Skip to content

Commit 48a30c3

Browse files
committed
clean up trivial symmetry bypassing overhead
1 parent 53c71e7 commit 48a30c3

2 files changed

Lines changed: 42 additions & 55 deletions

File tree

src/tensors/indexmanipulations.jl

Lines changed: 41 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -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
@@ -396,15 +384,8 @@ 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
410391
end
@@ -604,19 +585,38 @@ function unwrap_adjoints(tdst, tsrc, p::Index2Tuple, levels, conjsrc::Bool, α,
604585
return (tdst′, tsrc′, p″, levels′, conjsrc″, α′, β′)
605586
end
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)
614610
end
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)
620620
end
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
656650
end
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.

src/tensors/tensoroperations.jl

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ has_array_view(t) = has_array_view(typeof(t))
3939
has_array_view(::Type) = false
4040
has_array_view(::Type{T}) where {T <: TensorMap} = sectortype(T) === Trivial
4141
has_array_view(::Type{T}) where {T <: AdjointTensorMap} = has_array_view(parenttype(T))
42+
has_array_view(t, ts...) = has_array_view(t) && has_array_view(ts...)
4243

4344
# tensoradd!
4445
function TO.tensoradd!(
@@ -47,10 +48,6 @@ function TO.tensoradd!(
4748
α::Number, β::Number,
4849
backend, allocator
4950
)
50-
if has_array_view(C) && has_array_view(A)
51-
TO.tensoradd!(C[], A[], pA, conjA, α, β, backend, allocator)
52-
return C
53-
end
5451
tdst, tsrc, p, _, conjA′, α′, β′ = unwrap_adjoints(C, A, _canonicalize(pA, C), nothing, conjA, α, β)
5552
_braid!(tdst, tsrc, p, conjA′, allind(tsrc), α′, β′, backend, allocator)
5653
return C

0 commit comments

Comments
 (0)