【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, dtypebool) 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, dtypetorch.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, dtypetorch.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_len10 当成局部实际本地只有 4 个 token try: apply_token_lora_global(x, B, A, idx_global, seq_len4) 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_start4, local_end8) print(local 版输出形状:, tuple(out_local.shape)) print(只有本地 token1(全局5) 有非零增量:, out_local.abs().sum(dim-1).squeeze().tolist())运行后global 版抛IndexErrorlocal 版正确把全局索引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, dtypetorch.bool) active[trainable_sample_indices] True return active别把样本索引喂进 token 维掩码。修复 3保存时用全局索引命名加载再重映射# 保存参数名用全局索引保证各 rank 键一致 state {flora_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, dtypetorch.bool, devicedevice) 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(ftrainable_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, devicecpu): local_len local_end - local_start mask torch.zeros(local_len, dtypetorch.bool, devicedevice) 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, dtypetorch.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 增量应为 0token1 非零 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 才能在多卡下正确生效。