
【Bug已解决】TP support for Finetuning using LoRA and other PEFT techniques 解决方案一、现象长什么样Tensor ParallelismTP张量并行把单个大线性层沿维度切到多张卡上是大模型训练/推理的标配。但当你想用 LoRA或 PEFT 其它方法在 TP 模型上做微调时会遇到PEFT 的get_peft_model默认假设nn.Linear是完整、未切分的一旦底层权重被 TP 拆成DTensor/ColumnParallelLinear/RowParallelLinearLoRA 注入后A、B的维度与设备对不上报shape mismatch或device mismatch更隐蔽TP 把权重沿输出维column-parallel或输入维row-parallel切分而 LoRA 的增量Δ B·A·x里B ∈ ℝ^{out×r}、A ∈ ℝ^{r×in}必须沿匹配的维度切否则各 rank 算出的是错的切片前向结果在all-reduce/all-gather后数值不对训练时只有部分 rank 的 LoRA 参数收到有效梯度其余 rank 的lora_A/lora_B梯度为 0合并/保存时出现missing keys用device_map或tensor_parallel做推理时LoRA 权重没跟着切显存没省下来等于 TP 白做保存后再load_adapter发现 adapter 权重是“按未切分维度”存的加载回 TP 模型时维度对不上。根因PEFT 的 LoRA 注入假设权重是整体nn.Linear没有感知 TP 的切分维度LoRA 的A/B必须跟随基座的 TP 切分方式column-parallel 切 B 的输出维、row-parallel 切 A 的输入维才能正确组合。二、背景先明确 TP 怎么切一个线性层Y X · WᵀW ∈ ℝ^{out×in}Column Parallel按输出维切W沿out维切成[out/t, in]每个 rank 算Y_i X · W_iᵀ最后all-gather拼接成完整Y。对应nn.Linear的weight形状从[out, in]变成[out/t, in]。Row Parallel按输入维切W沿in维切成[out, in/t]每个 rank 用本地X_i已按输入维切好算Y_i X_i · W_iᵀ最后all-reduce求和。对应weight形状[out, in/t]。LoRA 增量Δ B·A·xB ∈ ℝ^{out×r}A ∈ ℝ^{r×in}。若基座是column-parallel沿 out 切则B必须沿out 维同样切成[out/t, r]A保持完整[r, in]每个 rank 都有完整 A因为x是完整输入切的是输出侧。若基座是row-parallel沿 in 切则A必须沿in 维切成[r, in/t]B保持完整[out, r]因为切的是输入侧x已按 in 切B 在输出侧需完整以便 all-reduce 求和。PEFT 默认不区分这两种统一把A、B当完整[r, in]/[out, r]建于是在 TP 模型上要么形状错、要么数值错。下面用单进程模拟两种切分演示正确的 LoRA-TP 维度分配。三、根因根因一句话LoRA 的A/B没有跟随基座的 TP 切分维度——column-parallel 应切B的输出维、row-parallel 应切A的输入维PEFT 默认把它们当完整矩阵导致形状/设备/数值三重错配。展开维度错配TP 把weight切成[out/t, in]或[out, in/t]而 LoRAB/A仍是完整尺寸矩阵乘维度对不上。设备错配切分后的weight在不同 rank 的不同 device 上LoRA 参数若建在默认 device前向时B A x报 device mismatch。数值错配即使形状巧合对上若A/B没沿正确维度切all-gather/all-reduce后合并出的Δ与真实B·A·x不等。修复方向实现 TP-aware 的 LoRA 层按基座切分方式切A/B并放到对应 rank 的 device 上。四、最小可运行复现下面用单进程模拟“column-parallel”与“row-parallel”下 LoRA 增量的正确切分不真正起多进程但把切分数学演示清楚。import torch import torch.nn as nn def reference_lora(W, A, B, x): 未切分的参考实现Y xWᵀ B A x。 return x W.T (B (A x.T)).T def column_parallel_lora(W_shards, A, B_shards, x, t2): Column-parallel: W、B 沿 out 维切成 t 份A 完整。 每个 rank 算本地 Y_i x W_iᵀ B_i A x最后 all-gather 拼接。 parts [] for i in range(t): Wi W_shards[i] # [out/t, in] Bi B_shards[i] # [out/t, r] yi x Wi.T (Bi (A x.T)).T parts.append(yi) return torch.cat(parts, dim-1) # 模拟 all-gather def row_parallel_lora(W_shards, A_shards, B, x_shards, t2): Row-parallel: W、A 沿 in 维切成 t 份B 完整。 每个 rank 用本地 x_i 算 Y_i x_i W_iᵀ B A_i x_i最后 all-reduce 求和。 acc None for i in range(t): Wi W_shards[i] # [out, in/t] Ai A_shards[i] # [r, in/t] xi x_shards[i] # [B, in/t] yi xi Wi.T (B (Ai xi.T)).T acc yi if acc is None else acc yi # 模拟 all-reduce return acc torch.manual_seed(0) out, in_f, r, t 8, 6, 3, 2 W torch.randn(out, in_f) A torch.randn(r, in_f) * 0.01 B torch.randn(out, r) x torch.randn(4, in_f) ref reference_lora(W, A, B, x) # column-parallel 切分 W_shards [W[i*out//t:(i1)*out//t] for i in range(t)] B_shards [B[i*out//t:(i1)*out//t] for i in range(t)] col column_parallel_lora(W_shards, A, B_shards, x, t) print(column-parallel 与参考一致:, torch.allclose(col, ref, atol1e-5)) # row-parallel 切分 W_shards_r [W[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] A_shards_r [A[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] x_shards [x[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] row row_parallel_lora(W_shards_r, A_shards_r, B, x_shards, t) print(row-parallel 与参考一致:, torch.allclose(row, ref, atol1e-5))运行后两者都True说明只要按 TP 切分方式正确切A/BLoRA 增量在切分后仍能还原成完整的B·A·x。五、解决方案第一层最小直接修复修复 1column-parallel 层 → 切B的输出维A保持完整class ColumnParallelLoraLinear(nn.Module): def __init__(self, base_weight, r, in_f, out_f, world_t, rank_idx): super().__init__() self.A nn.Parameter(torch.randn(r, in_f) * 0.01) # 完整 # B 沿输出维切到本 rank self.B nn.Parameter(base_weight[rank_idx*out_f//world_t: (rank_idx1)*out_f//world_t].new_zeros( out_f // world_t, r)) self.idx rank_idx self.world_t world_t self.out_f out_f def forward(self, x, W_local): # W_local: [out/t, in] 本 rank 的基座切片 base x W_local.T delta (self.B (self.A x.T)).T # B 已是 [out/t, r] return base delta # 外部 all-gather 拼接修复 2row-parallel 层 → 切A的输入维B保持完整class RowParallelLoraLinear(nn.Module): def __init__(self, base_weight, r, in_f, out_f, world_t, rank_idx): super().__init__() self.B nn.Parameter(torch.zeros(out_f, r)) # 完整 self.A nn.Parameter(base_weight[:, rank_idx*in_f//world_t: (rank_idx1)*in_f//world_t].new_zeros( r, in_f // world_t)) # 沿 in 切 self.idx rank_idx self.world_t world_t self.in_f in_f def forward(self, x_local, W_local): # x_local: [B, in/t] 本 rank 的输入切片 base x_local W_local.T delta (self.B (self.A x_local.T)).T return base delta # 外部 all-reduce 求和修复 3保存/加载时按“未切分”维度规整训练时各 rank 只持有切片保存前要把A/B用all-gather拼回完整[r, in]/[out, r]再save_pretrained否则加载回单卡/别的 TP 度时会维度错配import torch.distributed as dist def gather_full_B(B_local, world_t, out_f, r): # 假设 B 沿 out 维切all-gather 后再拼接 shards [torch.zeros(out_f // world_t, r) for _ in range(world_t)] dist.all_gather(shards, B_local) return torch.cat(shards, dim0)六、解决方案第二层结构性改进改进 1用torch.distributed.tensorDTensor声明切分而非手动切片DTensor 让你用DeviceMeshPlacement声明B是Colwise沿输出维切、A是Replicate完整框架自动处理通信from torch.distributed.tensor import DeviceMesh, DTensor, Shard, Replicate mesh DeviceMesh(cuda, list(range(world_t))) # B 沿输出维第 0 维切 B_dt DTensor.from_local(B_local, mesh, [Shard(0)]) # A 复制每 rank 完整 A_dt DTensor.from_local(A_local, mesh, [Replicate()]) delta B_dt (A_dt x.T) # 框架自动插入必要的 all-gather / all-reduce改进 2封装一个“自动识别 TP 类型”的 LoRA 工厂def make_tp_lora(base_linear, r, tp_mode): out_f, in_f base_linear.weight.shape if tp_mode column: return ColumnParallelLoraLinear(base_linear.weight, r, in_f, out_f, world_t, rank_idx) elif tp_mode row: return RowParallelLoraLinear(base_linear.weight, r, in_f, out_f, world_t, rank_idx) else: raise ValueError(f未知 tp_mode: {tp_mode})改进 3对其它 PEFT 方法同样处理不只 LoRAIA³的缩放向量、PrefixTuning的 prompt 张量、DoRA 的「方向向量」也要按对应 TP 维度切分通常沿输出维。原则是任何加性的、与权重同形状的适配项都跟随基座的切分方式与输入同形状的跟随输入切分。七、解决方案第三层断言 / CI 守护import torch import torch.nn as nn import pytest def reference_lora(W, A, B, x): return x W.T (B (A x.T)).T def column_parallel_lora(W_shards, A, B_shards, x, t2): parts [] for i in range(t): parts.append(x W_shards[i].T (B_shards[i] (A x.T)).T) return torch.cat(parts, dim-1) def row_parallel_lora(W_shards, A_shards, B, x_shards, t2): acc None for i in range(t): yi x_shards[i] W_shards[i].T (B (A_shards[i] x_shards[i].T)).T acc yi if acc is None else acc yi return acc def test_column_parallel_matches_reference(): torch.manual_seed(0) out, in_f, r, t 8, 6, 3, 2 W torch.randn(out, in_f); A torch.randn(r, in_f)*0.01; B torch.randn(out, r) x torch.randn(4, in_f) ref reference_lora(W, A, B, x) Ws [W[i*out//t:(i1)*out//t] for i in range(t)] Bs [B[i*out//t:(i1)*out//t] for i in range(t)] assert torch.allclose(column_parallel_lora(Ws, A, Bs, x, t), ref, atol1e-5) def test_row_parallel_matches_reference(): torch.manual_seed(1) out, in_f, r, t 8, 6, 3, 2 W torch.randn(out, in_f); A torch.randn(r, in_f)*0.01; B torch.randn(out, r) x torch.randn(4, in_f) ref reference_lora(W, A, B, x) Ws [W[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] As [A[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] xs [x[:, i*in_f//t:(i1)*in_f//t] for i in range(t)] assert torch.allclose(row_parallel_lora(Ws, As, B, xs, t), ref, atol1e-5) def test_lora_shard_dimensions(): out, in_f, r, t 8, 6, 3, 2 # column: B 切成 [out/t, r]A 完整 assert (out//t, r) (4, r) # row: A 切成 [r, in/t]B 完整 assert (r, in_f//t) (r, 3)这三个测试守护“column/row-parallel 的 LoRA 切片数学上等价于完整实现”“切分维度正确”。八、排查清单LoRA 在 TP 模型上出问题时按序查确认基座切分方式是 column-parallel沿 out还是 row-parallel沿 in看weight形状是[out/t, in]还是[out, in/t]。B跟随 out 切、A跟随 in 切column-parallel 切B输出维、row-parallel 切A输入维别搞反。设备对齐LoRA 参数必须和本地weight切片在同一 device。通信原语匹配column-parallel 输出all-gather拼接、row-parallel 输出all-reduce求和。保存前拼回完整用all-gather把A/B拼回[r,in]/[out,r]再存否则load_adapter维度错配。优先用 DTensor用Shard/Replicate声明切分让框架管通信比手写切片稳。其它 PEFT 方法同原则加性项跟随权重切分输入相关项跟随输入切分。验证数值单进程用本节的 reference 对比确认切片后all-gather/all-reduce能还原B·A·x。九、小结TP support for Finetuning using LoRA的根因是PEFT 的 LoRA 注入默认假设权重是完整nn.Linear没有感知 TP 的切分维度而 LoRA 的B输出侧和A输入侧必须分别跟随基座的 column/row 切分方式否则形状、设备、数值三重错配。最小修复是column-parallel 切B的输出维A 完整、row-parallel 切A的输入维B 完整保存前all-gather拼回完整维度结构性改进是用 DTensor 的Shard/Replicate声明切分让框架管通信、封装自动识别 TP 类型的 LoRA 工厂、把同原则推广到 IA³/PrefixTuning/DoRA最后用测试守护“column/row-parallel 切片数学等价于完整实现”。这样 LoRA 才能真正在张量并行模型上正确且高效地微调。