——“一个人用AI如何写出比PyTorch更快的自研深度学习框架”系列文章之十八
上一篇我们聊了 Tech-Renaissance 的多流并发架构:把本来串在一根绳上的任务拆到 TRANS、COMP_1/2/3、UPDATE 这几条流上,让计算、通信、传输有机会并行推进。但多流只是解决了“让 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/B、ZERO_GRAD、COMM、各种CAST、优化器更新等 shape 无关的图,全部指向 baseMemoryPlan,shape_id 设为kShapeInvariant;- 前向、反向、深层融合等 shape 相关的图,才使用各自变体的
MemoryPlan和ShapeId。
这样,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 捕获的是单一流上的操作序列,但我们的训练会用到多个流:TRANS、COMP_1、COMP_2、COMP_3、UPDATE。多流之间的依赖怎么办?
答案是在图内显式插入 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_1、COMP_2、COMP_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_2 和 COMP_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_table或g_range_op_table里的 CPU launch 函数; - 把函数指针和预先填好的
CpuOpContext一起存进cpu_ops_; - 运行时按顺序调用这些函数指针。
这保证了同一份高层训练代码,在 GPU 和 CPU 上都能得到一致的执行语义。性能上 CPU 路径当然无法和 CUDA Graph 相比,但它让调试、回归测试和没有 GPU 的环境都能复用同一套编译产物。
八、全捕获不是银弹:我们保留在图外的控制逻辑
必须诚实地说,CUDA Graph 全捕获带来巨大收益的同时,也施加了很多限制:
- 拓扑和内存地址必须稳定:捕获之后,图中 kernel 的参数、memcpy 的指针、事件对象都被固定。如果每次迭代张量的地址会变,图就需要重捕。我们的
MemoryPlan和ArenaKeeper通过静态分区保证了这一点。 - 形状变化需要不同的图: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 把静态图编译优势兑现为实际吞吐的关键一步。
它背后的逻辑链条是:
- 静态图编译让我们能在运行前看到完整拓扑;
MemoryPlan和DTensor让我们能在运行前锁定所有张量的地址和布局;GraphAtlas让我们能为不同形状变体复用同一份拓扑;pre_capture()把这些拓扑一次性编译成每张 GPU 上的cudaGraphExec_t;- 训练时,CPU 只需要按固定顺序
cudaGraphLaunch,GPU 自己就能跑完成百上千个 kernel。
如果说 PyTorch Eager 模式是“CPU 举着指挥棒,GPU 跟着一个个音符演奏”,那么 Tech-Renaissance 的运行时更像是“CPU 把整套乐谱一次性递给 GPU,然后站在一旁,只在乐章间隙翻页”。
下一篇,我们将进入训练算法的核心细节之一——AMP 自动混合精度训练,看看 FP16 的速度与 FP32 的精度是如何在这张已经被捕获好的图里共存的。
