Megakernel
2026/9/24 17:13:16 网站建设 项目流程

设计Megakernel是为了解决哪些问题?

  • 硬件层面:同 stream 相邻 kernel 之间有隐式 barrier。同一个 stream 上的两个 kernel,即使前一个只用了 1 个 SM、GPU 上还有 100 多个 SM 全空着,后一个也必须等它完全结束才能启动,例如算子的尾部效应。
  • 软件层面:生产者的输出没有精准地投递给消费者,消费者只能等所有生产者全部完成才开始消费。
  • 软件层面:减少launch开销(前向~50微秒,后向~100微秒)。减少launch开销还有另一种解决方案,就是CUDA Graph。但 CUDA Graph 是静态的:控制流、张量形状或数据依赖一旦发生变化,就必须重新捕获或修改已录制的图,因此难以应对模型推理中常见的动态负载。

解决方式

把模型推理中所有的计算和通信都融合进一个巨型 kernel(mega-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 bs=1. ===== 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 bs=1. 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::vector<LayerInfo> layers; // 顶点:每个 KN_CUSTOMIZED_OP 一个 std::vector<int> ordered_layers; // 拓扑序,决定 task 发射顺序 std::vector<EdgeInfo> edges; // 边:扁平存放,靠下标引用 std::vector<ForkGroupInfo> fork_groups; std::vector<JoinGroupInfo> join_groups; std::vector<EdgeInfo> 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 -> dim0

step (g) 逐边计算:

  • 边 F→A:event_dim[0] = gcd(4, 4) = 4,producer 侧last3.x = 4/4 = 1
  • 边 F→B:event_dim[0] = gcd(4, 2) = 2,producer 侧last3.x = 4/2 = 2

两条分支不一致(1 vs 2),F 没法用一个 event 同时服务两边。

step (h):lcm_last3.x = lcm(1, 2) = 2(能整除grid.x = 4,安全检查通过)。

  • 分支 A:scale = 2/1 = 2event_dim[0]: 4 → 2;consumer 侧last3.x = 4/2 = 2
  • 分支 B:scale = 2/2 = 1,不变;consumer 侧last3.x = 2/2 = 1

结果:两条边都是 2 个 event,F 的每个 event 覆盖 2 个 producer task(x∈{0,1} 和 x∈{2,3});第 0 个 event 触发 A 的 task 0-1 和 B 的 task 0。代价是 A 的同步粒度从 4 个 event 粗化到 2 个。

处理join consumer:(同个例子)

边 A→G:event_dim[0] = gcd(4, 4) = 4,consumer 侧 last3.x = 4/4 = 1
边 B→G:event_dim[0] = gcd(2, 4) = 2,consumer 侧 last3.x = 4/2 = 2
G 是 join-consumer,它的 task 只有一个 dependent_event 槽,所以两条入边必须落在 G 的 grid 上的同一套划分上——现在一个说切 4 份、一个说切 2 份,冲突。

step (i) join LCM(annotated_graph.cc:638-646):lcm_last3.x = lcm(1, 2) = 2,能整除 G.grid.x = 4 ✓

边 A→G:scale = 2/1 = 2 → event_dim[0]: 4 → 2;回推 producer 侧 last3.x = A.grid.x / 2 = 2
边 B→G:scale = 1,event_dim[0] 保持 2;producer 侧 last3.x = B.grid.x / 2 = 1
最终两个 join event:

join-event-0:num_triggers = 2 (A的 task 0,1) + 1 (B的 task 0) = 3,放行 G 的 task 0-1
join-event-1:num_triggers = 2 (A的 task 2,3) + 1 (B的 task 1) = 3,放行 G 的 task 2-3

gcd例子:

生产者切 128 份,消费者切 128 份 → gcd=128 → 128 个 event,每个 1→1,完全流水
生产者切 128 份,消费者切 64 份 → gcd=64 → 64 个 event,每个 2→1
生产者切 128 份,消费者切 1 份 → gcd=1 → 1 个 event,128→128,全屏障

生产者切128份,消费者切127份 → gcd=1 → 1 个 event,128→127,全屏障

搭建SM级DA

task信息包括: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_CODE:init_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.json

  • all_tasks:[task_type, inputs, outputs, dependent_event, trigger_event]我是什么任务、动哪块数据、等谁、完事通知谁
  • all_events:[num_triggers=1, first_task_id=4, last_task_id=100]等 1 个任务向我打卡,之后我就把 4~100 号工单派出去
  • first_tasks:不用等任何人,kernel 一起来就派它。

test.cu

  • construct_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_dir=folder_path) # python/mirage/mpk/persistent_kernel.py results = self.kn_graph.generate_task_graph(num_gpus=self.world_size, my_gpu_id=self.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::map<std::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_cast<type::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(input=x_dt, weight=w_dt, output=out_dt, grid_dim=(batch_size, 1, 1), block_dim=block_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::vector<int> 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::vector<int> const &params) { mirage::transpiler::CodeKeeper code; code.inc_indent(); code.e( "kernel::rms_norm_hopper_impl<bfloat16, $, $>(", 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 runtime

Megakernel的runtime指的是调度代码。In-kernel指的是把调度代码搬进了kernel里,不再依靠CPU逐个启动kernel。

为了把调度代码搬进kernel,Megakernel将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则是交错分发task:sched0 一叠连续的task、sched1 一叠连续的task、sched2 一叠连续的task、sched3 一叠连续的task,然后回到 sched0 继续。这个事件触发的task数量=DAG中所有task数量,远大于worker数量,所以不存在只有少数SM在干活儿的情况。

Worker

execute_workers算法: 死循环直到persistent kernel完成所有计算任务: 如果上一批计算任务已经完成: 从remote/local worker队列中,取一批计算任务 拿到下一个计算任务 如果计算任务有依赖的事件: 阻塞直到计算任务的依赖完成(通过轮询的方式检查依赖事件是否完成) _execute_task(task_desc, config) 完成计算任务,打卡下游事件 如果下游事件集齐全部打卡: 如果事件触发派发大量计算任务(≥8个): 将该下游事件加入广播队列 如果事件只触发少量计算任务: 找到该worker所属的Scheduler 将该下游事件加入该Scheduler的任务队列

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

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

立即咨询