跳转至

naive_rollout.py — Naive Rollout

文件概述

最简单的自回归生成实现,逐 token 生成,不使用任何推理引擎优化。

核心类

NaiveRollout(BaseRollout)

class NaiveRollout(BaseRollout):
    def generate_sequences(self, prompts: DataProto) -> DataProto:
        """纯 PyTorch 的逐 token 自回归生成

        for step in range(max_new_tokens):
            logits = model(input_ids)  # 前向传播
            next_token = sample(logits)  # 采样
            input_ids = cat([input_ids, next_token])  # 拼接
        """

适用场景

  • 教学和理解自回归生成原理
  • 极小规模测试
  • 不需要任何外部依赖

小结

NaiveRollout 是生成流程的最小实现,适合理解自回归生成的基本原理。