logitsLastToken method

Tensor logitsLastToken(
  1. Tensor hidden
)

Given decoder hidden [B, T, C], project only the last position into vocab logits [B, V] using the tied token-embedding matrix.

Implementation

Tensor logitsLastToken(Tensor hidden) {
  final b = hidden.shape[0];
  final t = hidden.shape[1];
  final c = hidden.shape[2];
  if (b < 1) throw ArgumentError('empty batch');
  // Slice the last position as a fresh [B, C] tensor on-device via
  // host round-trip (tiny: B*C floats).
  final all = hidden.toFloat32List();
  final last = Float32List(b * c);
  for (int bi = 0; bi < b; bi++) {
    for (int ci = 0; ci < c; ci++) {
      last[bi * c + ci] = all[bi * t * c + (t - 1) * c + ci];
    }
  }
  final lastT = Tensor.fromFloat32List([b, c], last, device: hidden.device);
  // [B, C] @ [C, V] = [B, V] on-device.
  return lastT.matmul(tokenEmbedding.weight.transpose());
}