Prhub

#24172 Fix flashinfer workspace OOM

原始 PR 作者 kpham-sgl 合并时间 2026-05-04 16:26 文件变更 3 提交数 4 评论 12 代码增减 +360 / -2

执行摘要

修复 FlashInfer 工作区 OOM 导致分布式测试挂起

在 CI 环境中,test_gpt_oss_120b.py 偶尔因 FlashInfer create_allreduce_fusion_workspace 触发 OOM,导致 NCCL watchdog 超时,测试进程挂起。作者经调研确认 FlashInfer 层面无更优解法,故从 SGLang 侧引入预检。

建议相关开发者精读分布式预检设计模式,重点关注 flashinfer_comm_fusion.py 中的内存分配计算和 all_reduce 投票逻辑。此变更为未来类似内存预检提供了可复用的模式。

讨论亮点

无公开 review 讨论。Fridge003 直接批准。作者在 PR body 中说明已花费时间研究 FlashInfer 内部,确认此修复是合理的且无其他方案。

实现拆解

  1. 添加 CUDA 驱动绑定函数:在 common.py 中新增 get_cuda_driver_bindings,统一 cuda.bindings.driver 和 cuda.cuda 的导入。
  2. 实现内存分配属性构造:在 flashinfer_comm_fusion.py 中新增 _make_flashinfer_workspace_allocation_prop,根据平台选择合适的句柄类型,填充 CUmemAllocationProp。
  3. 镜像 FlashInfer 大小计算:_flashinfer_trtllm_workspace_allocation_sizes 模拟 FlashInfer 内部 SymmDeviceMemory 逻辑,计算各缓冲区大小并考虑对齐和 multicast 粒度。
  4. 单 rank 探测函数 _probe_cumem_create_sequence:对每个分配大小尝试 cuMemCreate 和 cuMemRelease,返回成功或失败。
  5. 分布式预检与投票:_preflight_check_workspace_memory 内部调用探测函数,并将各 rank 的布尔结果通过 dist.all_reduce with BAND 聚合,仅当所有 rank 返回 True 时才允许创建融合 workspace。
  6. 修改初始化入口:在 FlashInferWorkspaceManager.initialize 中增加预检调用,失败时跳过 workspace 创建。
  7. 分布式测试:新增 test_flashinfer_fusion_preflight.py,使用 multiprocessing spawn 启动 2 进程模拟分布式环境,覆盖正常和显存耗尽场景。
文件 模块 状态 重要度
python/sglang/srt/layers/flashinfer_comm_fusion.py 通信融合层 modified 8.75
test/registered/distributed/test_flashinfer_fusion_preflight.py 分布式测试 added 7.62
python/sglang/srt/utils/common.py 工具函数 modified 5.48

关键符号

_make_flashinfer_workspace_allocation_prop _flashinfer_trtllm_workspace_allocation_sizes _probe_cumem_create_sequence _preflight_check_workspace_memory get_cuda_driver_bindings

关键源码片段

python/sglang/srt/layers/flashinfer_comm_fusion.py core-logic

核心变更文件,新增预检函数并修改 workspace 初始化流程

def _make_flashinfer_workspace_allocation_prop(cuda_driver):
    # 构建 CUmemAllocationProp,用于后续 cuMemCreate 探测
    if _should_force_posix_fd_transport():
        handle_type = (
            cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
        )
    else:
        from flashinfer.comm.mnnvl import is_mnnvl_fabric_supported
        # 优先使用 Fabric 句柄类型(若支持)
        if is_mnnvl_fabric_supported(torch.cuda.current_device()):
            handle_type = (
                cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
            )
        else:
            handle_type = (
                cuda_driver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
            )
    prop = cuda_driver.CUmemAllocationProp()
    prop.requestedHandleTypes = handle_type
    prop.type = cuda_driver.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
    prop.location = cuda_driver.CUmemLocation()
    prop.location.type = cuda_driver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
    prop.location.id = torch.cuda.current_device()
    prop.allocFlags.gpuDirectRDMACapable = 1
    return prop
test/registered/distributed/test_flashinfer_fusion_preflight.py test-coverage

新增的分布式测试,验证预检正常和 starvation 场景

def _run_rank(rank, world_size, port, scenario, result_q):
    held = None
    cuda_driver = None
    try:
        os.environ['MASTER_ADDR'] = '127.0.0.1'
        os.environ['MASTER_PORT'] = str(port)
        os.environ['RANK'] = str(rank)
        os.environ['WORLD_SIZE'] = str(world_size)
        os.environ['LOCAL_RANK'] = str(rank)
        torch.cuda.set_device(rank)
        import torch.distributed as dist
        dist.init_process_group(backend='gloo', rank=rank, world_size=world_size)
        cpu_group = dist.group.WORLD
​
        from sglang.srt.layers.flashinfer_comm_fusion import (
            _make_flashinfer_workspace_allocation_prop,
            _preflight_check_workspace_memory,
        )
        probe_kwargs = dict(
            world_size=8, max_token_num=2048, hidden_dim=12288,
            dtype=torch.bfloat16, cpu_group=cpu_group,
        )
        # 场景: rank0 饥饿(模拟显存不足)
        if scenario == 'rank0_starved' and rank == 0:
            cuda_driver = get_cuda_driver_bindings()
            prop = _make_flashinfer_workspace_allocation_prop(cuda_driver)
            free, _total = torch.cuda.mem_get_info(rank)
            target = max(free - (1 << 30), 0) # 预留 1GB
            err, gran = cuda_driver.cuMemGetAllocationGranularity(
                prop,
                cuda_driver.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED,
            )
            assert err == cuda_driver.CUresult.CUDA_SUCCESS, err
            aligned = (target // gran) * gran
            assert aligned > 0, 'not enough free memory'
            err, held = cuda_driver.cuMemCreate(aligned, prop, 0)
            assert err == cuda_driver.CUresult.CUDA_SUCCESS, (err, aligned)
​
        decision = _preflight_check_workspace_memory(**probe_kwargs)
        result_q.put((rank, 'ok', bool(decision)))
    except Exception as e:
        result_q.put((rank, 'err', repr(e)))
    finally:
        if held is not None:
            cuda_driver.cuMemRelease(held)
        try:
            if dist.is_initialized():
                dist.destroy_process_group()
        except Exception:
            pass

评论区精华

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

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

风险与影响

预检逻辑与 FlashInfer 内部分配算法紧密耦合,FlashInfer 升级后可能不同步导致预检失效或误判;AllReduce BAND 依赖 CPU group,若通信故障可能导致一致性错误;但预检失败仅跳过融合,不影响正常推理。

修复 CI 稳定性,减少因 OOM 导致的测试挂起;对用户透明,预检仅在初始化执行一次,无性能开销;若预检失败回退到非融合 allreduce,不影响结果正确性。

依赖 CUDA 驱动底层 API 镜像 FlashInfer 内部分配逻辑 AllReduce BAND 通信可能失败

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论