tiled_mlp.py — 分片 MLP 显存优化¶
文件路径¶
verl/models/transformers/tiled_mlp.py
文件概述¶
提供 TiledMLP 实现,通过将 MLP 的前向/反向传播按序列长度分片(chunking)来降低峰值显存。这是一种显存-计算权衡的优化策略,特别适合使用 FSDP2 训练大模型时。
关键代码讲解¶
核心思想¶
标准 MLP 前向计算 down_proj(act_fn(gate_proj(x)) * up_proj(x)) 会产生巨大的中间激活张量。TiledMLP 将输入按序列维度切分为多个小块,分批计算,从而降低同时驻留内存中的中间张量大小。
GradientAccumulator -- 梯度累积器¶
class GradientAccumulator:
def __init__(self, params, total_shards, dtype=None):
self.accumulated_grads = {}
for param in self.params:
self.accumulated_grads[param] = torch.zeros_like(param, dtype=self.grad_accumulation_dtype)
def install_hooks(self, is_last_shard):
def create_hook(param):
def hook(grad):
self.accumulated_grads[param] += grad.to(self.grad_accumulation_dtype)
if is_last_shard:
param.grad = None # 防止双重累积
return self.accumulated_grads[param].to(param.dtype)
return None # 非最后一个 shard,不返回梯度
return hook
for param in self.params:
if param.requires_grad:
hook = param.register_hook(create_hook(param))
关键设计:
- 每个参数的梯度在多个 shard 之间累积
- 只有处理最后一个 shard 时才返回最终梯度
- 使用 threading.Lock 保证线程安全
TiledMLP -- 自定义 Autograd 函数¶
前向传播:
class TiledMLP(torch.autograd.Function):
@staticmethod
def forward(ctx, fn, module, x, shards, compute_params):
ctx.save_for_backward(x)
# 按序列维度切分
x_shards = list(torch.chunk(x, chunks=shards, dim=-2))
with torch.no_grad():
output_shards = [fn(module, x_shard) for x_shard in x_shards]
output_unsharded = torch.cat(output_shards, dim=-2)
return output_unsharded
前向传播在 torch.no_grad() 下逐 shard 计算,只保存原始输入用于反向。
反向传播:
@staticmethod
def backward(ctx, *grads):
x = ctx.saved_tensors[0]
x_shards = list(torch.chunk(x, chunks=shards, dim=0))
grad_accumulator = GradientAccumulator(compute_params, shards)
for i, x_shard in enumerate(x_shards):
is_last_shard = i + 1 == shards
grad_accumulator.install_hooks(is_last_shard)
with torch.enable_grad():
output = fn(module, x_shard)
torch.autograd.backward(output, incoming_grad_shard)
grad_accumulator.cleanup()
return (None, None, x_grad, None, None)
反向传播重新计算每个 shard 的前向(recomputation),通过梯度钩子累积参数梯度。
MLP 前向函数¶
def _mlp_forward_fn(module, x):
"""LlamaMLP / Qwen2MLP / Qwen3MLP 通用前向"""
return module.down_proj(module.act_fn(module.gate_proj(x)) * module.up_proj(x))
这是 SwiGLU 风格的 MLP 计算:down_proj(SiLU(gate_proj(x)) * up_proj(x))。
Monkey Patch 入口¶
_MODEL_TYPE_TO_MLP_CLASS = {
"llama": ("transformers.models.llama.modeling_llama", "LlamaMLP"),
"qwen2": ("transformers.models.qwen2.modeling_qwen2", "Qwen2MLP"),
"qwen2_5": ("transformers.models.qwen2.modeling_qwen2", "Qwen2MLP"),
"qwen3": ("transformers.models.qwen3.modeling_qwen3", "Qwen3MLP"),
}
def apply_tiled_mlp_monkey_patch(num_shards=4, model_type=None):
"""必须在模型实例化之前调用"""
for mtype in types_to_patch:
module_path, class_name = _MODEL_TYPE_TO_MLP_CLASS[mtype]
module = importlib.import_module(module_path)
mlp_class = getattr(module, class_name)
_patch_mlp_class(mlp_class, _mlp_forward_fn, num_shards)
重要:这个 patch 必须在模型实例化 之前 调用,因为它修改的是类的方法,而非实例的方法。
核心类/函数列表¶
| 名称 | 作用 |
|---|---|
GradientAccumulator |
管理分片反向中的梯度累积 |
TiledMLP |
自定义 Autograd 函数,分片计算 MLP |
_mlp_forward_fn() |
SwiGLU MLP 的通用前向函数 |
apply_tiled_mlp_monkey_patch() |
对外的 monkey patch 入口 |
_patch_mlp_class() |
替换 MLP 类的 forward 方法 |
显存节省原理¶
标准 MLP:一次性计算所有 token 的中间激活
内存峰值 = batch_size * seq_len * intermediate_size
TiledMLP(4 shards):分4次计算
内存峰值 ≈ batch_size * (seq_len/4) * intermediate_size + 梯度累积
节省约 75% 的中间激活显存
与其他模块的关系¶
- 通过
transformers/__init__.py导出apply_tiled_mlp_monkey_patch - 被
monkey_patch.py的apply_monkey_patch()间接调用 - 支持 LLaMA、Qwen2、Qwen3 系列模型
小结¶
TiledMLP 是一种典型的"用计算换内存"的优化:通过分片计算 + 重计算(recomputation),将 MLP 的峰值显存降低到原来的 1/N(N 为分片数)。代价是反向传播时需要重新计算前向,增加了计算量。