(18) CUDA Graph全捕获:把训练循环变成一次GPU提交

——“一个人用AI如何写出比PyTorch更快的自研深度学习框架”系列文章之十八

上一篇我们聊了 Tech-Renaissance 的多流并发架构:把本来串在一根绳上的任务拆到 TRANSCOMP_1/2/3UPDATE 这几条流上,让计算、通信、传输有机会并行推进。但多流只是解决了“让 GPU 有事可做”的问题,它并没有解决另一个更根本的问题——

CPU 每次让 GPU 干点活,都要亲自发一次指令。

一张现代高端 GPU 做一次卷积或矩阵乘,真正计算时间可能只有几十微秒;而 CPU 通过 CUDA Runtime 提交一次 kernel,要经过参数校验、队列插入、驱动处理等一整套流程,通常要 4~10 微秒。如果模型由几百个小算子组成,CPU 发指令的时间就会和 GPU 计算的时间相当,GPU 不得不频繁停下来等 CPU 的“下一道口令”。训练迭代越短、kernel 越小,这个问题就越明显。

CUDA Graph 的目标,就是把这一整套口令一次性录下来,以后每次迭代只需要喊一声“按之前录好的来”。本文就来解释它的原理,以及 Tech-Renaissance 为什么敢说自己把“整个训练循环”都捕获进了 CUDA Graph。

一、CUDA Graph 是什么:把重复流水线录成一张唱片

CUDA Graph 是 NVIDIA 自 CUDA 10(2018 年)起提供的一种执行模式,现代 NVIDIA GPU 与较新的 CUDA 驱动栈普遍支持这一能力。它的核心思想很简单:

对于一段会被反复执行的 GPU 操作序列,与其每次让 CPU 逐条提交,不如在第一次把它录成一张“图”,以后整张图一次下发。

这张图不是神经网络意义上的计算图,而是 CUDA 级别的有向无环图:节点可以是 kernel 启动、异步内存拷贝、事件记录、集合通信调用等,边是它们之间的依赖关系。录制完成后,CUDA Driver 就已经知道了整张图的拓扑,可以把后续重放所需的调度信息预先准备好。

标准用法分为四步:

// 1. 把指定流置为捕获模式
cudaStreamBeginCapture(stream, cudaStreamCaptureModeThreadLocal);

// 2. 在这个流上正常提交要录制的 CUDA 工作——kernel、memcpy、event 等
my_kernel<<<grid, block, 0, stream>>>(...);
cudaMemcpyAsync(..., stream);

// 3. 结束捕获,得到 cudaGraph_t
cudaStreamEndCapture(stream, &graph);

// 4. 实例化为可执行图,之后反复重放
cudaGraphInstantiate(&exec, graph, ...);
cudaGraphLaunch(exec, stream);

注意,捕获期间提交的操作不会真正执行,只是被记录到图里。真正执行发生在 cudaGraphLaunch 之后。

CUDA Graph 并不让单个 kernel 跑得更快,也不改变 kernel 的计算量。它的收益在于:

  • 把 per-kernel 的 CPU 提交开销摊薄到几乎为零
  • 让 Driver 提前看到完整依赖,减少 GPU 流水线上的气泡
  • 对大量小 kernel、短延迟场景效果最明显,而这正是深度学习训练迭代中常见的局面。

举个粗略的例子:一次 ResNet-50 前向传播可能要 launch 超过两百个 kernel。如果每个 kernel 的 CPU 提交开销平均 5 微秒,那么一轮迭代光提交就要 1 毫秒以上;而捕获成图之后,一次 cudaGraphLaunch 的 CPU 开销只有亚微秒到几微秒级别。当单步迭代本身只有十几毫秒时,这省下来的 1 毫秒就意味着 5%~10% 的吞吐提升。

在推理领域,CUDA Graph 已经被广泛采用:TensorRT、ONNX Runtime、vLLM 都会把固定形状的推理路径固化成图。在训练领域,主流框架往往只 capture 前向或只 capture 推理子图,因为训练循环里还有反向传播、梯度同步、优化器更新、动态学习率、NaN 检测等“不规矩”的环节,完整捕获的工程复杂度要高得多。

Tech-Renaissance 的选择是:能 capture 的尽量 capture,不能 capture 的部分留在图外用最小开销控制。

二、从 Compiler 到 GraphAtlas:捕获之前的准备工作

要捕获 CUDA Graph,首先要有一张“施工图”。在我们的框架里,这张图就是 ComputationGraph

ComputationGraph 是一个零形状信息的纯拓扑容器。它把模型拆成两种节点:

  • COMPUTE 节点:DTensor 级别的算子,比如卷积、全连接、BN、激活,节点里只存算子类型和张量 ID;
  • RANGE 节点:Region 级别的批量操作,比如清零整个梯度区、AllReduce 某个桶、FP16↔FP32 整区转换,节点里存的是 (offset, size) 内存范围。

为什么形状信息要剥离?因为同一份网络结构会对应多种变体:正常 batch、最后一个不完整的 batch、低分辨率训练、验证分辨率等。如果每种变体都单独建图,图的数量会爆炸。我们把拓扑和形状解耦,同一份 ComputationGraph 就可以被多个 MemoryPlan 复用。

computation_graph.h 中,编译器定义了 33 个子图 ID,每个对应一个可独立捕获的语义单元:

enum class GraphId : uint8_t {
    TRANSFER_A, TRANSFER_B,
    FIRST_LAYER_FWD_A, FIRST_LAYER_FWD_B,
    DEEP_FWD_BWD,
    ZERO_GRAD,
    FIRST_LAYER_BWD_A, FIRST_LAYER_BWD_B,
    FIRST_COMM, DEEP_COMM,
    CAST_DEEP_GRAD_FP16_TO_FP32, CAST_FIRST_GRAD_FP16_TO_FP32,
    NAN_CHECK_AND_GRAD_SCALING,
    STATS_COMM, UPDATE_STATS,
    OPTIMIZER, EMA_UPDATE,
    INF_MAIN_A, INF_MAIN_B, INF_EMA_A, INF_EMA_B,
    CAST_MAIN_FP32_TO_FP16, CAST_EMA_FP32_TO_FP16,
    ACCUM_METRICS, ACCUM_METRICS_TRAIN_LAST, ACCUM_METRICS_VAL_LAST,
    VAL_RESULT_COMM, CLEAR_METRICS,
    SIMPLE_TASK_GRAPH,
    LARS_FC_OPT, LARS_FIRST_CONV_OPT, LARS_DEEP_CONV_OPT,
    UPDATE_BN_INF_PARAMS,
    COUNT              // = 33
};

真正决定“哪个变体用哪张图”的,是 GraphAtlas

GraphAtlas 是一张 6(变体)× 33(子图 ID)的表格。每个格子叫一个 Slot,里面填了:

  • cg:指向共享的 ComputationGraph
  • mp:该变体自己的 MemoryPlan
  • shape_id:去重键,shape 无关的图用 kShapeInvariant
  • stream_kind:这张子图默认跑在哪条流上;
  • captured_idx:Phase B 捕获完成后填回的索引。

DeepLearningTask::build_graph_atlas() 里,我们会遍历所有变体和所有 GraphId

  • TRANSFER_A/BZERO_GRADCOMM、各种 CAST、优化器更新等 shape 无关的图,全部指向 base MemoryPlan,shape_id 设为 kShapeInvariant
  • 前向、反向、深层融合等 shape 相关的图,才使用各自变体的 MemoryPlanShapeId

这样,Phase B 只要发现 (cg, gid, shape_id) 三元组相同,就会复用同一张 CapturedGraph,避免重复捕获。对于 8 卡训练,虽然每张卡都要有自己的 cudaGraphExec_t 句柄,但图的实例化过程仍然可以共享大部分元数据。

三、Phase B:pre_capture() 的六个步骤

捕获工作集中在 pre_capture() 里完成。把它拆开看,源码里大致是六个阶段:

PreCaptureResult pre_capture(const GraphAtlas& compile_atlas,
                              const std::vector<DeviceContext*>& contexts) {
    PreCaptureResult result;
    result.atlas = compile_atlas;

    // B1: 去重
    // B2: cuDNN 预热(每个 rank 串行)
    // B2.5: 识别含 NCCL 的图
    // B3: 普通图捕获(rank 0 串行 + rank 1~N-1 并行)
    // B3.5: NCCL 图协同捕获
    // B4: warmup launch(仅 rank 0,跳过 NCCL 图)
    return result;
}

B1:去重

pre_capture() 先遍历 GraphAtlas 的所有 slot,用 CapturedGraph::Key{cg, gid, shape_id} 做键,把重复的 slot 指向同一个 captured_idx

std::unordered_map<CapturedGraph::Key, int32_t, CapturedGraph::KeyHash> seen;

for (size_t vi = 0; vi < GraphAtlas::kMaxVariants; ++vi) {
    for (uint8_t gi = 0; gi < static_cast<uint8_t>(GraphId::COUNT); ++gi) {
        auto& slot = result.atlas.slot(vi, gi);
        if (!slot.cg || !slot.mp) continue;

        CapturedGraph::Key key{slot.cg, static_cast<GraphId>(gi), slot.shape_id};
        auto it = seen.find(key);
        if (it != seen.end()) {
            slot.captured_idx = it->second;   // 复用已有图
            ++result.reused;
        } else {
            slot.captured_idx = static_cast<int32_t>(seen.size());
            seen[key] = slot.captured_idx;    // 新增唯一图
            ++result.captured;
        }
        ++result.total_slots;
    }
}

同一个 shape 无关的图在 6 个变体里都会出现,它们在这里只会被真正捕获一次。

B2:cuDNN 预热

cuDNN Frontend 的 kernel/plan cache 是 per-device 的。如果在捕获阶段才让 cuDNN 第一次选计划,那么选计划本身也会被录进图里,污染图结构。所以我们在捕获之前,会让每个 rank 串行地把所有需要 warm-up 的 cuDNN 算子执行一遍,填满 cache。串行是为了避免 cuDNN 全局锁竞争。

B2.5:识别 NCCL 图

ncclAllReduce 的图不能按普通图那样各 rank 各自捕获。因为 NCCL 集合通信要求所有 rank 以相同顺序、同时进入同一 collective call,否则就会死锁。pre_capture() 会先扫描一遍,把所有含 NCCL RangeOp 的 key 标记出来。

B3:普通图捕获

对于不含 NCCL 集合通信的子图,我们采用“rank 0 串行 + rank 1~N-1 并行”的策略:

  • rank 0 在主线程捕获所有唯一图,负责初始化每张图的元数据;
  • 其他 rank 各自开线程并行捕获;
  • 每张图的 cudaGraphExec_t 句柄按 rank 填回 CapturedGraph::per_rank_execs_

B3.5:NCCL 协同捕获

标记为 NCCL 的图会在 capture_nccl_graph_coordinated() 中处理:

// Phase 1: 所有 rank 同时 BeginCapture
for (int r = 0; r < num_ranks; ++r) {
    cudaSetDevice(contexts[r]->device_id());
    cudaStreamBeginCapture(cap_streams[r], cudaStreamCaptureModeThreadLocal);
}

// Phase 2: ncclGroupStart → 重放所有 rank → ncclGroupEnd
ncclGroupStart();
for (int r = 0; r < num_ranks; ++r) {
    // 重放该 rank 的图节点,包括 ncclAllReduce
}
ncclGroupEnd();

// Phase 3: 所有 rank 同时 EndCapture 并各自 Instantiate
for (int r = 0; r < num_ranks; ++r) {
    cudaStreamEndCapture(cap_streams[r], &captured_graphs[r]);
    cudaGraphInstantiate(&exec[r], captured_graphs[r], ...);
}

B4:warmup launch

捕获完成后,在 rank 0 上对所有非 NCCL 图做一次 warmup launch,让 GPU 的各类内部状态处于“热”状态,后续 benchmark 测量更准确。含 NCCL 的图会被跳过——因为 warmup 只跑 rank 0,如果此时 launch 了 AllReduce,其他 rank 没参与,同样会死锁。

四、多流捕获里的依赖管理

CUDA Graph 捕获的是单一流上的操作序列,但我们的训练会用到多个流:TRANSCOMP_1COMP_2COMP_3UPDATE。多流之间的依赖怎么办?

答案是在图内显式插入 event barrier

MultiStreamCaptureState 管理最多 5 条活跃流:

struct PerStreamState {
    cudaStream_t stream = nullptr;
    cudaEvent_t  last_done_event = nullptr;
    bool         has_pending_work = false;
};

struct MultiStreamCaptureState {
    static constexpr int kMaxActiveStreams = 5;
    PerStreamState streams[kMaxActiveStreams] = {};
    int num_active = 0;
    cudaStream_t primary_stream = nullptr;
    int32_t output_stream_idx = -1;
};

捕获开始时:

  • 主捕获流被注册为 streams[0]
  • COMP_1COMP_2COMP_3 这三条计算流也会被预注册;
  • 在主捕获流上记录一个 event,让其他所有 secondary 流 cudaStreamWaitEvent,从而把它们引入同一个捕获上下文。

这里有一个关键细节:CUDA 不允许在捕获期间创建新的 event handle。因此所有 cudaEventCreateWithFlags 必须在 cudaStreamBeginCapture 之前完成,否则会在图中引入非法节点。源码里的 CaptureGuard 也是一个小保险:如果捕获中途抛异常,它会在析构时把未提交的图清理掉,防止句柄泄漏。

遍历节点时,insert_cross_op_barrier() 会根据下一个节点默认该跑哪条流,决定是否需要等待上一个节点的完成事件:

void insert_cross_op_barrier(const GraphNode& prev_node,
                              const GraphNode& next_node,
                              MultiStreamCaptureState& state,
                              const DeviceContext& ctx) {
    int out_idx = state.output_stream_idx;
    if (out_idx < 0) return;

    StreamKind target_sk;
    if (next_node.kind == GraphNode::Kind::COMPUTE) {
        target_sk = get_op_default_stream(next_node.compute_op);
    } else if (next_node.kind == GraphNode::Kind::RANGE) {
        switch (next_node.range_op) {
            case RangeOp::RANGE_ACCUM_METRICS:
            case RangeOp::RANGE_CLEAR:
            case RangeOp::RANGE_CAST_FP32_TO_FP16:
            case RangeOp::RANGE_CAST_FP16_TO_FP32:
            case RangeOp::RANGE_EMA_PARAM_UPDATE:
                target_sk = StreamKind::UPDATE; break;
            case RangeOp::RANGE_GRAD_SCALING:
            case RangeOp::RANGE_CHECK_NAN:
                target_sk = StreamKind::COMP_1; break;
            case RangeOp::RANGE_SUM_ALLREDUCE:
            case RangeOp::RANGE_MEAN_ALLREDUCE:
            case RangeOp::RANGE_BN_STATS_ALLREDUCE:
                target_sk = StreamKind::UPDATE; break;
            default:
                target_sk = StreamKind::COMP_1; break;
        }
    }

    cudaStream_t target_s = static_cast<cudaStream_t>(ctx.stream(target_sk));
    int target_idx = state.find_stream_index(target_s);
    if (target_idx >= 0 && target_idx != out_idx) {
        cudaStreamWaitEvent(target_s,
            state.streams[out_idx].last_done_event, 0);
    }
}

最后 finalize_cross_stream_barrier() 会让主捕获流等待所有有实际工作的 secondary 流,保证图提交出去的那一刻,整张图的语义是闭合的。

这套机制让多流并发和 CUDA Graph 全捕获可以共存:图内部依然存在并行执行的多个流,但 CPU 不需要在每次迭代里手动调度这些流,所有依赖都在捕获时固定下来。

五、运行时:从 CapturedGraph 到一次 cudaGraphLaunch

捕获完成后,就进入 Phase C:运行。

CapturedGraph 对外暴露的接口很简洁:

void launch(int rank, void* stream) const;

内部根据 is_cuda_ 分支:GPU 路径调用 cudaGraphLaunch(per_rank_execs_[rank], stream);CPU 路径则顺序执行预先收集好的 CpuOp 函数指针序列。

更高层有 GraphExecutor,它封装了训练/验证步骤的编排:

void run_train_step();
void run_val_step();
void launch(GraphId gid);
void launch_dual(GraphId gid1, GraphId gid2);

GraphExecutor::run_train_step() 把一次完整迭代拆成十几个 launch 调用,它是一个更通用、更同步的 API。但框架真正的性能路径在 DeepLearningTask::run_train_epoch_gpu() 里:为了把运行期开销压到最低,build_exec_table() 阶段就把每个 rank、每个变体、每个 GraphId 对应的 cudaGraphExec_t 句柄预先解析好,存进 gpu_exec_.variant_graphs

这里要注意一个容易混淆的地方:编译期用 GraphId(33 个),但运行期的 gpu_exec_.variant_graphs 用的是内部枚举 GraphSlot(32 个密集槽位)。GraphSlot 把 A/B 对合并、把语义相近的图聚合,让运行期数组更紧凑。build_exec_table() 负责把 GraphId 映射到 GraphSlot

// 示意:build_exec_table 中的一部分映射
auto resolve = [&](GraphId gid, int rank, size_t variant_idx) -> cudaGraphExec_t {
    int32_t idx = captured_result_.atlas.index(variant_idx, gid);
    if (idx < 0 || static_cast<size_t>(idx) >= captured_result_.graphs.size())
        return nullptr;
    return static_cast<cudaGraphExec_t>(
        captured_result_.graphs[idx].native_exec(rank));
};

// normal batch 变体
auto& g = gpu_exec_.variant_graphs[v][rank];
g[S(GraphSlot::XFER_A)]          = resolve(GraphId::TRANSFER_A, rank, v);
g[S(GraphSlot::XFER_B)]          = resolve(GraphId::TRANSFER_B, rank, v);
g[S(GraphSlot::FWD_BWD_DEEP_A)]  = resolve(GraphId::DEEP_FWD_BWD, rank, v);
g[S(GraphSlot::FIRST_LAYER_FWD_A)] = resolve(GraphId::FIRST_LAYER_FWD_A, rank, v);
// ... 其余 GraphId → GraphSlot 的映射

每个 GraphId 绑定到哪条流,由 DeepLearningTask::stream_for() 决定:

StreamKind DeepLearningTask::stream_for(GraphId gid) {
    switch (gid) {
        case GraphId::TRANSFER_A:
        case GraphId::TRANSFER_B:
            return StreamKind::TRANS;

        case GraphId::FIRST_LAYER_FWD_A:
        case GraphId::FIRST_LAYER_FWD_B:
        case GraphId::DEEP_FWD_BWD:
        case GraphId::FIRST_LAYER_BWD_A:
        case GraphId::FIRST_LAYER_BWD_B:
        case GraphId::LARS_FC_OPT:
            return StreamKind::COMP_1;

        case GraphId::LARS_FIRST_CONV_OPT:
            return StreamKind::COMP_2;

        case GraphId::LARS_DEEP_CONV_OPT:
            return StreamKind::COMP_3;

        case GraphId::INF_MAIN_A:
        case GraphId::INF_MAIN_B:
            return StreamKind::COMP_1;

        case GraphId::ZERO_GRAD:
        case GraphId::FIRST_COMM:
        case GraphId::DEEP_COMM:
        case GraphId::CAST_DEEP_GRAD_FP16_TO_FP32:
        case GraphId::CAST_FIRST_GRAD_FP16_TO_FP32:
        case GraphId::NAN_CHECK_AND_GRAD_SCALING:
        case GraphId::STATS_COMM:
        case GraphId::UPDATE_STATS:
        case GraphId::OPTIMIZER:
        case GraphId::EMA_UPDATE:
        case GraphId::CAST_MAIN_FP32_TO_FP16:
        case GraphId::ACCUM_METRICS:
        case GraphId::ACCUM_METRICS_TRAIN_LAST:
        case GraphId::ACCUM_METRICS_VAL_LAST:
        case GraphId::VAL_RESULT_COMM:
        case GraphId::CLEAR_METRICS:
        case GraphId::UPDATE_BN_INF_PARAMS:
            return StreamKind::UPDATE;

        default:
            return StreamKind::COMP_1;
    }
}

可以看到,传输图走 TRANS 流;首层/深层计算图和 LARS 的 FC 层走 COMP_1;LARS 的首层卷积和深层卷积分别走 COMP_2COMP_3;梯度通信、类型转换、NaN 检查、优化器更新、指标累积等全部走 UPDATE 流。这些映射在编译期就固定下来,运行期无需再查表。

run_train_epoch_gpu() 中,每个 rank 的线程会在 epoch 开头取出所有图句柄,然后在循环里只做轻量级的 cudaGraphLaunch

// epoch 级别:选择当前分辨率对应的变体
size_t v_base = at_begin_res ? 0 : 2;   // 正常 batch
size_t v_last = at_begin_res ? 1 : 3;   // 最后一个不完整 batch

const auto& g_n = gpu_exec_.variant_graphs[v_base][rank];
auto n_xfer_a  = g_n[S(GraphSlot::XFER_A)];
auto n_xfer_b  = g_n[S(GraphSlot::XFER_B)];
auto n_deep_a  = g_n[S(GraphSlot::FWD_BWD_DEEP_A)];
auto n_deep_b  = g_n[S(GraphSlot::FWD_BWD_DEEP_B)];
auto n_fwd_a   = g_n[S(GraphSlot::FIRST_LAYER_FWD_A)];
auto n_fwd_b   = g_n[S(GraphSlot::FIRST_LAYER_FWD_B)];
auto n_bwd_a   = g_n[S(GraphSlot::FIRST_LAYER_BWD_A)];
auto n_bwd_b   = g_n[S(GraphSlot::FIRST_LAYER_BWD_B)];
auto n_zg      = g_n[S(GraphSlot::ZERO_GRAD)];
auto n_dar     = g_n[S(GraphSlot::DEEP_ALLREDUCE)];
auto n_far     = g_n[S(GraphSlot::FIRST_LAYER_ALLREDUCE)];
auto n_wu      = g_n[S(GraphSlot::WEIGHT_UPDATE)];
auto n_cdg     = g_n[S(GraphSlot::CAST_DEEP_GRAD)];
auto n_cfg     = g_n[S(GraphSlot::CAST_FIRST_GRAD)];
auto n_ncg     = g_n[S(GraphSlot::NAN_CHECK_GRAD_SCALE)];
auto n_sc      = g_n[S(GraphSlot::STATS_COMM)];
auto n_us      = g_n[S(GraphSlot::UPDATE_STATS)];
auto n_cm      = g_n[S(GraphSlot::CAST_MAIN)];
auto n_accum   = g_n[S(GraphSlot::ACCUM_METRICS)];
auto n_lars_fc  = g_n[S(GraphSlot::LARS_FC_UPDATE)];
auto n_lars_fc2 = g_n[S(GraphSlot::LARS_FIRST_CONV_UPDATE)];
auto n_lars_dc  = g_n[S(GraphSlot::LARS_DEEP_CONV_UPDATE)];

// 普通 batch 的 launch 序列(示意,省略部分同步细节)
for (int batch = 0; batch < batches - 1; ++batch) {
    bool from_a   = (batch % 2 == 0);
    int  next_buf = from_a ? 1 : 0;
    auto g_fwd    = from_a ? n_fwd_a : n_fwd_b;
    auto g_deep   = from_a ? n_deep_a : n_deep_b;
    auto g_xfer_n = from_a ? n_xfer_b : n_xfer_a;  // 下一个 batch 的数据
    auto g_first  = from_a ? n_bwd_a : n_bwd_b;

    cudaGraphLaunch(n_zg, s_up);        // 清零梯度 + loss
    cudaGraphLaunch(g_fwd, s_c1);       // 首层前向
    cudaGraphLaunch(g_xfer_n, s_trans); // 下一 batch H2D 传输(与计算重叠)
    sync_comp(); sync_up();

    cudaGraphLaunch(g_deep, s_c1);      // 深层前向+反向
    sync_comp();

    if (!frozen) cudaGraphLaunch(g_first, s_c1);  // 首层反向
    cudaGraphLaunch(n_cdg, s_up);       // AMP 深层梯度转换
    cudaGraphLaunch(n_dar, s_up);       // 深层 AllReduce
    sync_up(); sync_comp();

    cudaGraphLaunch(n_cfg, s_up);       // AMP 首层梯度转换
    cudaGraphLaunch(n_far, s_up);       // 首层 AllReduce
    cudaGraphLaunch(n_accum, s_up);     // 累积 metrics
    cudaGraphLaunch(n_ncg, s_up);       // NaN 检查 + 梯度缩放
    cudaGraphLaunch(n_sc, s_up);        // BN 统计量通信
    cudaGraphLaunch(n_us, s_up);        // 更新 BN 统计量
    sync_up();

    cudaGraphLaunch(n_wu, s_up);        // 优化器更新权重
    cudaGraphLaunch(n_lars_fc,  s_c1);  // LARS FC 层
    cudaGraphLaunch(n_lars_fc2, s_c2);  // LARS 首层卷积
    cudaGraphLaunch(n_lars_dc,  s_c3);  // LARS 深层卷积
    sync_up(); sync_comp();

    cudaGraphLaunch(n_cm, s_up);        // AMP 主权重 FP32→FP16
    sync_up();
    sync_tr();                          // 等待下一 batch 数据传输完成
}

这个循环里有几个关键设计:

  • 双缓冲 A/B 交替:数据传输和首层计算各有 A/B 两套图,batch 0 用 A、batch 1 用 B,交替进行。当前 batch 还在计算时,下一个 batch 的数据已经在传输了。
  • 计算通信重叠g_xfer_n 在首层前向启动之后、同步之前被 launch,两者在不同流上并发执行,CPU 不等待任何一方完成。
  • 流同步的精确控制sync_comp()sync_up()sync_tr() 分别同步三条流。它们被精心放置在必须等待前序结果的位置,但不会过早同步。
  • 最后一个 batch 的特殊处理:最后一个 batch 的 batch size 可能不同,所以使用另一套变体图(v_last),确保在最后一个不完整 batch 上也能正确处理。

六、指针稳定性:为什么图能长期复用

CUDA Graph 里冻结的不只是算子顺序,还有 kernel 参数里的指针。如果每次迭代张量的设备地址都变,捕获的图就失效了。

Tech-Renaissance 能全捕获的前提,是前面几篇文章反复强调的静态内存模型:

  • MemoryPlan 在编译期为每个 DTensor 分配固定的 offset
  • ArenaKeeper 为每个 rank 维护一个统一的显存池;
  • DeviceContext::ptr_at(id) 只是 base_ptr + offset 的常量时间加法;
  • 不同变体之间,同一张量 ID 的 slot_bytes 保持一致,保证 offset 对齐。

这意味着,只要模型结构和超参数不变,每次训练迭代里所有 kernel 看到的指针都是稳定的。图捕获一次,就可以在成千上万个 batch 里反复 replay,不需要因为地址变化而重捕。

七、CPU 回退路径:没有 GPU 也能跑同一份图

Tech-Renaissance 不是只在 GPU 上才有效。当编译目标为 CPU 时,CapturedGraph::capture() 会走 capture_cpu() 路径:

  • 遍历同样的 ComputationGraph 节点;
  • 对每个节点找到 g_compute_op_tableg_range_op_table 里的 CPU launch 函数;
  • 把函数指针和预先填好的 CpuOpContext 一起存进 cpu_ops_
  • 运行时按顺序调用这些函数指针。

这保证了同一份高层训练代码,在 GPU 和 CPU 上都能得到一致的执行语义。性能上 CPU 路径当然无法和 CUDA Graph 相比,但它让调试、回归测试和没有 GPU 的环境都能复用同一套编译产物。

八、全捕获不是银弹:我们保留在图外的控制逻辑

必须诚实地说,CUDA Graph 全捕获带来巨大收益的同时,也施加了很多限制:

  • 拓扑和内存地址必须稳定:捕获之后,图中 kernel 的参数、memcpy 的指针、事件对象都被固定。如果每次迭代张量的地址会变,图就需要重捕。我们的 MemoryPlanArenaKeeper 通过静态分区保证了这一点。
  • 形状变化需要不同的图:batch size 或分辨率一旦改变,kernel 的 grid/block 配置也会变,必须换一张图。我们用 GraphAtlas 的多个变体来容纳正常 batch、last batch、低分辨率、验证分辨率等场景。
  • CPU 侧分支不能进图:NaN 检测、学习率更新、早停判断、A/B 双缓冲切换都发生在图外。它们的开销很小,而且只影响十几个 launch 决策,不会回到逐 kernel 提交的旧模式。
  • NCCL 图需要协同捕获:如前所述,集合通信图的所有 rank 必须一起捕获、一起 launch。

因此,Tech-Renaissance 的训练循环并不是“一个图打天下”,而是把能固化的全部固化成图,把必须动态决策的尽量推到图的最外层。这是一种务实的取舍。

九、为什么训练全捕获在主流框架里不多见

说到这里,可能有读者会问:既然 CUDA Graph 这么好,为什么 PyTorch、TensorFlow 不把整个训练循环都捕获进去?

答案不是不想,而是训练场景比推理场景复杂得多。

在推理场景,输入形状固定、权重不变、没有反向传播和通信,整个前向路径天然就是一张可以反复 replay 的图。TensorRT、ONNX Runtime 把它们 capture 下来,收益非常直接。

训练则不同。训练循环里至少包含:

  • 反向传播,它依赖前向保存的 mask 和中间特征;
  • 多卡梯度同步,需要 NCCL AllReduce 的跨 rank 协同;
  • 优化器状态更新,涉及 FP32 主权重、一阶/二阶动量、EMA、权重衰减;
  • 动态学习率、NaN 检测、梯度裁剪等控制逻辑;
  • 最后一个 batch 往往尺寸不同,需要另一套形状变体。

这些环节中的任何一个,如果在捕获时处理不好,都会导致图无效或死锁。主流框架为了保证通用性和易用性,通常只把部分子图交给 CUDA Graph。Tech-Renaissance 之所以能走得更远,是因为我们从一开始就选择了静态图 + 静态内存规划 + DTensor 跨 rank 一致布局的架构,这些前置条件把“全捕获”从不可能变成了工程问题。

十、小结

CUDA Graph 全捕获,是 Tech-Renaissance 把静态图编译优势兑现为实际吞吐的关键一步。

它背后的逻辑链条是:

  1. 静态图编译让我们能在运行前看到完整拓扑;
  2. MemoryPlanDTensor 让我们能在运行前锁定所有张量的地址和布局;
  3. GraphAtlas 让我们能为不同形状变体复用同一份拓扑;
  4. pre_capture() 把这些拓扑一次性编译成每张 GPU 上的 cudaGraphExec_t
  5. 训练时,CPU 只需要按固定顺序 cudaGraphLaunch,GPU 自己就能跑完成百上千个 kernel。

如果说 PyTorch Eager 模式是“CPU 举着指挥棒,GPU 跟着一个个音符演奏”,那么 Tech-Renaissance 的运行时更像是“CPU 把整套乐谱一次性递给 GPU,然后站在一旁,只在乐章间隙翻页”。

下一篇,我们将进入训练算法的核心细节之一——AMP 自动混合精度训练,看看 FP16 的速度与 FP32 的精度是如何在这张已经被捕获好的图里共存的。

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注

ICP备案号:京ICP备2025133467号-1