Skip to content

fix(autograd): gather/indexSelect backward and indexSelect forward reject rank>=2 indices - #995

Merged
michalharakal merged 2 commits into
developfrom
fix/gather-indexselect-backward-multidim-indices
Aug 16, 2026
Merged

michalharakal merged 2 commits into
developfrom
fix/gather-indexselect-backward-multidim-indices

Conversation

@michalharakal

Copy link
Copy Markdown
Contributor

Summary

Fixes #994 — gatherBackward/indexSelectBackward (DefaultExecutionTape.kt) and
indexSelect's forward pass (DefaultCpuOps.kt) all read the indices tensor via a single
flat index (indices.data[it]), which only works when indices is rank 1. Any rank >= 2
indices tensor — including a batched [B,T] token-embedding lookup, gather()'s primary
documented use case — throws:

java.lang.IllegalArgumentException: Number of indices (1) must match tensor dimensions (2)
	at sk.ainet.lang.graph.DefaultGradientTape.gatherBackward(DefaultExecutionTape.kt:1059)

Forward-pass gather() already read multi-dim indices correctly; only its backward, and both
directions of indexSelect, had this bug.

Fix

Read indices through a rank-agnostic path in all three spots:

  • gatherBackward / indexSelectBackward: TensorData.copyToFloatArray(), which already
    unravels flat positions into per-dimension coordinates correctly (see its own doc comment).
  • indexSelect forward: the same buffer/unravel dispatch gather()'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=1 boundary case, for both gather and indexSelect backward.
  • IndexSelectMultiDimIndicesTest (new): same rank-2/rank-3/boundary coverage for
    indexSelect's forward pass.
  • EmbeddingTableTrainingE2ETest (new): system-level — trains a tiny embedding table end to
    end (forward, backward, adamw step, 200 iterations) on a batched [B,T] lookup against a
    deterministic 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.

./gradlew :skainet-compile:skainet-compile-dag:jvmTest --tests "*OpsAutodiffBackwardTest*" --tests "*EmbeddingTableTrainingE2ETest*"
./gradlew :skainet-backends:skainet-backend-cpu:jvmTest --tests "*IndexSelectMultiDimIndicesTest*" --tests "*GatherRowDequantTest*"

…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).
@michalharakal
michalharakal requested a review from aharakal August 16, 2026 10:52
@michalharakal
michalharakal merged commit 6c7a5bf into develop Aug 16, 2026
13 checks passed
@michalharakal
michalharakal deleted the fix/gather-indexselect-backward-multidim-indices branch August 16, 2026 18:19
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.

gather/indexSelect backward (and indexSelect forward) throw for rank>=2 indices

2 participants