【Agentic RL / 强化学习框架】Molt 设计解读

【Agentic RL / 强化学习框架】Molt 设计解读

目录

0x00 概要

Molt 的定位是:能看懂、改得动、又不在规模上妥协的 LLM RL 训推框架。12K 行 Python, 55 个文件, PyTorch-native, Apache 2.0 — NVIDIA 出品。

molt

其主要信息如下:


0x01 从 RL 流程到 Molt 代码:一次完整映射

本文所有内容按三层结构组织:

  第一层: 标准 RL 训练流程 (8 步)         ← 你在任何 RL 教材里都见过的流程
  第二层: 每步对应 Molt 的哪个模块       ← 模块的文件路径
  第三层: 模块里的哪些类、它们做什么      ← 代码级,含类名、方法、谁调谁

读完这一章,我们就可以知道 Molt 的全貌。后续每章对应一个步骤深度展开。


1.1 第一层:标准 RL 训练流程

LLM 强化学习后训练的 8 个步骤:

molt-标准 RL 训练流程


1.2 第二层:Molt 模块对应关系

  步骤             Molt 模块                         文件路径
  ─────           ──────────                       ────────
  ① Dataset    datasets/prompts_dataset.py
               datasets/seqlen_balancing.py

  ② Rollout rollout/samples_generator.py        ← 调度入口
               rollout/router.py                ← Ray actor + vLLM 通信
               agents/base.py                   ← 循环控制器
               agents/chat_agent.py             ← 黑盒模式

  ③ Reward  agents/base.py → Env.step()         ← 你写的评分代码
               workers/reward_model_actor.py     ← 奖励模型(可选)

  ④ Exp      rollout/samples_generator.py        ← _process_response_into_experience()
               algorithm/experience.py           ← Experience 数据类

  ⑤ Adv      algorithm/advantage.py             ← GAE / GRPO

  ⑥ Policy  workers/policy_actor.py             ← FSDP 策略模型
               models/loss.py → PolicyLoss       ← PPO loss

  ⑦ Value   workers/critic_actor.py             ← 价值模型
               models/loss.py → ValueLoss        ← clipped value loss

  ⑧ Sync    trainer/fsdp/refit.py               ← flatten + broadcast
               trainer/vllm/vllm_worker_wrap.py  ← vLLM 加载

1.3 第三层:类定义 + 角色 + 交互

① Dataset

  类: PromptDataset              (datasets/prompts_dataset.py)
  角色: 从 JSONL 读 prompt, 按序列长度平衡采样

  方法: sample(batch_size)
    → 按 prompt_len 排序 → 分桶 → 每桶随机取 1 条
    → 返回 List[str] text prompts

  谁调用: SamplesGenerator.generate_samples()
    → _collect_prompt_batch(dataloader_iter, num_prompts)

② Rollout

  ┌─ SamplesGenerator (rollout/samples_generator.py:89)
  │  角色: 采样调度总入口
  │  方法: generate_samples()
  │    1. _collect_prompt_batch() → 从 dataset 拿 prompt
  │    2. _dispatch_to_agent_runners(prompts)
  │       → round-robin 分给 self.agent_runners[i] (AgentRunnerActor)
  │       → actor.run_group.remote(prompt, n_samples)  (Ray remote call)
  │    3. ray.wait() 等 rollout 完成 → 收集 Trajectory[]
  │    4. _process_response_into_experience(trajectories) → Experience[]
  │    5. replay_buffer.push(experiences)
  │
  ├─ AgentRunnerActor (rollout/router.py:283)
  │  角色: @ray.remote, 每个独立进程跑一个
  │  构造函数:
  │    self._runner = load_agent_runner(agent_path) ← 你的 StepEnvRunner
  │    self._client = RouterGenerateClient(http)     ← vLLM 客户端
  │  方法: run_group(prompt, n_samples)
  │    for _ in range(n_samples):
  │      traj = await self._runner.execute(llm_engine=self._client)
  │    return Trajectory[]
  │
  ├─ StepEnvRunner (agents/base.py:286)
  │  角色: 单个 rollout 的循环控制
  │  方法: execute(prompt, label, llm_engine, ...)
  │    1. env = self.env_cls()       ← 创建你的 Env
  │    2. env.reset(state)           ← 初始化
  │    3. tokens = tokenize(observation)
  │    4. loop (multi-turn):
  │         A) await llm_engine.generate(tokens, params)
  │              → HTTP POST → VllmRouterActor → vLLM Engine
  │              → return tokens + logprobs
  │         B) decode tokens → action_text
  │         C) result = await env.step({action_text, label, ...})
  │              → 你的 Env 评分 → Result(reward, observation, terminated)
  │         D) trajectory.append_action(tokens, logprobs)
  │         E) if terminated: break
  │            else: tokenize(result.observation) → goto A
  │    5. return Trajectory
  │
  ├─ VllmRouterActor (rollout/router.py:45)
  │  角色: 启动 Rust HTTP 负载均衡进程
  │  路由: consistent_hash(session_id) → 把请求送到固定 vLLM Engine
  │  通信: AgentRunnerActor 的 HTTP client → VllmRouterActor → Engine
  │
  └─ RouterGenerateClient (rollout/router.py:163)
      角色: HTTP client 包装
      方法: generate(token_ids, params) → POST /inference/v1/generate

③ Reward

  类: Env (你在 agent 文件里实现)          (agents/base.py)
  角色: 你的评分逻辑
  方法: step(state) → Result
    输入: state = {"action_text": str, "label": Any, ...}
    输出: Result(reward=torch.tensor(1.0), observation="...", terminated=True)

  类: RewardModelActor (可选)              (workers/reward_model_actor.py)
  角色: preference-based 奖励模型
  Trainer 调用: reward = ray.get(reward_actor.train.remote(experience))

④ Experience Assembly

  方法: _process_response_into_experience()  (samples_generator.py:461)
  角色: Trajectory → Experience
  步骤:
    1. trajectory_tokens → torch.tensor → sequences
    2. _build_action_token_mask() → action_mask
    3. 提取 action_log_probs, reward, scores
    4. 组装 Experience()

  数据类: Experience                        (algorithm/experience.py)
  字段:
    sequences        [batch, total_seq_len]
    attention_mask   [batch, total_seq_len]
    action_log_probs [batch, response_len]   ← rollout logprobs
    base_log_probs   [batch, response_len]   ← reference logprobs
    advantages       [batch, response_len]   ← GAE/GRPO 结果
    returns          [batch, response_len]   ← discounted return
    reward           [batch, 1]
    action_mask      [batch, response_len]   ← 哪些 token 参与 loss

⑤ Advantage

  函数: estimate(experiences)               (algorithm/advantage.py)
  角色: 从 Experience 计算 advantage
  方法: GAE (Generalized Advantage Estimation) 或 GRPO
  输出: advantages, returns (per-token)

  谁调用: RLTrainer._train()
    → batch = replay_buffer.sample()
    → advantages = estimate(batch)

⑥ Policy 更新

  类: PolicyModelActor                     (workers/policy_actor.py)
  角色: FSDP 包装的策略模型, Ray actor
  方法: train(experience)
    1. forward → logits → log_probs
    2. PolicyLoss.forward() → loss
    3. backward → FSDP all-reduce → optimizer.step()

  类: PolicyLoss                           (models/loss.py)
  角色: PPO-Clip + Importance Sampling
  公式:
    ratio = exp(current_logprobs - rollout_logprobs)
    surr1 = ratio * advantages
    surr2 = clamp(ratio, 1-eps, 1+eps) * advantages
    loss = -mean(min(surr1, surr2) * mask)

⑦ Value 更新

  类: CriticTrainer                        (workers/critic_actor.py)
  角色: 独立的价值模型 Ray actor
  方法: train(experience)
    1. forward → values
    2. ValueLoss.forward(values, returns) → v_loss
    3. backward → mean all-reduce (value head 不 FSDP wrap)

  类: ValueLoss                            (models/loss.py)
  公式: 0.5 * mean(max((v-r)^2, (clipped(v)-r)^2))

⑧ 权重同步

  类: RefitHelper                          (trainer/fsdp/refit.py)
  角色: FSDP shard → flatten → NCCL broadcast → vLLM
  步骤:
    1. strategy.unshard() → 完整权重
    2. flatten_params() → 1D 连续 tensor
    3. NCCL broadcast (rank 0 → ref model workers)
    4. NCCL broadcast (rank 0 → vLLM workers)

  类: MoltVLLMWorker                       (trainer/vllm/vllm_worker_wrap.py)
  角色: 接收 flat tensor, unpack, load_state_dict()

1.4 交互总图:谁调了谁

molt-交互总图

1.5 数据流

图的内容:

  • 单主链流程:CLI → RLTrainer(rollout_thread / train loop 两个线程)→ SamplesGenerator → AgentRunnerActor → StepEnvRunner
  • 分支 1:StepEnvRunner 同一 loop 内调 vLLM 链路(VllmRouterActor → vLLM Engine)和你的 Env,Trajectory 回传
  • 数据汇合:ExperienceMaker(_process_response_into_experience)→ ReplayBuffer
  • 训练链路:advantage.estimate → PolicyModelActor / CriticTrainer → PolicyLoss / ValueLoss
  • 权重同步:RefitHelper(NCCL broadcast)→ MoltVLLMWorker + ReferenceModelActor,⑱ 闭环回下一轮 rollout

molt-数据流

--- 后续每章对应 ” 第三层:类定义 + 角色 + 交互“ 的一个步骤展开 ---

0x02 第一部分:Dataset — 第 ① 步

2.1 在 RL 流程中的位置

RL 流程的第一步:从训练集中采样 prompt 文本,经过 tokenize 和长度均衡处理后交给 rollout。

2.2 模块与类

  文件: datasets/prompts_dataset.py

  类: PromptDataset
    方法: __getitem__(index) → Dict (prompt, label, images, ...)
    谁调用: PyTorch DataLoader → SamplesGenerator._collect_prompt_batch()
    输出: 原始文本 (未 tokenize)

  文件: datasets/utils.py
  函数: _tokenize_fn(batch, tokenizer)
    谁调用: PromptDataset 内部
    输出: tokenized tensors

  文件: datasets/seqlen_balancing.py
  函数: balance(prompts, batch_size)
    1. 按 prompt_len 排序
    2. 分成 len(dataset)/batch_size 个桶
    3. 每桶随机采样 1 条
    效果: 同 batch 内序列长度接近 → padding 少 → GPU 利用率高

2.3 数据流形状

  JSONL 文件
    │  {"prompt": "What is 2+2?", "label": "4"}
    │  {"prompt": "...", "label": "..."}
    ▼
  PromptDataset.__getitem__()
    │  Dict[prompt: str, label: Any, images: List[PIL]?, tools: List[Dict]?]
    ▼
  _collect_prompt_batch(dataloader_iter, num_prompts)
    │  List[str] prompts, List[Any] labels
    ▼
  → SamplesGenerator → AgentRunnerActor → StepEnvRunner
    (后续在 StepEnvRunner.execute() 内部 tokenize)

2.4 三种数据集的区分

数据集 文件 用途
PromptDataset prompts_dataset.py RL 训练时的 prompt 采样
SFTDataset sft_dataset.py SFT 阶段使用
RewardDataset utils/make_reward_dataset.py Reward Model 的 preference pairs

0x03 第二部分:Rollout — 第 ② 步(最复杂的环节)

3.1 在 RL 流程中的位置

这是 Molt 最复杂的部分。模型拿到 prompt 后生成 completion(action),如果有多轮环境则反复生成直到终止。涉及 5 个类和 2 个通信层。

3.2 SamplesGenerator — 调度总入口

  文件: rollout/samples_generator.py
  类: SamplesGenerator                          (第 89 行)

  角色: 管理整个采样管道

  构造函数参数:
    - prompts_dataloader: DataLoader            ← 数据源
    - agent_runners: List[AgentRunnerActor]     ← Ray actors 列表
    - tokenizer: PreTrainedTokenizer            ← 文本编解码

  核心方法: generate_samples(**kwargs) → List[Experience]
    流程:
      1. _collect_prompt_batch(dataloader_iter, num_prompts)
         → 从 DataLoader 取 up to num_prompts 条 prompt
      2. _dispatch_to_agent_runners(prompts, labels, **kwargs)
         → round-robin: for each prompt:
             actor = self.agent_runners[self._rr % N]
             actor.run_group.remote(prompt, label, n_samples, ...)
         → 返回 List[ObjectRef] (Ray remote call 的引用)
      3. ray.wait(refs, num_returns=1)
         → 等最先完成的 1 个 rollout
      4. _filter_group(finished_rollout, ...)
         → 检查过滤条件,去掉无效 rollout
      5. _process_response_into_experience(trajectory)
         → Trajectory → Experience
      6. 重复 3-5 直到收满 rollout.batch_size 条
      7. replay_buffer.push(experiences)

  关键设计: 异步流水线
    - 每次只等 1 个 rollout 完成(不是等全部)
    - 未完成的 rollout 留在 self._inflight_rollouts 里
    - 下次调用 generate_samples() 可以直接收它们
    - vLLM 不需要 drain → 每一步之间不间断

3.3 AgentRunnerActor — Ray 执行单元

  文件: rollout/router.py                       (第 283 行)
  类: @ray.remote class AgentRunnerActor

  角色: 每个独立进程跑一个的执行器

  构造函数:
    self._runner = load_agent_runner(agent_path)
      → 扫描你的 agent.py 文件
      → 找到 class AgentRunner(StepEnvRunner)
      → 实例化它 (e.g., StepEnvRunner(MathEnv))

    self._tokenizer = get_tokenizer(model_path)
    self._http = aiohttp.ClientSession(router_url)
    self._client = RouterGenerateClient(self._http)
      → HTTP client 连接 VllmRouterActor

  方法: run_group(prompt, label, images, sampling_params, max_length, n_samples)
    → 对同一 prompt 跑 n_samples 次 rollout:
      tasks = [
        self._runner.execute(
          llm_engine=self._client,     ← HTTP client
          hf_tokenizer=self._tokenizer,
          prompt=prompt, label=label,
          sampling_params=..., max_length=...,
        )
        for _ in range(n_samples)
      ]
      results = await asyncio.gather(*tasks)
      → 标记 group_id, rollout_id
      → 返回 Trajectory[]

  注意: n_samples > 1 时 (GRPO 模式),同批 rollout 并发执行

3.4 StepEnvRunner — 循环控制

  文件: agents/base.py                         (第 286 行)
  类: StepEnvRunner(Runner)

  角色: 单个 prompt 的完整 rollout 循环

  方法: execute(prompt, label, sampling_params, max_length,
                hf_tokenizer, llm_engine, images, tools)

  完整执行流程:

  ┌─ StepEnvRunner.execute() ──────────────────────────────────┐
  │                                                             │
  │  ① env = self.env_cls()     ← 创建你的 Env 的新实例         │
  │     reset = await env.reset({"observation": prompt,         │
  │                               "label": label})              │
  │     observation_text = reset["observation"]                 │
  │                                                             │
  │  ② obs_tokens, mm_inputs = tokenize(observation_text)      │
  │     trajectory = Trajectory(obs_tokens, ...)                │
  │                                                             │
  │  ③ 进入 while True 循环:                                    │
  │                                                             │
  │     A. 计算本轮 max_tokens = min(per_turn_cap, remaining)   │
  │        if max_tokens <= 0: break                            │
  │                                                             │
  │     B. await llm_engine.generate(                           │
  │          trajectory.observation_tokens,                      │
  │          turn_sp,                                           │
  │          multi_modal_data=...,                              │
  │          session_id=rollout_sid)                            │
  │          ↓                                                  │
  │        RouterGenerateClient → POST → VllmRouterActor        │
  │          → Rust consistent_hash routing → vLLM Engine       │
  │          ↓                                                  │
  │        tokens + logprobs + finish_reason + routed_experts   │
  │                                                             │
  │     C. action_text = decode(token_ids)                      │
  │                                                             │
  │     D. result = await env.step({                            │
  │          "observation_text": observation_text,              │
  │          "action_text": action_text,                        │
  │          "label": label,                                    │
  │          "sampling_params": turn_sp,                        │
  │        })                                                   │
  │          ↓                                                  │
  │        你的 Env.step() → Result(reward, observation,        │
  │                                   terminated, info)         │
  │                                                             │
  │     E. trajectory.append_action(action_tokens, logprobs)    │
  │        trajectory.reward += result.reward                   │
  │                                                             │
  │     F. if result.terminated: break                          │
  │        else:                                                │
  │          feedback_tokens = tokenize(result.observation)     │
  │          trajectory.append_feedback(feedback_text, tokens)  │
  │          → goto A                                           │
  │                                                             │
  │  ④ return Trajectory                                       │
  │                                                             │
  └─────────────────────────────────────────────────────────────┘

3.5 vLLM 通信层

  ┌─ VllmRouterActor (rollout/router.py:45) ───────────────────┐
  │                                                             │
  │  @ray.remote actor                                          │
  │  启动一个 Rust 子进程 (vllm_router.launch_router)           │
  │  作为 HTTP 代理放在所有 vLLM Engine 前面                    │
  │                                                             │
  │  路由策略: consistent_hash(session_id) → engine             │
  │  → 同一个 rollout_id 的多次 generate 永远去同一 engine      │
  │  → 多轮 KV cache 不被驱逐 + VLM render features 在同一节点  │
  │                                                             │
  │  端口: 30000 (默认)                                         │
  │  健康检查: ready() 等 router 绑定端口后返回 URL             │
  │                                                             │
  ├─ RouterGenerateClient (rollout/router.py:163) ─────────────┐
  │                                                             │
  │  aiohttp client 包装                                        │
  │  方法: generate(token_ids, sampling_params)                 │
  │    → 构造 body: {token_ids, sampling_params, features...}   │
  │    → POST /inference/v1/generate                            │
  │    → 解析 response: {token_ids, logprobs, finish_reason,   │
  │                       routed_experts}                       │
  │    → 返回 (request_output, off_policy_len)                  │
  │                                                             │
  └─────────────────────────────────────────────────────────────┘

  完整路径:
    StepEnvRunner.execute()
      → RouterGenerateClient.generate()     HTTP POST
        → VllmRouterActor (Rust proxy)      转发请求
          → vLLM Engine (worker)            实际推理

3.6 异步 Rollout 架构

传统同步模式的问题

  Step 1: rollout ──→ barrier (所有 GPU 等最慢的)
          ──→ sync weights (NCCL all-gather, 秒到分钟级)
          ──→ train
  Step 2: rollout ──→ barrier (又等一轮)
          ...

在 70B+ MoE 模型上,barrier + all-gather 可以消耗几十秒到几分钟。MoE 的 expert 负载不均——最慢的 expert 决定 rollout 延迟。

Molt 的异步解耦

  ┌─ Rollout Thread (独立线程) ───────────────────────────────┐
  │  while True:                                                │
  │    batch = dataset.sample()                                 │
  │    seqs = vllm.generate(batch)                              │
  │    traj = env_runner.run(seqs)                              │
  │    exp = experience_maker(traj)                             │
  │    replay_buffer.push(exp)                                  │
  │    if buffer is full: wait()    ← backpressure             │
  └────────────────────────────────────────────────────────────┘

  ┌─ Train Loop ───────────────────────────────────────────────┐
  │  while True:                                                │
  │    batch = replay_buffer.sample()                           │
  │    advantages = estimate(batch)                             │
  │    policy_loss = actor.train(batch)                         │
  │    if step % N == 0: sync_weights()                        │
  └────────────────────────────────────────────────────────────┘

Rollout 和 Train 之间没有 barrier。ReplayBuffer 是唯一的协同点。

异步的代价和补偿

异步意味着 rollout 用的策略权重可能比 train 的当前权重旧几轮(stale policy),用 importance sampling correction 补偿:

# models/loss.py
ratio = torch.exp(current_log_probs - rollout_log_probs)
# ↑ 训练时新 policy     ↑ rollout 时记录的旧 policy
# ratio 校正梯度权重: 新 policy 概率翻倍 → 梯度权重翻倍

3.7 两套 Agent API

为什么需要两套

"标准"RL 路径要求本地 vLLM 推理 + 完整的 per-token logprobs。但 Molt 也支持另一种场景:用 GPT-4/Claude 等远程 API 做 Agent,Molt 只负责收集 reward。

  ┌─ Agent 模式 (StepEnvRunner) ──────────────────────────┐
  │  模型 → 本地 vLLM(你有权重,FSDP 同步)                │
  │  logprobs → 精确 per-token 追踪                        │
  │  IS correction → 支持 stale policy 补偿                │
  │  性能 → GPU direct, 微秒级延迟                          │
  │  适用 → 高频实验、MoE 训推                              │
  └──────────────────────────────────────────────────────┘

  ┌─ Chat 模式 (ChatAgentRunner) ─────────────────────────┐
  │  模型 → 远程 API(只有 API key,无权重复制)            │
  │  token → session URL 追踪(数量,无 logprobs)         │
  │  IS correction → 不支持,纯 on-policy                 │
  │  性能 → 网络延迟,百毫秒级                              │
  │  适用 → 黑盒评测、第三方 API                            │
  └──────────────────────────────────────────────────────┘

ChatAgent 怎么工作

# agents/chat_agent.py
class ChatAgentRunner(Runner):
    """Runner 子类,用户不写 env.step(),而是写 agent.run()"""

class ChatAgent:
    """用户实现"""
    async def run(self, ctx: ChatContext) -> Result:
        client = AsyncOpenAI(base_url=ctx.base_url, api_key=ctx.api_key)
        resp = await client.chat.completions.create(
            model=ctx.model_name,
            messages=list(ctx.messages),
        )
        text = resp.choices[0].message.content or ""
        reward = grader.score_response(text, "", ctx.label or "")
        return Result(reward=torch.tensor(reward))

背后启动 FastAPI 服务器(_chat_server.py),做三件事:

  1. 暴露 /molt/chat 端点
  2. 代理请求到远程 API
  3. 追踪 session URL 获取 token 数量

为什么不合并

因为底层约束不同——StepEnvRunner 有模型权重所以能做 FSDP 同步 + per-token logprobs + IS correction;ChatAgentRunner 只有 API key 所以通通不能做。硬合并会在接口里塞满 if-else 分支。


0x04 第三部分:Reward — 第 ③ 步

4.1 在 RL 流程中的位置

模型生成 action 文本后,环境(Env)对 action 评分,返回 reward。这是你的算法最直接的介入点——你在 Env.step() 里定义评分规则。

4.2 Env 接口

# agents/base.py

class Env:
    async def reset(self, state):
        """可选:重置状态"""
        return {"observation": state["observation"]}

    async def step(self, state) -> Result:
        """你的评分逻辑在这里"""
        raise NotImplementedError

@dataclass
class Result:
    reward: torch.Tensor        # 标量奖励 (必须)
    observation: str = ""       # 下一轮的上下文 (多轮用)
    terminated: bool = False    # 是否终止
    truncated: bool = False     # 是否截断
    info: dict = {}             # 追踪用 metric
    score: Any = None           # 可选,默认等于 reward
    images: List = None         # VLM 多轮用
    sampling_params: Any = None # 覆盖下轮的采样参数

4.3 两个内置 Agent 的评分逻辑

单轮 (math.py)

class MathEnv(Env):
    async def step(self, state) -> Result:
        text = state["action_text"]
        result = _GRADER.score_response(text, "", state.get("label") or "")
        return Result(
            reward=torch.tensor(float(result.get("reward", 0.0))),
        )

多轮 (geo3k.py)

class GeoEnv(Env):
    async def step(self, state) -> Result:
        action = state["action_text"]
        label = state.get("label")
        self.turn += 1
        self.assistant_history.append(action)
        is_last_turn = self.turn >= _MAX_TURNS

        # 检查是否包含最终答案
        if self._has_answer(action):
            reward = self._final_reward(label)
            return Result(reward=reward, terminated=True)

        # 检查是否调用工具
        tool_call = _extract_tool_call(action)
        if tool_call:
            self.tool_call_count += 1
            obs_text = self._execute_tool(tool_call)
            return Result(reward=0.0, observation=obs_text, terminated=False)

        # 未知行为 → 终止并评分
        reward = self._final_reward(label)
        return Result(reward=reward, terminated=True)

4.4 Math Grader

  文件: examples/python/utils/math_grader.py

  支持的答案格式:
    <answer>42</answer>
    \boxed{42}
    \fbox{42}

  等价性:
    3 == 3.0 == 3.000
    1/2 == 0.5
    科学记数法: 1e5 == 100000

  输出:
    {"reward": 1.0, "score": 1.0, "missing_answer": 0.0}

4.5 RewardModelActor (可选)

  文件: workers/reward_model_actor.py

  角色: 用 preference pairs 训练奖励模型,替代规则评分

  使用: Trainer 里调用
    reward = ray.get(reward_actor.train.remote(experience))

  Molt 不强制使用 reward model——规则评分(grader)已经是 DAPO/GRPO 的核心思想

4.6 什么时候用规则 vs 奖励模型

场景 推荐
数学、代码、确定正确答案 规则评分 (grader)
对话、偏好、无绝对答案 Reward Model
混合 两者结合:grader 筛格式不对的,RM 评内容质量

0x05 第四部分:Experience 组装 — 第 ④ 步

Agent 返回的 Trajectory 包含原始 token_ids、logprobs、reward。Experience 组装把它转成 Trainer 能用的结构化张量。

5.1 Token-In-Trace-Out (TITO)

Molt 的核心理念:整个 pipeline 数据保持 token_id 形式,不转回文本

  "What is 2+2?" → tokenizer → [1432, 583, 17, 19]
                                     |
                                 vLLM 生成
                                     |
  [1432, 583, 17, 19,  ▒  ▒  ▒ ... ▒]
   ↑ prompt 部分 ↑     ↑ response 部分 ↑
                             |
                       每个 token:
                         - token_id (int64)
                         - logprob  (float32)
                         - [可选] value (float32)
                         - mask (bool, 是否参与 loss)

5.2 组装流程

  ┌─ Trajectory (agents/base.py) ───────────────────────────┐
  │  observation_tokens: List[int]      ← prompt + response  │
  │  action_ranges: List[Tuple[int,int]] ← response 位置     │
  │  action_logprobs: List[float]       ← rollout logprobs  │
  │  reward: float                      ← episode reward    │
  │  extra_logs: dict                   ← info 追踪         │
  │  mm_train_inputs: dict              ← VLM 输入          │
  │  routed_experts: ndarray            ← MoE routing 记录  │
  │  group_id, rollout_id: str          ← 全局唯一 ID       │
  └──────────────────────────────────────────────────────────┘
                            │
              _process_response_into_experience()
              (samples_generator.py:461)
                            │
                            ▼
  ┌─ Experience (algorithm/experience.py) ──────────────────┐
  │  sequences:        [batch, total_seq_len]    int64      │
  │  attention_mask:   [batch, total_seq_len]    int64      │
  │  action_log_probs: [batch, response_len]     float32    │
  │  base_log_probs:   [batch, response_len]     float32    │
  │  advantages:       [batch, response_len]     float32    │
  │  returns:          [batch, response_len]     float32    │
  │  reward:           [batch, 1]                float32    │
  │  action_mask:      [batch, response_len]     bool       │
  │  state_id:         List[str]                            │
  └──────────────────────────────────────────────────────────┘

5.3 _process_response_into_experience 做了什么

# samples_generator.py:461
def _process_response_into_experience(self, response, **kwargs):
    # 1. trajectory_tokens → torch.tensor → sequences
    trajectory_tokens = response.observation_tokens.copy()
    sequences = torch.tensor(trajectory_tokens, dtype=torch.long)

    # 2. 构造 attention_mask (全部有效)
    attention_mask = torch.ones(len(trajectory_tokens), dtype=torch.long)

    # 3. 构造 action_mask (哪些 token 是模型生成的)
    action_mask = _build_action_token_mask(
        len(trajectory_tokens),
        response.action_ranges,
        response.off_policy_action_lens,
    )

    # 4. 检查 VLM 截断安全
    if mm_train_inputs and len(trajectory_tokens) > max_len:
        if any image token truncated:
            return None, "vlm_truncation"  # 丢弃

    # 5. 提取 logprobs
    action_log_probs = response.extract_action_logprobs()

    # 6. 组装
    return Experience(
        sequences=sequences,
        attention_mask=attention_mask,
        action_log_probs=action_log_probs,
        action_mask=action_mask,
        reward=reward_val,
        ...
    ), None

5.5 数据形状变化汇总

阶段 表示 形状
Dataset {"prompt": str, "label": Any} 灵活
vLLM 生成 Sequences(prompt_ids + output_ids, logprobs) [seq_len] int64 + [seq_len, vocab] float32
Trajectory observation_tokens, action_ranges, logprobs 列表形式
Experience sequences, action_log_probs, advantages, masks [batch, seq_len] 张量
PolicyLoss loss, kl, approx_kl 标量

0x06 第五部分:训练 — 第 ⑤⑥⑦ 步

6.1 在 RL 流程中的位置

  第⑤步: Advantage 估计 — 从 Experience 计算优势函数
  第⑥步: Policy 更新    — PPO/GRPO loss, 策略梯度
  第⑦步: Value 更新     — Critic loss, 价值函数拟合

  这三步在 RLTrainer._train() 里连续执行:
    batch = replay_buffer.sample()
    advantages = estimate(batch)
    policy_actor.train.remote(batch)
    critic_actor.train.remote(batch)       # 可选

6.2 单 Actor 架构

Molt 只用 一个 PolicyModelActor(其他框架用多 actor 并行)。原因:在 70B+ 规模上,rollout 是瓶颈,不是 train。单 actor 避免了多 actor 的权重冲突和同步开销。

  你的 GPU:
    [FSDP shard 0]   [FSDP shard 1]   ...   [FSDP shard N-1]
    每卡只存 1/N 的参数量
    70B 在 32xH100 上每卡 ~2B 参数

6.3 Global Token Mean

分布式训练的关键设计:loss 除以 global token 数,不是 local token 数。

# 错误: DP 越多 loss 越小
loss = loss.sum() / local_num_tokens

# Molt: DP-invariant
global_num_tokens = all_reduce(local_num_tokens, op=SUM)
loss = all_reduce(loss.sum(), op=SUM) / global_num_tokens
# ↑ 增加 GPU 不改变 loss scale, 不重新调参

6.4 Advantage 估计 (algorithm/advantage.py)

  文件: algorithm/advantage.py

  两种模式:
    GAE (Generalized Advantage Estimation):
      delta_t = r_t + gamma * V(s_{t+1}) - V(s_t)
      advantage_t = sum((gamma * lambda)^k * delta_{t+k})
      需要 Critic 提供 V(s)

    GRPO (Group Relative Policy Optimization):
      advantage_i = (reward_i - mean(reward_group)) / std(reward_group)
      不需要 Critic, 同一 prompt 的 n_samples 之间比较

  谁调用: RLTrainer._train()

6.5 Policy 更新 (PolicyModelActor + PolicyLoss)

  文件: workers/policy_actor.py
  类: PolicyModelActor(BaseModelActor)
  角色: FSDP 包装的策略模型

  训练步骤:
    1. strategy.unwrap_model() → 获取原始模型
    2. forward(sequences, attention_mask) → logits
    3. 从 logits 提取 response 部分的 log_probs
    4. PolicyLoss.forward(log_probs, advantages, ...) → loss
    5. loss.backward() → FSDP reduce-scatter 梯度
    6. optimizer.step() → FSDP all-gather 更新参数
# models/loss.py
class PolicyLoss(nn.Module):
    def forward(self, data):
        # 1. 当前 policy 在 rollout tokens 上的 logprobs
        log_probs = compute_log_probs(data.sequences, data.response_ids)

        # 2. Importance Sampling ratio
        ratio = torch.exp(log_probs - data.action_log_probs)

        # 3. PPO-Clip
        surr1 = ratio * data.advantages
        surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * data.advantages
        policy_loss = -torch.mean(torch.min(surr1, surr2) * data.action_mask)

        # 4. KL penalty (对抗参考模型)
        kl_loss = kl_divergence(log_probs, data.base_log_probs)

        # 5. Entropy bonus (可选)
        entropy = -torch.mean(log_probs * data.action_mask)

        # 6. MoE auxiliary losses
        moe_aux_loss = load_balancing_loss(data) + z_loss(data)

        return {
            "loss": policy_loss + self.kl_weight * kl_loss
                    - self.entropy_weight * entropy + self.moe_weight * moe_aux_loss,
            "approx_kl": (data.base_log_probs - log_probs).mean().detach(),
        }

6.6 Value 更新 (CriticTrainer + ValueLoss)

  文件: workers/critic_actor.py
  类: CriticTrainer
  角色: 价值模型, 独立 Ray actor (可与 actor 同 GPU 或独立)
  特点:
    - Value Head 是标量输出 (不做 FSDP wrap, replicate)
    - 梯度: mean all-reduce over DP group
    - 支持比 actor 更多 epoch 的价值更新 (critic.max_epochs)
# models/loss.py
class ValueLoss(nn.Module):
    def forward(self, values, returns, old_values=None):
        if old_values is not None:
            value_clipped = old_values + torch.clamp(values - old_values, -self.clip_eps, self.clip_eps)
            v_loss1 = (values - returns).pow(2)
            v_loss2 = (value_clipped - returns).pow(2)
            value_loss = 0.5 * torch.mean(torch.max(v_loss1, v_loss2))
        else:
            value_loss = 0.5 * torch.mean((values - returns).pow(2))
        return value_loss

6.7 五种 Loss 信号

Loss 来源 作用
policy_loss PPO-Clip + IS ratio 策略梯度
kl_loss Reference Model 对比 约束策略偏移
entropy log_probs 均值 鼓励探索 (可选)
load_balancing_loss MoE expert 负载 防止 collapsed experts
z_loss MoE router logits 正则化门控输出

0x07 第六部分:权重同步 — 第 ⑧ 步

7.1 在 RL 流程中的位置

训练更新了 FSDP 中的模型参数,但 vLLM Engine 不知道。需要把 FSDP shard 还原成完整权重 → 广播到 vLLM worker。

7.2 三步流程

  Step 1: FSDP → 扁平化

    PolicyModelActor (FSDP)
      │ strategy.unshard() → 所有 GPU all-gather 自己的 shard
      │ 每 rank 拿到完整参数
      │
      │ refit.pack_shards():
      │   1. 遍历 model.state_dict()
      │   2. flatten 每个参数 → 1D
      │   3. 拼成一个连续 flat tensor
      │   4. 记录 meta (name, dtype, shape, offset)
      │
      ▼
    flat_tensor: [θ₁, θ₂, θ₃, ..., θₙ]  ← rank 0 上

  Step 2: NCCL broadcast

    rank 0 ──NCCL broadcast──→ ReferenceModelActor
          ──NCCL broadcast──→ vLLM Worker 0
          ──NCCL broadcast──→ vLLM Worker N
    70B 模型 ~140GB @ IB 400Gbps → ~3 秒

  Step 3: vLLM 加载

    MoltVLLMWorker:
      receive flat tensor + meta
      refit.unpack() → dict[name] = tensor
      load_state_dict(state_dict)
      → model 权重和 Actor 完全一致
      → vLLM Engine 就绪

7.3 为什么不用磁盘

  磁盘 I/O (NVMe):  70B × 2 bytes = 140GB @ 14GB/s → 10 秒 + 读回来又 10 秒
  NCCL broadcast: 140GB @ 400Gbps IB → ~3 秒

7.4 关键参数

参数 默认 说明
weights_sync_interval 0.5 每 0.5 个 RL step 同步一次
sync_using_ray False 用 Ray object store 替代 NCCL

0x08 第七部分:跨环节主题

以下三个主题不单独属于某个 RL 步骤,而是贯穿多步的设计。


8.1 训推一致性

问题

MoE 模型训推不一致的后果:

  场景: 训练 bf16, 推理 fp16 → router gate logits 差 0.1%
  结果: token 在训练时走 expert A, 推理时走 expert B
       → 梯度对齐了 expert B——但训练改的是 expert A 的参数

  场景: 训练 router 随梯度更新, 推理冻结
  结果: 训练最后几步 router 变了, rollout 用的是上一轮同步的 router
       → routing 决策不同, 整个 representation 分布不同

Molt 的六层保障

# 保障 代码位置 做了什么
重要性采样校正 models/loss.py ratio = exp(new - old) 补偿 stale policy
Router Freezing models/actor.py 训练冻结 MoE gate 参数,不更新
Router Replay models/actor.py 推理时记 routing → 训练回放
fp32 Router models/actor.py Gate logits 始终 fp32,不受混合精度影响
权重同步 fsdp/refit.py FSDP → NCCL broadcast → vLLM
enforce_eager trainer/vllm/vllm_engine.py vLLM 不建 CUDA graph,避免数值差异

六层如何协同

  Rollout:               Train:
    vLLM (最新权重)        FSDP (更新中)

    ① Router 决策          ② Router 冻结 (不更新)
       (gate fp32)

    ③ 记录 routing         ④ Router Replay: 回放 ③
       (token→expert)        (同一 token → 同一 expert)

    ⑤ 记录 logprobs        ⑥ Importance Sampling:
                             ratio = exp(new - old)

    ⑦ 权重同步: FSDP → NCCL broadcast → vLLM

    → 每层消除一类训推差异, 六层叠加 = MoE 训推一致性

8.2 MoE + Disaggregation

MoE 支持的四个层面

  模型层 (models/actor.py):
    MoEActorModel
      ├── freeze_router()         → 训练冻结 gate
      ├── replay_router()         → 回放推理时 routing
      ├── fp32 gate               → 不受混合精度影响
      └── expert + epilogue 各自 FSDP wrap

  Loss 层 (models/loss.py):
    load_balancing_loss(data)     → 防止 collapsed experts
    z_loss(data)                  → router logits 正则化

  vLLM 层 (trainer/vllm/vllm_engine.py):
    DeepEPMoEvLLMEngine           → NVIDIA DeepEP 实现 expert parallelism

  Placement 层 (trainer/placement.py):
    model_placement_strategy() → "PACK"
    → EP=8 / CP=8 全部在同一个 Node,不跨节点
    → 跨节点 deep_ep 会 illegal memory access

Disaggregated 部署

  ┌─ Node 1 (Actor Pool) ──────────────────────────────────┐
  │                                                         │
  │  [FSDP shard 0] [FSDP shard 1] [FSDP shard 2] ...      │
  │  [Reference Model co-located]                           │
  │                                                         │
  │  DP=N, CP=8, EP=8 (PACK 策略, 全在一个 Node)           │
  └──────────────────────┬──────────────────────────────────┘
                         │ NCCL / InfiniBand 权重同步
                         ▼
  ┌─ Node 2+ (Engine Pool) ────────────────────────────────┐
  │                                                         │
  │  [vLLM Engine 0] [vLLM Engine 1] ... [vLLM Engine N]   │
  │  Router Actor (HTTP 负载均衡)                           │
  └─────────────────────────────────────────────────────────┘

  Actor Pool 和 Engine Pool 独立扩展:
    Actor Pool:  固定大小 (每卡 ~2B 参数)
    Engine Pool: 横向扩展 (加 GPU = 加推理吞吐)

规模上限

模型 MoE GPU 数 关键
8B 4-16 FSDP full shard
70B 32-128 FSDP + CP=8
70B 16-64 DeepEP EP=8, freeze_router
405B 64-256 Disaggregated: 32 actor + 32+ engine
1T+ 256+ Disaggregated + DeepEP + CPU offload

8.3 对比 OpenRLHF

Molt 从 OpenRLHF 获取灵感(README 原文), 但代码库是完全重写的。

维度 OpenRLHF Molt 根因
Agent 抽象 固定 prompt/completion 格式 Env + Agent 通用接口 能看懂+改得动
Chat 模式 不支持 FastAPI + ChatAgentRunner 改得动(第三方 API)
Rollout 同步 barrier 异步 producer-consumer 不妥协规模
vLLM Router 固定映射 动态负载均衡 + cond_var 不妥协规模
MoE 训推 基础 freeze_router+replay+DeepEP 不妥协规模
Global Token Mean DP-invariant loss 不妥协规模
权重同步 直接 vLLM 加载 refit: FSDP→NCCL broadcast 不妥协规模
代码规模 40K+ LOC ~12K LOC 能看懂

移除了什么

  OpenRLHF 有而 Molt 没有:
    - Tensor Parallel (只用 FSDP)
    - Pipeline Parallel (只用 FSDP + CP)
    - 混合精度手动控制 (FSDP 自动)
    - 分布式 checkpoint 复杂逻辑 (依赖 HF + FSDP)

  移除理由: FSDP2 + CP + DeepEP 不需要
  megatron 风格的 TP/PP 手动切分。
  这些对"能看懂"是噪音,对"改得动"是障碍。

TransFormer-封面

0xFF 参考

Molt —— 轻量级、高性能的 Agentic RL Research 框架

posted @ 2026-08-31 20:55  罗西的思考  阅读(17)  评论(0)    收藏  举报