Skip to content

Reuse stride orders during dimension ordering - #89

Merged
lkdvos merged 7 commits into
mainfrom
ld-layout-planning
Oct 8, 2026
Merged

lkdvos merged 7 commits into
mainfrom
ld-layout-planning

Conversation

@lkdvos

@lkdvos lkdvos commented Oct 7, 2026 •

Copy link
Copy Markdown
Member

Compute stride orders once during dimension ordering, then move singleton axes last after fusion instead of recomputing importance. Use stable TupleTools.sortperm(...; rev=true) for both permutations.

Blocking and kernel execution remain unchanged. Validated 1,491 layout/blocking comparisons and 508 focused CPU checks with 1 and 4 threads.

@lkdvos
lkdvos marked this pull request as draft October 7, 2026 14:14
@lkdvos
lkdvos force-pushed the ld-promoteshape-overhead branch from 3434f93 to 2b50702 Compare October 7, 2026 14:26
Base automatically changed from ld-promoteshape-overhead to main October 7, 2026 14:26
@lkdvos

lkdvos commented Oct 7, 2026 •

Copy link
Copy Markdown
Member Author

Note: still investigating this one, seems to be some slowdowns depending on inlining behavior, I'm seeing if I can get it faster across the board

Seems to be fixed now!

@lkdvos lkdvos changed the title Reuse stride orders when planning mapreduce loops Reuse stride orders during dimension ordering Oct 7, 2026
@lkdvos
lkdvos force-pushed the ld-layout-planning branch from 64c01e1 to 19034da Compare October 7, 2026 14:36
@lkdvos
lkdvos marked this pull request as ready for review October 7, 2026 17:22
@lkdvos
lkdvos requested a review from Jutho October 7, 2026 17:22
Comment thread src/mapreduce.jl Outdated
Comment thread src/mapreduce.jl Outdated
# Each stride rank gets enough bits to hold all array votes without carries.
# The output array gets two votes; each input gets one.
function _importance(dims::NTuple{N, Int}, stride_orders::NTuple{M, NTuple{N, Int}}) where {N, M}
bits_per_rank = 8 * sizeof(Int) - leading_zeros(M + 1)

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.

Can you add the comment that this value is identical to ceil(Int, log2(M+2))

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.

Also, did you ask Claude whether this is really the best importance rating? There is very little justification for all of this. Or maybe there is a better sorting strategy all together?

@Jutho

Jutho commented Oct 7, 2026

Copy link
Copy Markdown
Member

I think this PR does not actually change the loop order, only the performance in computing it, so I of course have no objections. I do however wonder about the performance gain of the alternative sorting? The second change, i.e. doing a second permutation solely based on active, is clearly an improvement.

@lkdvos

lkdvos commented Oct 7, 2026

Copy link
Copy Markdown
Member Author

It of course depends on how long the tuples are, but it does seem to save another ~20ns which seems negligible but somehow adds up... since it isn't too complicated (as in I feel like the code is as unreadable as it was before ;)) i think it might be warranted ☺️

@codecov

codecov Bot commented Oct 8, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
src/mapreduce.jl 93.67% <100.00%> (+0.01%) ⬆️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@Jutho

Jutho commented Oct 8, 2026

Copy link
Copy Markdown
Member

But the question is whether it is actually faster? I cannot verify this:

julia> using Random, TupleTools, BenchmarkTools
julia> p = (randperm(6)...,)
(4, 2, 3, 1, 6, 5)
julia> function _sortperm_descending(values::NTuple{N, Int}) where {N}
           permutation = ntuple(identity, Val(N))
           issorted(values; rev = true) && return permutation

           @inbounds for source in 1:N
               value = values[source]
               position = 1
               @simd for other in 1:N
                   other_value = values[other]
                   position += (other_value > value) | ((other_value == value) & (other < source))
               end
               permutation = TupleTools.setindex(permutation, source, position)
           end
           return permutation
       end
[ Warning: setindex is defined in Base and is not public in TupleTools
_sortperm_descending (generic function with 1 method)
julia> @benchmark TupleTools.sortperm($p, rev=true)
BenchmarkTools.Trial: 10000 samples with 1000 evaluations per sample.
 Range (min … max):  7.417 ns … 24.250 ns  ┊ GC (min … max): 0.00% … 0.00%
 Time  (median):     8.125 ns              ┊ GC (median):    0.00%
 Time  (mean ± σ):   8.150 ns ±  0.204 ns  ┊ GC (mean ± σ):  0.00% ± 0.00%

                                                █             
  ▂▁▁▁▁▁▁▁▁▁▁▁▁▂▁▁▁▁▁▁▁▁▂▁▁▁▁▂▁▁▂▁▁▂▁▂▁▁▂▁▁▂▁▁▄▁█▁▁█▁▁▄▁▁▃▁▂ ▂
  7.42 ns        Histogram: frequency by time        8.29 ns <

 Memory estimate: 0 bytes, allocs estimate: 0.
julia> @benchmark _sortperm_descending($(p))
BenchmarkTools.Trial: 10000 samples with 997 evaluations per sample.
 Range (min … max):  12.914 ns … 28.084 ns  ┊ GC (min … max): 0.00% … 0.00%
 Time  (median):     13.039 ns              ┊ GC (median):    0.00%
 Time  (mean ± σ):   13.844 ns ±  1.550 ns  ┊ GC (mean ± σ):  0.00% ± 0.00%

  █▇▅▄▄▁  ▄▁▁   ▄▂▂         ▅▁▁           ▄▁                ▃ ▂
  ███████▄███▃▄▃███▃▁▁▃▁▁▃▁▁███▃▃▁▁▁▁▁▁▄▁▅███▄▃▁▁▁▁▁▃▃▁▁▃▁▁▃█ █
  12.9 ns      Histogram: log(frequency) by time      19.1 ns <

 Memory estimate: 0 bytes, allocs estimate: 0.

with similar results for tuples of length 4 to 8, with a tie around tuple length 9. But maybe this way of benchmarking is not sufficient. I also don't know if this is just constant propagation at this point; I forgot how to avoid this in benchmarking.

@Jutho

Jutho commented Oct 8, 2026 •

Copy link
Copy Markdown
Member

With randomized inputs, I get mixed results, with _sortperm_descending seemingly becoming worse for larger N

julia> N=4;

julia> @benchmark TupleTools.sortperm(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 997 evaluations per sample.
 Range (min … max):  19.349 ns … 496.532 ns  ┊ GC (min … max): 0.00% … 92.74%
 Time  (median):     19.725 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   20.622 ns ±  18.611 ns  ┊ GC (mean ± σ):  4.06% ±  4.29%

        ▁▃ ▄▆ ▇█ ▇▄ ▁                                           
  ▂▂▁▄▇▁██▁██▁██▁██▁█▅▄▅▁▄▃▁▃▂▁▂▂▁▂▂▁▂▂▁▂▁▁▂▁▂▂▁▂▁▁▂▂▁▂▁▁▂▂▁▂▂ ▃
  19.3 ns         Histogram: frequency by time           21 ns <

 Memory estimate: 48 bytes, allocs estimate: 1.

julia> @benchmark _sortperm_descending(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 997 evaluations per sample.
 Range (min … max):  16.257 ns … 478.142 ns  ┊ GC (min … max): 0.00% … 90.65%
 Time  (median):     18.974 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   19.891 ns ±  18.986 ns  ┊ GC (mean ± σ):  4.29% ±  4.30%

                              ▃██▅▄▆▆▃                          
  ▁▁▁▁▁▂▂▂▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▂▃▅████████▆▅▃▂▂▂▂▂▂▂▂▁▁▁▁▁▁▁▁▁▁▁▁▁ ▂
  16.3 ns         Histogram: frequency by time         21.3 ns <

 Memory estimate: 48 bytes, allocs estimate: 1.

julia> N = 5;

julia> @benchmark TupleTools.sortperm(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 996 evaluations per sample.
 Range (min … max):  22.590 ns … 493.558 ns  ┊ GC (min … max): 0.00% … 92.38%
 Time  (median):     23.553 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   24.614 ns ±  19.129 ns  ┊ GC (mean ± σ):  3.58% ±  4.35%

           ▁▅▁█▆▂                                               
  ▂▂▃▃▄▅▅▄▆████████▇▄▅▅▅▄▆▆▅▃▃▃▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▁▁▂▂▂▂▂▂▂▂▂▂▂▂▁▂ ▃
  22.6 ns         Histogram: frequency by time         26.9 ns <

 Memory estimate: 48 bytes, allocs estimate: 1.

julia> @benchmark _sortperm_descending(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 996 evaluations per sample.
 Range (min … max):  21.754 ns … 486.530 ns  ┊ GC (min … max): 0.00% … 89.67%
 Time  (median):     29.869 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   31.557 ns ±  20.203 ns  ┊ GC (mean ± σ):  3.06% ±  4.47%

                              ▂▅█▇▄                             
  ▃▂▂▁▁▁▁▂▂▁▁▁▁▁▁▁▁▂▁▂▂▂▂▂▂▃▄▇██████▅▄▃▃▃▃▄▅▆▆▆▅▄▃▂▂▂▂▂▂▂▂▂▂▂▂ ▃
  21.8 ns         Histogram: frequency by time           37 ns <

 Memory estimate: 48 bytes, allocs estimate: 1.

julia> N = 6;

julia> @benchmark TupleTools.sortperm(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 995 evaluations per sample.
 Range (min … max):  25.586 ns … 475.544 ns  ┊ GC (min … max): 0.00% … 90.23%
 Time  (median):     26.465 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   28.060 ns ±  23.100 ns  ┊ GC (mean ± σ):  5.11% ±  5.77%

  ▂▃▄▅▆▇▇▇█▇▅▂▁▁▂▁                                             ▂
  ███████████████████▅▆▅▅▆▃▇▅▆▆▆▇▆█▇▇▆▆▆▅▆▆▆▆▅▅▅▃▃▄▅▄▁▁▄▃▄▅▅▃▅ █
  25.6 ns       Histogram: log(frequency) by time      32.5 ns <

 Memory estimate: 64 bytes, allocs estimate: 1.

julia> @benchmark _sortperm_descending(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 997 evaluations per sample.
 Range (min … max):  18.764 ns … 463.473 ns  ┊ GC (min … max): 0.00% … 90.22%
 Time  (median):     28.084 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   29.560 ns ±  22.417 ns  ┊ GC (mean ± σ):  4.61% ±  5.63%

                                        █▆                      
  ▂▂▂▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▂▄▅██▆▃▂▂▂▂▂▂▂▂▂▂▂▂▂▁▂▂▂▂ ▂
  18.8 ns         Histogram: frequency by time         33.1 ns <

 Memory estimate: 64 bytes, allocs estimate: 1.

julia> N = 7;

julia> @benchmark TupleTools.sortperm(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 996 evaluations per sample.
 Range (min … max):  22.716 ns … 450.888 ns  ┊ GC (min … max): 0.00% … 89.89%
 Time  (median):     24.138 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   25.492 ns ±  22.152 ns  ┊ GC (mean ± σ):  5.21% ±  5.60%

             ▃   ▁▁█▄▁                                          
  ▂▂▂▂▂▃▄▅▆▇███▆▇█████▆▃▃▂▂▂▂▂▂▂▂▂▁▂▂▂▁▁▂▂▂▂▂▂▂▂▂▂▂▂▂▂▁▂▂▂▁▁▂▂ ▃
  22.7 ns         Histogram: frequency by time         28.2 ns <

 Memory estimate: 64 bytes, allocs estimate: 1.

julia> @benchmark _sortperm_descending(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 992 evaluations per sample.
 Range (min … max):  23.101 ns … 473.748 ns  ┊ GC (min … max): 0.00% … 88.83%
 Time  (median):     36.837 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   38.452 ns ±  22.893 ns  ┊ GC (mean ± σ):  3.70% ±  5.61%

                                          ▅██▅▃     ▁          ▂
  ▄▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁▁██████▆▅▆███▇▄▅▃▄▄▅▆ █
  23.1 ns       Histogram: log(frequency) by time      42.7 ns <

 Memory estimate: 64 bytes, allocs estimate: 1.

julia> N = 8;

julia> @benchmark TupleTools.sortperm(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 995 evaluations per sample.
 Range (min … max):  26.005 ns … 281.575 ns  ┊ GC (min … max): 0.00% … 82.03%
 Time  (median):     27.051 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   28.347 ns ±  15.697 ns  ┊ GC (mean ± σ):  4.09% ±  6.49%

        ▃█▅▃                                                    
  ▂▃▃▅▆█████▇▄▃▃▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▁▂▂▂▂▂▂▂▂▂▂▂▂▁▂▂▂▂▁▂▁▂ ▃
  26 ns           Histogram: frequency by time         33.9 ns <

 Memory estimate: 80 bytes, allocs estimate: 1.

julia> @benchmark _sortperm_descending(q) setup=(q=(randperm($N)...,))
BenchmarkTools.Trial: 10000 samples with 992 evaluations per sample.
 Range (min … max):  36.794 ns … 319.725 ns  ┊ GC (min … max): 0.00% … 80.65%
 Time  (median):     37.761 ns               ┊ GC (median):    0.00%
 Time  (mean ± σ):   39.102 ns ±  15.434 ns  ┊ GC (mean ± σ):  2.99% ±  6.31%

        ▆█▂                                                     
  ▂▃▃▃▄▅███▅▃▂▂▂▂▂▂▂▁▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▂▁▁▁▂▂▂▁▁▂▂▂▁▂▂▂▂▂▂▂▂▂▂▂▂ ▂
  36.8 ns         Histogram: frequency by time         44.7 ns <

 Memory estimate: 80 bytes, allocs estimate: 1.

@lkdvos
lkdvos force-pushed the ld-layout-planning branch from 28a4167 to 99f67e8 Compare October 8, 2026 12:18
@lkdvos

lkdvos commented Oct 8, 2026

Copy link
Copy Markdown
Member Author

Okay it turns out I was measuring against an earlier version of my own code that was doing some other sorting strategy, and you are right that TupleTools is just winning everywhere (Jutho/TupleTools.jl#27 needs fixing for larger tuples but that's kind of irrelevant here).

I'll investigate the importance strategy in a separate PR, since it does seem that Claude thinks there are some cases we are missing.

For example the following reduction:

output_strides = (1, 32, 0)
input_strides = (1, 32, 768)

will give (1, 3, 2) by the _importance, but you really want to fuse the first two axes which is then no longer possible.

@Jutho

Jutho commented Oct 8, 2026 •

Copy link
Copy Markdown
Member

With regards to sorting, I don't know how relevant this still is:
Jutho/TupleTools.jl#20

but SortingNetworks.jl seemed to be the fastest tuple sorting long time ago (not actually so long ago).

@lkdvos

lkdvos commented Oct 8, 2026

Copy link
Copy Markdown
Member Author

I think in any case it makes sense to do the sorting experiments in TupleTools, so this should be ready to merge

@lkdvos
lkdvos force-pushed the ld-layout-planning branch from 99f67e8 to 2e553df Compare October 8, 2026 15:00
@lkdvos
lkdvos enabled auto-merge (squash) October 8, 2026 15:27
@lkdvos
lkdvos disabled auto-merge October 8, 2026 19:10
@lkdvos
lkdvos merged commit 16454bf into main Oct 8, 2026
11 of 12 checks passed
@lkdvos
lkdvos deleted the ld-layout-planning branch October 8, 2026 19:11
lkdvos referenced this pull request Oct 8, 2026
* Skip layout planning for single-element arrays

* Skip stride ranking when the memory estimate fits (#92)

* Fix single-element GPU tests on Metal and CUDA

The `map!` closure captured the loop variable `T`, which is not isbits
and so cannot be passed to a GPU kernel; use `one(x)` instead. The
in-place reduction check compared GPU and CPU `cos` results exactly,
which can differ by an ulp; compare approximately.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants