Repository navigation
Speed up getindex and findindex for ProductSector SectorValues
#106
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
| 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 | ||
|
|
@@ -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 | ||
|
|
@@ -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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Do you need both here? Isn't the manhattan index
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
| # 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 | ||
There was a problem hiding this comment.
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
binomialfunction overflows also the recursive function would?There was a problem hiding this comment.
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
factorialso I thought I should catch it, but clearly I didn't think enough about this being useful or not.