You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Vicuna-7b离线Prompt Tuning脚本CUDA/CPU张量设备不匹配错误排查

RuntimeError: 张量设备不匹配(cuda:0 vs cpu)在Vicuna-7b离线Prompt Tuning中的分析与解决

错误成因分析

这个错误的核心是调用torch.isin时传入的两个张量分别位于CUDA和CPU设备,触发于transformers库的isin_mps_friendly函数。结合Vicuna-7b离线Prompt Tuning的场景,常见诱因包括:

  • 离线Prompt模板或部分训练数据未被正确转移到CUDA设备:比如手动构建的prompt张量、缓存的样本数据留在了CPU,而模型参数已经加载到cuda:0。
  • Transformers版本兼容问题:部分旧版本的isin_mps_friendly函数未处理跨设备张量的情况,在MPS友好逻辑中默认假设张量都在CPU,与GPU上的模型张量冲突。
  • 离线数据加载逻辑漏洞:离线预处理时保存的张量未记录设备信息,加载后默认留在CPU,未与模型设备同步。

调试方案

  • 定位冲突张量:在报错位置附近添加打印逻辑,输出两个张量的设备信息:
    # 临时在torch.isin调用前添加(可通过修改transformers源码或断点调试)
    print(f"First tensor device: {element.device}")
    print(f"Second tensor device: {test_elements.device}")
    
    或使用pdb断点调试,在isin_mps_friendly函数处暂停,直接查看传入张量的设备属性。
  • 验证模型与数据的设备一致性:
    • 确认模型设备:print(next(model.parameters()).device),确保模型运行在cuda:0。
    • 检查输入数据设备:遍历数据加载器的batch,打印每个张量的device属性,排查是否有遗漏的CPU张量。
  • 排查transformers版本:执行pip show transformers查看当前版本,对比官方release notes确认是否存在已知的跨设备兼容bug。

修复方案

  1. 强制同步所有张量到CUDA设备
    在数据加载、prompt模板构建完成后,显式将所有相关张量迁移到模型所在设备:
    # 将batch内所有张量移到模型设备
    model_device = next(model.parameters()).device
    batch = {key: tensor.to(model_device) for key, tensor in batch.items()}
    # 将自定义prompt张量移到CUDA
    prompt_tensor = prompt_tensor.to("cuda:0")
    
  2. 临时补丁transformers的isin_mps_friendly函数
    找到transformers库中该函数的实现路径(通常为transformers/utils/generic.py),添加设备同步逻辑:
    def isin_mps_friendly(element, test_elements):
        if isinstance(element, torch.Tensor) and isinstance(test_elements, torch.Tensor):
            # 新增:将test_elements同步到element的设备
            test_elements = test_elements.to(element.device)
        try:
            return torch.isin(element, test_elements)
        except TypeError:
            return False
    
  3. 调整transformers版本
    若为版本兼容问题,升级至最新稳定版(如>=4.30.0)或降级至无此bug的版本(如4.28.1):
    pip install --upgrade transformers==4.30.0
    
  4. 修复离线数据加载逻辑
    离线保存数据时,优先将张量转为numpy数组存储,加载后再转换为张量并迁移到目标设备;或保存时额外记录设备信息,加载时自动同步到对应设备。

内容的提问来源于stack exchange,提问作者Omega

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 01:43:11