已开启
【RFC】HyperParallel支持LLM推理Generate流程 #2101 #181
Moy创建于  6月1日
Moy
Moy
6月1日 创建

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 自动选择解码策略
  • 默认 Greedydo_sample=False 时为贪心搜索,do_sample=True 时切换到多项式采样
  • 参数体系max_new_tokens(控制生成长度)、temperature / top_k / top_p(控制采样)、repetition_penalty(惩罚重复)、eos_token_id(停止条件)
  • Batch 支持:通过 padding_side="left" + attention_mask 处理不等长 batch prompt
  • 扩展机制logits_processorstopping_criteria 可插拔自定义逻辑

HyperParallel generate 的设计目标是对齐 HF 的核心生成范式,同时适配分布式场景(TP/CP)的额外一致性需求。

核心价值:

  • 补齐 HyperParallel 推理生成能力,实现训推一体闭环
  • 对齐 HF generate 的核心使用范式,降低用户学习成本
  • 支持 Greedy / Top-K / Top-P 多种采样策略及 repetition_penalty
  • 内置 KV Cache 管理,避免 Decode 阶段重复计算历史 token
  • 分布式推理下保证单卡/多卡生成结果一致

功能描述

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 行为)

策略 触发条件 行为
Greedy do_sample=False(默认) argmax,确定性输出
Top-K do_sample=True, top_k>0 保留 top-k logits,softmax 后多项式采样
Top-P do_sample=True, top_p<1.0 核采样,累计概率 ≤ p 的最小 token 集合
Repetition Penalty repetition_penalty != 1.0 逐 batch item 独立惩罚已出现 token

3. KV Cache 管理

对标 HF 的 past_key_values 机制,存储每层 key/value 张量,避免 Decode 阶段重复计算。

  • 格式List[Tuple[Tensor, Tensor]],每层 (K, V),形状 (B, num_heads, seq_len, head_dim)
  • update:Prefill 后首次写入 / Decode 后替换为新值
  • merge:将历史 cache 与新 token 的 KV 沿 seq 维度拼接
  • detach:断开计算图(所有 generate 操作在 @torch.no_grad() 下)
  • clear:释放缓存

4. Batch 推理支持

对齐 HF 的 padding_side="left" + attention_mask 模式:

  • 接收 2D attention_mask (Batch_size, Seq_len),1=real token,0=padding
  • 自动推导 left-padding 下的 position_ids(padding 位置给 0,real token 从 0 递增)
  • 逐 batch item 独立追踪 EOS 停止状态(类似 HF 的 UnfinishedSequenceLogitsProcessor
  • 输出时逐 item 去 padding、拼接生成结果、pad 对齐到统一长度

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. 文件组织

hyper_parallel/generate/
├── __init__.py      # 导出 generate, GenerationConfig, KVCache, samplers
├── generation.py    # generate() — 核心 Prefill + Decode 主循环
├── sampler.py       # greedy_sample / top_k_sample / top_p_sample / _apply_repetition_penalty
├── kv_cache.py      # KVCache — update / merge / clear
├── utils.py         # GenerationConfig, build_causal_mask, build_position_ids
└── mixin.py         # GenerateMixin — model.generate() 便捷方法

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 的关键差异及设计考量:

  • 扩展机制(分阶段策略)
    • Phase 1:GenerationConfig预留 logits_processorstopping_criteria 字段(类型为 Optional[List[Callable]],默认 None),generate 循环内暂时忽略这两个字段。API 层面用户可见,文档标注"首期未实现,传值不生效"。generate 主循环内部通过 _apply_logits_processors_check_stopping_criteria 两个私有方法隔离采样/停止逻辑。
    • Phase 2+:实现 LogitsProcessor / StoppingCriteria 抽象基类,generate 循环读取 config.logits_processor / config.stopping_criteria 并应用。上层 API 不变(字段已在 Phase 1 定义),仅行为从"忽略"变为"生效"。
  • Beam search:Phase 1 聚焦 greedy 和 sampling 两种策略。Beam search 涉及多假设维护、KV Cache 结构变更(beam_size × batch_size)、分布式聚合逻辑变化,复杂度有数量级差异,有待考虑后续实现。

3. 生成主循环

generate(model, input_ids, generation_config, attention_mask=None)
  │
  ├─ Phase 1 — Prefill
  │   ├─ causal mask: (1, 1, S, S), 上三角 -inf
  │   ├─ position_ids: 感知 left-padding(padding→0, real→0,1,2,...)
  │   ├─ model(input_ids, position_ids, causal_mask, past_key_values=None)
  │   └─ cache.update(past_key_values)
  │
  ├─ Phase 2 — Decode (循环 max_new_tokens 次)
  │   ├─ logits[:, -1, :] → rep_penalty → greedy/top_k/top_p → next_tokens (B, 1)
  │   ├─ 逐 item EOS 检测 → 更新 is_finished
  │   ├─ 全部 finished → break
  │   ├─ position_ids = prompt_lengths + step  (B, 1)
  │   ├─ model(next_tokens, position_ids, attention_mask=None, past_key_values=cache.past_key_values)
  │   └─ cache.update(new_past_key_values)
  │
  └─ 输出: torch.LongTensor,shape (batch_size, max_total_len),
  			max_total_len =  max(prompt_len_i + generated_len_i),
            即 batch 中最长样本的总长度,不足此长度的样本在右侧用 pad_token_id 填充

4. 与模型侧的接口契约

Generate 模块不 import 任何具体模型类,通过行为契约对接(与 HF 的 GenerationMixin 设计理念一致:解耦生成逻辑和模型实现):

契约点 约定
模型前向签名 forward(input_ids, position_ids, attention_mask, past_key_values)dict{logits, past_key_values}
past_key_values List[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 位置给 0
logits (batch_size, seq_len, vocab_size),Decode 取 [:, -1, :]

attention_mask 不是纯 causal mask。Prefill 阶段 generate 模块负责将 causal mask 与 padding mask 合成为最终传给模型的 attention_mask:causal 部分(上三角 -inf)保证自回归约束,padding 部分(填充列 -inf)保证模型忽略无效位置。两者的合成逻辑由 generate 模块内部的 _build_combined_mask 完成。


实施计划

能力渐进叠加,每个 Phase 在前一个基础上增加一层能力,每个 Phase 完成后可独立验证。

Phase 1: Greedy 单 prompt          ← 最简,验证核心循环正确
  └─→ Phase 2: + KV Cache + 采样   ← 效率 + 多样性
       └─→ Phase 3: + Batch        ← 吞吐
            └─→ Phase 4: + TP 推理  ← HyperParallel 核心价值:单卡/多卡一致性
                 └─→ Phase 5: + CP 推理  ← Ring/All-to-All,复用 core/context_parallel/

Phase 1 — Greedy 单 prompt(基础闭环)

目标:在 qwen3.5 模型上跑通最简 greed 生成。无 KV Cache、无采样策略、单条 prompt。

Step 内容 产出
1.1 GenerationConfig dataclass + __post_init__ 校验 utils.py
1.2 build_causal_mask + build_position_ids utils.py
1.3 greedy_sample(logits)(batch_size, 1) sampler.py
1.4 generate() 最简循环:Prefill → 逐 step greedy 采样 → 拼接输出。全程 @torch.no_grad() generation.py

验证

CI 自动测试使用 stub 模型(CPU,秒级);本地手动验证使用 qwen3.5-0.8B(CPU,需 HF checkpoint)。二选一测试。

ID 检查项
UT-01 GenerationConfig 默认值 + 非法参数校验
UT-02 输入 (1, 4),max_new=8 → 输出 shape (1, 12)
UT-03 同一输入两次调用 → 输出完全一致(greedy 确定性)
UT-04 max_new_tokens=3 → 输出 ≤ prompt_len + 3
UT-05 eos 命中时提前停止

Phase 2 — KV Cache + 采样

目标:Decode 效率不随生成长度线性增长,支持非确定性采样。

Step 内容 产出
2.1 KVCache 类:update / merge / clear kv_cache.py
2.2 generate 接入 KV Cache:Prefill 后 cache.update,Decode 每步传 cache.past_key_values generation.py
2.3 top_k_sample + top_p_sample + _apply_repetition_penalty sampler.py

验证(CPU / stub 模型。KV Cache 与采样逻辑与模型无关,stub 模型即可充分验证):

ID 检查项
UT-06 KVCache 空/更新/merge/clear
UT-07 KV Cache 开/关,同一 prompt 生成结果完全一致
UT-08 Top-K/Top-P 输出合法 token
UT-09 Repetition penalty 生效

Phase 3 — Batch prompts

目标:支持 left-padded batch 输入,逐 item 独立 EOS 停止。

Step 内容 产出
3.1 generate() 新增 attention_mask 参数 generation.py
3.2 _build_prefill_position_ids:left-padding 感知 generation.py
3.3 逐 item EOS + 输出去 padding/pad 对齐 generation.py

验证(CPU / stub 模型。Batch 拼接与 EOS 逻辑与模型无关):

ID 检查项
BATCH-01 batch_size=2 等长输出 shape 正确
BATCH-02 left-padded 不等长 batch,greedy 两次一致
BATCH-03 不传 attention_mask 时行为与 Phase 2 一致

Phase 4 — TP 分布式推理

目标:Ascend A2 上 TP=2 单卡/双卡 Greedy 生成结果逐 token 一致。HyperParallel generate 区别于 HF 原生 generate 的核心能力

Step 内容 产出
4.1 TP 推理:all-gather logits → 全局采样 → 单 token forward generation.py(或 TP 模块中独立实现)
4.2 性能基线:Prefill latency + Decode tokens/s 性能数据

CP 推理见 Phase 5。

验证(qwen3.5-0.8B + MoE,Ascend A2,TP=2):

ID 检查项
DIST-01 TP=2 Greedy vs 单卡,逐 token 完全一致
DIST-02 Prefill latency + Decode tokens/s

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 需扩展分片管理能力。

Step 内容 说明
5.1 CP Prefill 适配:seq 分片下的 causal mask + position_ids 复用 CP 模块的 A2A 通信原语
5.2 CP Decode 适配:单 token(seq=1)天然不涉及切分 改动量小
5.3 KV Cache 分片管理:每 rank 只持有本地 seq 段的 KV 需扩展 KVCache 类

验证(qwen3.5-0.8B,Ascend A2,CP=2):

ID 检查项
DIST-03 CP=2 Greedy vs 单卡,结果一致
DIST-04 CP=2 Decode tokens/s

对外 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)

使用约束

  • 模型 forward 必须返回 dict 含 "logits",可选 "past_key_values"
  • generate 全程 @torch.no_grad(),不产生梯度图
  • Beam search 不在本次范围,计划通过独立 RFC 在后续版本实现(涉及 KV Cache 结构、分布式聚合逻辑的多处变更)
  • 扩展机制(logits_processor / stopping_criteria)Phase 1 通过私有方法预留扩展点,不暴露完整公共接口,后续可自然演进
  • temperature=0 不被允许,极端 greedy 用 do_sample=False
  • TP(Phase 4)和 CP(Phase 5)均为本次交付项

测试设计

单元测试(CPU / stub)

用例 ID 描述 期望
GEN-UT-01 GenerationConfig 默认值与参数校验 合法通过,非法抛异常
GEN-UT-02 Greedy 单 prompt 输出 shape (1, prompt+max_new)
GEN-UT-03 Greedy 确定性(同输入两次) 完全一致
GEN-UT-04 max_new_tokens 被遵守 输出 ≤ prompt + max_new
GEN-UT-05 EOS 触发提前停止 输出 < prompt + max_new
GEN-UT-06 KVCache 空/更新/合并/清理 操作正确
GEN-UT-07 KV Cache 开/关结果一致 完全一致
GEN-UT-08 Top-K/Top-P 输出合法性 值在 vocab 范围内

Batch 测试(CPU)

用例 ID 描述 期望
GEN-BATCH-01 batch_size=2 等长输出 shape (2, ≥prompt_len)
GEN-BATCH-02 left-padded batch greedy 确定性 两次一致
GEN-BATCH-03 不传 attention_mask 兼容 Phase 2 行为一致

模型验证(qwen3.5 + qwen3.5 MoE,Ascend A2)

用例 ID 描述 期望
GEN-QWEN-01 qwen3.5-0.8B base Greedy 生成合法文本 decode 后为可读文本
GEN-QWEN-02 qwen3.5 MoE Greedy 生成合法文本 decode 后为可读文本
GEN-QWEN-03 两模型 Top-K 输出多样性(两次调用不同) 两次输出不全等

分布式测试(Ascend A2)

用例 ID 描述 期望
GEN-DIST-01 TP=2 Greedy vs 单卡逐 token 一致 完全一致
GEN-DIST-02 TP=2 Prefill latency + Decode tokens/s 正常输出
GEN-DIST-03 CP=2 Greedy vs 单卡结果一致 完全一致
GEN-DIST-04 CP=2 Decode tokens/s 正常输出

回归测试

  • 训练链路不受影响(generate 与训练路径隔离)
  • Stub 模型测试可在 CPU CI 运行(确保 PR 独立可测)

规格 & 约束

  • 规格:实现 Prefill + Decode 基础 generate,支持 qwen3.5 / qwen3.5 MoE 两个模型可用
  • 额外功能:KV Cache、repetition_penalty、多采样策略、batch reasoning、GenerateMixin、stub 模型独立可测
  • 性能:NA(首期记录 Prefill latency + Decode tokens/s 基线)
  • 约束:首期不支持 beam search / logits_processor / stopping_criteria
  • 环境:Python 3.10 / PyTorch 2.6 / MindSpore>=2.8 / CANN 8.5(性能基线需 2×Ascend A2)

参考

likedislike
MoyMoy
6月1日 修改了issue 的描述
MoyMoy
6月2日 修改标题为 “【RFC】HyperParallel支持LLM推理Generate流程 #2101”,原标题为“【FRC】HyperParallel支持LLM推理Generate流程 #2101”