当前位置: 首页 > news >正文

【Bug已解决】GPT2 cannot be used with device_map=‘auto‘; Report “found at least two devices“ 解决方案

【Bug已解决】GPT2 cannot be used with device_map='auto'; Report "found at least two devices" 解决方案

一、现象长什么样

想用acceleratedevice_map="auto"把 GPT2 自动切到多张卡(或 CPU+GPU 混合)做大模型推理:

from transformers import GPT2LMHeadModel, AutoTokenizer model = GPT2LMHeadModel.from_pretrained( "gpt2", device_map="auto", torch_dtype="auto", ) tok = AutoTokenizer.from_pretrained("gpt2") out = model.generate(tok("Hello", return_tensors="pt").input_ids)

报错:

ValueError: You're trying to load a model that was saved with a different device mapping. ... Found at least two devices: cuda:0 and cpu (or cuda:1)

或者更常见的变体:

RuntimeError: GPT2LMHeadModel: weight wte.weight is on cuda:0 but lm_head.weight expects cpu — found at least two devices

最迷惑的是:用device_map="auto"本意是「让 accelerate 帮我自动放」,结果它反而因为「两个模块共享同一份权重却在两个设备」而拒绝。单卡model.cuda()完全正常,一上device_map就炸。

二、背景

GPT2 的wte(词嵌入)和lm_head(语言模型头)是**权重共享(tied)**的:lm_head.weight直接复用wte.weight,不单独存参数。

acceleratedevice_map="auto"做自动切分时,会逐个子模块决定放哪张卡。问题在于:它默认把wtelm_head当成两个独立模块分别规划设备。如果自动规划把wtecuda:0、把lm_headcpu(或cuda:1),就出现「同一份逻辑权重被要求同时在两个设备」的矛盾——生成时lm_head要拿wte的权重,但权重在另一张卡,于是报「found at least two devices」。

正确行为应该是:accelerate 知道lm_headwte是 tied,把它们强制放在同一设备。但 GPT2 早期的实现里,lm_head没有被显式标注为 tied(或_tied_weights_keys没登记),accelerate 无从得知,就各放各的,触发矛盾。

三、根因

根因一句话:GPT2 的lm_headwte是共享权重,但device_map="auto"的自动切分把它们当成独立模块分到不同设备,导致「同一权重需跨设备」的矛盾,触发 found-at-least-two-devices 错误。

三点展开:

  1. tied 未登记lm_head没在_tied_weights_keys/ tied 标记里登记,accelerate 不知道它和wte是同一份。
  2. 自动切分各放各的device_map="auto"wtelm_head分别规划到cuda:0/cpu(或不同卡),产生设备矛盾。
  3. 缺乏 co-locate 兜底:当检测到 tied 权重被分到不同设备时,没有「强制合并到同一设备」的兜底,直接抛出ValueError/RuntimeError

不是卡不够,是「权重共享关系没告诉切分器」。

四、最小可运行复现

不依赖真实大模型,模拟「tied 权重被分到两个设备」:

from dataclasses import dataclass from typing import List @dataclass class FakeModule: name: str device: str class FakeAutoMap: def __init__(self, modules): self.modules = modules def validate_tied(self, tied_pairs): # 检查 tied 的两个模块是否同设备 by_name = {m.name: m.device for m in self.modules} errors = [] for a, b in tied_pairs: if by_name[a] != by_name[b]: errors.append((a, by_name[a], b, by_name[b])) return errors # 模拟 device_map="auto" 把 wte 和 lm_head 分到不同设备 modules = [FakeModule("wte", "cuda:0"), FakeModule("lm_head", "cpu")] automap = FakeAutoMap(modules) errs = automap.validate_tied([("wte", "lm_head")]) print("tied 冲突:", errs) # [('wte','cuda:0','lm_head','cpu')] -> 报错

跑出来:wtecuda:0lm_headcpu,tied 冲突非空 → 触发 found-at-least-two-devices。这就是精确复现。

五、解决方案(第一层:最小直接修复)

最小修复:device_map="auto"知道lm_headwte是 tied,强制同设备。两条路:

路线 A(推荐):确保模型声明 tied,让 accelerate 自动 co-locate。

from transformers import GPT2LMHeadModel # 关键:显式登记 tied 权重,accelerate 会把它们放同一设备 model = GPT2LMHeadModel.from_pretrained( "gpt2", device_map="auto", torch_dtype="auto", # 若模型未自动登记,可手动: ) # 手动确保 tied(部分版本需要) if model.lm_head.weight is not model.wte.weight: model.lm_head.weight = model.wte.weight # 复用,而非独立参数 out = model.generate(...)

路线 B(稳妥):如果仍冲突,用显式device_map把共享权重的模块放同一设备,或干脆整模型放一张卡:

# 整模型一张卡,彻底避免跨设备 model = GPT2LMHeadModel.from_pretrained("gpt2", device_map={"": 0}) # 或显式指定:wte 和 lm_head 必须在同一设备 model = GPT2LMHeadModel.from_pretrained( "gpt2", device_map={"wte": 0, "lm_head": 0, "h": 0, "ln_f": 0}, )

要点:

  • tied 权重必须物理复用同一 Parameter 对象lm_head.weight = wte.weight),accelerate 才不会各放各的。
  • 若自动切分仍冲突,用显式device_map字典把共享模块锁在同一设备。
  • 纯多卡且必须切分时,确保所有引用同一权重的模块都在同一device_map条目下。

这一步单独就让device_map="auto"不再报 two-devices。

六、解决方案(第二层:结构性改进)

第一层是「在加载处补一行 tied」。但 GPT2 及所有 tied 权重的模型(OPT、BLOOM、GPT-Neo 等)都面临同样问题。更稳的做法把「tied 权重如何登记、如何强制同设备」收敛成单一策略对象。

from dataclasses import dataclass, field from typing import Dict, List, Tuple @dataclass class DeviceMapTieResolver: """解决 tied 权重在 device_map 下被分到多设备的单一策略。""" # tied 权重对:共享同一物理参数的模块路径 tied_pairs: List[Tuple[str, str]] = field(default_factory=list) def register(self, a: str, b: str): self.tied_pairs.append((a, b)) def co_located_map(self, auto_map: Dict[str, str]) -> Dict[str, str]: """把每个 tied 对的成员强制放到同一设备(取第一个出现的设备)。""" resolved = dict(auto_map) for a, b in self.tied_pairs: dev_a = resolved.get(a) dev_b = resolved.get(b) if dev_a is not None and dev_b is not None and dev_a != dev_b: # 强制 b 跟随 a 的设备 resolved[b] = dev_a elif dev_a is None and dev_b is not None: resolved[a] = dev_b elif dev_b is None and dev_a is not None: resolved[b] = dev_a return resolved def validate(self, final_map: Dict[str, str]) -> List[str]: errors = [] for a, b in self.tied_pairs: if final_map.get(a) != final_map.get(b): errors.append(f"tied 冲突: {a}={final_map.get(a)} vs {b}={final_map.get(b)}") return errors # 用法:针对 GPT2 登记 wte<->lm_head resolver = DeviceMapTieResolver() resolver.register("wte", "lm_head") auto = {"wte": "cuda:0", "lm_head": "cpu", "h": "cuda:0"} final = resolver.co_located_map(auto) print("修正后映射:", final) # lm_head 被拉回 cuda:0 print("冲突:", resolver.validate(final)) # []

结构收益:

  • 单一策略:所有 tied 对集中登记,device_map自动 co-locate。
  • 可校验validate在加载前抓出任何 tied 跨设备,避免运行时报错。
  • 可复用:OPT/BLOOM/GPT-Neo 等只需追加register即可。

七、解决方案(第三层:断言 / CI 守护)

写 pytest 守三条:(1) tied 对最终同设备;(2) 冲突被validate抓出;(3) 修正后无 two-devices。

import pytest from your_lib import DeviceMapTieResolver @pytest.fixture def resolver(): r = DeviceMapTieResolver() r.register("wte", "lm_head") return r def test_co_located(resolver): auto = {"wte": "cuda:0", "lm_head": "cpu", "h": "cuda:0"} final = resolver.co_located_map(auto) assert final["wte"] == final["lm_head"] == "cuda:0" def test_validate_catches_conflict(resolver): bad = {"wte": "cuda:0", "lm_head": "cuda:1"} errs = resolver.validate(bad) assert len(errs) == 1, "应抓出 tied 跨设备冲突" def test_no_conflict_after_fix(resolver): auto = {"wte": "cuda:0", "lm_head": "cpu"} final = resolver.co_located_map(auto) assert resolver.validate(final) == [], "修正后不应有冲突" def test_multiple_tied_pairs(): r = DeviceMapTieResolver() r.register("wte", "lm_head") r.register("shared1", "shared2") auto = {"wte": "cuda:0", "lm_head": "cpu", "shared1": "cuda:1", "shared2": "cpu"} final = r.co_located_map(auto) assert final["wte"] == final["lm_head"] assert final["shared1"] == final["shared2"]

CI 常驻跑这四条后,任何「tied 权重又被分到多设备」的回归都会立刻爆红。

八、排查清单

GPT2 / tied 模型上device_map报 two-devices 时,按顺序查:

  1. 先确认报错含found at least two deviceswte.weight ... lm_head.weight—— 是的话定位 tied 跨设备。
  2. 检查模型是否有 tied 权重:model.lm_head.weight is model.wte.weight应为True
  3. 若为False,手动model.lm_head.weight = model.wte.weight复用同一对象。
  4. 确认_tied_weights_keys"lm_head.weight"(或对应路径),让 accelerate 识别 tied。
  5. 若自动切分仍冲突,用显式device_map字典把 tied 模块锁同一设备。
  6. 多卡时,确认所有「共享同一物理参数」的模块都在同一device_map条目。
  7. 升级 transformers/accelerate 后,重跑一次device_map="auto"加载冒烟,断言 tied 同设备。

九、小结

GPT2 上device_map="auto"报 found-at-least-two-devices,根子是lm_headwte是 tied 共享权重,却没被登记进 tied 关系,accelerate 的自动切分把它们分到不同设备,造成「同一权重需跨设备」矛盾。修复三层次:第一层确保lm_head.weight物理复用wte.weight,或用显式device_map锁同设备;第二层用DeviceMapTieResolverdataclass 把 tied 对登记、自动 co-locate 并校验;第三层用 pytest 守「tied 同设备」「冲突被抓」「修正后无 two-devices」。

工程启示:凡是带 tied 权重的模型走device_map切分,都必须先把共享关系告诉切分器(登记_tied_weights_keys/ 物理复用 Parameter),否则自动切分必然把共享权重拆到多设备而拒绝加载。这是多卡/CPU-offload 推理最高频的坑之一。

http://www.jsqmd.com/news/1356943/

相关文章:

  • 浏览器中的微信:5分钟实现免安装终极工作沟通方案
  • word转图片用啥软件好?盘点这7款PDF格式转换工具,电脑端在线工具都有覆盖
  • Python、Java与C语言:主流编程语言对比与应用指南
  • 前端组件开发公众号技术社区的商业化路径,构建双赢的生态合作模式
  • 告别繁琐手动操作:semi-utils 让你的照片批量水印处理效率提升10倍
  • VISSIM交通仿真软件的核心技术与应用实践
  • M5Stack Cardputer ADV UART驱动安装与通信调试全攻略
  • Java多线程设计模式与并发编程实战指南
  • 小心数据手册的“乘法陷阱”:SPAD实际探测效率远低于标称的深层原因
  • 华为鸿蒙拒绝无效社交APP—小羊断交
  • 5分钟快速上手:用Python轻松获取NBA官方数据的终极指南
  • Selenium爬虫卡顿问题深度解析:从等待策略到实战调试
  • Power BI中Base64图片与超链接的实战应用
  • DLL 指纹库 采集、检索 TLSFOWARD tls指纹库
  • OpenClaw API密钥安全防护:从环境变量到纵深防御实战
  • 40K Star!支持2000+应用集成,开源版Zapier,自动化工作流神器
  • CVE-2026-8819 细节曝光:首个 AI 蠕虫 AgentWorm 如何利用 LangChain/AutoGPT 窃取 API Key
  • DolphinScheduler集成OIDC认证:企业级多租户身份管理方案
  • 软件工程基本功:超越AI热潮,构建可靠、可维护、可扩展的软件系统
  • 虚拟列表技术原理与性能优化实战
  • Vue 3 UI组件库从零搭建:Monorepo架构、按需加载与工程化实践
  • Unity Mono游戏逆向实战:Frida Hook绕过碰撞死亡判定
  • 交互式测试仪表盘:提升软件测试效率的关键工具
  • 微信自动化机器人开发指南与技术方案对比
  • VS Code集成GitHub Copilot全攻略与优化技巧
  • 从静态孪生到动态镜像:工业实时监管系统的架构演进与实践
  • Cocos Creator Shader源码解析:从特效原理到实战优化
  • COMSOL多物理场建模在地热能非均质储层开发中的应用
  • MyBatis N+1查询坑,百万数据下接口直接超时
  • React动态导入竞态问题与AI编程实践