Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -3171,9 +3171,25 @@ public open class DefaultCpuOpsBase(protected val dataFactory: TensorDataFactory
override fun <T : DType, V> indexSelect(input: Tensor<T, V>, indices: Tensor<DType, *>, dim: Int): Tensor<T, V> {
require(dim in 0 until input.rank) { "indexSelect: dim=$dim out of range for rank ${input.rank}" }
val numIndices = indices.volume
val indexList = IntArray(numIndices) { i ->
val v = indices.data[i]
(v as Number).toInt()
// The vararg element accessor requires one coordinate per dimension, so a flat
// `data[i]` throws for rank >= 2 indices — read the contiguous buffer when present,
// otherwise unravel the flat position into a coordinate (mirrors gather() above).
val idxData = indices.data
val indexList = when (idxData) {
is IntArrayTensorData<*> -> IntArray(numIndices) { idxData.buffer[it] }
is FloatArrayTensorData<*> -> IntArray(numIndices) { idxData.buffer[it].toInt() }
else -> {
val dims = indices.shape.dimensions
IntArray(numIndices) { flat ->
val coord = IntArray(dims.size)
var rem = flat
for (d in dims.indices.reversed()) {
coord[d] = rem % dims[d]
rem /= dims[d]
}
(idxData.get(*coord) as Number).toInt()
}
}
}

val resultDims = input.shape.dimensions.copyOf()
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
package sk.ainet.exec.tensor.ops

import sk.ainet.context.DirectCpuExecutionContext
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.Tensor
import sk.ainet.lang.types.FP32
import sk.ainet.lang.types.Int32
import kotlin.test.Test
import kotlin.test.assertContentEquals
import kotlin.test.assertEquals

/**
* `ops.indexSelect` used to read its `indices` tensor via a single flat index (`indices.data[i]`),
* which only matches a rank-1 indices tensor — it threw for any indices rank != 1, even though
* indexSelect only ever cares about the flat, row-major sequence of index values (the `dim` axis
* size becomes `indices.volume` regardless of the indices tensor's own shape).
*/
class IndexSelectMultiDimIndicesTest {

@Test
fun indexSelectAcceptsMultiDimensionalIndices() {
val ctx = DirectCpuExecutionContext.create()
// x: [3,4], row r = [4r, 4r+1, 4r+2, 4r+3]
val x = ctx.fromFloatArray<FP32, Float>(Shape(3, 4), FP32::class, FloatArray(12) { it.toFloat() })
// rank-2 indices [[0,2],[2,0]] along dim=1 -> flat sequence [0,2,2,0]
val ids = ctx.fromIntArray<Int32, Int>(Shape(2, 2), Int32::class, intArrayOf(0, 2, 2, 0))

@Suppress("UNCHECKED_CAST")
val out = ctx.ops.indexSelect(x, ids as Tensor<sk.ainet.lang.types.DType, *>, dim = 1)

// indexSelect keeps input rank, replacing dim's size with indices.volume: [3,4].
assertEquals(listOf(3, 4), out.shape.dimensions.toList())
assertContentEquals(
floatArrayOf(
0f, 2f, 2f, 0f,
4f, 6f, 6f, 4f,
8f, 10f, 10f, 8f,
),
out.data.copyToFloatArray(),
)
}

@Test
fun indexSelectAcceptsRank3Indices() {
val ctx = DirectCpuExecutionContext.create()
val x = ctx.fromFloatArray<FP32, Float>(Shape(4), FP32::class, floatArrayOf(10f, 20f, 30f, 40f))
// rank-3 indices [2,2,2] (volume 8) along dim=0 -> a [batch, group, seq]-shaped lookup.
val ids = ctx.fromIntArray<Int32, Int>(Shape(2, 2, 2), Int32::class, intArrayOf(0, 1, 2, 3, 3, 2, 1, 0))

@Suppress("UNCHECKED_CAST")
val out = ctx.ops.indexSelect(x, ids as Tensor<sk.ainet.lang.types.DType, *>, dim = 0)

assertEquals(listOf(8), out.shape.dimensions.toList())
assertContentEquals(floatArrayOf(10f, 20f, 30f, 40f, 40f, 30f, 20f, 10f), out.data.copyToFloatArray())
}

@Test
fun indexSelectHandlesSingleIndex() {
// numIndices=1 boundary, still rank 2 (not rank 1) indices.
val ctx = DirectCpuExecutionContext.create()
val x = ctx.fromFloatArray<FP32, Float>(Shape(3, 2), FP32::class, floatArrayOf(1f, 2f, 3f, 4f, 5f, 6f))
val ids = ctx.fromIntArray<Int32, Int>(Shape(1, 1), Int32::class, intArrayOf(2))

@Suppress("UNCHECKED_CAST")
val out = ctx.ops.indexSelect(x, ids as Tensor<sk.ainet.lang.types.DType, *>, dim = 0)

assertEquals(listOf(1, 2), out.shape.dimensions.toList())
assertContentEquals(floatArrayOf(5f, 6f), out.data.copyToFloatArray())
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -1056,7 +1056,12 @@ public class DefaultGradientTape(
val indices = inputs[1]
val gradInput = zerosLike(input)
val numIndices = indices.volume
val indexList = IntArray(numIndices) { (indices.data[it] as Number).toInt() }
// indices.data[it] requires exactly one coordinate per dimension (arity check in
// calcFlatIndex), so it only works when indices is rank 1. copyToFloatArray() unravels
// the flat position for us and reads correctly for any indices rank (e.g. a batched
// [B,T] token-id lookup).
val indexFloats = indices.data.copyToFloatArray()
val indexList = IntArray(numIndices) { indexFloats[it].toInt() }
fun rowOf(outIdx: IntArray): Int {
val flatIdx = if (outIdx.size == 2) outIdx[0] else {
var flat = 0
Expand Down Expand Up @@ -1090,7 +1095,9 @@ public class DefaultGradientTape(
val dim = (attributes["dim"] as? Int) ?: 0
val gradInput = zerosLike(input)
val numIndices = indices.volume
val indexList = IntArray(numIndices) { (indices.data[it] as Number).toInt() }
// See gatherBackward above: indices.data[it] only supports rank-1 indices.
val indexFloats = indices.data.copyToFloatArray()
val indexList = IntArray(numIndices) { indexFloats[it].toInt() }
val upDims = upstream.shape.dimensions
val outIdx = IntArray(upDims.size)
val srcIdx = IntArray(upDims.size)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
package sk.ainet.exec.autograd

import kotlin.math.ln
import kotlin.random.Random
import kotlin.test.Test
import kotlin.test.assertTrue
import sk.ainet.context.Phase
import sk.ainet.exec.tensor.ops.DefaultCpuOps
import sk.ainet.lang.graph.DefaultGradientTape
import sk.ainet.lang.graph.DefaultGraphExecutionContext
import sk.ainet.lang.nn.loss.CrossEntropyLoss
import sk.ainet.lang.nn.loss.Reduction
import sk.ainet.lang.nn.optim.adamw
import sk.ainet.lang.nn.topology.ModuleParameter
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.Tensor
import sk.ainet.lang.tensor.data.DenseTensorDataFactory
import sk.ainet.lang.tensor.withRequiresGrad
import sk.ainet.lang.types.DType
import sk.ainet.lang.types.FP32
import sk.ainet.lang.types.Int32

/**
* System-level regression for #994. `gather()`'s primary documented use case — a batched
* `[B,T]` token-embedding lookup — must not just avoid throwing during backward but actually
* train: forward, backward, optimizer step, repeated, driving the loss down on a real (if
* tiny) synthetic task. This is deliberately close to a minimal "bigram" next-token model —
* the shape that surfaced the bug in the first place — rather than a synthetic op-level check.
*/
class EmbeddingTableTrainingE2ETest {

private fun ctx(): DefaultGraphExecutionContext {
val dataFactory = DenseTensorDataFactory()
return DefaultGraphExecutionContext(
baseOps = DefaultCpuOps(dataFactory),
phase = Phase.TRAIN,
tensorDataFactory = dataFactory,
createTapeFactory = { _ -> DefaultGradientTape(true) },
)
}

@Test
fun batched_embedding_lookup_trains_without_throwing_and_reduces_loss() {
val ctx = ctx()
val rng = Random(42)
val vocabSize = 6
val batchSize = 4
val blockSize = 3

val tableData = ctx.tensorDataFactory.randn<FP32, Float>(Shape(vocabSize, vocabSize), FP32::class, 0f, 0.02f, rng)
val table = ctx.fromData(tableData, FP32::class).withRequiresGrad(true)
val param = ModuleParameter.WeightParameter("table", table, true)

val optimizer = adamw(lr = 0.1)
optimizer.addParameter(param)
val lossFn = CrossEntropyLoss()

// A tiny fixed pattern (next id = id+1 mod vocabSize) so there's something real to
// learn — a bigram table can fit this exactly, unlike i.i.d. random targets.
fun batch(): Pair<Tensor<Int32, Int>, Tensor<Int32, Int>> {
val x = IntArray(batchSize * blockSize) { rng.nextInt(vocabSize) }
val y = IntArray(x.size) { (x[it] + 1) % vocabSize }
return ctx.fromIntArray<Int32, Int>(Shape(batchSize, blockSize), Int32::class, x) to
ctx.fromIntArray<Int32, Int>(Shape(batchSize, blockSize), Int32::class, y)
}

@Suppress("UNCHECKED_CAST")
fun step(): Float {
val (x, y) = batch()
ctx.startRecording()
val loss = try {
val logits = ctx.ops.gather(param.value, x as Tensor<DType, *>, dim = 0) // [B,T,vocab]
val flatLogits = ctx.ops.reshape(logits, Shape(batchSize * blockSize, vocabSize))
val flatTargets = ctx.ops.reshape(y, Shape(batchSize * blockSize))
lossFn.forward(flatLogits, flatTargets, ctx, Reduction.MEAN)
} finally {
ctx.stopRecording()
}
ctx.backward(targets = listOf(loss), sources = listOf(param.value))
optimizer.step()
optimizer.zeroGrad()
return loss.data.get()
}

val firstLoss = step()
var lastLoss = firstLoss
repeat(199) { lastLoss = step() }

assertTrue(lastLoss < firstLoss, "loss should decrease over training: first=$firstLoss last=$lastLoss")
val chanceLoss = ln(vocabSize.toFloat())
assertTrue(
lastLoss < 0.25f * chanceLoss,
"a perfectly learnable deterministic next-token pattern should converge well below " +
"chance loss ln(vocabSize)=$chanceLoss, got $lastLoss",
)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,49 @@ class OpsAutodiffBackwardTest {
}
}

@Test
fun gather_backward_scatter_adds_rows_with_rank2_indices() {
// table [vocab=4, emb=3], indices [B=2,T=2] = [[0,2],[2,1]] -> same gather counts as the
// rank-1 case above (1,1,2,0), just batched — the shape a token-embedding lookup uses.
// Regression: gatherBackward read indices via a single flat index, which only matches
// rank-1 indices tensors and threw for any indices rank != 1.
assertGradMatchesFiniteDiff(Shape(4, 3), FloatArray(12) { (it - 6) * 0.1f }) { c, x ->
val idx = intTensor(c, Shape(2, 2), intArrayOf(0, 2, 2, 1))
x.ops.gather(x, idx, dim = 0)
}
}

@Test
fun indexSelect_backward_scatter_adds_along_dim_with_rank2_indices() {
// x [3,4], dim=1, indices [2,2] = [[0,2],[2,0]] -> flat index sequence [0,2,2,0],
// batched. Same regression as gather above.
assertGradMatchesFiniteDiff(Shape(3, 4), FloatArray(12) { (it - 6) * 0.1f }) { c, x ->
val idx = intTensor(c, Shape(2, 2), intArrayOf(0, 2, 2, 0))
x.ops.indexSelect(x, idx, dim = 1)
}
}

@Test
fun gather_backward_scatter_adds_rows_with_rank3_indices() {
// table [vocab=4, emb=2], indices [2,2,2] (volume 8) -> exercises the generic
// multi-dim branch of gatherBackward's rowOf() (upstream rank 4, not the outIdx.size==2
// fast path the rank-1/rank-2 cases above hit), e.g. a [batch, group, seq] lookup shape.
assertGradMatchesFiniteDiff(Shape(4, 2), FloatArray(8) { (it - 4) * 0.1f }) { c, x ->
val idx = intTensor(c, Shape(2, 2, 2), intArrayOf(0, 3, 1, 3, 2, 0, 3, 1))
x.ops.gather(x, idx, dim = 0)
}
}

@Test
fun gather_backward_handles_single_index() {
// numIndices=1 boundary: the smallest possible indices tensor, shape (1,1) not (1,) —
// still rank >= 2, still must not throw and must scatter to exactly one row.
assertGradMatchesFiniteDiff(Shape(3, 2), FloatArray(6) { (it - 3) * 0.1f }) { c, x ->
val idx = intTensor(c, Shape(1, 1), intArrayOf(1))
x.ops.gather(x, idx, dim = 0)
}
}

@Test
fun unfold_backward_folds_overlapping_windows() {
// x [6], size 3, step 1 -> 4 windows; each element's grad = number of windows covering it.
Expand Down
Loading