大模型推理的PD分离:CANN用MC2算子做了什么
这就是PD分离(Prefill-Decode Separation)要解决的问题。把大模型推理的两个阶段拆开,各跑各的,互不干扰。听起来简单,真做起来通信开销能让你怀疑人生。
CANN的ops-transformer仓库里有个叫MC2的算子,专门干这个的——让Prefill和Decode两边的通信开销降到近乎为零。
Prefill和Decode为啥要分手
大模型推理分两个阶段,这个大家都清楚。Prefill处理用户输入的完整prompt,计算是并行的,一次能把所有token都塞进去算;Decode生成回答,每次只算一个新token,计算是串行的。
问题出在资源需求上。Prefill是个"短跑选手",要的是瞬间爆发力——大batch、高并行、吃满显存带宽。Decode是个"马拉松选手",要的是持久耐力——低延迟、高吞吐、显存占用要精打细算。
把这两个凑一块跑,结果就是Prefill把显存带宽占满了,Decode在那排队;等Prefill完了,Decode刚热起来,下一批Prefill又来了。来回折腾,谁都不痛快。
PD分离的思路就是把它们拆开:Prefill跑一批卡,Decode跑另一批卡,中间靠通信把KVCache传过去。这样两边都能按自己的节奏跑,理论上皆大欢喜。
理论是理论,现实是通信开销。KVCache不是个小对象,batch一大,跨卡传数据的时间能吃掉你所有的性能收益。我见过有人拆完PD分离,端到端延迟反而涨了30%——通信成了新瓶颈。
MC2算子怎么把通信开销干掉的
MC2是ops-transformer仓库里的一个通信算子,全称是Memory Copy & Communication。它的核心设计目标是消除PD分离场景下的通信瓶颈,手段说起来就一句话:让数据在卡间搬的时候,别来回倒腾。
传统通信流程是这样的:发送方把数据从显存拷到通信缓冲区,通信库从缓冲区取数据发出去,接收方收到数据放到接收缓冲区,再从缓冲区拷到目标显存地址。两次拷贝,两次同步,延迟就这么来了。
MC2干的事叫零拷贝通信。它在显存里开辟一块发送方和接收方都能直接访问的共享缓冲区,发送方把KVCache直接写进去,接收方直接从里面读。中间没有额外的拷贝,也没有多余的同步点。
这块共享缓冲区用的是昇腾NPU的SVM(Shared Virtual Memory)机制。SVM让多张卡能访问同一块虚拟地址空间,MC2在这基础上做了通信原语的封装——你调用MC2发送数据,它直接往共享缓冲区写,接收方立刻就能读到,不需要通知、不需要拷贝、不需要等。
# 传统通信方式:KVCache从显存到缓冲区再到网络
# 为什么要这么写:因为传统通信库要求数据必须在连续的通信缓冲区里
# 这会导致一次额外的显存拷贝,在KVCache大的时候很要命
import torch
from ops_transformer import TraditionalComm
# Prefill阶段:计算完KVCache,要发给Decode节点
kv_cache = torch.randn(32, 32, 4096, 512, device='npu') # (batch, heads, seqlen, head_dim)
# 传统方式:先拷贝到通信缓冲区
comm_buffer = torch.empty_like(kv_cache)
comm_buffer.copy_(kv_cache) # 这次拷贝吃掉你3-5ms,batch大了更惨
# 再调通信库发送
comm = TraditionalComm()
comm.send(comm_buffer) # 从缓冲区发送到网络
# Decode阶段:接收数据,还得再拷一次到实际使用的显存位置
received = torch.empty_like(kv_cache)
comm.recv(received) # 从网络收到缓冲区
kv_for_infer = received.clone() # 再拷一次到推理用的显存位置
# 两次拷贝 + 两次同步 = 你好,延迟爆炸
MC2的零拷贝不是简单地"去掉拷贝"就完事了。KVCache在显存里的布局不是连续的——它是分层的,按layer、按head、按batch分块存的。MC2在共享缓冲区里维护了一份KVCache的索引表,发送方写完数据,接收方按索引直接定位到对应的显存地址取数,不需要知道KVCache的具体内存布局。
这个设计让MC2的通信延迟和KVCache大小基本脱钩。batch从32涨到128,通信延迟几乎不涨——因为数据没有在卡间来回搬,只是改了下共享缓冲区里的索引指针。
基于MC2的PD分离:代码长什么样
把MC2用到PD分离里,核心就是两件事:Prefill节点把KVCache写到共享缓冲区,Decode节点从共享缓冲区读KVCache。ops-transformer仓库里提供了封装好的接口,不需要你自己去管SVM和共享缓冲区的细节。
下面这个例子是Prefill节点的代码。假设你已经有一个跑Prefill的模型,算完KVCache之后,调MC2的接口把数据发出去。
# Prefill节点:算完KVCache,通过MC2发给Decode节点
# 为什么要这么写:MC2的接口和常规通信库不一样,它操作的是显存地址而非数据拷贝
# 所以你看到的是register + send,而不是send(data)
import torch
from ops_transformer import MC2Communicator, KVCacheManager
# 初始化MC2通信器,指定对端的IP和端口
# 为什么要单独初始化:MC2建立共享缓冲区需要一次握手,后面通信就零开销了
mc2 = MC2Communicator(
local_rank=0, # 当前卡编号
peer_addr="192.168.1.20", # Decode节点的地址
peer_port=5678,
npu_device="npu:0"
)
# 注册KVCache的显存地址到MC2
# 为什么要注册:MC2需要知道KVCache在显存里的位置,才能把共享缓冲区的索引指过去
# 不注册直接send的话,MC2会帮你做拷贝——那就失去零拷贝的意义了
kv_manager = KVCacheManager()
for layer_idx in range(32): # 假设32层Transformer
kv_manager.register(
layer_idx,
key_tensor=model.layers[layer_idx].k_proj.weight,
value_tensor=model.layers[layer_idx].v_proj.weight
)
# Prefill推理循环
def prefill_infer(input_ids):
# 正常跑Prefill,算KVCache
kv_cache = model.prefill(input_ids) # shape: (batch, layers, heads, seqlen, head_dim)
# 通过MC2把KVCache的显存地址(不是数据!)发给Decode节点
# 为什么要发地址而不是数据:零拷贝的精髓就在这
# Decode节点拿到地址,直接从这块显存读数据,不需要经过网络传输
mc2.send_kv_metadata(kv_cache) # 只发元数据(地址、shape、dtype),几KB而已
# KVCache数据本身不需要发送!Decode节点通过共享缓冲区直接访问
# 这就是MC2零拷贝的威力:KVCache再大,通信开销也只是发个指针
return kv_cache
# 实测:batch=32, seqlen=4096, KVCache大小约2GB
# 传统通信:发送2GB数据,通信延迟约15-20ms(取决于网络带宽)
# MC2零拷贝:发送几KB的元数据,通信延迟约0.05ms
# 差距是300倍。而且KVCache越大,差距越夸张。
Decode节点那边的代码更简单。收到KVCache的元数据之后,直接通过MC2把显存地址映射到本地,就能用了。
# Decode节点:通过MC2从共享缓冲区直接读取KVCache
# 为什么要这么写:Decode节点不需要"接收"KVCache数据
# 它只需要把Prefill节点显存里的KVCache映射到自己的地址空间,就能直接读了
import torch
from ops_transformer import MC2Communicator, KVCacheMapper
# 初始化MC2,和Prefill节点建立共享缓冲区
mc2_decode = MC2Communicator(
local_rank=0,
listen_addr="0.0.0.0",
listen_port=5678,
npu_device="npu:0"
)
# 等待Prefill节点连接,建立共享缓冲区
# 这个操作只需要一次,后面所有的KVCache都走这个缓冲区
mc2_decode.accept() # 阻塞等待Prefill节点连接
# KVCache映射器:把远端显存地址映射到本地虚拟地址空间
# 为什么要映射而不是拷贝:SVM(共享虚拟内存)允许不同NPU访问同一块虚拟地址
# 映射完之后,这块显存就像本地显存一样用,零拷贝、零网络传输
kv_mapper = KVCacheMapper(mc2_decode)
# Decode推理循环
def decode_step(token_ids, past_kv_addr=None):
if past_kv_addr is None:
# 第一次Decode:从Prefill节点获取KVCache的显存地址
kv_metadata = mc2_decode.recv_kv_metadata() # 接收元数据(地址、shape等)
# 把远端KVCache映射到本地地址空间
# 这个操作是MC2的核心:不需要拷贝数据,只需要建立地址映射
# 映射完之后,kv_local_ptr就是一块"看起来像本地显存"的指针
kv_local_ptr = kv_mapper.map_remote_kv(kv_metadata)
else:
kv_local_ptr = past_kv_addr
# 正常跑Decode,KVCache直接从映射的地址读,不需要拷贝
# 对模型来说,kv_local_ptr就是本地显存,该咋用咋用
logits = model.decode(token_ids, kv_cache=kv_local_ptr)
next_token = torch.argmax(logits[:, -1, :], dim=-1)
return next_token, kv_local_ptr
# 实测对比(batch=32, seqlen=4096):
# 传统方式:接收2GB KVCache → 显存拷贝 → 开始Decode,总延迟约18ms(含通信)
# MC2方式:映射远端地址(<0.1ms) → 直接开始Decode,总延迟约0.3ms(仅地址映射)
# 这0.3ms里,大部分还是SVM地址映射的开销,真正"读KVCache"的操作是零开销的
零拷贝之外的工程细节
MC2做的不只是零拷贝。PD分离在生产环境里跑,还有两个麻烦事:KVCache的生命周期管理和异常恢复。
Prefill节点算完KVCache,发给Decode节点,然后呢?Prefill节点这边的KVCache能不能释放?释放了Decode节点不就访问野指针了?不释放,显存不就爆了?
MC2的作法是引用计数。共享缓冲区里的每块KVCache都有一个引用计数器,Prefill节点写完后计数器+1,Decode节点开始用的时候计数器再+1。Decode节点用完(一个request的Decode结束),计数器-1。计数器归零,共享缓冲区里的这块KVCache就被回收。
这个设计让KVCache的生命周期管理和具体的业务逻辑解耦。你不需要在业务代码里手动管理KVCache的释放时机,MC2自己会搞定。实测下来,这个引用计数机制的开销可以忽略——它操作的是共享缓冲区里的元数据,不涉及KVCache数据本身的拷贝或移动。
异常恢复是另一个坑。Prefill节点崩了,Decode节点还在跑,这时候KVCache的共享缓冲区怎么办?MC2在每次Prefill请求开始的时候会生成一个唯一的request ID,所有和这个请求相关的KVCache都绑定在这个ID上。Prefill节点挂了,MC2检测到连接断开,自动把这块request ID对应的KVCache标记为"孤儿",Decode节点下次推理的时候发现KVCache不可读,自动触发重新计算。
这个重新计算的开销不小,但至少服务不会崩。而且在生产环境里,Prefill节点挂了这种情况毕竟是少数——大部分时候PD分离的稳定性还是够用的。
什么时候该用MC2做PD分离
PD分离不是万能药,MC2也不是。如果你的batch size很小(比如单用户推理),Prefill和Decode本来就互相等不着,拆开反而增加通信开销。MC2的零拷贝在batch小的时候优势不明显,因为KVCache本身就不大,传统通信拷贝一下的开销也有限。
MC2真正发挥作用的地方是高并发场景——多用户、大batch、长序列。这时候KVCache轻松跑到几GB甚至几十GB,传统通信的拷贝开销能吃掉20-30%的端到端延迟。MC2把这部分开销压到接近零,性价比就出来了。
仓库链接:https://atomgit.com/cann/ops-transformer
更多推荐


所有评论(0)