@dataclassclassGenerationConfig:
# ── Phase 1 实现 ──
max_new_tokens: int = 256
temperature: float = 1.0
top_k: int = 50
top_p: float = 1.0
do_sample: bool = False
eos_token_id: int = 2
pad_token_id: int = 0
repetition_penalty: float = 1.0# ── 预留扩展字段(Phase 1 预定义接口) ──
logits_processor: Optional[List[Callable]] = None
stopping_criteria: Optional[List[Callable]] = Nonedef__post_init__(self):
ifself.temperature <= 0:
raise ValueError("temperature must be > 0")
ifself.top_k < 0:
raise ValueError("top_k must be >= 0")
ifnot0 < self.top_p <= 1.0:
raise ValueError("top_p must be in (0, 1]")
ifself.logits_processor isnotNone:
import warnings
warnings.warn(
"logits_processor is reserved for a future release and is ""currently ignored. Custom logits processors will take effect ""once the LogitsProcessor interface is implemented.",
FutureWarning,
)
ifself.stopping_criteria isnotNone:
import warnings
warnings.warn(
"stopping_criteria is reserved for a future release and is ""currently ignored. Custom stopping criteria will take effect ""once the StoppingCriteria interface is implemented.",
FutureWarning,
)
RFC: HyperParallel 支持 LLM 推理 Generate 流程
需求背景 & 价值
开源实习任务地址:https://gitcode.com/mindspore/community/issues/2101
当前 HyperParallel 的能力重心在分布式训练,缺乏推理侧 Generate 流程的端到端链路。随着训推一体(RLHF / PPO / GRPO)和 RL 训练场景的兴起,框架需要在训练态和推理态之间无缝切换——训练完成后直接在同一分布式环境下进行模型推理生成,无需导出到第三方推理框架。
业界参考:HuggingFace Transformers generate()
HF 的 generate() 是业界事实标准,其核心设计要点:
model.generate(**inputs),内部根据GenerationConfig自动选择解码策略do_sample=False时为贪心搜索,do_sample=True时切换到多项式采样max_new_tokens(控制生成长度)、temperature/top_k/top_p(控制采样)、repetition_penalty(惩罚重复)、eos_token_id(停止条件)padding_side="left"+attention_mask处理不等长 batch promptlogits_processor、stopping_criteria可插拔自定义逻辑HyperParallel generate 的设计目标是对齐 HF 的核心生成范式,同时适配分布式场景(TP/CP)的额外一致性需求。
核心价值:
功能描述
1. 自回归生成流程
实现 Prefill + Decode 两阶段生成,对标 HF
model.generate():Prefill 阶段:根据 prompt_len 构造 causal attention mask;构造 position_ids;一次性调用
model.forward()编码全部输入 token;将返回的past_key_values写入 KV Cache。Decode 阶段(循环
max_new_tokens次):取logits[:, -1, :]作为当前步预测分数;应用 repetition_penalty;按采样策略选择下一个 token;遇eos_token_id或达max_new_tokens则停止;以单 token(seq_len=1)调用model.forward();将新的past_key_values增量合并到 KV Cache。2. 采样策略(对齐 HF 行为)
do_sample=False(默认)do_sample=True, top_k>0do_sample=True, top_p<1.0repetition_penalty != 1.03. KV Cache 管理
对标 HF 的
past_key_values机制,存储每层 key/value 张量,避免 Decode 阶段重复计算。List[Tuple[Tensor, Tensor]],每层 (K, V),形状(B, num_heads, seq_len, head_dim)@torch.no_grad()下)4. Batch 推理支持
对齐 HF 的
padding_side="left"+attention_mask模式:attention_mask(Batch_size, Seq_len),1=real token,0=paddingUnfinishedSequenceLogitsProcessor)5. TP/CP 分布式推理
HF 原生
generate()不具备分布式并行推理能力。TP 推理(Phase 4):各 rank 持有部分 logits(vocab 维度分片),采样前 all-gather 聚合完整 logits,全局执行 Greedy/Top-K 采样。TP 模型层内部已处理权重切分与激活聚合,单 token forward 无需额外修改。
CP 推理(Phase 5):上下文并行涉及 Ring Attention 或 All-to-All 通信(参考
core/context_parallel/的 Ulysses / Colossal AI 实现),KV Cache 的 seq 维度在 rank 间分片,causal mask 适配本地 seq 范围。复用 CP 模块现有通信原语,需扩展 KVCache 支持分片管理。验证目标:单卡与双卡(TP=2 / CP=2)生成结果逐 token 完全一致。
设计方案
1. 文件组织
2. GenerationConfig
@dataclass class GenerationConfig: # ── Phase 1 实现 ── max_new_tokens: int = 256 temperature: float = 1.0 top_k: int = 50 top_p: float = 1.0 do_sample: bool = False eos_token_id: int = 2 pad_token_id: int = 0 repetition_penalty: float = 1.0 # ── 预留扩展字段(Phase 1 预定义接口) ── logits_processor: Optional[List[Callable]] = None stopping_criteria: Optional[List[Callable]] = None def __post_init__(self): if self.temperature <= 0: raise ValueError("temperature must be > 0") if self.top_k < 0: raise ValueError("top_k must be >= 0") if not 0 < self.top_p <= 1.0: raise ValueError("top_p must be in (0, 1]") if self.logits_processor is not None: import warnings warnings.warn( "logits_processor is reserved for a future release and is " "currently ignored. Custom logits processors will take effect " "once the LogitsProcessor interface is implemented.", FutureWarning, ) if self.stopping_criteria is not None: import warnings warnings.warn( "stopping_criteria is reserved for a future release and is " "currently ignored. Custom stopping criteria will take effect " "once the StoppingCriteria interface is implemented.", FutureWarning, )与 HF 的关键差异及设计考量:
GenerationConfig中预留logits_processor和stopping_criteria字段(类型为Optional[List[Callable]],默认None),generate 循环内暂时忽略这两个字段。API 层面用户可见,文档标注"首期未实现,传值不生效"。generate 主循环内部通过_apply_logits_processors和_check_stopping_criteria两个私有方法隔离采样/停止逻辑。LogitsProcessor/StoppingCriteria抽象基类,generate 循环读取config.logits_processor/config.stopping_criteria并应用。上层 API 不变(字段已在 Phase 1 定义),仅行为从"忽略"变为"生效"。beam_size × batch_size)、分布式聚合逻辑变化,复杂度有数量级差异,有待考虑后续实现。3. 生成主循环
4. 与模型侧的接口契约
Generate 模块不 import 任何具体模型类,通过行为契约对接(与 HF 的
GenerationMixin设计理念一致:解耦生成逻辑和模型实现):forward(input_ids, position_ids, attention_mask, past_key_values)→dict{logits, past_key_values}past_key_valuesList[Tuple[Tensor, Tuple]],K/V 形状(batch_size, num_heads, seq_len, dim)attention_mask(batch_size, 1, seq_len, seq_len),causal + padding 的合成 mask。上三角为-inf(causal),padding 列也设置为-inf。Decode 阶段单 token 时传 None(注意力天然可见全部历史 KV)position_ids(batch_size, seq_len),Prefill 从 0 开始,left-padding 位置给 0logits(batch_size, seq_len, vocab_size),Decode 取[:, -1, :]实施计划
能力渐进叠加,每个 Phase 在前一个基础上增加一层能力,每个 Phase 完成后可独立验证。
Phase 1 — Greedy 单 prompt(基础闭环)
目标:在 qwen3.5 模型上跑通最简 greed 生成。无 KV Cache、无采样策略、单条 prompt。
GenerationConfigdataclass +__post_init__校验utils.pybuild_causal_mask+build_position_idsutils.pygreedy_sample(logits)→(batch_size, 1)sampler.pygenerate()最简循环:Prefill → 逐 step greedy 采样 → 拼接输出。全程@torch.no_grad()generation.py验证:
Phase 2 — KV Cache + 采样
目标:Decode 效率不随生成长度线性增长,支持非确定性采样。
KVCache类:update / merge / clearkv_cache.pygeneration.pytop_k_sample+top_p_sample+_apply_repetition_penaltysampler.py验证(CPU / stub 模型。KV Cache 与采样逻辑与模型无关,stub 模型即可充分验证):
Phase 3 — Batch prompts
目标:支持 left-padded batch 输入,逐 item 独立 EOS 停止。
generate()新增attention_mask参数generation.py_build_prefill_position_ids:left-padding 感知generation.pygeneration.py验证(CPU / stub 模型。Batch 拼接与 EOS 逻辑与模型无关):
Phase 4 — TP 分布式推理
目标:Ascend A2 上 TP=2 单卡/双卡 Greedy 生成结果逐 token 一致。HyperParallel generate 区别于 HF 原生 generate 的核心能力
generation.py(或 TP 模块中独立实现)验证(qwen3.5-0.8B + MoE,Ascend A2,TP=2):
Phase 5 — CP 分布式推理
目标:Ascend A2 上 CP=2 单卡/多卡 Greedy 生成结果一致。CP 涉及 Ring Attention 或 All-to-All 通信(参考
core/context_parallel/的 Ulysses / Colossal AI 实现),KV Cache 的 seq 维度在各 rank 间分片,causal mask 需适配本地 seq 范围,KVCache 需扩展分片管理能力。验证(qwen3.5-0.8B,Ascend A2,CP=2):
对外 API
from hyper_parallel.generate import generate, GenerationConfig # 返回值: torch.LongTensor, shape (batch_size, max_total_len) # 每个 batch item = [去 padding 的 real_prompt | generated_tokens | pad 对齐] # Phase 1/2: 单 prompt config = GenerationConfig(max_new_tokens=128, do_sample=True, top_k=50) output = generate(model, input_ids, config) # Phase 3: Batch prompts (left-padded) output = generate(model, batch_ids, config, attention_mask=attn_mask) # Phase 4: TP 推理(模型需先完成 TP 切分) # generate() 内部自动 all-gather logits 后采样,用户无需额外操作 output = generate(tp_model, input_ids, config) # Phase 5: CP 推理(模型需先完成 CP 切分) # generate() 内部适配 seq 分片的 mask 构造与 KVCache 管理 output = generate(cp_model, input_ids, config) # Mixin 便捷方式(对齐 HF 的 model.generate() 调用习惯) from hyper_parallel.generate.mixin import GenerateMixin class MyModel(GenerateMixin, nn.Module): ... output = model.generate(input_ids, config)使用约束
"logits",可选"past_key_values"@torch.no_grad(),不产生梯度图logits_processor/stopping_criteria)Phase 1 通过私有方法预留扩展点,不暴露完整公共接口,后续可自然演进temperature=0不被允许,极端 greedy 用do_sample=False测试设计
单元测试(CPU / stub)
Batch 测试(CPU)
模型验证(qwen3.5 + qwen3.5 MoE,Ascend A2)
分布式测试(Ascend A2)
回归测试
规格 & 约束
参考
hyper_parallel/core/context_parallel/hyper_parallel/core/tensor_parallel/hyper_parallel/integration/llamafactory/