generate method

List<double> generate(
  1. List<double> prompt, {
  2. required int maxNewTokens,
  3. double temperature = 1.0,
  4. int? topK,
  5. Random? rng,
  6. bool useCache = true,
})

Autoregressive sampling. Same signature and behaviour as GPT.generate.

Implementation

List<double> generate(
  List<double> prompt, {
  required int maxNewTokens,
  double temperature = 1.0,
  int? topK,
  math.Random? rng,
  bool useCache = true,
}) {
  if (prompt.isEmpty) {
    throw ArgumentError('PythiaModel.generate: prompt must be non-empty');
  }
  if (useCache && prompt.length > config.maxCtx) {
    throw ArgumentError(
      'PythiaModel.generate(useCache: true): prompt length '
      '${prompt.length} exceeds maxCtx ${config.maxCtx}. Pass '
      'useCache: false to enable sliding-window truncation.',
    );
  }
  final r = rng ?? math.Random();
  final wasTraining = training;
  eval();
  try {
    return Tensor.noGrad(
      () => useCache
          ? _generateCached(prompt, maxNewTokens, temperature, topK, r)
          : _generateNoCache(prompt, maxNewTokens, temperature, topK, r),
    );
  } finally {
    if (wasTraining) train();
  }
}