news 2026/7/25 4:11:30

【Bug已解决】trainable_token_indices of LoraConfig not working when using more than 1 NVIDIA GPU 解决方案

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【Bug已解决】trainable_token_indices of LoraConfig not working when using more than 1 NVIDIA GPU 解决方案

【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被理解成“跨样本的全局索引”时。

更常见的真实踩坑场景:

  1. 索引是按“样本维”给的,但代码按“序列维”用:用户想训练第 0、1、2 个样本,却把索引喂进了序列维掩码,单卡序列短碰巧不越界,多卡序列分片后越界。
  2. FSDP/DeepSpeed 对长序列做分片:序列维被切到多卡,trainable_token_indices仍用全局索引,本地分片只有[start, end)区间的 token,全局索引落在别处就IndexError
  3. 保存时各 rank 索引键不一致:token adapter 把索引写进参数名(如...trainable_tokens_delta.default),各 rank 本地索引不同 →state_dict键集合不同 → 聚合/保存报KeyError

下面用最小可运行代码演示“全局索引在本地分片上越界”以及如何重映射。

三、根因

根因一句话:trainable_token_indices是全局序列/样本索引语义,多卡训练时每个进程只持有本地分片(序列维或样本维被切),但掩码仍用全局索引构造,导致越界或错位。

展开:

  1. 索引未重映射:本地分片 token 范围是[local_start, local_end),全局索引i需减local_start才是本地合法索引;不重映射就IndexError
  2. 索引语义混淆:把“样本索引”误当“序列索引”喂给 token 掩码。
  3. 多卡键不一致: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多卡失效时按序查:

  1. 确认索引语义:是序列位置索引还是样本索引?别混用。
  2. 检查本地分片范围:序列维被 SP/FSDP 切分时,本地 token 是[local_start, local_end),全局索引要减偏移。
  3. 重映射为本地相对索引:只对落在本地范围内的全局索引做i - local_start,其余丢弃。
  4. 保存用全局命名:参数名用全局索引,避免各 rank 键不一致导致KeyError
  5. rank0 做全局校验validate_token_indices在全局序列长上校验越界,提前报错。
  6. 核对多卡 loss 一致:固定 seed,对比单卡与多卡在“等价为单卡切分”下的 loss,索引错位会表现为 loss 随卡数漂移。
  7. DDP 样本切分不影响序列索引:若只是DistributedSampler按样本切,序列维完整,trainable_token_indices无需重映射(但仍要确认代码没误把它当样本索引)。
  8. 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 才能在多卡下正确生效。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/25 4:11:22

Physics of AI:从物理规律探索通用人工智能新路径

1. 专访背景与核心观点解读最近MIT研究员刘子鸣提出的"Physics of AI"研究路径在人工智能领域引发广泛讨论。这位年轻科学家主张跳出当前主流的大模型规模竞赛&#xff0c;转而从物理学的底层规律出发探索AGI&#xff08;通用人工智能&#xff09;的实现路径。这种&q…

作者头像 李华
网站建设 2026/7/25 4:11:21

函数式编程与游戏引擎融合:Haskell绑定Godot开发实践

1. 项目概述&#xff1a;当函数式编程遇上游戏引擎如果你和我一样&#xff0c;既着迷于Haskell那种纯粹、优雅的函数式编程范式&#xff0c;又被Godot引擎的轻量、高效与节点化设计所吸引&#xff0c;那么“Godot-Haskell”这个项目对你来说&#xff0c;可能就像发现了一座宝藏…

作者头像 李华
网站建设 2026/7/25 4:10:05

前端项目骨架模板化:从 Create React App 到定制化脚手架

前端项目骨架模板化&#xff1a;从 Create React App 到定制化脚手架CRA 给你一个项目&#xff0c;定制化脚手架给你的是一整个团队的工程共识。一、场景痛点 新项目启动&#xff0c;npx create-react-app 一把梭。然后开始删代码&#xff1a;删 App.css、删 logo.svg、删测试文…

作者头像 李华
网站建设 2026/7/25 4:07:15

AI人格化技术解析与实战指南

1. 项目概述&#xff1a;当AI人格化成为现象级话题上周三凌晨&#xff0c;一张疑似GPT-5.3系统生成的对话截图突然在各大技术社区刷屏。图中AI不仅准确预判了用户的隐藏需求&#xff0c;还展现出类似人类的情感共鸣能力——这直接引爆了关于"AI人格化"的技术伦理讨论…

作者头像 李华
网站建设 2026/7/25 4:06:21

CNN结合时频分析与注意力机制的信号分类模型

1. 项目概述&#xff1a;当CNN遇见时频分析与注意力机制这个项目实现了一个融合三种核心技术的分类预测模型&#xff1a;卷积神经网络&#xff08;CNN&#xff09;负责提取局部特征&#xff0c;S变换&#xff08;Stockwell Transform&#xff09;提供信号的时频表示&#xff0c…

作者头像 李华
网站建设 2026/7/25 4:06:11

LinkSwift:九大网盘直链解析工具,免费解锁高速下载的完整指南

LinkSwift&#xff1a;九大网盘直链解析工具&#xff0c;免费解锁高速下载的完整指南 【免费下载链接】Online-disk-direct-link-download-assistant 一个基于 JavaScript 的网盘文件下载地址获取工具。基于【网盘直链下载助手】修改 &#xff0c;支持 百度网盘 / 阿里云盘 / 中…

作者头像 李华