RopeCache constructor

RopeCache({
  1. required int maxCtx,
  2. required int headDim,
  3. int? rotaryDim,
  4. double base = 10000.0,
  5. Device device = Device.CPU,
})

Build tables for positions [0, maxCtx) and head-dim headDim (must be even). base is the geometric-progression base for the frequency schedule (10000 for GPT-NeoX / Pythia; the same in LLaMA/Mistral). rotaryDim defaults to the full headDim; pass a smaller even number to only rotate a prefix of each head (GPT-NeoX rotary_pct convention).

Implementation

factory RopeCache({
  required int maxCtx,
  required int headDim,
  int? rotaryDim,
  double base = 10000.0,
  Device device = Device.CPU,
}) {
  if (headDim.isOdd) {
    throw ArgumentError('RopeCache: headDim ($headDim) must be even');
  }
  final rDim = rotaryDim ?? headDim;
  if (rDim.isOdd) {
    throw ArgumentError('RopeCache: rotaryDim ($rDim) must be even');
  }
  if (rDim <= 0 || rDim > headDim) {
    throw ArgumentError(
      'RopeCache: rotaryDim ($rDim) must be in (0, $headDim]',
    );
  }
  final halfR = rDim ~/ 2;

  // Inverse frequencies: shape [halfR]. Note the divisor is
  // `rDim` (not headDim) — matches HF GPT-NeoX which passes
  // rotary_ndims as the "dim" arg to RotaryEmbedding.
  final invFreq = List<double>.generate(
    halfR,
    (i) => 1.0 / math.pow(base, 2.0 * i / rDim),
  );

  // For each position m, build a [headDim] vector of cos/sin values
  // laid out as:
  //   cos: [cos(m*f_0)..cos(m*f_{halfR-1}) | cos(m*f_0)..cos(m*f_{halfR-1}) | 1..1]
  //   sin: [sin(m*f_0)..sin(m*f_{halfR-1}) | sin(m*f_0)..sin(m*f_{halfR-1}) | 0..0]
  // The trailing `headDim - rDim` slots ensure the pass-through
  // dims are multiplied by 1 (cos) / 0 (sin), leaving them
  // unrotated in the formula q*cos + rotate_half(q)*sin.
  final cosPerPos = <Tensor>[];
  final sinPerPos = <Tensor>[];
  for (int m = 0; m < maxCtx; m++) {
    final cRow = Float64List(headDim);
    final sRow = Float64List(headDim);
    for (int j = 0; j < halfR; j++) {
      final ang = m * invFreq[j];
      final c = math.cos(ang);
      final s = math.sin(ang);
      cRow[j] = c;
      cRow[j + halfR] = c;
      sRow[j] = s;
      sRow[j + halfR] = s;
    }
    for (int j = rDim; j < headDim; j++) {
      cRow[j] = 1.0;
      // sRow[j] already 0.0 from Float64List default.
    }
    cosPerPos.add(
      Tensor.fromList([1, headDim], cRow.toList(), device: device),
    );
    sinPerPos.add(
      Tensor.fromList([1, headDim], sRow.toList(), device: device),
    );
  }

  // rotate_half([a | b | pass]) = [-b | a | 0]. Encode as a matmul:
  // P has shape [headDim, headDim] where P[i, j] tells the output
  // column j to source input column i.
  //   * for j in [0, halfR):      out_j = -in_{j + halfR}
  //     -> P[j + halfR, j] = -1
  //   * for j in [halfR, rDim):   out_j = in_{j - halfR}
  //     -> P[j - halfR, j] = 1
  //   * for j in [rDim, headDim): out_j = 0
  // The last block is left zero (multiplied by sin==0 anyway).
  final pData = Float64List(headDim * headDim);
  for (int j = 0; j < halfR; j++) {
    pData[(j + halfR) * headDim + j] = -1.0;
  }
  for (int j = halfR; j < rDim; j++) {
    pData[(j - halfR) * headDim + j] = 1.0;
  }
  final rotateHalfP = Tensor.fromList(
    [headDim, headDim],
    pData.toList(),
    device: device,
  );

  return RopeCache._(
    maxCtx,
    headDim,
    rDim,
    base,
    cosPerPos,
    sinPerPos,
    rotateHalfP,
    device,
  );
}