ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Megakernel

Megakernel 设计Megakernel是为了解决哪些问题硬件层面同 stream 相邻 kernel 之间有隐式 barrier。同一个 stream 上的两个 kernel即使前一个只用了 1 个 SM、GPU 上还有 100 多个 SM 全空着后一个也必须等它完全结束才能启动例如算子的尾部效应。软件层面生产者的输出没有精准地投递给消费者消费者只能等所有生产者全部完成才开始消费。软件层面减少launch开销前向~50微秒后向~100微秒。减少launch开销还有另一种解决方案就是CUDA Graph。但 CUDA Graph 是静态的控制流、张量形状或数据依赖一旦发生变化就必须重新捕获或修改已录制的图因此难以应对模型推理中常见的动态负载。解决方式把模型推理中所有的计算和通信都融合进一个巨型 kernelmega-kernel也叫 persistent kernel。在这种设计下系统只启动一个 GPU kernel由它来跑完整个模型。性能论文号称比SGLang提升1-1.7x但实际性能不如SGLang 0.5.12。run_tgx没开mtp# python3 scripts/plot_fig9_h100.py qwen3-30b-a3b (metric: per-request TPOT, ms/token) bs mirage sglang mirage/sglang winner 1 6.646 4.741 1.40 sglang 2 7.903 5.872 1.35 sglang 4 9.478 7.343 1.29 sglang 8 12.441 8.792 1.41 sglang 16 16.846 10.011 1.68 sglang - Mirage loses its latency edge starting at bs1. qwen3-8b (metric: per-request TPOT, ms/token) bs mirage sglang mirage/sglang winner 1 6.869 6.506 1.06 sglang 2 7.038 6.755 1.04 sglang 4 7.279 6.738 1.08 sglang 8 7.504 6.949 1.08 sglang 16 8.197 7.221 1.14 sglang - Mirage loses its latency edge starting at bs1. Wrote results/H100/fig9_latency.png Wrote results/H100/fig9_throughput.png如何使用Megakernel目前要求用户手写一遍整个模型的计算图并且每一层都要自己算grid_dim/block_dim比如demo/qwen3/demo.py是881 行并不像论文中声称的用户只需要写几行代码就可以将整个模型的算子融合为一个大算子。Megakernel实现方式把整个模型的计算和跨 GPU 通信拆成一个个「以单个 SM 为执行单位的小任务」用一张记录任务间细粒度依赖的图tGraph把它们串起来然后由一个常驻 GPU 的内核按依赖关系自行调度这些任务——而不是像传统方式那样每个算子启动一次占满全 GPU 的 kernel。MPK Compiler把一个模型用到的所有算子组成的DAG图拆成SM级别的算子DAG图。换句话说把模型用到的每个算子切成一堆小task每个task小到能由一个 SM 独立完成记录到task_graph.json里面生成test.cu调nvcc将test.cu编成.so拆图过程拼模型计算图-generate_task_graph - register_mugraph - print_task_graph/* 标注拆图的信息 */ struct AnnotatedGraph { std::vectorLayerInfo layers; // 顶点每个 KN_CUSTOMIZED_OP 一个 std::vectorint ordered_layers; // 拓扑序决定 task 发射顺序 std::vectorEdgeInfo edges; // 边扁平存放靠下标引用 std::vectorForkGroupInfo fork_groups; std::vectorJoinGroupInfo join_groups; std::vectorEdgeInfo stripped_residual_edges; };细粒度流水线# src/kernel/annotated_graph.cc auto prod_part build_partition(prod_op-bgraph.grid_dim, e.output_map); auto cons_part build_partition(cons_op-bgraph.grid_dim, e.input_map); for (int d 0; d (int)mirage::config::MAX_TENSOR_DIMS; d) { e.event_dim[d] std::gcd(prod_part[d], cons_part[d]); }切event例子grid(128,1,1) map(1,-1,-1) part [1, 128, 1, 1]map的原型是map(grid.x切tensor第几维, grid.y切tensor第几维, grid.z切tensor第几维)。没有被任何grid维度切的填1例子中的map显示grid.0切的是input tensor第一维也就是tensor的第一维被切成128份。总event数 event_dim[·]各维度的乘积。处理fork producer:F grid(4,1,1)/ \A B A grid(4,1,1), B grid(2,1,1)\ /G grid(4,1,1)A: grid_dim (4,1,1), input_map: x - dim0 B: grid_dim (2,1,1), input_map: x - dim0step (g) 逐边计算边 F→Aevent_dim[0] gcd(4, 4) 4producer 侧last3.x 4/4 1边 F→Bevent_dim[0] gcd(4, 2) 2producer 侧last3.x 4/2 2两条分支不一致1 vs 2F 没法用一个 event 同时服务两边。step (h)lcm_last3.x lcm(1, 2) 2能整除grid.x 4安全检查通过。分支 Ascale 2/1 2→event_dim[0]: 4 → 2consumer 侧last3.x 4/2 2分支 Bscale 2/2 1不变consumer 侧last3.x 2/2 1结果两条边都是 2 个 eventF 的每个 event 覆盖 2 个 producer taskx∈{0,1} 和 x∈{2,3}第 0 个 event 触发 A 的 task 0-1 和 B 的 task 0。代价是 A 的同步粒度从 4 个 event 粗化到 2 个。处理join consumer:同个例子边 A→Gevent_dim[0] gcd(4, 4) 4consumer 侧 last3.x 4/4 1边 B→Gevent_dim[0] gcd(2, 4) 2consumer 侧 last3.x 4/2 2G 是 join-consumer它的 task 只有一个 dependent_event 槽所以两条入边必须落在 G 的 grid 上的同一套划分上——现在一个说切 4 份、一个说切 2 份冲突。step (i) join LCMannotated_graph.cc:638-646lcm_last3.x lcm(1, 2) 2能整除 G.grid.x 4 ✓边 A→Gscale 2/1 2 → event_dim[0]: 4 → 2回推 producer 侧 last3.x A.grid.x / 2 2边 B→Gscale 1event_dim[0] 保持 2producer 侧 last3.x B.grid.x / 2 1最终两个 join eventjoin-event-0num_triggers 2 (A的 task 0,1) 1 (B的 task 0) 3放行 G 的 task 0-1join-event-1num_triggers 2 (A的 task 2,3) 1 (B的 task 1) 3放行 G 的 task 2-3gcd例子生产者切 128 份消费者切 128 份 → gcd128 → 128 个 event每个 1→1完全流水生产者切 128 份消费者切 64 份 → gcd64 → 64 个 event每个 2→1生产者切 128 份消费者切 1 份 → gcd1 → 1 个 event128→128全屏障生产者切128份消费者切127份 → gcd1 → 1 个 event128→127全屏障搭建SM级DAtask信息包括task type, variant id, input切分output切分trigger_event, dependent_event最后两个创建时不填遍历op级DAG的linearization中每个节点(layer): 按角色分四种情形: first layer(无入边)只按bid字典序把tasks塞进all_tasks记入first_tasks不发event fork bundle(head)遍历producer侧的event维度 join consumer遍历consumer侧的event维度 chain layer递归遍历event_dim各维叶子即一个event 对每个event索引: 创建一个event event.first_task_id all_tasks.size() 把这个event该触发的下游tasks创建出来塞进all_tasks // 连续 event.last_task_id all_tasks.size() 遍历上游producer的bid子范围 该task.trigger_event 当前event的id event.num_triggers all_events.push_back(event) 所有task和event构建完毕遍历所有event的下游task统一设置所有task的dependent_event task { task type variant id, input map, output map, trigger event, dependent event }完成all_tasks, all_events, first_tasks构建输出为test.cu和task_graph.json两个文件。将all_events、all_tasks、first_tasks写入task_graph.json 将编译、运行函数写入test.cu HARD_CODEinit_func由mpk.compile()调用 _init_persistent_kernel 由init_func调用分配显存地址 Construct_task_graph由_init_persistent_kernel调用从task_graph.json中读出 all_tasks, all_events, first_tasks _execute_task由worker调用task_graph.jsonall_tasks[task_type, inputs, outputs, dependent_event, trigger_event]我是什么任务、动哪块数据、等谁、完事通知谁all_events[num_triggers1, first_task_id4, last_task_id100]等 1 个任务向我打卡之后我就把 4~100 号工单派出去first_tasks不用等任何人kernel 一起来就派它。test.cuconstruct_task_graph()反序列化task_graph.json读出all_tasks, all_events, first_tasks_init_persistent_kernel()分配显存地址_execute_task()这个函数会根据task_type, variant_id找到对应的CUDA kernel调用代码all_task_variants存着这张图所有要用到的算子的调用和调用参数代码把task_desc-input_ptrs/output_ptrs传进去跑# tests/runtime_python/test_mode/test_rmsnorm_testmode.py pk.compile(output_dirfolder_path) # python/mirage/mpk/persistent_kernel.py results self.kn_graph.generate_task_graph(num_gpusself.world_size, my_gpu_idself.mpi_rank) # python/mirage/kernel.py def generate_task_graph(self, num_gpus: int, my_gpu_id: int): return self.cygraph.generate_task_graph(num_gpus, my_gpu_id)/* src/kernel/runtime.cc */ TaskGraphResult Graph::generate_task_graph(int _num_gpus, int _my_gpu_id) { /* 一共有哪些task、每个task读写哪块数据、谁等谁 产出三个C数组all_tasks, all_events, first_tasks */ register_mugraph(...); /* 把上面三个数组序列化成 task_graph.json再拼出 test.cu */ print_task_graph(...); } TaskGraphResult print_task_graph(...) { /* 名字显存地址*/ if (use_json_format) { code.e(std::mapstd::string, void* all_tensors;); } for (auto const iter : io_configs) { IODesc desc iter.second; switch (desc.type) { /* 其他case */ case IODesc::CUDAMallocTensor: { code.e(void *$;, desc.name); size_t size mirage::type::get_datatype_size( static_casttype::DataType(desc.tensor.data_type)); for (int i 0; i desc.tensor.num_dims; i) { size * desc.tensor.dim[i]; } /* 生成test.cu代码: 现场malloc一个地址 */ code.e(CUDA_CHECK(cudaMalloc($, $));, desc.name, size); if (use_json_format) { code.e(all_tensors[\$\] $;, desc.name, desc.name); } break; } /* 其他case */ } if (use_json_format) { // Add nullptr for tensors set as None code.e(all_tensors[\nullptr\] nullptr;); /* 这个函数会将JSON反序列化将JSON文件包括显存地址读回内存变成C对象 */ code.e(construct_task_graph(num_gpus, my_gpu_id, all_tasks, all_events, first_tasks, all_tensors);); } else { code.e(tgbody.to_string()); }task如何找对应的实现pk PersistentKernel对象PersistentKernel对象是一个建造者builder。TaskRegister是一个单例对象整个程序只有一个TaskRegister对象。# tests/runtime_python/test_mode/test_rmsnorm_testmode.py # 搭建模型计算图 pk.rmsnorm_layer(inputx_dt, weightw_dt, outputout_dt, grid_dim(batch_size, 1, 1), block_dimblock_dim) # python/mirage/mpk/persistent_kernel.py def rmsnorm_layer( self, input: DTensor, weight: DTensor, output: DTensor, grid_dim: tuple, block_dim: tuple, ): self.kn_graph.register_task(tb_graph, rmsnorm_hopper if self.target_cc 90 else rmsnorm) # src/kernel/graph.cc void Graph::register_task(char const *task_type, std::vectorint params) { else if (name rmsnorm_hopper) { int variant_id task_register-register_rmsnorm_hopper_task(customized-bgraph, params); task_config[op] std::make_tuple(2, 1, TASK_RMS_NORM_HOPPER, variant_id); } } int TaskRegister::register_rmsnorm_hopper_task(threadblock::Graph const bgraph, std::vectorint const params) { mirage::transpiler::CodeKeeper code; code.inc_indent(); code.e( kernel::rms_norm_hopper_implbfloat16, $, $(, batch_size, hidden_dim); code.e( task_desc-input_ptrs[0],); code.e( task_desc-input_ptrs[1],); code.e( task_desc-output_ptrs[0],); code.e( 1e-6f);); return register_task_variant(TASK_RMS_NORM_HOPPER, code.to_string()); }# include/mirage/persistent_kernel/tasks/hopper/rmsnorm_hopper.cuh namespace kernel { template typename T, int BATCH_SIZE, int HIDDEN_DIM, int NUM_THREADS 256 __device__ __forceinline__ void rms_norm_hopper_impl(void const *input_ptr, void const *weight_ptr, void *output_ptr, float eps) {...}register_task_variant会将生成的调用和调用参数放到all_task_variants[type]下。generate_task_graph会遍历 all_task_variants 生成一个巨大的分派函数 _execute_task。这些算子并不完全是mirage团队自己开发的有些来自FlashInfer有些来自DeepGEMM。In-kernel parallel runtimeMegakernel的runtime指的是调度代码。In-kernel指的是把调度代码搬进了kernel里不再依靠CPU逐个启动kernel。为了把调度代码搬进kernelMegakernel将GPU的SM分成worker和scheduler两大角色。运行过程prefill跟SGLang--chunked-prefill-size默认为8192不同Megakernel的chunked prefill size被设定为跟batch size是一样大因为Megakernel的prefill和decode用的是同一张计算图这张计算图在compile时就已经焊死之后不能改。其实也就是说Megakernel并没有chunked prefill size这个概念如果实在要说那就是把batch size当成chunked prefill size。Scheduler在Scheduler SM上每个warp的0号线程当一个Scheduler每个SM上4个Scheduler。Scheduler负责哪些worker以B200为例sched 0 → worker [0, 9)sched 1 → [9, 18)sched 2 → [18, 27)sched 3 → [27, 36)sched 4 → [36, 45)sched 5 → [45, 54)sched 6 → [54, 63)sched 7 → [63, 72)sched 8 → [72, 81)sched 9 → [81, 90)sched 10 → [90, 99)sched 11 → [99, 108)sched 12 → [108, 117)sched 13 → [117, 126)sched 14 → [126, 135)sched 15 → [135, 144)execute_schedulers算法: 死循环直到persistent kernel完成所有计算任务 轮询取一个event在local和广播队列间来回切换 如果是termination event即没有更多的request 给所有Scheduler发一个termination event 每个Scheduler收到后给自己的worker发taskid 0 如果是一次计算图计算结束 派TASK_BEGIN_TASK_GRAPH任务给下一个worker 如果是计算图根节点EVENT_LAUNCH_DEPENDENT_TASKS 交错分发task 如果是EVENT_LAUNCH_MASSIVE_TASKS 均分task后派发[first_task_id, last_task_id) 如果是普通event 派发[first_task_id, last_task_id)重要的Task/Event类型TASK_BEGIN_TASK_GRAPH: DAG的根节点这个任务没有任何计算量作用仅在于把EVENT_LAUNCH_DEPENDENT_TASKS加入广播队列然后触发EVENT_LAUNCH_DEPENDENT_TASKS。如何保持负载均衡EVENT_LAUNCH_MASSIVE_TASKS触发≥8个任务和EVENT_LAUNCH_DEPENDENT_TASKS触发DAG所有task会被塞进广播队列。每个Scheduler都会去读广播队列每个Scheduler用一个私有的pointer去遍历广播队列里的event。EVENT_LAUNCH_MASSIVE_TASKS会按照Scheduler数量均分任务区间[first_task_id, last_task_id)每个Scheduler拿到的任务区间是均分过的。EVENT_LAUNCH_DEPENDENT_TASKS则是交错分发tasksched0 一叠连续的task、sched1 一叠连续的task、sched2 一叠连续的task、sched3 一叠连续的task然后回到 sched0 继续。这个事件触发的task数量DAG中所有task数量远大于worker数量所以不存在只有少数SM在干活儿的情况。Workerexecute_workers算法: 死循环直到persistent kernel完成所有计算任务 如果上一批计算任务已经完成 从remote/local worker队列中取一批计算任务 拿到下一个计算任务 如果计算任务有依赖的事件 阻塞直到计算任务的依赖完成通过轮询的方式检查依赖事件是否完成 _execute_task(task_desc, config) 完成计算任务打卡下游事件 如果下游事件集齐全部打卡 如果事件触发派发大量计算任务≥8个 将该下游事件加入广播队列 如果事件只触发少量计算任务 找到该worker所属的Scheduler 将该下游事件加入该Scheduler的任务队列
返回列表