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 是生成流程的最小实现,适合理解自回归生成的基本原理。