Prhub

#32984 [MLX] Upgrade to Torch 2.13/MLX 0.32+ and redesign the Torch-MLX tensor bridge

原始 PR 作者 yeahdongcn 合并时间 2026-08-22 09:51 文件变更 36 提交数 27 评论 18 代码增减 +1695 / -343

执行摘要

升级 Torch 2.13/MLX 0.32,重设计零拷贝张量桥

PR body 给出三层动机:其一,MLX 0.32 可以直接导入 Torch MPS 分配,而 Torch 2.13 改变了运行时与 DLPack 边界语义,旧桥不再适配;其二,旧桥把 MPS 张量先搬到 CPU 再经 NumPy 中转,双向都付出 payload 大小的拷贝,且 borrowed-vs-owned 生命周期被隐式处理,参数替换或 GC 后容易引用失效存储;其三,diffusion 的 SGLANG_USE_MLX norm override 在生产图中实际不可达——CustomOp.dispatch_forward() 对 MPS 路由到 forward_native,Sana 仪器化确认每 DiT forward 的 58 个 norm 零调用,强制全部穿过 Torch-MLX 桥反而使 denoise step 中位数 +62.5%。因此「删除不可达路径、收紧版本门控」是数据驱动的决策。

值得精读,尤其是 python/sglang/srt/utils/tensor_bridge.py 的所有权契约(owned vs borrowed 三个层次)与负 stride 物化设计:它把「双框架共用同一块 Metal 内存」的竞态、生命周期问题显式化,并用 RLock 串行化完整穿越、用 DLPack 消除 payload 拷贝。对任何做 Torch + 其他运行时张量桥的团队都有直接借鉴意义;runtime.py 的版本门控矩阵(拒 rc/dev、Torch 锁系列、MLX 只设下限)也是可复用的 fail-fast 模式。

讨论亮点

评审核心围绕桥的边界语义与数值精度:

  • noob-se7en 发现 CPU float64 导出缺陷:「My local test failed here because MLX CPU float64 array is exported here as a PyTorch float64 MPS tensor. (PyTorch MPS does not support float64)」。yeahdongcn 修复为 CPU float64 在 mx.cpu stream 构造,并在 _prepare_mlx_export() 对 MPS 目标显式拒绝 float64,配套新增 test_mlx_call_multi_preserves_cpu_float64
  • alexnails 质疑:「there can be race condition / deadlock on MLX vs Torch borrow on exit?」yeahdongcn 复现并修复,补充了借用期存活与负 stride 物化的处理。
  • alexnails 对 fallback_torch.py 做精度审计:bias 被加到已窄化回输入 dtype 的结果上,旧代码是先 fp32 累加再一次性舍入;实测 bf16+fp32 参数时 max err 从 1.5e-2 升到 2.7e-2。yeahdongcn 修复后由 test_norm_infer_preserves_input_dtype_with_fp32_parameters 锁定。
  • noob-se7en 指出 CI 依赖集冲突:「The MLX Stage A workflow installs srt_mps,test, but this test imports diffusion modules requiring diffusers, causing a clean-install failure.」扩散导入冒烟被拆到 test_diffusion_import_isolation.py,Stage A 依赖集保持不变。
  • alexnails 质疑新增依赖:「we are adding torchcodec?」yeahdongcn 解释其用于多模态音视频解码,与 cudaxpu extra 对齐。
  • noob-se7en 整体认可所有权模型:「The distinction between MLX-owned storage in torch_to_mlx() and scoped borrowing in mlx_call() is much clearer.」,alexnails 最终批准:「all my remaining comments are minor, so approving」。

实现拆解

  1. 依赖升级与运行时提前门控python/pyproject_other.tomlsrt_mps extra 将 torch==2.11.0 提升到 torch==2.13.0mlx 改为 mlx>=0.32.0,并新增 torchcodec==0.15.0torchvision==0.28.0;版本继续钉在 srt_mps 内,all_mps 继承,CPU 与独立 diffusion_mps 依赖列表不变。新增 python/sglang/srt/hardware_backend/mlx/runtime.py_validate_runtime()packaging.version 同时校验 stable Torch 2.13.x、stable MLX>=0.32.0、torch.backends.mps.is_available()mlx.metal.is_available(),失败信息统一带「reinstall with the srt_mps extra」的可操作指引;server_args.py_run_resolution_pipeline() 中新增 _handle_hardware_runtime_validation(),位于模型路径解析、权重下载与 dummy-model 短路之前;use_mlx() 使用 lru_cache,未设置 SGLANG_USE_MLX 时不导入 mlx(由子进程测试验证 sys.modules)。

  2. 张量桥所有权模型重构tensor_bridge.py 删除全局 _MLX_AVAILABLE 探测与 NumPy/CPU 中转,改用 _mlx_core() 懒加载与 _BRIDGE_LOCK(RLock)经 _serialized_bridge 装饰器串行化所有穿越。torch_to_mlx() 语义改为「独立 MLX 拷贝」:MPS 输入 mx.asarray(copy=True) 后立即 mx.eval 物化;CPU 输入永远构造 MLX-owned 存储,float64mx.cpu stream、complex128 显式报错。新增 MlxTensorView(同时持有 detached Torch tensor 与 MLX array 两个强引用,含 matches() 供参数替换后校验)与 borrow_torch_tensors()(批量借用、单次 fence)。mlx_call() / mlx_call_multi() 先校验目标设备(cuda 等在开跑前拒绝),一次 torch.mps.synchronize() 后零拷贝导入,所有输出共享一次 mx.eval 边界,经 DLPack 导出(CPU 目标用 dl_device=(1,0) 显式视图);_has_negative_stride() 预检负 stride 输出并物化,因为 Torch 2.13 的 DLPack 导入器遇负 stride 会 abort 进程而非抛 Python 异常。

  3. diffusion MPS 路径收敛:删除 python/sglang/kernels/ops/diffusion/common/fallback_mps.py(122 行 MLX 加速实现);fallback_torch.pynorm_infer_native() 改用 F.rms_norm / F.layer_norm,在 fp32 下累加 bias 后一次性舍入回输入 dtype,保留 optional-bias 与「fp16/bf16 输入 + fp32 参数」契约;rotary、scale/shift 原本就是 Torch 别名直接清理。移除 diffusion 侧 SGLANG_USE_MLX 运行时处理,文档明确该变量只作用于 SRT 服务。

  4. Torch 2.13 兼容桩与配套迁移_platform_stubs.py 新增 _KernelInterface_OutOfResources_PTXASError_CompiledKernel_ASTSource_GPUTarget 类符号以及 triton.compiler.compiler.triton_key,避免 catch-all meta-path finder 把类名物化为模块,保证 MPS 上 Inductor import 成功。rebase 后把 VAE/Transformer loader 迁到 upstream 的 unified residency API;scheduler.pyone_batch.pyoverrides.pyprofiler.pykv_cache_builder.py 等仅调整 use_mlx 导入来源。

  5. 测试与 CI 配套:新增 4 个注册测试——test_tensor_bridge.py(603 行:owned/borrowed/persistent 生命周期、源变异与删除、单次 fence、单次 eval、负 stride abort-safe 子进程、CPU float64、无效目标 fail-fast、autograd 分离);test_runtime.py(版本门控矩阵、缺失 MLX/MPS/Metal 的可操作报错、dummy 模型前中止);test_diffusion_torch_fallback.py(同时注册 MLX 与 CPU/ARM64 套件,锁定 fp16/bf16+fp32 参数精度);test_diffusion_import_isolation.py(未启用时不导入 mlx)。4 个文件全部注册进 stage-a-unit-test-mlx,diffusion fallback 测试额外注册进 base-a-test-cpubase-b-test-cpubase-b-test-cpu-arm64

文件 模块 状态 重要度
python/sglang/srt/utils/tensor_bridge.py 张量桥 modified 8.84
python/sglang/srt/hardware_backend/mlx/runtime.py 运行时门控 added 8.54
test/registered/unit/utils/test_tensor_bridge.py 桥接测试 added 7.91
python/sglang/_platform_stubs.py 平台桩 modified 7.88
test/registered/unit/hardware_backend/mlx/test_runtime.py 门控测试 added 7.7
python/sglang/kernels/ops/diffusion/common/fallback_torch.py 扩散回退 modified 5.64
python/sglang/kernels/ops/diffusion/common/fallback_mps.py 扩散回退 removed 6.85
python/sglang/srt/server_args.py 启动参数 modified 6.39

关键符号

mlx_call mlx_call_multi torch_to_mlx mlx_to_torch borrow_torch_tensors MlxTensorView _torch_to_mlx _serialized_bridge _validate_runtime use_mlx _handle_hardware_runtime_validation norm_infer_native

关键源码片段

python/sglang/srt/hardware_backend/mlx/runtime.py runtime-gate

新增运行时门控:集中校验 stable Torch 2.13.x、MLX>=0.32、MPS 与 MLX Metal 可用性,是 fail-fast 策略的执行主体。

# python/sglang/srt/hardware_backend/mlx/runtime.py(新增,核心逻辑)
# 目标:把版本 / 设备校验前移到 ServerArgs 构造的最早阶段,
# 在模型路径解析、权重下载甚至 dummy-model 短路之前就 fail fast。from functools import lru_cacheimport torch
from packaging.version import InvalidVersion, Versionfrom sglang.srt.environ import envs_MIN_MLX_VERSION = Version("0.32.0")
_SUPPORTED_TORCH_SERIES = (2, 13)
​
​
def _is_stable_series(raw_version: object, series: tuple[int, int]) -> bool:
    """只接受 stable 系列:拒绝 rc/dev,也拒绝 2.14 等未验证的大版本。"""
    try:
        version = Version(str(raw_version))
    except InvalidVersion:
        return False
    return not version.is_prerelease and (version.major, version.minor) == series
​
​
def _is_stable_at_least(raw_version: object, minimum: Version) -> bool:
    """MLX 只设下限:0.32.0 之后的 stable 小版本(含 0.33)都放行。"""
    try:
        version = Version(str(raw_version))
    except InvalidVersion:
        return False
    return not version.is_prerelease and version >= minimum
​
​
@lru_cache(maxsize=1)
def _validate_runtime() -> None:
    try:
        import mlx.core as mx
    except ImportError:
        raise RuntimeError(
            "SGLANG_USE_MLX requires stable Torch 2.13.x and MLX >= 0.32.0, "
            "but MLX is not installed; reinstall with the srt_mps extra"
        ) from None
    mlx_version = getattr(mx, "__version__", None)
    torch_version = getattr(torch, "__version__", None)
    if not _is_stable_series(
        torch_version, _SUPPORTED_TORCH_SERIES
    ) or not _is_stable_at_least(mlx_version, _MIN_MLX_VERSION):
        raise RuntimeError(
            "SGLANG_USE_MLX requires stable Torch 2.13.x and MLX >= 0.32.0; "
            f"found Torch {torch_version or 'unknown'} + MLX {mlx_version or 'unknown'}; "
            "reinstall with the srt_mps extra"
        )
​
    # 版本达标后还要确认 MPS 设备与 MLX Metal 设备真实可用,
    # 给出可操作报错而不是让后续在随机位置崩溃。
    mps_backend = getattr(torch.backends, "mps", None)
    is_mps_available = getattr(mps_backend, "is_available", None)
    if not callable(is_mps_available) or not is_mps_available():
        raise RuntimeError("SGLANG_USE_MLX requires an available PyTorch MPS device")
​
    metal = getattr(mx, "metal", None)
    is_available = getattr(metal, "is_available", None)
    if not callable(is_available) or not is_available():
        raise RuntimeError("SGLANG_USE_MLX requires an available MLX Metal device")
​
​
@lru_cache(maxsize=1)
def use_mlx() -> bool:
    """显式启用才触发校验;未设置 SGLANG_USE_MLX 时保持懒导入,不碰 mlx。"""
    enabled = bool(envs.SGLANG_USE_MLX.get())
    if enabled:
        _validate_runtime()
    return enabled

评论区精华

CPU float64 导出到 MPS 目标失败 正确性

noob-se7en 本地测试发现 MLX CPU float64 数组被导出成 PyTorch float64 MPS 张量,而 PyTorch MPS 不支持 float64。

结论:yeahdongcn 修复:CPU float64 输入在 mx.cpu stream 构造,_prepare_mlx_export() 对 MPS 目标显式拒绝 float64,并新增 test_mlx_call_multi_preserves_cpu_float64 锁定。 · 已解决

桥退出路径的竞态 / 死锁疑虑 正确性

alexnails 询问「there can be race condition / deadlock on MLX vs Torch borrow on exit?」,关注借用期与进程退出时的双框架交互。

结论:yeahdongcn 复现并修复,补充借用对象存活保持与负 stride 物化处理。 · 已解决

norm fallback 的 bias 累加精度回归 正确性

alexnails 实测:bias 在结果窄化回输入 dtype 之后才加,bf16 + fp32 参数场景 max err 从旧实现的 1.5e-2 升到 2.7e-2。

结论:yeahdongcn 修复为 fp32 累加、末尾一次性舍入,由 test_norm_infer_preserves_input_dtype_with_fp32_parameters 覆盖。 · 已解决

Stage A CI 依赖集与 diffusion 导入冲突 测试

noob-se7en 指出 stage-a-unit-test-mlx 只安装 srt_mps,test,而 test_runtime.py 导入 diffusion 模块需要 diffusers,干净安装会失败。

结论:diffusion 导入冒烟拆到 test_diffusion_import_isolation.py,Stage A 依赖集保持不变。 · 已解决

新增 torchcodec 依赖的疑问 question

alexnails 询问「we are adding torchcodec?」。

结论:yeahdongcn 说明 torchcodec 用于多模态音视频解码路径,与 cuda、xpu extra 对齐。 · 已解决

桥所有权模型设计评审 设计

noob-se7en 认可 owned(torch_to_mlx)与 scoped borrow(mlx_call)的区分更清晰,并说明正以本 PR 为参考评估 SRT/LLM serving 侧 Torch ModelRunner + 可选 MLX 算子的对比。

结论:无需改动;设计获两位 reviewer 认可后合并。 · 已解决

风险与影响

崩溃级风险(DLPack 负 stride):Torch 2.13 的 DLPack 导入器遇负 stride 视图会 abort 整个进程而非抛异常,tensor_bridge.py_has_negative_stride() 预检 + _export_evaluated_mlx()mx.contiguous 物化是唯一防线;已有 abort-safe 子进程测试覆盖,但未来新增输出路径若绕过 mlx_call_multi 的共享物化边界仍可能漏检。版本硬锁定的升级窗口SGLANG_USE_MLX=1 只接受 stable Torch 2.13.x 与 stable MLX>=0.32.0,Torch 2.12/2.14 和 rc/dev 全部拒绝;MLX 后续 minor 只有形式化放行、缺每版本回归,Torch 2.14 发布后用户会立即遇到硬失败。双框架流同步边界_serialized_bridge 的 RLock 只覆盖函数内完成的桥工作,无法串行化函数外对同一 MPS 张量的并发使用/修改,契约依赖调用方自律。diffusion fallback 数值变化norm_infer_native 从手工 fp32 分解改为 F.rms_norm / F.layer_norm,与 Triton 内核存在 1e-5 量级差异(测试容忍 2e-5),严格位级一致场景需留意;该代码路径同时服务 CPU Torch 2.11/2.12。依赖面扩大srt_mps 新增 torchcodec==0.15.0,安装体积与解析复杂度上升。

用户:Apple Silicon 上启用 SGLANG_USE_MLX 的用户必须升级到 Torch 2.13 / MLX 0.32+;未启用用户的 import 路径完全不受影响(懒加载 + 子进程测试保证 sys.modules 无 mlx);diffusion 用户不再有 MLX norm 加速路径,但该路径本就不可达。系统:零拷贝桥的主要收益将在堆叠 PR yeahdongcn/sglang#7 落地的 Torch 权重与算子穿越中兑现;当前 SRT MLX runner 并未被替换,SGLANG_USE_MLX=1 仍走旧 runner,但运行时门控已收紧。团队:本 PR 确立了 Torch-MLX 桥的 ownership/lifetime 契约,成为 Apple Silicon 路线图(issue #19137)上消除 MlxModelRunnerStub drift 的公共基础;review 环节的精度审计与 CI 依赖集把关提升了该路径的工程门槛。

崩溃级:DLPack 负 stride 需预检物化 依赖版本硬锁定 Torch 2.13.x 双框架流同步边界依赖调用方自律 扩散 fallback 数值行为变化 新增 torchcodec 依赖 CI 依赖集需保持精简

关联 Issue

#7 [MPS] Run SRT with Torch ModelRunner and MLX/Metal operators

完整报告

参与讨论