@@ -221,12 +221,6 @@ See also [`permute`](@ref) for creating a new tensor.
221221 )
222222 @boundscheck spacecheck_transform (permute, tdst, tsrc, p)
223223 @timeit_debug GLOBAL_TIMER " permute!/braid!" begin
224- if has_array_view (tdst) && has_array_view (tsrc)
225- @timeit_debug GLOBAL_TIMER " dense: tensoradd" TO. tensoradd! (
226- tdst[], tsrc[], p, false , α, β, backend, allocator
227- )
228- return tdst
229- end
230224 tdst′, tsrc′, p′, _, conjsrc, α′, β′ = unwrap_adjoints (tdst, tsrc, p, nothing , false , α, β)
231225 @inbounds _braid! (tdst′, tsrc′, p′, conjsrc, allind (tsrc′), α′, β′, backend, allocator)
232226 end
@@ -317,12 +311,6 @@ See also [`braid`](@ref) for creating a new tensor.
317311 )
318312 @boundscheck spacecheck_transform (braid, tdst, tsrc, p, levels)
319313 @timeit_debug GLOBAL_TIMER " permute!/braid!" begin
320- if has_array_view (tdst) && has_array_view (tsrc)
321- @timeit_debug GLOBAL_TIMER " dense: tensoradd" TO. tensoradd! (
322- tdst[], tsrc[], p, false , α, β, backend, allocator
323- )
324- return tdst
325- end
326314 tdst′, tsrc′, p′, levels′, conjsrc, α′, β′ = unwrap_adjoints (tdst, tsrc, p, levels, false , α, β)
327315 @inbounds _braid! (tdst′, tsrc′, p′, conjsrc, levels′, α′, β′, backend, allocator)
328316 end
396384 )
397385 @boundscheck spacecheck_transform (transpose, tdst, tsrc, p)
398386 @timeit_debug GLOBAL_TIMER " transpose!" begin
399- if has_array_view (tdst) && has_array_view (tsrc)
400- @timeit_debug GLOBAL_TIMER " dense: tensoradd" TO. tensoradd! (
401- tdst[], tsrc[], p, false , α, β, backend, allocator
402- )
403- return tdst
404- end
405387 tdst′, tsrc′, p′, _, conjsrc, α′, β′ = unwrap_adjoints (tdst, tsrc, p, nothing , false , α, β)
406- transformer = treetransposer (tdst′, tsrc′, p′, conjsrc)
407- @inbounds add_transform! (tdst′, tsrc′, p′, conjsrc, transformer, α′, β′, backend, allocator)
388+ @inbounds _transpose! (tdst′, tsrc′, p′, conjsrc, α′, β′, backend, allocator)
408389 end
409390 return tdst
410391end
@@ -604,19 +585,38 @@ function unwrap_adjoints(tdst, tsrc, p::Index2Tuple, levels, conjsrc::Bool, α,
604585 return (tdst′, tsrc′, p″, levels′, conjsrc″, α′, β′)
605586end
606587
607- # shared by `permute!`, `braid!` and `TO.tensoradd!` after the adjoints have been unwrapped
588+ # dense transform that bypasses overhead
589+ function _dense_transform! (tdst, tsrc, p:: Index2Tuple , conjsrc:: Bool , α, β, backend, allocator)
590+ p2 = (linearize (p), ()) # only the linear permutation matters for the array kernels
591+ @timeit_debug GLOBAL_TIMER " dense: tensoradd" TO. tensoradd! (
592+ tdst[], tsrc[], p2, conjsrc, α, β, backend, allocator
593+ )
594+ return tdst
595+ end
596+
597+ # space check for `tdst = permutedims(conjsrc ? conj(tsrc) : tsrc, p)`
598+ function spacecheck_transform (f, tdst:: AbstractTensorMap , tsrc:: AbstractTensorMap , p:: Index2Tuple , conjsrc:: Bool )
599+ Vsrc′, p′ = transform_source (space (tsrc), p, conjsrc)
600+ return spacecheck_transform (f, space (tdst), Vsrc′, p′)
601+ end
602+
608603@propagate_inbounds function _braid! (
609604 tdst, tsrc, p:: Index2Tuple , conjsrc:: Bool , levels:: IndexTuple , α, β, backend, allocator
610605 )
611606 @boundscheck spacecheck_transform (permute, tdst, tsrc, p, conjsrc)
607+ has_array_view (tdst, tsrc) && return _dense_transform! (tdst, tsrc, p, conjsrc, α, β, backend, allocator)
612608 transformer = treebraider (tdst, tsrc, p, conjsrc, levels)
613609 return @inbounds add_transform! (tdst, tsrc, p, conjsrc, transformer, α, β, backend, allocator)
614610end
615611
616- # space check for `tdst = permutedims(conjsrc ? conj(tsrc) : tsrc, p)`
617- function spacecheck_transform (f, tdst:: AbstractTensorMap , tsrc:: AbstractTensorMap , p:: Index2Tuple , conjsrc:: Bool )
618- Vsrc′, p′ = transform_source (space (tsrc), p, conjsrc)
619- return spacecheck_transform (f, space (tdst), Vsrc′, p′)
612+ # counterpart of `_braid!` for `transpose!`; the cyclicity of `p` is checked by the caller
613+ @propagate_inbounds function _transpose! (
614+ tdst, tsrc, p:: Index2Tuple , conjsrc:: Bool , α, β, backend, allocator
615+ )
616+ @boundscheck spacecheck_transform (permute, tdst, tsrc, p, conjsrc)
617+ has_array_view (tdst, tsrc) && return _dense_transform! (tdst, tsrc, p, conjsrc, α, β, backend, allocator)
618+ transformer = treetransposer (tdst, tsrc, p, conjsrc)
619+ return @inbounds add_transform! (tdst, tsrc, p, conjsrc, transformer, α, β, backend, allocator)
620620end
621621
622622"""
@@ -635,40 +635,30 @@ of `tsrc`, using the fusion tree transformation encoded in `transformer` (see [`
635635 add! (tdst, tsrc, α, β)
636636 else
637637 p2 = (linearize (p), ()) # only the linear permutation matters for the array kernels
638- if has_array_view (tdst) && has_array_view (tsrc)
639- @timeit_debug GLOBAL_TIMER " dense: tensoradd" TO. tensoradd! (
640- tdst[], tsrc[], p2, conjsrc, α, β, backend, allocator
641- )
638+ ntasks = use_threaded_transform (tdst, transformer) ? get_num_transformer_threads () : 1
639+ # resolve the conjugation flag into the view type here, with a statically typed call per branch
640+ if conjsrc
641+ dst, src = _transform_subblocks (tdst, tsrc, transformer, conj)
642+ add_transform_kernel! (dst, src, p2, transformer, α, β, backend, allocator, ntasks)
642643 else
643- ntasks = use_threaded_transform (tdst, transformer) ? get_num_transformer_threads () : 1
644- # resolve the conjugation flag into the view type here, with a statically typed call per branch
645- if conjsrc
646- dst, src = _transform_subblocks (tdst, tsrc, transformer, conj)
647- add_transform_kernel! (dst, src, p2, transformer, α, β, backend, allocator, ntasks)
648- else
649- dst, src = _transform_subblocks (tdst, tsrc, transformer, identity)
650- add_transform_kernel! (dst, src, p2, transformer, α, β, backend, allocator, ntasks)
651- end
644+ dst, src = _transform_subblocks (tdst, tsrc, transformer, identity)
645+ add_transform_kernel! (dst, src, p2, transformer, α, β, backend, allocator, ntasks)
652646 end
653647 end
654648
655649 return tdst
656650end
657651
658652# TensorMaps address their flat data directly, other tensor types go through `subblock`
659- function _transform_subblocks (tdst:: TensorMap , tsrc:: TensorMap , transformer, op)
660- return StridedSubblocks (tdst, transformer. structure_dst), StridedSubblocks (tsrc, transformer. structure_src, op)
661- end
662- function _transform_subblocks (tdst:: AbstractTensorMap , tsrc:: AbstractTensorMap , transformer, op)
663- return TreeSubblocks (tdst), TreeSubblocks (tsrc, op)
664- end
665-
666- function use_threaded_transform (t:: TensorMap , transformer)
667- return get_num_transformer_threads () > 1 && length (t. data) > Strided. MINTHREADLENGTH
668- end
669- function use_threaded_transform (t:: AbstractTensorMap , transformer)
670- return get_num_transformer_threads () > 1 && dim (space (t)) > Strided. MINTHREADLENGTH
671- end
653+ _transform_subblocks (tdst:: TensorMap , tsrc:: TensorMap , transformer, op) =
654+ StridedSubblocks (tdst, transformer. structure_dst), StridedSubblocks (tsrc, transformer. structure_src, op)
655+ _transform_subblocks (tdst:: AbstractTensorMap , tsrc:: AbstractTensorMap , transformer, op) =
656+ TreeSubblocks (tdst), TreeSubblocks (tsrc, op)
657+
658+ use_threaded_transform (t:: TensorMap , transformer) =
659+ get_num_transformer_threads () > 1 && length (t. data) > Strided. MINTHREADLENGTH
660+ use_threaded_transform (t:: AbstractTensorMap , transformer) =
661+ get_num_transformer_threads () > 1 && dim (space (t)) > Strided. MINTHREADLENGTH
672662
673663# The kernel operates on the subblocks addressed by position, so that for `TensorMap`s this only
674664# depends on `numind`, `eltype` and the transformer data, not on the sectortype.
0 commit comments