Skip to content

Segmented radix for sort(A; dims) via onesweep #120

Description

@shreyas-omkar

Follow-up to #117.

sort(A; dims) currently uses composite (slice_id, value) keys with a custom comparator, which forces merge sort and lands at a uniform ~0.65 ms on a 1M-element array regardless of shape, trailing CUDA's dedicated segmented sorter by roughly 1.2x to 1.9x. Widening the key so plain radix applies is worse, since it multiplies the radix passes.

The performant path is a segmented radix, where the segmentation lives in the histogram and scatter (per-segment digit counts) rather than in the key, so keys stay native width and radix throughput is preserved. This composes with onesweep: reset the decoupled-lookback carry at each segment boundary and reuse the existing per-tile histogram plus lookback machinery.

Plan:

  1. Land onesweep (unblocked now that the DL device fence is merged and validated).
  2. Extend it to a segmented variant with per-segment boundary resets for sort(A; dims).

Until then, #117 provides a correct, feature-complete generic fallback.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions