0 前言
按照官方代码中学习手册里的顺序对源码进行学习,再根据自己的技术方向加相应的feature进行学习。
1 qwen3 0.6B模型结构
下载的模型参数是safetensor类型,以pt文件储存,打印模型结构:
Qwen3ForCausalLM
|- model: Qwen3Model
|- embed_tokens: VocabParallelEmbedding [weight(151936, 1024)]
|- layers: ModuleList
|- 0: Qwen3DecoderLayer × 27
|- input_layernorm: LayerNorm [weight(1024,)]
|- self_attn: Qwen3Attention
|- qkv_projection: QKVColumnParallelLinear [weight(4096, 1024)]
|- q_norm: LayerNorm [weight(128,)]
|- k_norm: LayerNorm [weight(128,)]
|- rotary_emb: RotaryEmbedding
|- attention: Attention
|- o_proj: RowParallelLinear [weight(1024, 2048)]
|- post_attention_layernorm: LayerNorm [weight(1024,)]
|- mlp: Qwen3MLP
|- gate_up: MergedColumnParallelLinear [weight(6144, 1024)]
|- activation: SiluAndMul
|- down_proj: RowParallelLinear [weight(1024, 3072)]
|- norm: LayerNorm [weight(1024,)]
|- lm_head: ParallelLMHead [weight(151936, 1024)]

该项目使用dense架构模型,可以看到上面打印出来的模型layer参数可以和下面流程图对应上。
2 layers
2.1 activation
最简单的layer实现,将激活函数silu和乘法合在一起:
class SiluAndMul(nn.Module):
"""
A custom activation layer that applies the SiLU (Sigmoid Linear Unit) activation
function followed by element-wise multiplication with the input tensor.
"""
def __init__(self):
super().__init__()
@torch.compile
def forward(self, x: torch.Tensor) -> torch.Tensor:
x, y = x.chunk(2, -1)
return F.silu(x) * y
引擎把两次gate和up合并起来减少一次nn.Linear计算在原来的实现中将y和x分为两路去执行,输入进来的x形状是两个支路拼接的形状,所以在前向中要将x chunk回两支路,并对gate支路的做silu激活。
2.2 RMSnorm
负责层归一化:
RMSNorm(x) = (x / sqrt(mean(x²) + ε)) ⊙ γ
class LayerNorm(torch.nn.Module):
def __init__(self, gamma: torch.Tensor, eps: float = 1e-5):
super().__init__()
# Use nn.Parameter to make gamma learnable and loadable from checkpoints
self.weight = torch.nn.Parameter(gamma.detach().clone())
self.eps = eps
@property
def gamma(self):
"""Backward compatibility: gamma alias for weight"""
return self.weight
@torch.compile
def rms_forward(self, x: torch.Tensor) -> torch.Tensor:
# RMSNorm(x) = (x / sqrt(mean(x²) + ε)) ⊙ γ
variance = x.pow(2).mean(dim=-1, keepdim=True) + self.eps
sqrt_variance = variance.sqrt()
x_norm = (x / sqrt_variance * self.weight)
return x_norm
def residual_rms_forward(self, x: torch.Tensor, residual: torch.Tensor) -> torch.Tensor:
x = x + residual
return self.rms_forward(x), x
def forward(self, x: torch.Tensor, residual: torch.Tensor | None = None) -> torch.Tensor:
if residual is not None:
return self.residual_rms_forward(x, residual)
else:
return self.rms_forward(x)
函数类实现了带残差和不带残差的forward,可以自己选用这个功能。
主要实现在rms_forward中,代码也很简单,根据公式返回rms(x)即可。
2.3 linear
实现了支持张量并行的线性层和权重的加载。
linear.py实现了什么:
LinearBase # 基类:权重形状、weight_loader 挂载机制
├── ReplicatedLinear # 不切分,完整复制(做参考基准)
├── ColumnParallelLinear # 列并行:按输出维度切,forward 无通信
│ ├── MergedColumnParallelLinear # QKV 三个矩阵合并成一个大 Linear,分段加载
│ └── QKVColumnParallelLinear # 按 head 粒度切分,适配 GQA
└── RowParallelLinear # 行并行:按输入维度切,forward 末尾 all_reduce
先看基类:基类不实际实现任何功能,只定义本rank的切片形状,登记TP信息等
class LinearBase(nn.Module):
"""
A base class for linear layers.
"""
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = True,
tp_dim: int | None = None
):
super().__init__()
# set tp_dim, tp_rank, tp_world_size for tensor parallelism
self.tp_dim = tp_dim
self.tp_rank = dist.get_rank()
self.tp_size = dist.get_world_size()
# create weight parameter with custom weight loader
self.weight = nn.Parameter(torch.empty(output_size, input_size))
self.weight.weight_loader = self.weight_loader
# create bias parameter
if bias:
self.bias = nn.Parameter(torch.zeros(output_size))
self.bias.weight_loader = self.weight_loader
else:
self.register_parameter('bias', None)
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor):
raise NotImplementedError("Subclasses should implement this method.")
def forward(self, x: torch.Tensor) -> torch.Tensor:
raise NotImplementedError("Subclasses should implement this method.")
其中
# create weight parameter with custom weight loader
self.weight = nn.Parameter(torch.empty(output_size, input_size))
self.weight.weight_loader = self.weight_loader
第一句是将Tensor注册为模块参数,以便后续某些函数在遍历时能扫描到它、
第二句self.weight.weight_loader对象动态附加属性,注释是这么写的:
"""
these functions are for is that we deploy a maybe randomly initialized model on GPU using some tensor/pipeline parallel method
then we wanna load a saved model checkpoint to it
for name, param in model.named_parameters():
if name in checkpoint:
loaded_weight = checkpoint[name] # full model parameter (4096, 4096)
# check if the parameter has a custom weight_loader
if hasattr(param, 'weight_loader'):
# call custom weight_loader
param.weight_loader(param, loaded_weight)
# weight_loader will automatically:
# 1. extract the shard corresponding to the current GPU
# 2. copy it to param.data
else:
# default: copy directly
param.data.copy_(loaded_weight)
"""
在每个rank中的权重都是经过切分的,也就是原本权重大小的1/tp,在填充权重时直接使用copy函数进行填充每个切片对应不上原始权重的大小,程序就会报错,为了省去一堆判断分支,代码为对象附加一个动态属性,直接把属性挂在参数上,就不用再去写代码判断了:
if hasattr(param, 'weight_loader'):
param.weight_loader(param, loaded_weight) # 参数自己知道怎么切
else:
param.data.copy_(loaded_weight) # 普通层(如 embedding)直接复制
基础的线性层实现:
代码比较简单,没有什么花里胡哨的地方,copy时继承了前面父类参数挂载动态模式的接口。
# the simpliest Linear layer: ReplicatedLinear(LinearBase)
# where we simply copy the weight as the weight_loader
# and run the forward as a normal linear layer
class ReplicatedLinear(LinearBase):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = True
):
super().__init__(input_size, output_size, bias)
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor):
param.data.copy_(loaded_weights)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return nn.functional.linear(x, self.weight, self.bias)
在定义这个子类时直接将完整权重传进类中,不进行任何操作,单卡测试时使用的就是这个类函数。
列并行线性层:
class ColumnParallelLinear(LinearBase):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = True,
):
tp_size = dist.get_world_size()
assert output_size % tp_size == 0, "Output size must be divisible by tensor parallel size."
super().__init__(input_size, output_size//tp_size, bias, tp_dim=0)
# param: parameter after tensor parallelism
# loaded_weights: the original full parameter to be loaded into param
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor):
param_data = param.data
# full_dim on the output column
full_data_output_size = loaded_weights.size(0)
# dim size after sharding
shard_size = full_data_output_size // self.tp_size
assert shard_size == param_data.size(0), "Shard size does not match parameter size."
# starting index
start_index = self.tp_rank * shard_size
slided_weight = loaded_weights.narrow(0, start_index, shard_size)
param_data.copy_(slided_weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return nn.functional.linear(x, self.weight, self.bias)
每个rank都持有x的完整副本,计算权重结果时每个rank持有按列切分后的权重,x与权重矩阵相乘后每个rank的权重结果进行列拼接得到完整结果。
tp_size = dist.get_world_size()
这一行代码查询torchrun启动的进程数量,每个进程分在不同的device上,因此进程数就作为切分数量tp。
# starting index
start_index = self.tp_rank * shard_size
slided_weight = loaded_weights.narrow(0, start_index, shard_size)
param_data.copy_(slided_weight)
这三行代码决定了每个rank分到的是哪个切片,按rank索引对原始矩阵进行切分,最后存入每个rank中。
在计算linear时,wq,wk,wv是三个独立的矩阵,而在计算每个x的qkv时要分别用x乘这三个矩阵,如果用上面的这个类去实现的话需要启动三次kernel,所以下一个类函数将三个矩阵合并成一个大矩阵,一次GEMM算完,省去开销。
MergedColumnParallelLinear:多个矩阵合并成一个大 Linear 的列并行层
class MergedColumnParallelLinear(ColumnParallelLinear):
def __init__(
self,
input_size: int,
output_sizes: list[int], # e.g. merge QKV matrices to compute MM together and then split
bias: bool = True,
):
self.output_sizes = output_sizes
super().__init__(input_size, sum(output_sizes), bias)
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor, loaded_weight_id: int):
"""
checkpoint = {
'q_proj.weight': torch.randn(4096, 4096),
'k_proj.weight': torch.randn(4096, 4096),
'v_proj.weight': torch.randn(4096, 4096),
}
load to
merged_layer = Linear(
input_size=4096,
output_sizes=sum([4096, 4096, 4096]), # Q, K, V
) which is also sharded by tp_size
"""
param_data = param.data
# compute offset
offset = sum(self.output_sizes[:loaded_weight_id]) // self.tp_size
# compute size
shard_size = self.output_sizes[loaded_weight_id] // self.tp_size
# find the correct slice to be loaded in the sharded parameter
param_data = param_data.narrow(0, offset, shard_size)
# shard the original full weight
loaded_weights_start_index = self.tp_rank * shard_size
shard_weights = loaded_weights.narrow(0, loaded_weights_start_index, shard_size)
param_data.copy_(shard_weights)
在计算索引时有区别,按wqkv排列顺序再进行分割。
在GQA中kv的头数少于q,上面这种完全均分的策略就不行了,于是有为GQA策略特化的合并线性层:
class QKVColumnParallelLinear(ColumnParallelLinear):
def __init__(
self,
input_size: int,
head_size: int,
num_heads: int,
num_kv_heads: int | None = None,
bias: bool = False,
):
self.tp_size = dist.get_world_size()
num_kv_heads = num_kv_heads or num_heads
self.head_size = head_size
self.num_heads = num_heads // self.tp_size
self.num_kv_heads = num_kv_heads // self.tp_size
# Calculate per-GPU output size
self.output_size = head_size * (self.num_heads + 2 * self.num_kv_heads)
# Pass TOTAL output size to parent (it will divide by tp_size)
total_output_size = head_size * (num_heads + 2 * num_kv_heads)
super().__init__(input_size, total_output_size, bias=bias)
# load_weight_id: q, k, v
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor, load_weight_id: str):
# batch_size * num_heads * num_token * head_size
param_data = param.data
# loaded_weights: batch_size * num_token * (head_size*num_heads)
assert load_weight_id in ['q', 'k', 'v'], "load_weight_id must be one of 'q', 'k', 'v'"
# compute offset
if load_weight_id == 'q':
offset = 0
shard_size = self.head_size * self.num_heads
elif load_weight_id == 'k':
offset = self.head_size * self.num_heads
shard_size = self.head_size * self.num_kv_heads
elif load_weight_id == 'v':
offset = self.head_size * self.num_heads + self.head_size * self.num_kv_heads
shard_size = self.head_size * self.num_kv_heads
else:
raise ValueError(f"Unknown load_weight_id: {load_weight_id}")
param_data = param_data.narrow(0, offset, shard_size)
# shard the original full weight
loaded_weights_start_index = self.tp_rank * shard_size
shard_weights = loaded_weights.narrow(0, loaded_weights_start_index, shard_size)
param_data.copy_(shard_weights)
Q头和KV头分开来算。每个rank分到的Q头数和KV头数不相等,切分时策略也不同,但逻辑也好理解,按headdim和头数进行索引确定切分点即可。
以上两个类都是列并行线性层的子类,共用其forward。
最后是行并行的实现:
class RowParallelLinear(LinearBase):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = True,
):
tp_size = dist.get_world_size()
assert input_size % tp_size == 0, "Input size must be divisible by tensor parallel size."
super().__init__(input_size // tp_size, output_size, bias, tp_dim=1)
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor):
param_data = param.data
# full_dim on the input row
full_data_input_size = loaded_weights.size(1)
# dim size after sharding
shard_size = full_data_input_size // self.tp_size
assert shard_size == param_data.size(1), "Shard size does not match parameter size."
# starting index
start_index = self.tp_rank * shard_size
slided_weight = loaded_weights.narrow(1, start_index, shard_size)
param_data.copy_(slided_weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
result = nn.functional.linear(x, self.weight, self.bias)
if self.tp_size > 1:
dist.all_reduce(result, op=dist.ReduceOp.SUM)
return result
行并行在attn得分计算完后的混合线性层使用,前面计算时使用列切分wqkv,本rank中得到的注意力得分矩阵拿去混合线性层中做计算,列并行输入得到的输出刚好可以在本rank内被行并行作为输入,可以省去中间的allredece操作,假设tp=2:
┌────────────── 每卡状态:完整 ──────────────┐
x [B,N,4096](完整复制)
│
▼ ① qkv_projection 列并行 权重 [6144,4096] → 本卡 [3072,4096]
qkv [B,N,3072](= [Q_i|K_i|V_i]) ← 状态转为分片
│ split
Q_i[B,N,16,128] K_i,V_i[B,N,4,128]
│ ② attention 本卡算(GQA 广播) ← 零通信,head 间独立
z_i [B,N,2048]
│ ③ o_proj 行并行 权重 [4096,4096] → 本卡 [4096,2048]
部分和 y_i = z_i·W_o,i^T
│ ④ all_reduce ← 第 1 次通信:部分和 → 完整
Y [B,N,4096](每卡都完整) ← 状态回到完整
│ +residual, RMSNorm ← 需要完整输入
│
▼ ⑤ gate_up 列并行 权重 [11008,4096] → 本卡 [5504,4096]
│ ⑥ SiLU·mul 本卡逐元素 ← 零通信
act [B,N,5504]
│ ⑦ down_proj 行并行 权重 [4096,11008] → 本卡 [4096,5504]
部分和
│ ⑧ all_reduce ← 第 2 次通信
Y [B,N,4096]
│ +residual, RMSNorm
▼
下一层(同样的循环)… → lm_head
2.4 词表embedding
在现代大模型中,使用分词器提前提取词表,在完整词表中每个片段(分词器可能不会将一个完整单词作为一个单独token,而是将一个单词拆开成几个片段)都对应有一个token序号,可学习的词embedding层就负责将这些token的语义联系起来,这里实现的前向过程是查找输入token序号所属的embedding行数,反向过程实现embedding层被多次decoder加工后返回回来最终含有语义的x,将x与本卡分到的embbeding行做点积得到相似度得分矩阵。
按行并行的实现:
class VocabParallelEmbedding(nn.Module):
def __init__(self, num_embeddings: int, embedding_dim: int):
super().__init__()
self.tp_size = dist.get_world_size()#获取设备个数决定tp分配数量
self.tp_rank = dist.get_rank()#获取本卡的rank标签
# keep the original num_embeddings
self.num_embeddings = num_embeddings0#记录真实词表大小 后续判断时使用
# pad to make it divisible by tp_size
self.padded_num_embeddings = (num_embeddings + self.tp_size - 1) // self.tp_size * self.tp_size
# this is the num_embeddings per partition in this current GPU
self.num_embeddings_per_partition = self.padded_num_embeddings // self.tp_size
self.embedding_dim = embedding_dim
self.weight = nn.Parameter(torch.empty(self.num_embeddings_per_partition, embedding_dim))
self.weight.weight_loader = self.weight_loader
def weight_loader(self, param: nn.Parameter, loaded_weights: torch.Tensor):
param_data = param.data
offset = self.tp_rank * self.num_embeddings_per_partition
shard_size = self.num_embeddings_per_partition
# calculate how much of the original vocab falls in this partition
actual_start = min(offset, self.num_embeddings)
actual_end = min(offset + shard_size, self.num_embeddings)
actual_size = max(0, actual_end - actual_start)
if actual_size > 0:
# load the actual weights
sharded_weights = loaded_weights.narrow(0, actual_start, actual_size)
param_data[:actual_size].copy_(sharded_weights)
# pad the rest with zeros if needed
if actual_size < shard_size:
param_data[actual_size:].zero_()
def forward(self, x: torch.Tensor) -> torch.Tensor:
# mask for tokens in this partition's range and within original vocab size
mask = (x >= self.tp_rank * self.num_embeddings_per_partition) & \
(x < (self.tp_rank + 1) * self.num_embeddings_per_partition) & \
(x < self.num_embeddings)
x = mask * (x - self.tp_rank * self.num_embeddings_per_partition)
output = F.embedding(x, self.weight)
if dist.get_world_size() > 1:
# need to mask again, otherwise the embedding for the out-of-range ids will be the embedding of id 0
output = mask.unsqueeze(1) * output
dist.all_reduce(output, op=dist.ReduceOp.SUM)
return output
可以看出,加载方式和分卡方式与前面介绍过的层基本完全一样,在前向中为了准确实现功能,需要mask掉不属于本卡分到的词和剔除掉前面padding的空格,确保查到的是所要的embedding行。
反向:
# weight tying with embedding layer
class ParallelLMHead(VocabParallelEmbedding):
def __init__(self, num_embeddings: int, embedding_dim: int):
super().__init__(num_embeddings, embedding_dim)
# x: [batch_size, seq_len, hidden_size]
# weight: [vocab_size_per_partition, hidden_size]
def forward(self, x: torch.Tensor) -> torch.Tensor:
context = get_context()
if context.is_prefill:
# cu_seqlens_q = [0, 5, 8, 12]
# last_indices = [5, 8, 12] - 1 = [4, 7, 11]
last_token = context.cu_seqlens_q[1:] - 1 # exclude the first element which is 0
x = x[last_token].contiguous()
# logits: [batch_size, seq_len, vocab_size_per_partition]
# F.linear automatically transpose the weight
logits = torch.nn.functional.linear(x, self.weight)
if self.tp_size > 1:
# prepare for all_gather only for GPU 0 which is the main GPU
all_logits = [torch.empty(logits.size(), device=logits.device) for _ in range(self.tp_size)] if self.tp_rank == 0 else None
# dist.gather collects the logits from all GPUs to GPU 0
dist.gather(logits, gather_list=all_logits, dst=0)
# concatenate
if self.tp_rank == 0:
# [batch_size, seq_len, padded_vocab_size]
logits = torch.cat(all_logits, dim=-1)
# trim to original vocab size
logits = logits[..., :self.num_embeddings]
return logits
在prefill阶段,需要取出提示词序列里的最后一个token计算得分去做后续预测,decode阶段输入的就只有一个token因此不需要这样操作。
不同卡分到的词表最后直接拼接不需要其他任何操作是因为此处词表的行是按设备序号递增的,cat回来后顺序与原来是一样的。
2.5 ROPE旋转位置编码
我们知道,attn中只是对矩阵进行相乘计算得分,这个得分结果不关注token的排列顺序,比如“我爱你”和“你爱我”两句话得到的结果会是相同,所以我们需要对每个token进行位置编码,让模型知道token的排序。普通的位置编码只能让token获取绝对位置,而不同token之间是有语义联系的,embbeding层中语义相近的词旋转编码后距离就近,使用旋转位置编码可以获得其相对位置,Rope还有许多好处,这里不一一赘述。
前向过程对查询到的token形成的q和k张量做旋转编码,按行数和该行每个元素的列数进行旋转:
cos_sin_cache(示意,base=10000)
对0 对1 对63
θ=1.0 θ=0.87 θ=0.0001
第0行: cos(0×1.0) cos(0×0.87) ... cos(0×0.0001)
第1行: cos(1×1.0) cos(1×0.87) ... cos(1×0.0001)
第57行: cos(57×1.0) cos(57×0.87) ... cos(57×0.0001) ← decode 第57步查这行
完整类函数:
class RotaryEmbedding(nn.Module):
def __init__(
self,
base:int,
rotary_embedding: int,
max_position: int = 2048,
is_llama3: bool = False,
# the following params are only used in llama3.2
llama3_rope_factor: float = 32.0,
llama3_rope_high_freq_factor: float = 4.0,
llama3_rope_low_freq_factor: float = 1.0,
llama3_rope_original_max_position_embeddings: int = 8192,
):
super().__init__()
self.base = base
# how many dimensions to apply rotary embedding
self.rotary_embedding = rotary_embedding
# max position that the long context can reach
self.max_position = max_position
self.inv_freq = 1/(base ** (torch.arange(0, self.rotary_embedding, 2)/self.rotary_embedding))
if is_llama3:
# specifically for llama3.2
import math
inv_freq = self.inv_freq
# no smooth if low_freq_factor == high_freq_factor
wave_len = 2 * math.pi / inv_freq
if llama3_rope_low_freq_factor == llama3_rope_high_freq_factor:
inv_freq = torch.where(
wave_len < llama3_rope_original_max_position_embeddings / llama3_rope_high_freq_factor,
inv_freq,
inv_freq / llama3_rope_factor,
)
else:
delta = llama3_rope_high_freq_factor - llama3_rope_low_freq_factor
smooth = (llama3_rope_original_max_position_embeddings / wave_len - llama3_rope_low_freq_factor) / delta
smooth = torch.clamp(smooth, 0, 1)
factor = (1 - smooth) / llama3_rope_factor + smooth
inv_freq = factor * inv_freq
self.inv_freq = inv_freq
positions = torch.arange(self.max_position).float()
# (max_position, rotary_embedding/2)
freqs = torch.einsum("i,j -> ij", positions, self.inv_freq)
cos = torch.cos(freqs)
sin = torch.sin(freqs)
# (max_position, rotary_embedding)
cos_sin_cache = torch.cat([cos, sin], dim=-1)
self.register_buffer("cos_sin_cache", cos_sin_cache)
@torch.compile
# tell the position index of the token
# apply rotary embedding to query and key
def forward(self, positions, query, key):
cos_sin = self.cos_sin_cache[positions] # (seq_len, rotary_embedding)
cos, sin = cos_sin.chunk(2, dim=-1)
return (
apply_rotary_pos_emb(query, cos, sin),
apply_rotary_pos_emb(key, cos, sin)
)
2.6 Attention
搬运算子,将前向过程计算出来的的k和v数组从slot mapping搬运至分页kv cache的固定槽位,也就是所谓block槽:
@triton.jit
def store_kvcache_kernel(
key_ptr, # pointer to what we want to store
value_ptr,
k_cache_ptr, # pointer to where we want to store
v_cache_ptr,
slot_mapping_ptr,
num_kv_heads: tl.constexpr,
head_dim: tl.constexpr,
block_size: tl.constexpr
):
"""
Store keys and values into paged KV cache.
Each token is mapped to a slot via slot_mapping.
Grid layout: (num_tokens, num_kv_heads)
Cache layout: (num_blocks, block_size, num_kv_heads, head_dim)
"""
# thread ID, in dimension 0
token_idx = tl.program_id(0) # each GPU thread processes one token
# slot ID, where in cache to store this token
slot_idx = tl.load(slot_mapping_ptr + token_idx)
if slot_idx == -1:
return
# Calculate which block and position within block
block_idx = slot_idx // block_size
block_offset = slot_idx % block_size
# Process each head
# program_id(0) = which token
# program_id(1) = which head
head_idx = tl.program_id(1)
# it creates a vector [0, 1, ..., head_dim-1]
# Load key and value for this token and head
head_offsets = tl.arange(0, head_dim)
# Input: (num_tokens, num_kv_heads, head_dim)
# example: input_offset = 5 * (8 * 128) + 3 * 128 + [0, 1, 2, ..., 127]
# = 5120 + 384 + [0, 1, 2, ..., 127]
# = [5504, 5505, 5506, ..., 5631]
input_offset = (token_idx * num_kv_heads * head_dim + # skip previous tokens
head_idx * head_dim + # skip previous heads
head_offsets)
# Cache: (num_blocks, block_size, num_kv_heads, head_dim)
cache_offset = (block_idx * block_size * num_kv_heads * head_dim + # skip previous blocks
block_offset * num_kv_heads * head_dim + # skip previous positions in block
head_idx * head_dim + # skip previous heads
head_offsets)
# load key and value value floats from the pointers's memory
key = tl.load(key_ptr + input_offset)
value = tl.load(value_ptr + input_offset)
# store into cache
tl.store(k_cache_ptr + cache_offset, key)
tl.store(v_cache_ptr + cache_offset, value)
逻辑/物理解耦 :读地址用 token_idx (连续),写地址用 slot_idx(分页)——序列的 KV 在逻辑上连续、物理上散落在任意块里,这一步搬运是解耦的落点,之后 decode 的 attention kernel 靠 block_tables 沿物理块读回
对于变长序列varlen的flash处理:
参考