mindspore-ai / mindspore-ai/hyper-parallel
【RFC】HyperParallel支持LLM推理Generate流程 #2101
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
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自动选择解码策略 - 默认 Greedy:
do_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_processor、stopping_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_processor和stopping_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 定义),仅行为从"忽略"变为"生效"。
- 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)
参考
- HF generate 教程:https://huggingface.co/docs/transformers/main/en/llm_tutorial
- HF GenerationConfig:https://huggingface.co/docs/transformers/main/en/main_classes/text_generation#transformers.GenerationConfig
- 上下文并行:
hyper_parallel/core/context_parallel/ - 张量并行:
hyper_parallel/core/tensor_parallel/ - LLaMAFactory 集成:
hyper_parallel/integration/llamafactory/ - SIG 仓库:https://atomgit.com/mindspore/community/tree/master/sigs/parallel_training_system
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 181
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/181
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the proposed hyper_parallel/generate/ layout, especially utils.py, sampler.py, kv_cache.py, and generation.py, then inspect core/context_parallel/ for the later CP phase. Begin with the Phase 1 stub-model checks UT-01 through UT-05; completion spans the staged generation, cache, batching, TP, and CP milestones with their listed validations.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning, testing-qa
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 25/100