Prhub

#36991 [Diffusion] Add exact component precision overrides

原始 PR 作者 mickqian 合并时间 2026-08-31 11:12 文件变更 25 提交数 13 评论 1 代码增减 +385 / -27

执行摘要

新增组件级精度覆盖,贯通加载与驻留阶段

PR body 说明了目标:为扩散模型组件提供精确的精度覆盖能力,覆盖原生文本和图像编码器、标准 VAE 组件(包括 audio_vae)以及原生 plain-state 组件,并让该覆盖从加载阶段贯穿到驻留阶段执行 dtype;同时必须拒绝不支持的组件覆盖,而不是接受一个加载/执行阶段并不会兑现的 dtype。这本质上是把此前散落在各 pipeline_config 中的默认精度,升级为可按组件名显式覆盖的统一机制。

值得精读。核心看点有三个:一是 fail-closed 的组件能力声明模式(component_load_precision 基类默认拒绝、plain-state 子类放行、load() 模板方法统一调用),二是精度解析的优先级设计(显式覆盖 > pipeline_config 默认,且 residency 阶段只读 resolve_component_precision_override 以保持加载 dtype),三是 _extract_dynamic_component_map 的泛化复用(关闭别名后缀与路径展开)——这套模式可直接迁移到其他组件级配置键。

讨论亮点

本 PR 没有人工 review 评论,唯一的 issue 评论来自 mintlify[bot],仅用于文档预览部署通知(Mintlify Previews,lmsysorg-codex-component-encoder-precision.mintlify.site)。设计取舍主要体现为 13 个 commit 的渐进迭代:从最初的 encoder 精度覆盖,逐步扩展到 VAE、audio_vae、plain-state 组件,两次合并 main 时解决了 server_args.py、component_loader.py、cli.mdx 等文件的冲突,最终以「Fix transformer fallback test fixture」收尾,说明实现过程中对 fallback 路径的 fixture 也做了同步修正。

实现拆解

本 PR 按「配置入口 → 解析链路 → 加载器接入 → 驻留阶段 → 测试与文档」五步落地,全部改动位于 sglang/multimodal_gen 子系统。

  1. 配置入口(server_args.py):在 ServerArgs 上新增 component_precisions: dict[str, str] 字段;新增 _normalize_component_precisions 做白名单校验(精度值必须落在 PRECISION_TO_TYPE 内,组件名连字符转下划线);新增 _extract_component_precisions 从 CLI 未知参数中提取 --component-precisions.<component>--component_precisions.<component> 两种写法,并在 from_cli_args 中并入 provided_args。为了让精度值不被当作路径展开,把 _extract_dynamic_component_map 泛化为支持 alias_suffix=Noneexpand_values=False

  2. 解析链路(utils/precision.py)resolve_precisionresolve_decode_precisionresolve_component_precision 三处解析都优先查询 server_args.component_precisions 覆盖,未命中才回退到 pipeline_config 的默认精度;新增独立的 resolve_component_precision_override,供驻留阶段只读「显式覆盖」。这样既保证加载阶段总能得到可执行 dtype,又让运行阶段能区分「用户显式覆盖」与「默认精度」。

  3. 加载器接入(component_loader.py 及各 loader):基类 ComponentLoader 新增 component_load_precision,默认实现直接抛 ComponentCheckpointUnsupportedError(fail-closed),并在 load() 模板方法入口处调用;PlainStateDictComponentLoader 覆写为放行配置值。text_encoder_loader.py 覆写为「override 优先,否则按 _extract_encoder_index 索引 text_encoder_precisions」;image_encoder_loader.py 覆写为「override 优先,否则用 image_encoder_precision」;vae_loader.py 通过继承 PlainStateDictComponentLoader 获得放行能力;sound_tokenizer_loader.py 改用 resolve_component_precision 解析 dtype,并把模型构造从 set_default_torch_dtype 隐式控制改为显式 .to(device=..., dtype=...)

  4. 驻留阶段(pipelines_core/stages)image_encoding.pytext_encoding.pycomponent_uses 为每个 ComponentUse 传入 target_dtype=resolve_component_precision_override(...),使显式精度覆盖能进入驻留调度;无覆盖时传 None,保持模型加载时的 dtype,避免多余的类型转换。

  5. 测试与文档配套:新增/更新约 10 个单元测试文件,覆盖 VAE 与 audio_vae 的精度准入(test_vae_loader.py)、image encoder 加载与驻留(test_image_encoder_loader.py、test_component_residency.py)、CLI 动态参数提取(test_server_args.py)、VAE 加载与解码默认精度(test_precision_consistency.py)、文本编码缓存与解码并行度(test_text_encoding_cache.py、test_decoding_stage_parallelism.py)、plain loader 准入(test_component_quantization_admission.py)以及 disagg/ideogram4 的 fixture 修补;文档侧更新了 docs/docs/sglang-diffusion/api/cli.mdx,说明能力边界。

文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/server_args/server_args.py 参数配置 modified 7.69
python/sglang/multimodal_gen/runtime/utils/precision.py 精度解析 modified 7.19
python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py 组件加载器 modified 6.79
python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py 编码器加载 modified 6.44
python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py 图像编码阶段 modified 6.19
python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py 声音分词加载 modified 5.83
python/sglang/multimodal_gen/test/unit/test_vae_loader.py VAE 测试 modified 6.08
python/sglang/multimodal_gen/test/unit/test_component_residency.py 驻留测试 modified 5.75

关键符号

_normalize_component_precisions _extract_component_precisions resolve_component_precision_override resolve_component_precision resolve_precision resolve_decode_precision component_load_precision component_uses

关键源码片段

python/sglang/multimodal_gen/runtime/server_args/server_args.py core-logic

配置入口:新增 component_precisions 字段、规范化校验函数与 CLI 提取逻辑,并将 _extract_dynamic_component_map 泛化以支持精度值不做路径展开。

def _normalize_component_precisions(value: object) -> dict[str, str]:
    """规范化组件精度映射:组件名连字符转下划线,精度值小写并做白名单校验。    白名单来自 `PRECISION_TO_TYPE`,因此 CLI 上传入的任意字符串会在
    `ServerArgs` 构造阶段就被拒绝,而不是等到加载期才暴露问题。
    """
    if not isinstance(value, dict):
        raise ValueError("component_precisions must be a mapping")
​
    normalized: dict[str, str] = {}
    for component, precision in value.items():
        component_name = str(component).strip().replace("-", "_")
        precision_name = str(precision).strip().lower()
        # 精度必须是 PRECISION_TO_TYPE 白名单内的写法,否则直接拒绝,
        # 防止用户在 CLI 上传入一个加载器无法识别的 dtype 字符串。
        if not component_name or precision_name not in PRECISION_TO_TYPE:
            raise ValueError(
                "Component precision entries require a component and one of "
                f"{sorted(PRECISION_TO_TYPE)}"
            )
        normalized[component_name] = precision_name
    return normalized
​
​
@classmethod
def _extract_component_precisions(
    cls,
    unknown_args: list[str],
) -> tuple[dict[str, str], list[str]]:
    """从未知 CLI 参数中提取 `--component-precisions.<component>` 形式的覆盖。    复用 `_extract_dynamic_component_map`,但关闭别名后缀(alias_suffix=None)
    与路径展开(expand_values=False):精度值(如 bf16、fp32)不是文件路径,
    不应被 `os.path.expanduser` 改写。
    """
    return cls._extract_dynamic_component_map(
        unknown_args,
        option_prefixes=("--component-precisions.", "--component_precisions."),
        alias_suffix=None,
        expand_values=False,
    )
python/sglang/multimodal_gen/runtime/utils/precision.py core-logic

精度解析核心:所有解析函数统一优先查 component_precisions 覆盖,新增独立的 resolve_component_precision_override 供驻留阶段使用。

def resolve_component_precision_override(
    server_args, module_name: str
) -> Optional[torch.dtype]:
    """只解析显式组件精度覆盖;无覆盖时返回 None,让调用方保持加载时 dtype。    驻留(residency)阶段的 `ComponentUse` 依赖此函数决定是否覆盖目标 dtype:
    没有显式覆盖就保持模型加载时的精度,避免一次多余的类型转换。
    """
    exact_precision = server_args.component_precisions.get(module_name)
    if exact_precision is None:
        return None
    return precision_to_dtype(exact_precision, f"component_precisions.{module_name}")
​
​
def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]:
    # 优先级:显式组件覆盖 > pipeline_config 默认精度。
    # 加载阶段用本函数保证总能得到可执行的 dtype;而驻留阶段的
    # resolve_component_precision_override 只在用户显式指定时返回非 None。
    exact_precision = resolve_component_precision_override(server_args, module_name)
    if exact_precision is not None:
        return exact_precision
​
    pipeline_config = server_args.pipeline_config
    if module_name in ("audio_vae", "vocoder"):
        precision_attr = "audio_vae_precision"
    elif module_name in ("vae", "video_vae", "diffusion_decoder"):
        precision_attr = "vae_precision"
    elif module_name in (
        "transformer",
        "transformer_2",
        "audio_dit",
        "video_dit",
        "connectors",
        "dual_tower_bridge",
    ):
        precision_attr = "dit_precision"
    elif module_name == "image_encoder":
        precision_attr = "image_encoder_precision"
    elif module_name == "text_encoder" or module_name.startswith("text_encoder_"):
        # 文本编码器支持多个实例,按 text_encoder_2 后缀映射到
        # pipeline_config.text_encoder_precisions 的对应下标。
        precisions = getattr(pipeline_config, "text_encoder_precisions", None)
        if not precisions:
            return None
        suffix = module_name.removeprefix("text_encoder")
        index = 0 if suffix == "" else int(suffix.removeprefix("_")) - 1
        if index < 0 or index >= len(precisions):
            raise ValueError(
                f"No configured precision for {module_name!r}; "
                f"text_encoder_precisions has {len(precisions)} entries"
            )
        return precision_to_dtype(precisions[index], f"text_encoder_precisions[{index}]")
    else:
        return None
python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py core-logic

fail-closed 的核心落点:基类 component_load_precision 默认拒绝覆盖,PlainStateDictComponentLoader 覆写为放行,load() 模板方法统一调用。

# 基类 ComponentLoader:默认拒绝任何精度覆盖,保证 fail-closed。
# 加载入口 load() 会先调用本方法,因此不支持覆盖的组件一旦被用户
# 配置 --component-precisions.<name> 就会立刻报错,而不是静默接受
# 一个加载器并不会兑现的 dtype。
def component_load_precision(
    self, server_args: ServerArgs, component_name: str
) -> str | None:
    """Return an exact precision override or reject an unsupported one."""
    precision = server_args.component_precisions.get(component_name)
    if precision is not None:
        raise ComponentCheckpointUnsupportedError(
            f"{component_name!r} does not support an exact component precision "
            "override"
        )
    return None
​
​
# 原生 plain-state 加载器按声明 dtype 物化权重,真正兑现精度覆盖,
# 因此子类覆写为直接放行配置值。
class PlainStateDictComponentLoader(ComponentLoader):
    def component_load_precision(
        self, server_args: ServerArgs, component_name: str
    ) -> str | None:
        return server_args.component_precisions.get(component_name)
​
​
# 加载主流程统一入口:无论哪个 loader,加载前先做一次精度能力校验,
# 之后再走 direct-GPU / quantization 等其他能力检查。
def load(
    self,
    component_model_path: str,
    server_args: ServerArgs,
    component_name: str,
    transformers_or_diffusers: str,
) -> tuple[AutoModel, float]:
    self._native_load_manages_placement = False
    self.component_load_precision(server_args, component_name)
    ...

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  1. resolve_precision 依赖 component_precisions 属性utils/precision.pyresolve_precisionresolve_component_precision_override 直接访问 server_args.component_precisions,若外部调用方或测试传入未初始化该属性的伪 ServerArgs,会抛 AttributeError;本 PR 已同步修补了 test_component_residency.py 等 fixture,但第三方扩展代码仍可能遗漏。
  2. fail-closed 行为变化ComponentLoader.component_load_precision 默认拒绝任何精度覆盖,意味着对不支持的组件(如 GenericComponentLoader、DiT 系列 loader)传入 --component-precisions.xxx 会直接启动失败。这是有意设计,但属于行为变更,可能影响已有启动脚本。
  3. sound_tokenizer 加载路径微调:从 set_default_torch_dtype 隐式控制改为显式 .to(device=..., dtype=...),整体更确定,但可能改变个别 checkpoint 的加载结果(例如原本依赖模块默认 dtype 的权重)。
  4. text_encoder_loader 边界component_load_precision 在 pipeline_config 缺失 text_encoder_precisions 时仍可能 IndexError,属原有行为,但新方法把该路径包装得更隐晦,错误信息不友好。

影响范围限定在 sglang/multimodal_gen 的 Diffusion 生成链路,涉及配置解析(server_args)、组件加载器(component_loader 及四个具体 loader)、运行阶段(image_encoding / text_encoding 的 component_uses)三个分层,对 SRT 核心推理路径无影响。用户侧收益:可对单个组件做显式精度覆盖(如 audio_vae 用 fp32、encoder 用 fp16),并获得 fail-closed 的明确报错;团队侧价值:确立了「组件能力声明 + 默认拒绝 + 子类显式放行」的统一模式,为后续按组件控制显存与质量打开了扩展空间。

配置入口变更 fail-closed 行为变更 加载路径行为微调 无人工 review

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论