greedyDecode method

List<int> greedyDecode({
  1. required List<int> startTokens,
  2. required int eot,
  3. int maxLen = 100,
  4. List<int> initialSuppress = const [],
})

Greedy decode from startTokens up to maxLen, stopping when eot is emitted. Caller must have called primeCrossAttn first.

Implementation

List<int> greedyDecode({
  required List<int> startTokens,
  required int eot,
  int maxLen = 100,
  List<int> initialSuppress = const [],
}) {
  final tokens = List<int>.of(startTokens);
  while (tokens.length < maxLen) {
    final tensor = Tensor.fromList(
      [1, tokens.length],
      List<double>.generate(tokens.length, (i) => tokens[i].toDouble()),
      device: device,
    );
    final hidden = forward(tensor);
    final logitsT = logitsLastToken(hidden);
    final logits = logitsT.toFloat32List();

    final isFirstSample = tokens.length == startTokens.length;
    double best = -double.infinity;
    int bestId = -1;
    for (int v = 0; v < vocabSize; v++) {
      if (isFirstSample && initialSuppress.contains(v)) continue;
      final lv = logits[v];
      if (lv > best) {
        best = lv;
        bestId = v;
      }
    }
    if (bestId == eot) break;
    tokens.add(bestId);
  }
  return tokens;
}