Prhub

#22817 [diffusion] Extract post-training weight APIs into mixins and add tensor update/checker paths

原始 PR 作者 MikukuOvO 合并时间 2026-06-09 13:57 文件变更 8 提交数 30 评论 7 代码增减 +816 / -84

执行摘要

提取扩散后训练权重 API 到 mixins 并添加 Tensor 更新 / 检查路径

在本次变更之前,扩散后训练权重操作分散在调度器/worker实现中,并且tensor更新/验证流程没有通过与磁盘更新路径相同的后训练mixin结构路由。本PR使后训练API表面更一致且更易扩展,同时保持运行时行为明确。

值得精读。该PR展示了如何通过mixin模式将分散的治理逻辑解耦到独立模块,同时为后续扩展后训练特性提供了清晰的架构参考。

讨论亮点

Rockdu在gpu_worker_post_training_mixin.pyTYPE_CHECKING导入上建议确认是否可移除以避免循环导入;gxlvera解释了fsdp_load.py中的修改是必要的,因为该代码块无论FSDP是否启用都会执行。

实现拆解

  1. 提取Worker Mixin:新建python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py,定义GPUWorkerPostTrainingMixin,将原本直接写在GPUWorker中的update_weights_from_diskget_weights_checksum、新加的update_weights_from_tensorupdate_weights_from_tensor_checker方法集中到mixin中,通过组合方式引入GPUWorker
  2. 提取Scheduler Mixin:新建python/sglang/multimodal_gen/runtime/post_training/scheduler_post_training_mixin.py,定义SchedulerPostTrainingMixin,集中_handle_update_weights_from_disk_handle_get_weights_checksum以及新增的_handle_update_weights_from_tensor_handle_update_weights_from_tensor_checker,Scheduler类改为继承该mixin。
  3. 新增Tensor权重更新与检查API:在weights_api.py中添加两个HTTP端点POST /update_weights_from_tensorPOST /update_weights_from_tensor_checker,接收序列化的tensor数据或SHA-256校验值。WeightsUpdater类新增update_weights_from_tensor方法,支持模块级payload解析、flattened_bucket重建和weight_loader感知加载。
  4. 实现Tensor校验检查器:新建tensor_update_checker.py,实现TensorUpdateChecker类,提供verify_across_tp方法:在TP环境下通过gather_object收集各rank的本地tensor,在root rank上计算SHA-256并与期望值比较,支持DTensor分片重建校验。
  5. 调整调度器与worker的实例化:修改python/sglang/multimodal_gen/runtime/managers/scheduler.pypython/sglang/multimodal_gen/runtime/managers/gpu_worker.py,移除内联的方法实现,改为通过mixin继承自动获得。同时更新io_struct.py添加UpdateWeightFromTensorReqInputUpdateWeightFromTensorCheckerReqInput请求结构体。
文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/post_training/tensor_update_checker.py 校验器 added 9.25
python/sglang/multimodal_gen/runtime/post_training/weights_updater.py 权重更新器 renamed 9.23
python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py Worker Mixin added 9.08
python/sglang/multimodal_gen/runtime/post_training/scheduler_post_training_mixin.py 调度器 Mixin added 8.76
python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py API 入口 modified 8.08
python/sglang/multimodal_gen/runtime/managers/scheduler.py 调度器 modified 7.95
python/sglang/multimodal_gen/runtime/managers/gpu_worker.py Worker modified 7.93
python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py 数据结构 modified 6.8

关键符号

GPUWorkerPostTrainingMixin.update_weights_from_tensor GPUWorkerPostTrainingMixin.update_weights_from_tensor_checker SchedulerPostTrainingMixin._handle_update_weights_from_tensor SchedulerPostTrainingMixin._handle_update_weights_from_tensor_checker WeightsUpdater.update_weights_from_tensor TensorUpdateChecker.verify_across_tp compute_tensor_sha256 build_named_tensor_sha256

关键源码片段

python/sglang/multimodal_gen/runtime/post_training/weights_updater.py rename-or-move

Core 逻辑文件:从 loader 目录搬迁到 post_training,新增 update_weights_from_tensor 系列方法,支持模块级 tensor 更新和 flattened_bucket 重建。

# 部分关键新增:update_weights_from_tensor 方法
class WeightsUpdater:
    # ... 已有代码
​
    def update_weights_from_tensor(
        self,
        named_tensors: dict[str, dict[str, torch.Tensor]],
        load_format: str | None = None,
        target_modules: list[str] | None = None,
    ) -> tuple[bool, str]:
        """从内存中的 tensor 字典更新模型权重,支持模块级粒度和扁平桶重建。"""
        try:
            # 1. 将外层 dict 按模块拆分:{module_name: {param_name: tensor}}
            module_payloads = self._resolve_module_payloads(named_tensors)
            if not module_payloads:
                return False, "No module payloads resolved"
​
            # 2. 对每个模块应用权重更新
            for module_name, payload in module_payloads.items():
                module = self.pipeline.get_module(module_name)
                if module is None:
                    return False, f"Module {module_name} not found"
​
                # 3. 重建 flattened bucket(若 payload 是扁平格式)
                weights_iter = self._materialize_weights_iter(
                    payload, load_format
                )
​
                # 4. 将权重实际加载到模型,处理 offload 与 DTensor
                self._load_weights_into_module(module, weights_iter)
​
            return True, "Weights updated successfully"
        except Exception as e:
            return False, str(e)
python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py dependency-wiring

Worker Mixin 新文件:聚合所有后训练权重操作(disk/tensor/checker/checksum),并处理 TP rank 范围的 payload 选择和 SP gather。

# GPUWorkerPostTrainingMixin 的关键方法:update_weights_from_tensor
from __future__ import annotations
from typing import TYPE_CHECKINGfrom sglang.multimodal_gen.runtime.distributed import get_tp_rank, get_tp_world_size
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
from sglang.multimodal_gen.runtime.post_training.tensor_update_checker import TensorUpdateChecker
from sglang.multimodal_gen.runtime.post_training.weights_updater import WeightsUpdater, get_updatable_modules
from sglang.srt.utils import MultiprocessingSerializer
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductionsif TYPE_CHECKING:
    from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
        UpdateWeightFromTensorCheckerReqInput,
        UpdateWeightFromTensorReqInput,
    )class GPUWorkerPostTrainingMixin:
    # ... 其他方法
​
    def update_weights_from_tensor(
        self,
        req: UpdateWeightFromTensorReqInput,
    ) -> tuple[bool, str]:
        if not self.pipeline:
            return False, "Pipeline is not initialized"
​
        # 根据 TP rank 选择对应的序列化 payload
        payload, error = self._select_rank_scoped_payload(
            payloads=req.serialized_named_tensors,
            field_name="serialized_named_tensors",
        )
        if error is not None:
            return False, error
​
        monkey_patch_torch_reductions() # 避免序列化时触发 dist 操作
        try:
            named_tensors = MultiprocessingSerializer.deserialize(payload)
        except Exception as e:
            return False, f"Failed to deserialize serialized_named_tensors: {e}"
​
        updater = WeightsUpdater(self.pipeline)
        return updater.update_weights_from_tensor(
            named_tensors=named_tensors,
            load_format=req.load_format,
            target_modules=req.target_modules,
        )
​
    def update_weights_from_tensor_checker(
        self,
        req: UpdateWeightFromTensorCheckerReqInput,
    ) -> tuple[bool, str]:
        if not self.pipeline:
            return False, "Pipeline is not initialized"
​
        checker = TensorUpdateChecker(self.pipeline)
        result = checker.verify_across_tp(
            target_module=req.target_module,
            expected_named_tensors_sha256=req.expected_named_tensors_sha256,
            tp_rank=get_tp_rank(),
            tp_world_size=get_tp_world_size(),
            tp_cpu_group=self.tp_cpu_group,
            tp_root_rank=self.tp_group.first_rank,
        )
​
        # 如果启用 SP(序列并行),还需要跨 SP group gather 结果
        if self.sp_group.world_size > 1:
            import torch
            is_sp_root = self.sp_group.rank_in_group == 0
            gathered_results = [None] * self.sp_group.world_size if is_sp_root else None
            torch.distributed.gather_object(
                result,
                gathered_results,
                dst=self.sp_group.first_rank,
                group=self.sp_cpu_group,
            )
            # 在 SP root 上合并结果并 broadcast 回去
            if is_sp_root:
                failures = [(r, m) for r, (s, m) in enumerate(gathered_results) if not s]
                if failures:
                    result = (False, failures[0][1])
            result_holder = [result]
            torch.distributed.broadcast_object_list(
                result_holder,
                src=self.sp_group.first_rank,
                group=self.sp_cpu_group,
            )
            return result_holder[0]
        return result

评论区精华

TYPE_CHECKING 导入的必要性 style

Rockdu 评论:如果不会导致循环导入,可以移除 TYPE_CHECKING 块。

结论:没有后续响应,但代码保留了 TYPE_CHECKING,说明可能仍有循环导入风险。 · 已解决

修改 fsdp_load.py 的必要性 正确性

Rockdu 提问是否有必要修改 FSDP 代码。gxlvera 回应:该代码块无论 FSDP 开关如何都会执行,所以必须修改。

结论:接受修改,已确认必要性。 · 已解决

风险与影响

风险较低:该PR不修改模型前向计算或推理核,仅影响后训练权重管理路径。但需要注意的是,新增的tensor更新路径会涉及跨TP rank的gather_object和序列化/反序列化操作,若在大型TP集群上频繁调用可能引入通信开销。另PS:fsdp_load.py的修改虽然必要,但如果FSDP配置特殊,可能引入边缘情况。

对用户:提供了基于Tensor的权重更新和验证途径,便于强化学习/后训练工作流中的热更新和校验;对系统:重构后训练权重管理,代码更集中,但现有磁盘更新API保持不变;对团队:新mixin结构降低了添加后训练功能的门槛。

后训练路径变更 跨 TP 通信开销 FSDP 兼容性微调

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论