diff --git a/src/TensorKit.jl b/src/TensorKit.jl index 87a3a2380..929dd804e 100644 --- a/src/TensorKit.jl +++ b/src/TensorKit.jl @@ -44,7 +44,7 @@ export infimum, supremum, isisomorphic, ismonomorphic, isepimorphic export sectortype, sectors, hassector export unit, rightunit, leftunit, allunits, isunit, otimes, deligneproduct, timereversed export Nsymbol, Fsymbol, Rsymbol, Bsymbol, frobenius_schur_phase, frobenius_schur_indicator, twist, fusiontensor -export sectorscalartype, fusionscalartype, braidingscalartype +export sectorscalartype, fusionscalartype, braidingscalartype, dimscalartype # Export methods for fusion trees export fusiontrees, braid, permute, transpose diff --git a/src/auxiliary/dicts.jl b/src/auxiliary/dicts.jl index 626599c01..dbce97087 100644 --- a/src/auxiliary/dicts.jl +++ b/src/auxiliary/dicts.jl @@ -89,21 +89,12 @@ end Base.empty(::SortedVectorDict, ::Type{K}, ::Type{V}) where {K, V} = SortedVectorDict{K, V}() Base.empty!(d::SortedVectorDict) = (empty!(d.keys); empty!(d.values); return d) -# _searchsortedfirst(v::Vector, k) = searchsortedfirst(v, k) -function _searchsortedfirst(v::Vector, k) - i = 1 - @inbounds while i <= length(v) && isless(v[i], k) - i += 1 - end - return i -end - function Base.delete!(d::SortedVectorDict{K}, k) where {K} key = convert(K, k) if !isequal(k, key) return d end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) if i <= length(d) && isequal(d.keys[i], key) deleteat!(d.keys, i) deleteat!(d.values, i) @@ -118,7 +109,7 @@ function Base.haskey(d::SortedVectorDict{K}, k) where {K} if !isequal(k, key) return false end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) return (i <= length(d) && isequal(d.keys[i], key)) end function Base.getindex(d::SortedVectorDict{K}, k) where {K} @@ -126,7 +117,7 @@ function Base.getindex(d::SortedVectorDict{K}, k) where {K} if !isequal(k, key) throw(KeyError(k)) end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) @inbounds if (i <= length(d) && isequal(d.keys[i], key)) return d.values[i] else @@ -138,7 +129,7 @@ function Base.setindex!(d::SortedVectorDict{K}, v, k) where {K} if !isequal(k, key) throw(ArgumentError("$k is not a valid key for type $K")) end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) if i <= length(d) && isequal(d.keys[i], key) d.values[i] = v else @@ -153,7 +144,7 @@ function Base.get(d::SortedVectorDict{K}, k, default) where {K} if !isequal(k, key) return default end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) @inbounds begin return (i <= length(d) && isequal(d.keys[i], key)) ? d.values[i] : default end @@ -163,7 +154,7 @@ function Base.get(f::Union{Function, Type}, d::SortedVectorDict{K}, k) where {K} if !isequal(k, key) return f() end - i = _searchsortedfirst(d.keys, key) + i = searchsortedfirst(d.keys, key) @inbounds begin return (i <= length(d) && isequal(d.keys[i], key)) ? d.values[i] : f() end @@ -186,6 +177,64 @@ function Base.:(==)(d1::SortedVectorDict, d2::SortedVectorDict) return true end +# merge over two SectorDicts +# the intersect case for infimum is kind of tricky, so there's an extra bool +# to indicate keeping keys that are only present in one of the two dicts +# zero results are dropped, matching how GradedSpace never stores an explicit zero dimension +function _sortedmerge( + combine, ::Val{keepunique}, d1::SortedVectorDict{K, V}, d2::SortedVectorDict{K, V} + ) where {keepunique, K, V} + k1, v1 = d1.keys, d1.values + k2, v2 = d2.keys, d2.values + n1, n2 = length(k1), length(k2) + ks, vs = Vector{K}(), Vector{V}() + sizehint!(ks, keepunique ? n1 + n2 : min(n1, n2)) + sizehint!(vs, keepunique ? n1 + n2 : min(n1, n2)) + i, j = 1, 1 + @inbounds while i <= n1 && j <= n2 + if k1[i] == k2[j] + d = combine(v1[i], v2[j]) + if !iszero(d) + push!(ks, k1[i]) + push!(vs, d) + end + i += 1 + j += 1 + elseif k1[i] < k2[j] + if keepunique + push!(ks, k1[i]) + push!(vs, v1[i]) + end + i += 1 + else + if keepunique + push!(ks, k2[j]) + push!(vs, v2[j]) + end + j += 1 + end + end + if keepunique + @inbounds while i <= n1 + push!(ks, k1[i]) + push!(vs, v1[i]) + i += 1 + end + @inbounds while j <= n2 + push!(ks, k2[j]) + push!(vs, v2[j]) + j += 1 + end + end + return SortedVectorDict{K, V}(ks, vs) +end + +Base.mergewith(combine, d1::SortedVectorDict{K, V}, d2::SortedVectorDict{K, V}) where {K, V} = + _sortedmerge(combine, Val(true), d1, d2) + +_sortedintersect(combine, d1::SortedVectorDict{K, V}, d2::SortedVectorDict{K, V}) where {K, V} = + _sortedmerge(combine, Val(false), d1, d2) + """ Hashed(value, hashfunction = Base.hash, isequal = Base.isequal) diff --git a/src/factorizations/factorizations.jl b/src/factorizations/factorizations.jl index fbe87a63a..b16a28e9b 100644 --- a/src/factorizations/factorizations.jl +++ b/src/factorizations/factorizations.jl @@ -6,9 +6,9 @@ module Factorizations export copy_oftype, factorisation_scalartype, one!, truncspace using ..TensorKit -using ..TensorKit: AdjointTensorMap, SectorDict, SectorVector, +using ..TensorKit: AdjointTensorMap, SectorDict, SectorVector, findindex, blocktype, foreachblock, one!, - similar_diagonal, similarstoragetype + similar_diagonal, similarstoragetype, sectorstoragetype using LinearAlgebra: LinearAlgebra, BlasFloat, Diagonal, svdvals, svdvals!, eigen, eigen!, diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index b910a184c..68ac57bb4 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -24,16 +24,22 @@ for pullback! in (:qr_null_pullback!, :lq_null_pullback!) return Δt end end -_notrunc_ind(t) = SectorDict(c => Colon() for c in blocksectors(t)) +function _notrunc_ind(t) + I = sectortype(t) + return _builddensemap(sectorstoragetype(I), I, blocks(t), Colon) do _, _ + Colon() + end +end for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) @eval function MAK.$pullback!( Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, inds = _notrunc_ind(t); kwargs... ) + Isec = sectortype(t) foreachblock(Δt, t) do c, (Δb, b) - haskey(inds, c) || return nothing - ind = inds[c] + ind = _denseget(inds, Isec, c) + isnothing(ind) && return nothing Fc = block.(F, Ref(c)) ΔFc = block.(ΔF, Ref(c)) MAK.$pullback!(Δb, b, Fc, ΔFc, ind; kwargs...) diff --git a/src/factorizations/truncation.jl b/src/factorizations/truncation.jl index 2facfd2e9..d2e455ebd 100644 --- a/src/factorizations/truncation.jl +++ b/src/factorizations/truncation.jl @@ -34,32 +34,101 @@ _blocklength(ax, ind) = length(ax[ind]) _blocklength(ax::Base.OneTo, ind::AbstractVector{<:Integer}) = length(ind) _blocklength(ax::Base.OneTo, ind::AbstractVector{Bool}) = count(ind) +# TODO: it quacks like a duck, just define a subtype of AbstractDict? +# represent the sector-index mapping as Vector{Union{Nothing, V}} where V is the type of the index +# mapping is indexed through findindex +# the type V is needed because the concrete type of ind depends on the strategy (except for intersect/union) +_densenew(::Type{I}, ::Type{V}) where {I <: Sector, V} = + Vector{Union{Nothing, V}}(nothing, length(values(I))) + +function _denseset!(v::Vector, ::Type{I}, c::I, val) where {I <: Sector} + v[findindex(values(I), c)] = val + return v +end +_denseget(v::Vector, ::Type{I}, c::I) where {I <: Sector} = v[findindex(values(I), c)] +function _densepairs(v::Vector, ::Type{I}) where {I <: Sector} + vals = values(I) + return (vals[i] => x for (i, x) in enumerate(v) if !isnothing(x)) +end +_densekeys(v::Vector, ::Type{I}) where {I <: Sector} = (c for (c, _) in _densepairs(v, I)) + +# fallbacks to catch SectorVector/SectorDict, even for NTuple sectorstoragetype +_denseget(v, ::Type{I}, c::I) where {I <: Sector} = get(v, c, nothing) +_densekeys(v, ::Type{I}) where {I <: Sector} = keys(v) +_densepairs(v, ::Type{I}) where {I <: Sector} = pairs(v) + +# builds either a dense Vector or SectorDict based on sectorstoragetype +# mapping each (c, v) pair's sector c to f(c, v) +# so every `findtruncated` method shares one output-construction path +# `pairsiter` are c => v pairs, can be c => nothing for NoTruncation/TruncationIntersection/TruncationUnion +function _builddensemap(f, ::Type{D}, ::Type{I}, pairsiter, ::Type{V}) where {D <: Tuple, I <: Sector, V} + d = _densenew(I, V) + for (c, v) in pairsiter + _denseset!(d, I, c, f(c, v)) + end + return d +end +function _builddensemap(f, ::Type{D}, ::Type{I}, pairsiter, ::Type{V}) where {D <: SectorDict, I <: Sector, V} + return SectorDict(c => f(c, v) for (c, v) in pairsiter) # V unused +end + function truncate_space(V::ElementarySpace, inds) - return spacetype(V)(c => _blocklength(dim(V, c), ind) for (c, ind) in pairs(inds)) + @assert !isdual(V) + I = sectortype(V) + @assert I == Trivial + return spacetype(V)(c => _blocklength(dim(V, c), ind) for (c, ind) in _densepairs(inds, I)) +end +function truncate_space(V::GradedSpace{I, NTuple{N, Int}}, inds) where {I <: Sector, N} + @assert !isdual(V) + vals = values(I) + newdims = zeros(Int, N) + for (c, ind) in _densepairs(inds, I) + d = dim(V, c) + n_write = findindex(vals, c) + newdims[n_write] = _blocklength(d, ind) + end + return typeof(V)(NTuple{N, Int}(newdims), false) +end +function truncate_space(V::GradedSpace{I, <:SectorDict}, inds) where {I <: Sector} + @assert !isdual(V) + ks, vs = Vector{I}(), Vector{Int}() # accumulate and sort once at the end + for (c, ind) in pairs(inds) + d = dim(V, c) + len = _blocklength(d, ind) + if !iszero(len) + push!(ks, c) + push!(vs, len) + end + end + perm = sortperm(ks) + return typeof(V)(SectorDict{I, Int}(ks[perm], vs[perm]), false) end function truncate_domain!(tdst::AbstractTensorMap, tsrc::AbstractTensorMap, inds) + Isec = sectortype(tdst) for (c, b) in blocks(tdst) - I = get(inds, c, nothing) - @assert !isnothing(I) + I = _denseget(inds, Isec, c) + @assert !isnothing(I) # kept for safety, but should be guaranteed by _densepairs b′ = block(tsrc, c) b .= view(b′, :, I) end return tdst end function truncate_codomain!(tdst::AbstractTensorMap, tsrc::AbstractTensorMap, inds) + Isec = sectortype(tdst) for (c, b) in blocks(tdst) - I = get(inds, c, nothing) - @assert !isnothing(I) + I = _denseget(inds, Isec, c) + @assert !isnothing(I) # kept for safety, but should be guaranteed by _densepairs b′ = block(tsrc, c) b .= view(b′, I, :) end return tdst end function truncate_diagonal!(Ddst::DiagonalTensorMap, Dsrc::DiagonalTensorMap, inds) + Isec = sectortype(Ddst) for (c, b) in blocks(Ddst) - I = get(inds, c, nothing) - @assert !isnothing(I) + I = _denseget(inds, Isec, c) + @assert !isnothing(I) # kept for safety, but should be guaranteed by _densepairs diagview(b) .= view(diagview(block(Dsrc, c)), I) end return Ddst @@ -114,7 +183,10 @@ end function MAK.truncate( ::typeof(left_null!), (U, S)::NTuple{2, AbstractTensorMap}, strategy::NoTruncation ) - ind = SectorDict(c => (size(b, 2) + 1):size(b, 1) for (c, b) in blocks(S)) + I = sectortype(S) + ind = _builddensemap(sectorstoragetype(I), I, blocks(S), UnitRange{Int}) do _, b + (size(b, 2) + 1):size(b, 1) + end V_truncated = truncate_space(space(S, 1), ind) Ũ = similar(U, codomain(U) ← V_truncated) truncate_domain!(Ũ, U, ind) @@ -123,7 +195,10 @@ end function MAK.truncate( ::typeof(right_null!), (S, Vᴴ)::NTuple{2, AbstractTensorMap}, strategy::NoTruncation ) - ind = SectorDict(c => (size(b, 1) + 1):size(b, 2) for (c, b) in blocks(S)) + I = sectortype(S) + ind = _builddensemap(sectorstoragetype(I), I, blocks(S), UnitRange{Int}) do _, b + (size(b, 1) + 1):size(b, 2) + end V_truncated = truncate_space(dual(space(S, 2)), ind) Ṽᴴ = similar(Vᴴ, V_truncated ← domain(Vᴴ)) truncate_codomain!(Ṽᴴ, Vᴴ, ind) @@ -160,7 +235,10 @@ function MAK.findtruncated_svd(values::SectorVector, strategy::TruncationStrateg end function MAK.findtruncated(values::SectorVector, ::NoTruncation) - return SectorDict(c => Colon() for c in keys(values)) + I = sectortype(values) + return _builddensemap(sectorstoragetype(I), I, (c => nothing for c in keys(values)), Colon) do _, _ + Colon() + end end # Need to select the first k values here after sorting across blocks, weighted by quantum dimension @@ -206,18 +284,29 @@ MAK.findtruncated_svd(values::SectorVector, strategy::TruncationByOrder) = MAK.findtruncated(values, strategy) function MAK.findtruncated(values::SectorVector, strategy::TruncationByFilter) - return SectorDict(c => findall(strategy.filter, d) for (c, d) in pairs(values)) + I = sectortype(values) + return _builddensemap(sectorstoragetype(I), I, pairs(values), Vector{Int}) do _, v + findall(strategy.filter, v) + end end function MAK.findtruncated(values::SectorVector, strategy::TruncationByValue) + I = sectortype(values) atol = rtol_to_atol(values, strategy.p, strategy.atol, strategy.rtol) strategy′ = trunctol(; atol, strategy.by, strategy.keep_below) - return SectorDict(c => MAK.findtruncated(d, strategy′) for (c, d) in pairs(values)) + V = Base.promote_op(MAK.findtruncated, valtype(values), typeof(strategy′)) + return _builddensemap(sectorstoragetype(I), I, pairs(values), V) do _, v + MAK.findtruncated(v, strategy′) + end end function MAK.findtruncated_svd(values::SectorVector, strategy::TruncationByValue) + I = sectortype(values) atol = rtol_to_atol(values, strategy.p, strategy.atol, strategy.rtol) strategy′ = trunctol(; atol, strategy.by, strategy.keep_below) - return SectorDict(c => MAK.findtruncated_svd(d, strategy′) for (c, d) in pairs(values)) + V = Base.promote_op(MAK.findtruncated_svd, valtype(values), typeof(strategy′)) + return _builddensemap(sectorstoragetype(I), I, pairs(values), V) do _, v + MAK.findtruncated_svd(v, strategy′) + end end # Need to select the first k values here after sorting by error across blocks, @@ -259,14 +348,24 @@ MAK.findtruncated_svd(values::SectorVector, strategy::TruncationByError) = MAK.findtruncated(values, strategy) function MAK.findtruncated(values::SectorVector, strategy::TruncationSpace) - sectortype(values) == sectortype(strategy) || throw(SectorMismatch("sectortype of truncation strategy does not match values")) + I = sectortype(values) + I == sectortype(strategy) || throw(SectorMismatch("sectortype of truncation strategy does not match values")) blockstrategy(c) = truncrank(dim(strategy.space, c); strategy.by, strategy.rev) - return SectorDict(c => MAK.findtruncated(d, blockstrategy(c)) for (c, d) in pairs(values)) + Vstrategy = Base.promote_op(blockstrategy, I) + V = Base.promote_op(MAK.findtruncated, valtype(values), Vstrategy) + return _builddensemap(sectorstoragetype(I), I, pairs(values), V) do c, v + MAK.findtruncated(v, blockstrategy(c)) + end end function MAK.findtruncated_svd(values::SectorVector, strategy::TruncationSpace) - sectortype(values) == sectortype(strategy) || throw(SectorMismatch("sectortype of truncation strategy does not match values")) + I = sectortype(values) + I == sectortype(strategy) || throw(SectorMismatch("sectortype of truncation strategy does not match values")) blockstrategy(c) = truncrank(dim(strategy.space, c); strategy.by, strategy.rev) - return SectorDict(c => MAK.findtruncated_svd(d, blockstrategy(c)) for (c, d) in pairs(values)) + Vstrategy = Base.promote_op(blockstrategy, I) + V = Base.promote_op(MAK.findtruncated_svd, valtype(values), Vstrategy) + return _builddensemap(sectorstoragetype(I), I, pairs(values), V) do c, v + MAK.findtruncated_svd(v, blockstrategy(c)) + end end # The implementations below assume that the `SectorDict` always contains an entry for every block sector @@ -274,40 +373,40 @@ end # This is always the case in the implementations above. function MAK.findtruncated(values::SectorVector, strategy::TruncationIntersection) + I = sectortype(values) inds = map(Base.Fix1(MAK.findtruncated, values), strategy.components) - @assert TensorKit._allequal(keys, inds) "missing blocks are not supported right now" - sectors = keys(first(inds)) - vals = map(keys(first(inds))) do c - mapreduce(Base.Fix2(getindex, c), MatrixAlgebraKit._ind_intersect, inds) + @assert TensorKit._allequal(v -> collect(_densekeys(v, I)), inds) "missing blocks are not supported right now" + sectors = collect(_densekeys(first(inds), I)) + return _builddensemap(sectorstoragetype(I), I, (c => nothing for c in sectors), Any) do c, _ + mapreduce(v -> _denseget(v, I, c), MatrixAlgebraKit._ind_intersect, inds) end - return SectorDict{eltype(sectors), eltype(vals)}(sectors, vals) end function MAK.findtruncated_svd(values::SectorVector, strategy::TruncationIntersection) + I = sectortype(values) inds = map(Base.Fix1(MAK.findtruncated_svd, values), strategy.components) - @assert TensorKit._allequal(keys, inds) "missing blocks are not supported right now" - sectors = keys(first(inds)) - vals = map(keys(first(inds))) do c - mapreduce(Base.Fix2(getindex, c), MatrixAlgebraKit._ind_intersect, inds) + @assert TensorKit._allequal(v -> collect(_densekeys(v, I)), inds) "missing blocks are not supported right now" + sectors = collect(_densekeys(first(inds), I)) + return _builddensemap(sectorstoragetype(I), I, (c => nothing for c in sectors), Any) do c, _ + mapreduce(v -> _denseget(v, I, c), MatrixAlgebraKit._ind_intersect, inds) end - return SectorDict{eltype(sectors), eltype(vals)}(sectors, vals) end function MAK.findtruncated(values::SectorVector, strategy::TruncationUnion) + I = sectortype(values) inds = map(Base.Fix1(MAK.findtruncated, values), strategy.components) - @assert TensorKit._allequal(keys, inds) "missing blocks are not supported right now" - sectors = keys(first(inds)) - vals = map(keys(first(inds))) do c - mapreduce(Base.Fix2(getindex, c), MatrixAlgebraKit._ind_union, inds) + @assert TensorKit._allequal(v -> collect(_densekeys(v, I)), inds) "missing blocks are not supported right now" + sectors = collect(_densekeys(first(inds), I)) + return _builddensemap(sectorstoragetype(I), I, (c => nothing for c in sectors), Any) do c, _ + mapreduce(v -> _denseget(v, I, c), MatrixAlgebraKit._ind_union, inds) end - return SectorDict{eltype(sectors), eltype(vals)}(sectors, vals) end function MAK.findtruncated_svd(values::SectorVector, strategy::TruncationUnion) + I = sectortype(values) inds = map(Base.Fix1(MAK.findtruncated_svd, values), strategy.components) - @assert TensorKit._allequal(keys, inds) "missing blocks are not supported right now" - sectors = keys(first(inds)) - vals = map(keys(first(inds))) do c - mapreduce(Base.Fix2(getindex, c), MatrixAlgebraKit._ind_union, inds) + @assert TensorKit._allequal(v -> collect(_densekeys(v, I)), inds) "missing blocks are not supported right now" + sectors = collect(_densekeys(first(inds), I)) + return _builddensemap(sectorstoragetype(I), I, (c => nothing for c in sectors), Any) do c, _ + mapreduce(v -> _denseget(v, I, c), MatrixAlgebraKit._ind_union, inds) end - return SectorDict{eltype(sectors), eltype(vals)}(sectors, vals) end # Truncation error @@ -315,7 +414,8 @@ end MAK.truncation_error(values::SectorVector, ind) = MAK.truncation_error!(copy(values), ind) function MAK.truncation_error!(values::SectorVector, ind) - for (c, ind_c) in pairs(ind) + Isec = sectortype(values) + for (c, ind_c) in _densepairs(ind, Isec) v = values[c] v[ind_c] .= zero(eltype(v)) end diff --git a/src/spaces/gradedspace.jl b/src/spaces/gradedspace.jl index 53f48dafd..facb63811 100644 --- a/src/spaces/gradedspace.jl +++ b/src/spaces/gradedspace.jl @@ -30,17 +30,17 @@ end sectortype(::Type{<:GradedSpace{I}}) where {I <: Sector} = I function GradedSpace{I, NTuple{N, Int}}(dims; dual::Bool = false) where {I, N} - d = ntuple(n -> 0, N) - isset = ntuple(n -> false, N) + d = zeros(Int, N) + isset = falses(N) # see if this is still needed if we're restricting to small N for (c, dc) in dims k = convert(I, c) i = findindex(values(I), k) - k = dc < 0 && throw(ArgumentError(lazy"Sector $k has negative dimension $dc")) + dc < 0 && throw(ArgumentError(lazy"Sector $k has negative dimension $dc")) isset[i] && throw(ArgumentError(lazy"Sector $c appears multiple times")) - isset = TupleTools.setindex(isset, true, i) - d = TupleTools.setindex(d, dc, i) + isset[i] = true + d[i] = dc end - return GradedSpace{I, NTuple{N, Int}}(d, dual) + return GradedSpace{I, NTuple{N, Int}}(NTuple{N, Int}(d), dual) end function GradedSpace{I, NTuple{N, Int}}(dims::Pair; dual::Bool = false) where {I, N} return GradedSpace{I, NTuple{N, Int}}((dims,); dual = dual) @@ -89,9 +89,9 @@ GradedSpace(g::AbstractDict; dual::Bool = false) = GradedSpace(g...; dual = dual field(::Type{<:GradedSpace}) = ℂ InnerProductStyle(::Type{<:GradedSpace}) = EuclideanInnerProduct() -function dim(V::GradedSpace) - init = 0 * dim(first(allunits(sectortype(V)))) - return sum(c -> dim(c) * dim(V, c), sectors(V); init = init) +function dim(V::GradedSpace{I}) where {I <: Sector} + init = zero(dimscalartype(I)) + return sum(((c, d),) -> dim(c) * d, blockdims(V); init) end function dim(V::GradedSpace{I, <:AbstractDict}, c::I) where {I <: Sector} return get(V.dims, isdual(V) ? dual(c) : c, 0) @@ -102,9 +102,9 @@ end Base.axes(V::GradedSpace) = Base.OneTo(dim(V)) function Base.axes(V::GradedSpace{I}, c::I) where {I <: Sector} offset = 0 - for c′ in sectors(V) + for (c′, d′) in blockdims(V) c′ == c && break - offset += dim(c′) * dim(V, c′) + offset += dim(c′) * d′ end return (offset + 1):(offset + dim(c) * dim(V, c)) end @@ -115,9 +115,9 @@ isdual(V::GradedSpace) = V.dual isconj(V::GradedSpace) = isdual(V) function flip(V::GradedSpace{I}) where {I <: Sector} return if isdual(V) - typeof(V)(c => dim(V, c) for c in sectors(V)) + typeof(V)(blockdims(V)) else - typeof(V)(dual(c) => dim(V, c) for c in sectors(V))' + typeof(V)(dual(c) => d for (c, d) in blockdims(V))' end end @@ -126,54 +126,88 @@ function unitspace(S::Type{<:GradedSpace{I}}) where {I <: Sector} end zerospace(S::Type{<:GradedSpace}) = S() -# TODO: the following methods can probably be implemented more efficiently for -# `FiniteGradedSpace`, but we don't expect them to be used often in hot loops, so -# these generic definitions (which are still quite efficient) are good for now. -function ⊕(V₁::GradedSpace{I}, V₂::GradedSpace{I}) where {I <: Sector} +function ⊕(V₁::GradedSpace{I, <:SectorDict}, V₂::GradedSpace{I, <:SectorDict}) where {I <: Sector} + dual1 = isdual(V₁) + dual1 == isdual(V₂) || throw(SpaceMismatch("Direct sum of a vector space and a dual space does not exist")) + return typeof(V₁)(mergewith(+, V₁.dims, V₂.dims), dual1) +end +function ⊕(V₁::GradedSpace{I, <:Tuple}, V₂::GradedSpace{I, <:Tuple}) where {I <: Sector} dual1 = isdual(V₁) dual1 == isdual(V₂) || throw(SpaceMismatch("Direct sum of a vector space and a dual space does not exist")) - dims = SectorDict{I, Int}() - for c in union(sectors(V₁), sectors(V₂)) - cout = ifelse(dual1, dual(c), c) - dims[cout] = dim(V₁, c) + dim(V₂, c) - end - return typeof(V₁)(dims; dual = dual1) + newdims = map(+, V₁.dims, V₂.dims) + return typeof(V₁)(newdims, dual1) end -function ⊖(V::GradedSpace{I}, W::GradedSpace{I}) where {I <: Sector} - dual = isdual(V) - V ≿ W && dual == isdual(W) || - throw(SpaceMismatch("$(W) is not a subspace of $(V)")) - return typeof(V)(c => dim(V, c) - dim(W, c) for c in sectors(V); dual) +function ⊖(V::GradedSpace{I, <:Tuple}, W::GradedSpace{I, <:Tuple}) where {I <: Sector} + dualV = isdual(V) + V ≿ W && dualV == isdual(W) || throw(SpaceMismatch("$(W) is not a subspace of $(V)")) + newdims = map(-, V.dims, W.dims) + return typeof(V)(newdims, dualV) +end +function ⊖(V::GradedSpace{I, <:SectorDict}, W::GradedSpace{I, <:SectorDict}) where {I <: Sector} + dualV = isdual(V) + V ≿ W && dualV == isdual(W) || throw(SpaceMismatch("$(W) is not a subspace of $(V)")) + return typeof(V)(mergewith(-, V.dims, W.dims), dualV) end -function fuse(V₁::GradedSpace{I}, V₂::GradedSpace{I}) where {I <: Sector} - dims = SectorDict{I, Int}() - for a in sectors(V₁), b in sectors(V₂) +function fuse(V₁::GradedSpace{I, <:SectorDict}, V₂::GradedSpace{I, <:SectorDict}) where {I <: Sector} + acc = Dict{I, Int}() # SectorDict `get` within the double for loop accumulates O(N^2) `findindex` calls -> sort afterwards + for (a, da) in blockdims(V₁), (b, db) in blockdims(V₂) + dab = da * db for c in a ⊗ b - dims[c] = get(dims, c, 0) + Nsymbol(a, b, c) * dim(V₁, a) * dim(V₂, b) + acc[c] = get(acc, c, 0) + Nsymbol(a, b, c) * dab + end + end + ks0 = collect(keys(acc)) + vs0 = collect(values(acc)) + perm = sortperm(ks0) + return typeof(V₁)(SectorDict{I, Int}(ks0[perm], vs0[perm]), false) +end +function fuse(V₁::GradedSpace{I, NTuple{N, Int}}, V₂::GradedSpace{I, NTuple{N, Int}}) where {I <: Sector, N} + vals = values(I) + dual1, dual2 = isdual(V₁), isdual(V₂) + newdims = zeros(Int, N) + @inbounds for na in 1:N + da = V₁.dims[na] + iszero(da) && continue + a₀ = vals[na] # avoid call to sectors(V₁) + a = dual1 ? dual(a₀) : a₀ + for nb in 1:N + db = V₂.dims[nb] + iszero(db) && continue + b₀ = vals[nb] # idem for V₂ + b = dual2 ? dual(b₀) : b₀ + dab = da * db + for c in a ⊗ b + nc = findindex(vals, c) + newdims[nc] += Nsymbol(a, b, c) * dab + end end end - return typeof(V₁)(dims) + return typeof(V₁)(NTuple{N, Int}(newdims), false) end -function infimum(V₁::GradedSpace{I}, V₂::GradedSpace{I}) where {I <: Sector} +function infimum(V₁::GradedSpace{I, <:Tuple}, V₂::GradedSpace{I, <:Tuple}) where {I <: Sector} Visdual = isdual(V₁) - Visdual == isdual(V₂) || - throw(SpaceMismatch("Infimum of space and dual space does not exist")) - return typeof(V₁)( - (Visdual ? dual(c) : c) => min(dim(V₁, c), dim(V₂, c)) - for c in intersect(sectors(V₁), sectors(V₂)); dual = Visdual - ) + Visdual == isdual(V₂) || throw(SpaceMismatch("Infimum of space and dual space does not exist")) + newdims = map(min, V₁.dims, V₂.dims) + return typeof(V₁)(newdims, Visdual) end -function supremum(V₁::GradedSpace{I}, V₂::GradedSpace{I}) where {I <: Sector} +function infimum(V₁::GradedSpace{I, <:SectorDict}, V₂::GradedSpace{I, <:SectorDict}) where {I <: Sector} Visdual = isdual(V₁) - Visdual == isdual(V₂) || - throw(SpaceMismatch("Supremum of space and dual space does not exist")) - return typeof(V₁)( - (Visdual ? dual(c) : c) => max(dim(V₁, c), dim(V₂, c)) - for c in union(sectors(V₁), sectors(V₂)); dual = Visdual - ) + Visdual == isdual(V₂) || throw(SpaceMismatch("Infimum of space and dual space does not exist")) + return typeof(V₁)(_sortedintersect(min, V₁.dims, V₂.dims), Visdual) +end +function supremum(V₁::GradedSpace{I, <:Tuple}, V₂::GradedSpace{I, <:Tuple}) where {I <: Sector} + Visdual = isdual(V₁) + Visdual == isdual(V₂) || throw(SpaceMismatch("Supremum of space and dual space does not exist")) + newdims = map(max, V₁.dims, V₂.dims) + return typeof(V₁)(newdims, Visdual) +end +function supremum(V₁::GradedSpace{I, <:SectorDict}, V₂::GradedSpace{I, <:SectorDict}) where {I <: Sector} + Visdual = isdual(V₁) + Visdual == isdual(V₂) || throw(SpaceMismatch("Supremum of space and dual space does not exist")) + return typeof(V₁)(mergewith(max, V₁.dims, V₂.dims), Visdual) end hassector(V::GradedSpace{I}, s::I) where {I <: Sector} = dim(V, s) != 0 @@ -186,6 +220,24 @@ function sectors(V::GradedSpace{I, NTuple{N, Int}}) where {I <: Sector, N} end end +""" + blockdims(V::GradedSpace) + +Return an iterator over the non-zero blocks of the graded space `V`. +These blocks contain the `Sector`s and their corresponding degeneracy, +i.e. the number of times the sector appears in the direct sum decomposition of `V`. +""" +function blockdims(V::GradedSpace{I, <:AbstractDict}) where {I <: Sector} + return ((isdual(V) ? dual(c) : c) => d for (c, d) in V.dims) +end +function blockdims(V::GradedSpace{I, NTuple{N, Int}}) where {I <: Sector, N} + vals = values(I) + return ( + (isdual(V) ? dual(vals[n]) : vals[n]) => V.dims[n] + for n in 1:N if !iszero(V.dims[n]) + ) +end + Base.hash(V::GradedSpace, h::UInt) = hash(V.dual, hash(V.dims, h)) function Base.:(==)(V₁::GradedSpace, V₂::GradedSpace) return sectortype(V₁) == sectortype(V₂) && (V₁.dims == V₂.dims) && V₁.dual == V₂.dual @@ -218,7 +270,7 @@ function Base.show(io::IO, V::GradedSpace) cls = ")" end - v = [c => dim(V, c) for c in sectors(V)] + v = collect(blockdims(V)) # logic stolen from Base.show_vector limited = get(io, :limit, false)::Bool @@ -249,7 +301,7 @@ function Base.show(io::IO, ::MIME"text/plain", V::GradedSpace) # print detailed sector information - hijack Base.Vector printing print(io, ":\n") isdual(V) && (V = dual(V)) - print_data = [c => dim(V, c) for c in sectors(V)] + print_data = collect(blockdims(V)) ioc = IOContext(io, :typeinfo => eltype(print_data)) Base.print_matrix(ioc, print_data) @@ -267,13 +319,24 @@ specify `D`. const Vect = SpaceTable() Base.getindex(::SpaceTable) = ComplexSpace Base.getindex(::SpaceTable, ::Type{Trivial}) = ComplexSpace -function Base.getindex(::SpaceTable, I::Type{<:Sector}) +Base.getindex(::SpaceTable, I::Type{<:Sector}) = GradedSpace{I, sectorstoragetype(I)} + +# based on Julia tuple unrolling range +const _ntuple_storage_threshold = 32 + +""" + sectorstoragetype(I::Type{<:Sector}) -> Type + +The storage type `D` used for the `dims` field of `GradedSpace{I, D}`. +This is `NTuple{N,Int}` with `N = length(values(I))` if `I` has a finite, known length +of at most `$_ntuple_storage_threshold`, or `SectorDict{I,Int}` otherwise. +""" +Base.@assume_effects :foldable function sectorstoragetype(::Type{I}) where {I <: Sector} if Base.IteratorSize(values(I)) isa Union{HasLength, HasShape} N = length(values(I)) - return GradedSpace{I, NTuple{N, Int}} - else - return GradedSpace{I, SectorDict{I, Int}} + N <= _ntuple_storage_threshold && return NTuple{N, Int} end + return SectorDict{I, Int} end Base.getindex(::ComplexNumbers, I::Type{<:Sector}) = Vect[I]