Skip to content

fix(scatter_add): accumulate duplicate destination indices via unique_indices flag - #11

Open
ryan-williams wants to merge 1 commit into
aws-neuron:mainfrom
Open-Athena:fix/scatter-add-dup-indices
Open

fix(scatter_add): accumulate duplicate destination indices via unique_indices flag#11
ryan-williams wants to merge 1 commit into
aws-neuron:mainfrom
Open-Athena:fix/scatter-add-dup-indices

Conversation

@ryan-williams

@ryan-williams ryan-williams commented Jul 31, 2026

Copy link
Copy Markdown

scatter_add silently drops contributions when index repeats a destination row inside the
same 128-row tile. That is the common case for embedding gradients, so
grad_table.scatter_add_(0, token_ids, grads) currently returns wrong gradients for any
batch containing a repeated token.

This PR adds an opt-in unique_indices=False path that accumulates correctly. The default is
unchanged, so existing callers and the existing fast path are untouched.

The bug

scatter_add does a tile-wide gather → add → scatter over 128-row tiles. Within a tile the
pre-image is gathered once, so duplicate destination rows all read the same pre-image and the
scatter back is last-write-wins — every duplicate contribution but one is discarded, with no
error.

index=[0, 0, 1, 2, 2, 2], src=ones:

row scatter_add torch.scatter_add_
0 1 2
1 1 1
2 1 3
3 0 0

The docstring noted the precondition in passing ("indices within a tile of 128 rows should be
unique for correctness"), but nothing enforced or detected the violation.

Why the existing test does not catch it: test_scatter_add.py's input generator builds
per-tile permutations (np.random.permutation(bs_slen)), so every tile is unique by
construction and the duplicate path is never exercised.

The fix

Add unique_indices: bool = True. When False, set k_tile_size = 1 so each index becomes its
own sequential tile; nl.sequential_range then makes each add observe prior writes, so
duplicates accumulate.

  • Default True preserves today's behavior and performance — no existing caller changes.
  • unique_indices=False is correct but slower (1 row per tile instead of 128). We chose the
    minimal change that makes correct results reachable; a faster approach (sorted segments,
    atomics) would be a better long-term fix if you want one, and we would be happy to follow
    your lead there.
  • The kwarg is mirrored on scatter_add_torch_ref for test-harness signature parity.
  • Docstring updated to state that duplicates are silently dropped under the default.

Tests

Adds test/integration/nkilib/experimental/misc/test_scatter_add_dup_indices.py, following the
conventions of the existing test_scatter_add.py (same Orchestrator / UnitTestFramework /
torch_ref_wrapper structure), with indices sampled with replacement so duplicates occur
both within and across tile boundaries:

bs_slen dim_size src_rows dtype covers
16 512 64 float32 all duplicates inside one <128-row tile
16 512 64 bfloat16 same, reduced precision
10 256 200 float32 duplicates crossing the 128-row tile boundary
64 1024 512 float32 larger shape

Validation

  • CPU simulation (--test-mode simulation): existing test_scatter_add.py plus the new
    dup-index tests, 10 passed / 4 skipped.
  • Trainium2 hardware (trn2.3xlarge, neuronx-cc 2.26.6360.0): the fixed kernel called from
    JAX accumulates duplicate indices correctly, matching a reference implementation.
  • End-to-end: used as the backward half of an embedding-lookup jax.custom_vjp (NKI gather
    forward + this kernel backward), which differentiates correctly under jax.grad on device. With
    the stock kernel the same test produces incorrect gradients on batches with repeated tokens.

Notes

Found during Trainium bring-up for a JAX/Equinox training stack, where the embedding backward is
the first place duplicate indices show up. Happy to adjust naming, defaults, or the test layout to
match your conventions.

…ices flag

The tile-wide gather/accumulate/scatter path is last-write-wins for repeated
destination indices within a 128-row tile, silently dropping all but one
contribution. This is the natural embedding-gradient backward, where repeats
(common tokens, BOS/EOS) are ubiquitous. The existing test never caught it: its
input generator builds per-tile permutations, so tiles are unique by construction.

Add unique_indices: bool = True. When False, k_tile_size=1 makes each index its
own sequential tile, so nl.sequential_range accumulates duplicates correctly.
Mirror the kwarg on scatter_add_torch_ref (harness signature parity) and add a
dup-index regression test (sampled with replacement, within- and cross-tile).
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.

1 participant