跳转至

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 为分片数)。代价是反向传播时需要重新计算前向,增加了计算量。