【Bug已解决】trainable_token_indices of LoraConfig not working when using more than 1 NVIDIA GPU 解决方案
一、现象长什么样
PEFT 的LoraConfig支持一个相对少有人用的参数trainable_token_indices:它让 LoRA只作用在序列里的特定 token 位置(例如只对 padding 之外的某些 token、或只对 prompt 位置生效),其余位置走基座权重。这是 token-level / 位置级微调的关键能力。
但当你把训练从单卡切到**多卡(DDP / FSDP / DeepSpeed)**时,会出现:
- 单卡下
trainable_token_indices=[0,1,2]工作正常,多卡下报IndexError: index ... is out of bounds; - 或者更隐蔽:多卡不报错,但 LoRA 实际作用到了错误的 token 位置(因为每个 rank 只持有 batch 的一个分片,而
trainable_token_indices是按全局序列索引算的,没映射到本地分片); KeyError: '...trainable_tokens...'出现在保存 checkpoint 时(多卡下各 rank 的索引集合不一致,聚合 state_dict 时键对不上);- 多卡 loss 和单卡对不上,且随 GPU 数量变化(因为索引错位导致每个 rank 应用的 token 不同);
grad_accum/no_sync场景下,只有 rank0 的索引有效,其它 rank 把 LoRA 应用到了全部 token。
核心症状:trainable_token_indices是“全局序列索引”语义,但多卡训练时每个进程只看到本地分片,索引没有被重映射到本地范围,于是越界或错位。
二、背景
先理解trainable_token_indices在 forward 里怎么用。PEFT 在应用 LoRA 增量时,会用一张布尔/索引掩码决定哪些 token 位置叠加B·A·x:
# 伪代码:token-level LoRA active = torch.zeros(seq_len, dtype=bool) active[trainable_token_indices] = True delta = (B @ (A @ x.T)).T # [B, seq, out] delta = delta * active.unsqueeze(-1) # 只保留指定 token 位置 out = base + delta这里trainable_token_indices假设x的序列维长度是完整的。
多卡训练时(以 DDP 的DistributedSampler为例):每个 rank 拿到的是整个 batch 的一个子集,但序列长度seq_len通常保持不变(按样本切,不按 token 切)。这种情况下索引其实不会越界——真正出问题的是按 token/序列维度切分的并行(如序列并行 SP、或某些 FSDP 对长序列的分片),以及trainable_token_indices被理解成“跨样本的全局索引”时。
更常见的真实踩坑场景:
- 索引是按“样本维”给的,但代码按“序列维”用:用户想训练第 0、1、2 个样本,却把索引喂进了序列维掩码,单卡序列短碰巧不越界,多卡序列分片后越界。
- FSDP/DeepSpeed 对长序列做分片:序列维被切到多卡,
trainable_token_indices仍用全局索引,本地分片只有[start, end)区间的 token,全局索引落在别处就IndexError。 - 保存时各 rank 索引键不一致:token adapter 把索引写进参数名(如
...trainable_tokens_delta.default),各 rank 本地索引不同 →state_dict键集合不同 → 聚合/保存报KeyError。
下面用最小可运行代码演示“全局索引在本地分片上越界”以及如何重映射。
三、根因
根因一句话:trainable_token_indices是全局序列/样本索引语义,多卡训练时每个进程只持有本地分片(序列维或样本维被切),但掩码仍用全局索引构造,导致越界或错位。
展开:
- 索引未重映射:本地分片 token 范围是
[local_start, local_end),全局索引i需减local_start才是本地合法索引;不重映射就IndexError。 - 索引语义混淆:把“样本索引”误当“序列索引”喂给 token 掩码。
- 多卡键不一致:token adapter 把索引编入参数名,各 rank 本地索引集合不同,聚合 state_dict 时键对不上 →
KeyError。
修复方向:在 forward 里把全局trainable_token_indices重映射为本地分片内的相对索引(并对落在本地范围外的索引直接丢弃),保存时统一用全局索引命名。
四、最小可运行复现
下面演示“全局 token 索引在本地序列分片上越界”与“重映射后正确”。
import torch def apply_token_lora_global(x, B, A, idx_global, seq_len): """错误示范:用全局索引直接构造序列掩码。""" active = torch.zeros(seq_len, dtype=torch.bool) active[idx_global] = True # 若本地 seq_len 更小 -> IndexError delta = (B @ (A @ x.T)).T return delta * active.unsqueeze(-1).to(delta.dtype) def apply_token_lora_local(x, B, A, idx_global, local_start, local_end): """正确:把全局索引重映射到本地分片 [local_start, local_end)。""" local_len = local_end - local_start active = torch.zeros(local_len, dtype=torch.bool) for i in idx_global: if local_start <= i < local_end: # 只保留落在本地范围内的 active[i - local_start] = True # 减偏移变成本地相对索引 delta = (B @ (A @ x.T)).T return delta * active.unsqueeze(-1).to(delta.dtype) torch.manual_seed(0) r, out, in_f = 3, 8, 6 B = torch.randn(out, r) A = torch.randn(r, in_f) * 0.01 x = torch.randn(1, 10, in_f) # 序列长 10,但本 rank 只持有 [4, 8) idx_global = [0, 1, 5, 9] # 全局想作用的 token # 错误:把完整 seq_len=10 当成局部,实际本地只有 4 个 token try: apply_token_lora_global(x, B, A, idx_global, seq_len=4) print("global 版未报错(异常)") except IndexError as e: print("global 版越界(符合预期):", e) # 正确:本地 [4, 8),全局 5 映射到本地 1,其余丢弃 out_local = apply_token_lora_local(x, B, A, idx_global, local_start=4, local_end=8) print("local 版输出形状:", tuple(out_local.shape)) print("只有本地 token1(全局5) 有非零增量:", out_local.abs().sum(dim=-1).squeeze().tolist())运行后:global 版抛IndexError,local 版正确把全局索引5映射为本地索引1,且只对该位置叠加增量,范围外的索引被安全丢弃。
五、解决方案(第一层:最小直接修复)
修复 1:在 forward 里把全局索引重映射为本地相对索引
如apply_token_lora_local:对每个全局索引i,仅当local_start <= i < local_end时,以i - local_start写入本地掩码;范围外的直接忽略。
修复 2:明确索引语义——是序列索引还是样本索引
# 如果是“只对前 N 个样本训练”,应在样本维处理,而非序列掩码 def sample_mask(batch_size, trainable_sample_indices): active = torch.zeros(batch_size, dtype=torch.bool) active[trainable_sample_indices] = True return active别把样本索引喂进 token 维掩码。
修复 3:保存时用全局索引命名,加载再重映射
# 保存:参数名用全局索引,保证各 rank 键一致 state = {f"lora_A.trainable_tokens.{i}.default.weight": p for i, p in zip(global_indices, params)} torch.save(state, "adapter.bin") # 加载:本地 rank 只取自己范围内的全局索引,重映射为本地这样无论多少卡,参数名都基于全局索引,不会因本地分片差异而KeyError。
六、解决方案(第二层:结构性改进)
改进 1:封装一个“多卡安全的 token-LoRA 掩码”工具
def build_local_token_mask(idx_global, local_start, local_end, device): local_len = local_end - local_start mask = torch.zeros(local_len, dtype=torch.bool, device=device) for i in idx_global: if local_start <= i < local_end: mask[i - local_start] = True return mask # 用法 mask = build_local_token_mask(trainable_token_indices, local_start, local_end, x.device) delta = (B @ (A @ x.T)).T * mask.unsqueeze(0).unsqueeze(-1)改进 2:用 Dist 获取本地序列偏移
import torch.distributed as dist def local_seq_range(world_size, rank, seq_len): # 均匀切分序列维 base = seq_len // world_size rem = seq_len % world_size start = rank * base + min(rank, rem) end = start + base + (1 if rank < rem else 0) return start, end rank = dist.get_rank() if dist.is_initialized() else 0 ws = dist.get_world_size() if dist.is_initialized() else 1 local_start, local_end = local_seq_range(ws, rank, seq_len) mask = build_local_token_mask(trainable_token_indices, local_start, local_end, x.device)改进 3:把 token 索引校验固化进配置加载
def validate_token_indices(idx_global, seq_len): bad = [i for i in idx_global if i < 0 or i >= seq_len] if bad: raise ValueError(f"trainable_token_indices 含越界值 {bad},序列长 {seq_len}") return True validate_token_indices(trainable_token_indices, seq_len) # 全局校验在 rank0 做七、解决方案(第三层:断言 / CI 守护)
import torch import pytest def build_local_token_mask(idx_global, local_start, local_end, device="cpu"): local_len = local_end - local_start mask = torch.zeros(local_len, dtype=torch.bool, device=device) for i in idx_global: if local_start <= i < local_end: mask[i - local_start] = True return mask def test_global_index_out_of_local_bounds_is_dropped(): # 全局 [0,1,5,9],本地 [4,8) -> 仅 5 命中,映射为本地 1 mask = build_local_token_mask([0, 1, 5, 9], 4, 8) assert mask.shape == (4,) assert mask.tolist() == [False, True, False, False] def test_no_indexerror_on_local_slice(): # 本地长度只有 4,但全局索引最大 9,不应抛错 mask = build_local_token_mask([0, 1, 5, 9], 4, 8) assert mask.sum().item() == 1 def test_sample_vs_token_semantics(): # 样本索引不应进 token 掩码 batch = 8 sample_active = torch.zeros(batch, dtype=torch.bool) sample_active[[0, 2, 3]] = True assert sample_active.sum().item() == 3 # token 掩码是另一个维度,互不干扰 token_mask = build_local_token_mask([1], 0, 10) assert token_mask.sum().item() == 1 def test_mask_applied_only_to_active_tokens(): r, out, in_f = 3, 8, 6 torch.manual_seed(0) B = torch.randn(out, r); A = torch.randn(r, in_f) * 0.01 x = torch.randn(1, 4, in_f) mask = build_local_token_mask([1], 0, 4) delta = (B @ (A @ x.T)).T * mask.unsqueeze(0).unsqueeze(-1) # token0/2/3 增量应为 0,token1 非零 assert delta[0, 0].abs().sum() == 0 assert delta[0, 2].abs().sum() == 0 assert delta[0, 1].abs().sum() > 0这四个测试守护“全局索引越界被丢弃、本地不抛错、样本/序列语义分离、掩码只作用于活跃 token”。
八、排查清单
trainable_token_indices多卡失效时按序查:
- 确认索引语义:是序列位置索引还是样本索引?别混用。
- 检查本地分片范围:序列维被 SP/FSDP 切分时,本地 token 是
[local_start, local_end),全局索引要减偏移。 - 重映射为本地相对索引:只对落在本地范围内的全局索引做
i - local_start,其余丢弃。 - 保存用全局命名:参数名用全局索引,避免各 rank 键不一致导致
KeyError。 - rank0 做全局校验:
validate_token_indices在全局序列长上校验越界,提前报错。 - 核对多卡 loss 一致:固定 seed,对比单卡与多卡在“等价为单卡切分”下的 loss,索引错位会表现为 loss 随卡数漂移。
- DDP 样本切分不影响序列索引:若只是
DistributedSampler按样本切,序列维完整,trainable_token_indices无需重映射(但仍要确认代码没误把它当样本索引)。 no_sync/ 梯度累积:确保trainable_token_indices在每个 micro-batch 都用正确的本地掩码,而非只在第一步计算。
九、小结
trainable_token_indices not working with >1 GPU的根因是:trainable_token_indices是全局序列/样本索引语义,多卡训练时每个进程只持有本地分片(序列维或样本维被切),但掩码仍用全局索引构造,导致越界或错位;再加上把样本索引误当序列索引、各 rank 把索引编入参数名导致保存时键不一致,问题被放大。
最小修复是在 forward 里把全局索引重映射为本地相对索引(范围外丢弃)、明确索引语义、保存时用全局索引命名;结构性改进是封装多卡安全的 token 掩码工具、用local_seq_range获取本地偏移、在 rank0 做全局校验;最后用测试守护“越界索引被丢弃、本地不抛错、样本/序列语义分离、掩码只作用活跃 token”。这样 token-level LoRA 才能在多卡下正确生效。