sampleNext method

int sampleNext(
  1. Float64List logits,
  2. List<int> recent
)

Implementation

int sampleNext(Float64List logits, List<int> recent) {
  // Apply repetition penalty
  if (cfg.repetitionPenalty > 1.0 && recent.isNotEmpty) {
    final set = recent.length > cfg.penaltyWindow
        ? recent.sublist(recent.length - cfg.penaltyWindow)
        : recent;
    final seen = Set<int>.from(set);
    for (final id in seen) {
      if (id >= 0 && id < logits.length) {
        logits[id] /= cfg.repetitionPenalty; // simple variant
      }
    }
  }

  // Temperature scaling
  final temp = cfg.temperature.clamp(1e-6, 1000.0);
  for (var i = 0; i < logits.length; i++) {
    logits[i] /= temp;
  }

  // Build candidate list (id, logit)
  final List<_Tok> toks = [
    for (var i = 0; i < logits.length; i++) _Tok(i, logits[i])
  ];

  // Top-K filter
  if (cfg.topK > 0 && cfg.topK < toks.length) {
    toks.sort((a, b) => b.logit.compareTo(a.logit));
    toks.removeRange(cfg.topK, toks.length);
  }

  // Convert to probs with softmax on remaining
  final maxv = toks.fold<double>(-double.infinity, (m, t) => m > t.logit ? m : t.logit);
  var sum = 0.0;
  for (final t in toks) {
    t.prob = math.exp(t.logit - maxv);
    sum += t.prob;
  }

  // Top-P nucleus
  if (cfg.topP < 1.0) {
    toks.sort((a, b) => b.prob.compareTo(a.prob));
    var cum = 0.0;
    final kept = <_Tok>[];
    for (final t in toks) {
      kept.add(t);
      cum += t.prob / sum;
      if (cum >= cfg.topP) break;
    }
    // renormalize
    sum = kept.fold<double>(0.0, (s, t) => s + t.prob);
    for (final t in kept) {
      t.prob /= sum;
    }
    return _draw(kept);
  }

  // Renormalize for general case
  for (final t in toks) {
    t.prob /= sum;
  }
  return _draw(toks);
}