mindspore-ai / mindspore-ai/hyper-parallel

【RFC】HyperParallel支持LLM推理Generate流程 #2101

Open
#709 0 comments 0 reactions 0 assignees View on GitHub

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 自动选择解码策略
  • 默认 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)

参考

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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reviewing the proposed hyper_parallel/generate/ files and the model forward contract, then inspect core/context_parallel/ for the existing Ulysses and Colossal AI communication patterns. Use a CPU stub model to validate the Phase 1 and Phase 2 checks listed as UT-01 through UT-09 before investigating batch, TP, and CP behavior. Done means the phased generation flow and its listed validation checks work, including single-card and distributed consistency.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.