☰
DeepJIT融合内核实战:拆掉TensorRT串行小核墙
2026/10/6 14:55:24 网站建设 项目流程

第一次在 Nsight Systems 里看到自己 TensorRT 引擎的完整时间轴时,我整个人是懵的:一长串两三微秒的小核排着队顺序执行,CUDA kernel 的启动开销占掉整帧耗时将近三分之一。这就是标题里说的"串行小核墙"。为了拆掉它,我试了 DeepJIT 的路子——手写 CUDA 内核,在运行期用 NVRTC 现场编译、动态加载,把被优化器拆散的串行小核重新焊成大核。

这篇文章把我完整踩过的过程记录下来:怎么定位问题、为什么 TensorRT 合不上、DeepJIT 最小骨架怎么写、融合内核分几层写,以及最后到底提了多少、有哪些绕不过去的坑。适合两类人:一类是正在做低延迟推理、被一堆 elementwise 小核卡帧率的人;另一类是想搞清楚"动态编译 + 手写 kernel"到底能在工程里解决什么问题的同学。我会尽量把原理讲明白,也会给出可以直接抄的代码骨架。

1. 先定位问题:那一堵"串行小核墙"是怎么量出来的

1.1 三微秒的核,排两百个队

很多同学以为 GPU 上 kernel 就是"一次调用一次执行",快得很。实际上一次 kernel launch 从 CPU 侧提交到 GPU 真正开始干活,中间要经过命令缓冲区、驱动校验、上下文切换,通常要 3~10 微秒。这还没算 GPU 侧启动新 kernel 时的流水线排空和重新灌满的时间。

如果你的 kernel 本身要跑 50 微秒,3 微秒的 launch overhead 无所谓;但如果一个 kernel 只要跑 2~3 微秒,那 launch overhead 就和计算时间一样长了。

我那个检测模型,ONNX 导出来有几百个节点,TensorRT 优化后主干的卷积基本都被融合掉了,但检测头的部分保留了大量 elementwise 小算子:sigmoid、exp、乘法、加法、slice、concat、cast……它们被拆成了一百多个小 kernel,每个都在 2~4 微秒左右,整串跑下来 GPU 真正干活的时间只有零点几毫秒,剩下的全耗在"排队启动"上。

更亏的是内存带宽。每个小 kernel 都要把张量从全局内存读一遍、算完再写一遍。一个 4MB 的中间张量被十个串行小核轮一遍,就是 80MB 的读写流量——哪怕不启动排队,光带宽也浪费掉了。所以拆串行小核墙,本质上是同时解决两个问题:启动开销和无效内存往返。

1.2 用 Nsight Systems 复现"GPU 空转"现场

定位这件事别靠猜,直接上 profiler。我的标准操作是:

nsys profile --trace=cuda -o deepjit_trace python run_engine.py nsys stats --report=cuda_gpu_trace deepjit_trace.nsys-rep

看 CUDA GPU trace 的时候,我一般盯两个东西:一是 kernel 时长分布,如果出现一大片 1~5 微秒的短条,基本可以断定是串行小核墙;二是 CPU 侧的 CUDA API 提交时间和 GPU 侧实际执行时间之间的差值,如果 CPU 侧的cuLaunchKernel一批一批地排队,而 GPU timeline 上 kernel 之间有明显空档,那就说明启动开销在支配整个推理。

再用--report=gpu_memops看一下全局内存读写总量,会发现比理论最小值高好几倍——这就是小核各自读各自写的代价。

一句话总结定位方法:短 kernel 数量多 + GPU timeline 空档多 + 内存流量虚高,三条同时成立,DeepJIT 就值得试。

2. TensorRT 想帮你又帮不上的地方:融合边界的真相

2.1 TensorRT 最擅长和最不擅长的两张表

TensorRT 确实做了大量图优化,但它不是万能的。它最擅长的是把 CNN 里"标配三件套"焊在一起:

TensorRT 擅长的融合例子
Conv + Bias + Activation一个 kernel 搞定
BatchNorm 折叠推理期 BN 变成 scale+bias 直接并进 conv
常见 elementwise 短链连续几层 add / mul / activation 有时能合

它不擅长的场景,恰恰是串行小核墙的来源:

TensorRT 不擅长的场景结果
自定义算子 / Plugin 两侧融合链在这里断开,前后各剩一堆小核
动态 shape 下的通用策略宁可下发单个小 kernel,也不冒险做激进融合
layout 转换、transpose、split/concat经常被降级成纯 copy kernel,一次一次拷
数据依赖的控制流图上有 if/while,优化器只能保守处理

换句话说,TensorRT 的融合是"模式匹配式"的,它只在认识且能证明安全的模式上合。一旦碰到它不认识的边界,哪怕只是一个小算子,整条融合链也会断掉,断口两侧的算子全变成独立小核。

2.2 三个最容易拆散小核的元凶

我这次项目里,罪魁祸首有三个,你可以对照自己的模型排查:

第一,自定义算子。检测头里我用了几个 ONNX 里就有的 op,但 TensorRT 对这些 op 的支持不完整,比如某个自定义的归一化逻辑,它干脆不认识。于是以它为分界,前面一堆 op、后面一堆 op,全被拆成小核。

第二,动态 shape。我在部署时用了动态 batch。动态环境下 TensorRT 不敢做太激进的融合,因为同一个 kernel 可能要服务多种 shape,它宁可选一个"通用但低效"的实现。我后来固定 batch=1 重新构建引擎,小核数量立刻少了一批。

第三,slice/concat/transpose 这一类张量搬运算子。这些 op 本身没啥计算量,但会被实现成内存拷贝 kernel。如果模型里有 NMS 前处理、多分支拼接这类结构,你会看到大量 copy kernel 排在时间轴上。

搞清楚这三个元凶之后,你的选择就清晰了:要么改模型结构绕开它们(有时候可行),要么在融合链断裂的地方自己上融合 kernel——这正是 DeepJIT 的位置。

3. DeepJIT 的核心动作:运行期拼源码、NVRTC 现场编译、动态加载

3.1 为什么选择"运行期生成源码"而不是预编译

最早我的想法是直接写一个 TensorRT Plugin,把融合 kernel 预编译进 .so 里。做了一轮之后发现三个不舒服的地方:一是要维护一堆 CRUD 代码,Plugin 的 create/configure/serialize/enqueue 全是样板;二是每个 shape 变体都得靠运行时参数传,编译器没法专门优化;三是调试一次要重新编 .so、重新对齐 TRT 版本,来回很慢。

DeepJIT 的思路完全反过来:kernel 的源码在运行期拼出来,用 NVRTC 现场编译成 PTX,再用 CUDA driver API 加载。因为源码是运行期生成的,我可以把张量大小、向量化宽度、是否处理尾块这些信息直接写死进代码里,编译器能看到常量,就能做循环展开、向量化、用立即数,这些都是预编译插件给不了的。

还有一层对比是 CUDA Graph。很多人遇到小核多,第一反应是上 CUDA Graph 录一波。这个办法只解决了 launch 开销,并没有减少 kernel 数量,每个小核照样各自读一遍、写一遍全局内存。DeepJIT 是真正把 N 个小核算法合并到一次 kernel 里,内存流量按比例降下来。两者不冲突,可以在融合之后再用 CUDA Graph 保一层,但别指望 Graph 替代融合。

3.2 一个最小可跑的 DeepJIT 骨架

先给一个最精简的编译加载管线,你可以直接抄。

#include <nvrtc.h> #include <cuda.h> #include <cuda_runtime.h> #include <string> #include <stdexcept> #define NVRTC_CHECK(x) \ do { \ nvrtcResult r = (x); \ if (r != NVRTC_SUCCESS) { \ throw std::runtime_error(std::string("nvrtc: ") + \ nvrtcGetErrorString(r)); \ } \ } while (0) #define CU_CHECK(x) \ do { \ CUresult r = (x); \ if (r != CUDA_SUCCESS) { \ const char* msg = nullptr; \ cuGetErrorString(r, &msg); \ throw std::runtime_error(std::string("cuda: ") + msg); \ } \ } while (0) // 一个简单的融合 kernel:add -> relu -> scale const char* kFusedSource = R"( extern "C" __global__ void fused_add_relu_scale( const float* a, const float* b, float* out, float scale, int n) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < n) { float t = a[i] + b[i]; t = t > 0.f ? t : 0.f; out[i] = t * scale; } } )"; CUfunction DeepJitLoad(const std::string& src, const std::string& name, int sm_arch) { nvrtcProgram prog; NVRTC_CHECK(nvrtcCreateProgram(&prog, src.c_str(), "fused.cu", 0, nullptr, nullptr)); std::string arch = "--gpu-architecture=sm_" + std::to_string(sm_arch); const char* opts[] = {arch.c_str(), "-use_fast_math", "--std=c++17"}; nvrtcResult res = nvrtcCompileProgram(prog, 3, opts); if (res != NVRTC_SUCCESS) { size_t logSize = 0; nvrtcGetProgramLogSize(prog, &logSize); std::string log(logSize, '\0'); nvrtcGetProgramLog(prog, log.data()); nvrtcDestroyProgram(&prog); throw std::runtime_error("nvrtc compile failed:\n" + log); } size_t ptxSize = 0; NVRTC_CHECK(nvrtcGetPTXSize(prog, &ptxSize)); std::string ptx(ptxSize, '\0'); NVRTC_CHECK(nvrtcGetPTX(prog, ptx.data())); NVRTC_CHECK(nvrtcDestroyProgram(&prog)); CUmodule mod; CU_CHECK(cuModuleLoadData(&mod, ptx.c_str())); CUfunction fn; CU_CHECK(cuModuleGetFunction(&fn, mod, name.c_str())); return fn; }

注意一点:调用 driver API 之前最好先把 primary context 建好,避免和 PyTorch / CUDA runtime 抢上下文。常用姿势是:

cuInit(0); CUdevice dev; cuDeviceGet(&dev, 0); CUcontext ctx; cuDevicePrimaryCtxRetain(&ctx, dev); cuCtxSetCurrent(ctx);

然后启动 kernel:

void LaunchFused(CUfunction fn, const float* d_a, const float* d_b, float* d_out, float scale, int n, cudaStream_t stream) { int threads = 256; int blocks = (n + threads - 1) / threads; void* args[] = {&d_a, &d_b, &d_out, &scale, &n}; CU_CHECK(cuLaunchKernel(fn, blocks, 1, 1, threads, 1, 1, 0, stream, args, nullptr)); }

这套骨架跑通之后,你就有了一个"运行期生成 CUDA 内核"的最小闭环。剩下的问题全在源码生成策略和 kernel 写法上。

顺带提醒一句部署环境的事:NVRTC 的头文件和库是随 CUDA Toolkit 一起装的,在 Ubuntu 上别只装 runtime 版。我遇到过一次nvrtc.h找不到,查了半天才发现是装 TensorRT 时顺手装了个运行时,没有完整 toolkit,后来把/usr/local/cuda/include/nvrtc.h确认存在才跑通编译。

3.3 Shape 硬化:让编译器替你展开循环

NVRTC 编译一次大约要几十到几百毫秒,挺贵的,所以源码生成时尽量把一切能确定的都变成常量。比如上面的 kernel,n是运行期参数,编译器没法针对它做太多优化。如果我把n直接宏替换进源码:

#define N 1048576 extern "C" __global__ void fused_add_relu_scale_fixed( const float* a, const float* b, float* out, float scale) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < N) { float t = a[i] + b[i]; t = t > 0.f ? t : 0.f; out[i] = t * scale; } }

N 变成编译期常量后,编译器知道循环边界,能把条件判断优化成无分支甚至完全展开。同理,#define VEC 1开向量化、#define TAIL 0关尾块处理,都会让生成的 SASS 更精简。

代价是编译缓存的 key 变多了:每个 shape 组合对应一份源码。我的做法是维护一个unordered_map<string, CUfunction>,key 用hash(source + arch + nvrtcOptions)生成,命中就直接跳过 NVRTC。上线前把常见的 batch 尺寸、分辨率都预编译一遍,warmup 时全部塞进缓存,线上推理零编译。

4. 手写融合内核的三个层次:元素级、向量化、归约类

4.1 第一层:把一条 elementwise 链焊成一个核

融合 kernel 的起点是 elementwise 链。比如检测头里常见的(a + b) -> relu -> scale -> add c,原来是三四个小核,焊成一个:

extern "C" __global__ void fused_chain(const float* a, const float* b, const float* c, float* out, float scale, int n) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < n) { float t = fmaxf(a[i] + b[i], 0.f) * scale + c[i]; out[i] = t; } }

别看代码不长,收益是实打实的:原来 (a+b) 写一次临时张量、relu 读一次写一次、scale 又读一次写一次、加 c 再读两次写一次,同一个数据在全局内存里被来回折腾四五趟;融合后只读 a/b/c 一次、写 out 一次,内存流量直接砍到三分之一以下。GPU 算力通常不是瓶颈,带宽才是,所以这种融合在 elementwise 密集的模型上效果最明显。

写这种 kernel 我习惯直接用if (i < n)而不是 grid-stride loop,因为推理场景 kernel 是一次性调度,起一个恰好覆盖 n 的 grid 就够了。grid-stride loop 适合常驻 kernel 或需要动态限流的情况,这里用不上。

4.2 第二层:float4 向量化,把带宽榨干

elementwise kernel 有个天然优势:每个线程处理的数据彼此独立,可以一次性读四个 float,用float4向量化。向量化的本质是让单个线程产生更多的内存级并行,减少指令数和访问次数,对带宽敏感型 kernel 非常有效。

extern "C" __global__ void fused_chain_vec4(const float4* a, const float4* b, const float4* c, float4* out, float scale, int n4) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < n4) { float4 va = a[i], vb = b[i], vc = c[i]; float4 t; t.x = fmaxf(va.x + vb.x, 0.f) * scale + vc.x; t.y = fmaxf(va.y + vb.y, 0.f) * scale + vc.y; t.z = fmaxf(va.z + vb.z, 0.f) * scale + vc.z; t.w = fmaxf(va.w + vb.w, 0.f) * scale + vc.w; out[i] = t; } }

这里有两个必须注意的坑:

第一,指针要对齐。float4*要求 16 字节对齐。cudaMalloc出来的显存没问题,但 TensorRT 引擎里拿到的 binding 指针未必对齐,尤其是经过一些 layout 优化之后。我吃过的亏是:插件里拿到 TRT 给的指针直接按 float4 读,结果非法地址错误排查了半天。正确做法是先检查指针地址和n4的整除性,不对齐就退回标量 kernel。

第二,尾块。n不一定被 4 整除,剩下的 1~3 个元素要用标量路径处理。我的习惯是源码生成时把n/4作为常量写进主循环,剩下的尾块单独生成一个 if 分支,甚至单独一个标量 kernel,别在主循环里加"最后一块处理一下"的判断,那样会干扰编译器向量化。

4.3 第三层:归约类融合,处理 LayerNorm / Softmax

elementwise 之外,另一个串行小核重灾区是归约类算子。TensorRT 处理 LayerNorm 往往会拆成均值 kernel、方差 kernel、归一化 kernel 三个小核,每个都对张量做一次完整读写。融合的思路是:一个 block 负责一行,先在 shared memory 里把均值和方差算出来,再做归一化,一次内存遍历解决问题。

extern "C" __global__ void fused_layernorm(const float* x, const float* gamma, const float* beta, float* y, float eps, int cols) { int row = blockIdx.x; int tid = threadIdx.x; const float* row_in = x + row * cols; float* row_out = y + row * cols; __shared__ float s_sum[256], s_sq[256]; __shared__ float s_mean, s_rstd; float sum = 0.f, sq = 0.f; for (int i = tid; i < cols; i += blockDim.x) { float v = row_in[i]; sum += v; sq += v * v; } s_sum[tid] = sum; s_sq[tid] = sq; __syncthreads(); if (tid == 0) { float ts = 0.f, tq = 0.f; for (int i = 0; i < blockDim.x; ++i) { ts += s_sum[i]; tq += s_sq[i]; } s_mean = ts / cols; float var = tq / cols - s_mean * s_mean; s_rstd = rsqrtf(var + eps); } __syncthreads(); for (int i = tid; i < cols; i += blockDim.x) { row_out[i] = (row_in[i] - s_mean) * s_rstd * gamma[i] + beta[i]; } }

这个版本是"一个 block 处理一行"的经典结构,适合 cols 在几百到一两千的场景。实际部署时我还会做两处升级:一是行内归约用 warp shuffle 替代 shared memory 数组,二是当一行太长时拆成多个 block 做 split-K 归约,再用 atomics 合并。不过对大多数推理模型来说,上面的版本已经能跑赢 TensorRT 拆出来的三个小核了。

注意__syncthreads()一定不能少,否则第二个循环里读s_mean、s_rstd时可能还是旧值。这个 bug 我写过一次,症状是输出里偶尔有 NaN,时有时无,特别阴间。

5. 把 DeepJIT 塞进推理链路的三种姿势

有了能跑的融合 kernel,接下来是想办法让它进到真正的推理链路里。按工程成本从高到低,有三种姿势。

5.1 姿势一:写成 TensorRT Plugin,最正统也最重

在 TensorRT 里注册自定义算子,走IPluginV2DynamicExt或新版的IPluginV3。plugin 的enqueue里调用你缓存的CUfunction,JIT 编译放在引擎加载阶段完成。

这个姿势的优点是融合真的发生在 TensorRT 图里,引擎序列化后整体一致;缺点是样板代码巨多,要处理 input/output 的格式协商、维度推断、序列化反序列化,还要和具体 TRT 版本绑定。我的建议是:如果团队有精力维护这套代码,选它;否则先别碰。

就算走这条路线,也别把 NVRTC 编译过程塞进 plugin 的 enqueue。正确做法是:plugin 构造函数或 configure 阶段编译好 kernel 并缓存,enqueue 只做cuLaunchKernel。推理热路径上任何编译动作都是灾难。

5.2 姿势二:在引擎外面替换瓶颈段(推荐先试)

我这次实际落地的就是这条路线。检测头的串行小核区域在计算图上是一整块连续区域:主干输出 -> 一堆 elementwise -> 后处理。我把它从 TensorRT 网络里拆掉,让引擎只跑主干,输出一个中间张量;后处理部分用 DeepJIT 生成的几个融合 kernel 在同一条 CUDA stream 上继续跑。

这样做的好处是:TensorRT 引擎保持"标准构建",不用碰 plugin API;中间张量的 shape 和 layout 是固定的,源码生成时能把所有 shape 硬化到极致;调试也简单,前面引擎跑完,后面核函数逐个打日志。

代价是你要保证引擎输出和后续 kernel 的输入格式完全一致,包括 channel order、dtype、stride。如果 TensorRT 引擎为了内部效率给输出定了 NHWC 或者半精度,你的融合 kernel 要么跟着改,要么在引擎输出处加一层转换。我的做法是在 ONNX 导出时就把尾部分割出来,让主干引擎输出强制为 NCHW + FP32,后处理 kernel 全部按这个约定写,省心很多。

5.3 姿势三:干脆自己组一小段执行图

最后一招是彻底绕开 TensorRT 来管那些病态的小算子。如果你的模型里有那么一段图,TensorRT 怎么动都是拆成小核,那你可以不把它放进引擎,而是在推理代码里直接按顺序 launch 你的 JIT kernel,甚至用cudaStreamBeginCapture把这几个 launch 录成 CUDA Graph。

我自己试过之后觉得,融合 + 手写执行序的体验很像在用"微型的 TVM",自由度很高——你可以任意调节 kernel 的 grid/block、shared memory 大小、换用不同的向量化策略——代价是失去了引擎的统一管理,shape 一变所有 launch 参数都要重新算。好在有 JIT 编译兜底,shape 变了重新生成源码重新编译就行,这套组合拳配合好之后其实挺稳的。

6. 实测数据与避坑清单

6.1 一组可复现的 benchmark 数据

交代一下测试环境:RTX 3060 12G,CUDA 11.8,TensorRT 8.5,模型是一个带复杂检测头的检测模型,batch=1,FP32。我主要优化的是检测头那一段串行小核区域,主干引擎保持不变。

指标原始 TensorRT只加 CUDA GraphDeepJIT 融合
尾部 region 的 kernel 数1321324
kernel 平均执行时长2.3 μs2.3 μs11 μs
估算 launch 开销约 370 μs约 10 μs约 10 μs
尾部内存往返次数7 次左右7 次左右2 次
尾部端到端耗时约 0.68 ms约 0.36 ms约 0.09 ms
整帧延迟(含主干)2.31 ms1.99 ms1.72 ms

数字只看比例就好,不同模型差异很大。关键结论是三条:CUDA Graph 能砍掉 launch 开销,但内存往返省不了,所以它把 0.68ms 降到 0.36ms;DeepJIT 融合把计算和内存一起省了,尾部直接降到 0.09ms;整帧从 2.31ms 降到 1.72ms,提升大约 26%,对延迟敏感场景是很可观的量级。

6.2 第一条坑:NVRTC 编译慢,忘了做缓存

NVRTC 编译一次五十毫秒起步,如果线上每次推理都编译,你的服务早就超时了。我的做法是三层缓存:源码字符串做 key 的unordered_map是最基本的;第二层是 PTX 落盘,进程重启后直接从磁盘读;第三层是上线前 warmup,把常见 shape 组合全部编译一遍。千万不要在 enqueue 或者热路径里触发新的编译。

6.3 第二条坑:-use_fast_math 把精度吃掉了

-use_fast_math会把除法、平方根、倒数都变成近似指令,对 elementwise 融合通常没问题,但对归一化、softmax 这类算子,误差可能大到肉眼可见。我的策略是:默认不开-use_fast_math,只在某一个 kernel 确实需要且验证过误差在可接受范围时才单独开。更细的做法是源码里只对特定表达式用__fdividef、__frsqrt_rn这类 intrinsic,而不是全局开。

6.4 第三条坑:cuModuleUnload 和异步流的恩怨

kernel 还在 GPU 上排队执行,你这边cuModuleUnload把包含它的模块卸载了,轻则随机的非法地址,重则直接 CUDA_ERROR_ILLEGAL_ADDRESS,而且这错误还不好复现,因为它是时序相关的。我的规矩是:除非确定所有相关流都同步过,否则 module 不主动卸载,让它和进程同生共死。一个推理进程里模块数量本来就不多,内存又不是问题,没必要在这省。

6.5 第四、五条坑:对齐假设与指针所有权

第四坑是内存对齐。前面在 float4 那节提过,这里再强调一次:从 TRT binding、PyTorch tensor、自建显存池拿到的指针,对齐情况各不相同。写融合 kernel 前先写个 assertion 检查指针地址和步长,别默认人家帮你对齐了。

第五坑是指针所有权。如果载体是 PyTorch,你从tensor.data_ptr()拿到的指针只有在 tensor 存活时有效,融合 kernel 是异步的,kernel 还没跑完,tensor 就被释放回显存池,你读到的就是脏数据。稳妥做法是在 launch 前tensor.record_stream(stream),或者干脆让张量跨过 kernel 的生命周期再释放。

7. 边界感:哪些模型不值得上 DeepJIT

7.1 大算子为主时,收益约等于零

DeepJIT 解决的问题非常具体:大量短 kernel 串行执行 + 全局内存反复读写。如果你的模型主力是卷积、GEMM、attention 这类大算子,TensorRT 早就把它们融合得很好了,你再手写 kernel 纯属给自己找活。判断标准很简单:profiler 里短 kernel 占比高不高,内存流量是不是远超理论值。两者都不满足,就不要折腾。

7.2 动态 shape 下的维护成本

动态 shape 是 DeepJIT 最头疼的敌人。每一个新的 shape 组合都可能触发一次新的源码生成和编译,缓存体积膨胀、warmup 枚举不完、线上偶发编译超时。我的经验是把连续 shape 量化为有限的几个桶:batch 只允许 1/2/4/8,分辨率只允许预设的几种,一旦命中桶就直接 pad 到桶内固定大小,用 mask 或尾块处理掉多出来的部分。这样编译次数从无限收敛到几十次,完全可控。

最后说点个人感受。做完这个项目之后,我最大的变化是再也不把推理优化当成"换引擎"或者"调参数"了。先看 profiler,确认瓶颈是启动开销还是带宽,再决定上 CUDA Graph 还是 DeepJIT 融合。前者是消除排队,后者是减少干活次数,两者解决的问题维度不一样,配合起来才是完整的优化思路。希望这篇尝鲜记录也能帮你少走几步弯路。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询