fix(autograd): gather/indexSelect backward and indexSelect forward reject rank>=2 indices - #995
Merged
michalharakal merged 2 commits intoAug 16, 2026
Conversation
…ject rank>=2 indices gatherBackward and indexSelectBackward (DefaultExecutionTape.kt) read the indices tensor via `indices.data[it]`, a single flat index. The vararg element accessor requires exactly one coordinate per dimension, so this only works when indices is rank 1 — it throws "Number of indices (1) must match tensor dimensions (N)" during backward for any indices tensor of rank >= 2, e.g. a batched [B,T] token-id lookup into an embedding table, which is the primary documented use case for gather(). Forward-pass gather already handles multi-dim indices correctly (see GatherRowDequantTest); backward never did, and had no test coverage for it. indexSelect's forward (DefaultCpuOps.kt) had the identical bug in its own index-reading loop, so a rank>=2 indices tensor throws before backward is even reached. Fix: read indices through a rank-agnostic path in all three spots — TensorData.copyToFloatArray() (which already unravels flat positions correctly per its own doc comment) for the two backward functions, and the same buffer/unravel dispatch gather()'s forward already uses for indexSelect's forward. Adds regression coverage: rank-2 indices variants of the existing gather_backward/indexSelect_backward finite-difference tests, plus a forward-pass indexSelect multi-dim indices test mirroring GatherRowDequantTest's gatherAcceptsMultiDimensionalIndices. All three fail on the pre-fix code and pass after.
…for #994 Rounds out the gather/indexSelect rank>=2 indices fix with coverage beyond the original rank-2 regression tests: - gather backward: rank-3 indices (exercises rowOf()'s generic multi-dim branch, not just the outIdx.size==2 fast path) and a numIndices=1 boundary case. - indexSelect forward: rank-3 indices and a numIndices=1 boundary case. - System-level: EmbeddingTableTrainingE2ETest trains a tiny embedding table end to end (forward, backward, adamw step, repeated 200 steps) on a batched [B,T] token lookup against a deterministic next-token pattern, asserting the loss actually converges — not just that nothing throws. This is the shape (and the real failure mode) that surfaced the bug: a minimal bigram-style language model. All five new tests verified to fail on the pre-fix code and pass on the fix (checked out DefaultExecutionTape.kt/DefaultCpuOps.kt at HEAD~1, confirmed failures, restored).
aharakal
approved these changes
Aug 16, 2026
michalharakal
deleted the
fix/gather-indexselect-backward-multidim-indices
branch
August 16, 2026 18:19
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Fixes #994 —
gatherBackward/indexSelectBackward(DefaultExecutionTape.kt) andindexSelect's forward pass (DefaultCpuOps.kt) all read theindicestensor via a singleflat index (
indices.data[it]), which only works whenindicesis rank 1. Any rank >= 2indices tensor — including a batched
[B,T]token-embedding lookup,gather()'s primarydocumented use case — throws:
Forward-pass
gather()already read multi-dim indices correctly; only its backward, and bothdirections of
indexSelect, had this bug.Fix
Read
indicesthrough a rank-agnostic path in all three spots:gatherBackward/indexSelectBackward:TensorData.copyToFloatArray(), which alreadyunravels flat positions into per-dimension coordinates correctly (see its own doc comment).
indexSelectforward: the same buffer/unravel dispatchgather()'s forward already uses.Test plan
OpsAutodiffBackwardTest: rank-2 (the batched-lookup shape that broke) and rank-3(exercises the generic multi-dim branch, not just the rank-2 fast path) indices, plus a
numIndices=1boundary case, for bothgatherandindexSelectbackward.IndexSelectMultiDimIndicesTest(new): same rank-2/rank-3/boundary coverage forindexSelect's forward pass.EmbeddingTableTrainingE2ETest(new): system-level — trains a tiny embedding table end toend (forward, backward,
adamwstep, 200 iterations) on a batched[B,T]lookup against adeterministic next-token pattern, asserting the loss actually converges, not just that
nothing throws. This mirrors the real scenario that surfaced the bug: a minimal bigram-style
language model.
All new tests verified to fail against the pre-fix code and pass with the fix.