Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/TensorKitSectors.jl
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ end

# imports
# -------
using Base: SizeUnknown, HasLength, IsInfinite
using Base: SizeUnknown, HasLength, HasShape, IsInfinite
using Base: HasEltype, EltypeUnknown
using Base.Iterators: product, filter
using Base: @assume_effects
Expand Down
62 changes: 59 additions & 3 deletions src/auxiliary.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,16 +19,31 @@ function _kron_promote(A₁, B₁, sz₁, sz₂)
end

# Manhattan based distance enumeration: I is supposed to be one-based index
# TODO: is there any way to make this faster?

# forward mapping from multidimensional to single Manhattan index
@inline function num_manhattan_points(d::Int, sz::Dims{1})
return Int(sz[1] > d)
end

# when sz[k] > d for every k, then x_k < sz[k] - 1 is never violated by any split of d units among N axes (one x_k is maximally d, and sz[k] - 1 ≥ d covers that)
# so the bound drops -> calculate ways to write d as an ordered sum of N nonnegative integers,
# which is the number of compositions of d into N parts
@inline function _compositions(d::Int, ::Val{N}) where {N}
try
return binomial(d + N - 1, N - 1)
catch e
e isa OverflowError || rethrow()
return nothing # fallback to recursion

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Did you hit this case somewhere specifically? I did not look at this in detail, but I somehow would have expected that if the binomial function overflows also the recursive function would?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah you're right, which really means that previously you could've potentially summed up incorrectly in the recursion, and now I'm catching it in a subset of cases. Although I think realistically this can't be reached, as this requires a crazy amount of product sectors, so maybe I can just remove this?

Btw I didn't hit this myself, it's something I read in the docstring of factorial so I thought I should catch it, but clearly I didn't think enough about this being useful or not.

end
end
@inline function num_manhattan_points(d::Int, sz::Dims{N}) where {N}
d == 0 && return 1
if N ≥ 3 && all(>(d), sz)
c = _compositions(d, Val(N))
isnothing(c) || return c
end
num = 0
for i in 1:min(sz[1], d + 1)
for i in 1:min(sz[1], d + 1) # for N = 2 the loop is already linear
num += num_manhattan_points(d - i + 1, Base.tail(sz))
end
return num
Expand Down Expand Up @@ -79,7 +94,7 @@ function manhattan_to_multidimensional_index(index::Int, sz::Dims{N}) where {N}
index == 1 && return ntuple(one, Val(N))
offset = 1
d = 1
while true
while true # reason for the boundscheck
currentlayer = num_manhattan_points(d, sz)
if index <= offset + currentlayer
break
Expand All @@ -92,3 +107,44 @@ function manhattan_to_multidimensional_index(index::Int, sz::Dims{N}) where {N}
index -= offset
return invertlocaloffset(d, index - 1, sz)
end

# tabulation of the mapping to avoid repeated calls to manhattan_to_multidimensional_index
# and to_manhattan_index
struct MTable{N}
multi::Vector{NTuple{N, Int}} # Manhattan index -> multi-index
lin::Array{Int, N} # multi-index -> Manhattan index
Comment on lines +114 to +115

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you need both here? Isn't the manhattan index 1:n, so as long as you store the multi index in sorted order by their linear index you could fold both into a single array, using the array index as the key.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I guess the tradeoffs occurring here as that you save a bit of memory removing lin (not much compared to multi, though), but looking for a multi-index given the Manhattan index requires a small binary search (which google tells me is log(n)). Is that worth it?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I also realized that I wasn't really following what was going on here anyways, and I think my suggestion is wrong and that would defeat the entire point, the ordering is important 🙃

end

function build_table(sz::Dims{N}) where {N}
n = prod(sz)
multi = Vector{NTuple{N, Int}}(undef, n)
k = 0
for J in CartesianIndices(sz)
k += 1
multi[k] = Tuple(J)
end
sort!(multi; by = J -> (sum(J), J)) # sort first by distance, matches isless on ProductSector
lin = Array{Int, N}(undef, sz)
for i in 1:n
lin[CartesianIndex(multi[i])] = i
end
return MTable{N}(multi, lin)
end

const TABLES = IdDict{DataType, Any}()
const TABLE_LOCK = ReentrantLock() # dictionaries aren't thread-safe for concurrent mutation
Comment on lines +134 to +135

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wonder if it might be worth it to try something similar to TensorKit's caches for this, using an LRU which is both threadsafe as well as avoids keeping too many entries around.

# within TensorKit no problem: OhMyThreads allows concurrent calls into e.g. sectors(V)

@inline function mtable(::Type{I}, sz::Dims{N}) where {I, N}
t = get(TABLES, I, nothing)
if t === nothing
t = @lock TABLE_LOCK get!(() -> build_table(sz), TABLES, I)
end
return t::MTable{N} # re-establish concrete typing for downstream
end

# refuse to tabulate absurdly large label sets
# GradedSpace NTuple storage is only used for finite sectors
# not sure what a reasonable limit is, might be edited to smaller value
# depending on findings of NTuple vs Dict performance tests
const MAXTABLE = 2^20
23 changes: 20 additions & 3 deletions src/product.jl
Original file line number Diff line number Diff line change
Expand Up @@ -39,12 +39,29 @@ end
function _size(::SectorValues{I}) where {I <: ProductSector}
return map(s -> _length(values(s)), _sectors(I))
end
function Base.getindex(P::SectorValues{I}, i::Int) where {I <: ProductSector}
inds = manhattan_to_multidimensional_index(i, _size(P))
# finite <=> the product iterator has a shape, which is exactly the condition
# under which Vect[I] uses NTuple storage (pre potential cutoff)
@inline function tabulate(P::SectorValues, sz::Dims)
return Base.IteratorSize(P) isa Union{HasLength, HasShape} && prod(sz) ≤ MAXTABLE
end
Base.@propagate_inbounds function Base.getindex(P::SectorValues{I}, i::Int) where {I <: ProductSector}
sz = _size(P)
@boundscheck checksectorindex(P, i)
inds = if tabulate(P, sz)
@inbounds mtable(I, sz).multi[i]
else
manhattan_to_multidimensional_index(i, sz)
end
return I(getindex.(values.(_sectors(I)), inds))
end
function findindex(P::SectorValues{I}, c::I) where {I <: ProductSector}
return to_manhattan_index(findindex.(values.(_sectors(I)), Tuple(c)), _size(P))
J = findindex.(values.(_sectors(I)), Tuple(c))
sz = _size(P)
return if tabulate(P, sz)
@inbounds mtable(I, sz).lin[CartesianIndex(J)]
else
to_manhattan_index(J, sz)
end
end

function Base.iterate(P::SectorValues{I}, i = 1) where {I <: ProductSector}
Expand Down
19 changes: 10 additions & 9 deletions src/sectors.jl
Original file line number Diff line number Diff line change
Expand Up @@ -74,16 +74,17 @@ Base.IteratorEltype(::Type{<:SectorValues}) = HasEltype()
Base.eltype(::Type{SectorValues{I}}) where {I <: Sector} = I
Base.values(::Type{I}) where {I <: Sector} = SectorValues{I}()

Base.@propagate_inbounds function Base.getindex(
v::SectorValues{I}, i::Int
) where {I <: Sector}
@boundscheck begin
if Base.IteratorSize(v) === HasLength()
1 ≤ i ≤ length(v) || throw(BoundsError(v, i))
else
1 ≤ i || throw(BoundsError(v, i))
end
@inline function checksectorindex(v::SectorValues, i::Integer)
if Base.IteratorSize(v) isa Union{HasLength, HasShape}
1 ≤ i ≤ length(v) || throw(BoundsError(v, i))
else
1 ≤ i || throw(BoundsError(v, i))
end
return nothing
end

Base.@propagate_inbounds function Base.getindex(v::SectorValues{I}, i::Int) where {I <: Sector}
@boundscheck checksectorindex(v, i)
for (j, c) in enumerate(v)
j == i && return c
end
Expand Down
Loading