执行摘要
- 一句话:提取扩散后训练权重API到mixins并添加Tensor更新/检查路径
- 推荐动作:值得精读。该PR展示了如何通过mixin模式将分散的治理逻辑解耦到独立模块,同时为后续扩展后训练特性提供了清晰的架构参考。
功能与动机
在本次变更之前,扩散后训练权重操作分散在调度器/worker实现中,并且tensor更新/验证流程没有通过与磁盘更新路径相同的后训练mixin结构路由。本PR使后训练API表面更一致且更易扩展,同时保持运行时行为明确。
实现拆解
- 提取Worker Mixin:新建
python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py,定义GPUWorkerPostTrainingMixin,将原本直接写在GPUWorker中的update_weights_from_disk、get_weights_checksum、新加的update_weights_from_tensor和update_weights_from_tensor_checker方法集中到mixin中,通过组合方式引入GPUWorker。
- 提取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。
- 新增Tensor权重更新与检查API:在
weights_api.py中添加两个HTTP端点POST /update_weights_from_tensor和POST /update_weights_from_tensor_checker,接收序列化的tensor数据或SHA-256校验值。WeightsUpdater类新增update_weights_from_tensor方法,支持模块级payload解析、flattened_bucket重建和weight_loader感知加载。
- 实现Tensor校验检查器:新建
tensor_update_checker.py,实现TensorUpdateChecker类,提供verify_across_tp方法:在TP环境下通过gather_object收集各rank的本地tensor,在root rank上计算SHA-256并与期望值比较,支持DTensor分片重建校验。
- 调整调度器与worker的实例化:修改
python/sglang/multimodal_gen/runtime/managers/scheduler.py和python/sglang/multimodal_gen/runtime/managers/gpu_worker.py,移除内联的方法实现,改为通过mixin继承自动获得。同时更新io_struct.py添加UpdateWeightFromTensorReqInput和UpdateWeightFromTensorCheckerReqInput请求结构体。
关键文件:
python/sglang/multimodal_gen/runtime/post_training/tensor_update_checker.py(模块 校验器;类别 source;类型 core-logic;符号 _materialize_local_tensor, compute_tensor_sha256, build_named_tensor_sha256, TensorUpdateChecker): 核心新增文件:实现基于SHA-256的Tensor校验检查器,支持DTensor感知的跨TP校验,是tensor更新路径的验证保障。
python/sglang/multimodal_gen/runtime/post_training/weights_updater.py(模块 权重更新器;类别 source;类型 rename-or-move;符号 _build_module_weight_name_mapper, map_name, load_weights_into_model, _iter_module_weight_updates): Core逻辑文件:从loader目录搬迁到post_training,新增update_weights_from_tensor系列方法,支持模块级tensor更新和flattened_bucket重建。
python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py(模块 Worker Mixin;类别 source;类型 dependency-wiring;符号 GPUWorkerPostTrainingMixin, update_weights_from_disk, update_weights_from_tensor, update_weights_from_tensor_checker): Worker Mixin新文件:聚合所有后训练权重操作(disk/tensor/checker/checksum),并处理TP rank范围的payload选择和SP gather。
python/sglang/multimodal_gen/runtime/post_training/scheduler_post_training_mixin.py(模块 调度器Mixin;类别 source;类型 core-logic;符号 SchedulerPostTrainingMixin, _handle_update_weights_from_disk, _handle_update_weights_from_tensor, _handle_update_weights_from_tensor_checker): Scheduler Mixin新文件:定义事件处理函数,将输入请求调度到worker的mixin方法,并在tensor更新路径上添加TP barrier。
python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py(模块 API入口;类别 source;类型 entrypoint;符号 update_weights_from_tensor, update_weights_from_tensor_checker): 入口文件:新增两个HTTP端点/update_weights_from_tensor和/update_weights_from_tensor_checker,对外暴露tensor更新和校验能力。
python/sglang/multimodal_gen/runtime/managers/scheduler.py(模块 调度器;类别 source;类型 core-logic;符号 Scheduler, _handle_update_weights_from_disk, _handle_get_weights_checksum): 调度器主文件:移除内联的后训练处理函数,改为继承SchedulerPostTrainingMixin,并注册新请求类型到路由表。
python/sglang/multimodal_gen/runtime/managers/gpu_worker.py(模块 Worker;类别 source;类型 core-logic;符号 GPUWorker, update_weights_from_disk, get_weights_checksum): Worker主文件:移除内联的update_weights_from_disk和get_weights_checksum实现,改为组合GPUWorkerPostTrainingMixin。
python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py(模块 数据结构;类别 source;类型 core-logic;符号 UpdateWeightFromTensorReqInput, UpdateWeightFromTensorCheckerReqInput): 数据结构文件:新增UpdateWeightFromTensorReqInput和UpdateWeightFromTensorCheckerReqInput请求体定义。
关键符号: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
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
Worker Mixin新文件:聚合所有后训练权重操作(disk/tensor/checker/checksum),并处理TP rank范围的payload选择和SP gather。
# GPUWorkerPostTrainingMixin 的关键方法:update_weights_from_tensor
from __future__ import annotations
from typing import TYPE_CHECKING
from 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_reductions
if 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
评论区精华
Rockdu在gpu_worker_post_training_mixin.py的TYPE_CHECKING导入上建议确认是否可移除以避免循环导入;gxlvera解释了fsdp_load.py中的修改是必要的,因为该代码块无论FSDP是否启用都会执行。
- TYPE_CHECKING导入的必要性 (style): 没有后续响应,但代码保留了TYPE_CHECKING,说明可能仍有循环导入风险。
- 修改fsdp_load.py的必要性 (correctness): 接受修改,已确认必要性。
风险与影响
- 风险:风险较低:该PR不修改模型前向计算或推理核,仅影响后训练权重管理路径。但需要注意的是,新增的tensor更新路径会涉及跨TP rank的gather_object和序列化/反序列化操作,若在大型TP集群上频繁调用可能引入通信开销。另PS:
fsdp_load.py的修改虽然必要,但如果FSDP配置特殊,可能引入边缘情况。
- 影响:对用户:提供了基于Tensor的权重更新和验证途径,便于强化学习/后训练工作流中的热更新和校验;对系统:重构后训练权重管理,代码更集中,但现有磁盘更新API保持不变;对团队:新mixin结构降低了添加后训练功能的门槛。
- 风险标记:后训练路径变更, 跨TP通信开销, FSDP兼容性微调
关联脉络
- PR #18306 [Feature] Implement update_weights_from_disk for SGLang-D: 基础PR:首次实现磁盘权重更新,本次重构将其提取到mixin。
- PR #20464 Add update_weights_from_tensor pipeline to Diffusion: 增量PR:添加tensor更新管线,本次完善并集成到mixin结构。
- PR #21106 [diffusion] Add update_weights_from_tensor checker: 增量PR:添加tensor检查器,本次重构校验逻辑并整合进tensor_update_checker模块。
参与讨论