手动映射设备后Hugging Face Trainer仍报编码器与分类头设备不一致错误
多GPU环境下Hugging Face Trainer设备不匹配问题排查与解决
问题场景
在多GPU环境运行Hugging Face Trainer时出现错误:RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cuda:1!
使用T5模型仅提取编码器,分片到两个设备,用LoRA封装后附加分类头,模型代码如下:
class ProtT5ForClassification(nn.Module): def __init__(self, encoder, device_map): super().__init__() self.encoder = encoder # already sharded hidden = self.encoder.config.d_model # create classifier but don’t push it to a device yet self.classifier = nn.Linear(hidden, 1, bias=True).to(torch.float16) # dispatch classifier to follow the encoder device map # simplest: put it entirely on the last shard (cuda:1 here) dispatch_model(self.classifier, device_map={"" : "cuda:1"}) self.loss_fn = nn.BCEWithLogitsLoss() def masked_mean_pool(self, hidden_states, attention_mask): mask = attention_mask.unsqueeze(-1).type_as(hidden_states) summed = (hidden_states * mask).sum(dim=1) denom = mask.sum(dim=1).clamp(min=1e-9) return summed / denom def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): # IMPORTANT: do not pass anything other than encoder-expected args to encoder enc_out = self.encoder(input_ids=input_ids, attention_mask=attention_mask, return_dict=True) last_hidden = enc_out.last_hidden_state pooled = self.masked_mean_pool(last_hidden, attention_mask) logits = self.classifier.to(pooled.device)(pooled).squeeze(-1) loss = None if labels is not None: labels = labels.float().view(-1) loss = self.loss_fn(logits, labels) return SequenceClassifierOutput(loss=loss, logits=logits)
已尝试将分类头映射到编码器最后分片所在设备(cuda:1),但错误仍存在。
排查与解决步骤
1. 确认各张量与模块的实际设备
在forward函数中添加打印语句,明确各组件的设备信息,避免假设错误:
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): enc_out = self.encoder(input_ids=input_ids, attention_mask=attention_mask, return_dict=True) last_hidden = enc_out.last_hidden_state # 打印设备信息 print("last_hidden device:", last_hidden.device) print("classifier param device:", next(self.classifier.parameters()).device) pooled = self.masked_mean_pool(last_hidden, attention_mask) print("pooled device:", pooled.device) # ...后续代码
2. 修正分类头与张量的设备迁移逻辑
你当前在forward中动态移动分类头到pooled的设备,这会和dispatch_model的固定映射冲突,改为将pooled张量移到分类头所在设备:
# 替换原logits行 classifier_device = next(self.classifier.parameters()).device logits = self.classifier(pooled.to(classifier_device)).squeeze(-1)
3. 同步标签与logits的设备
计算损失时,标签可能默认在cuda:0,需手动移到logits所在设备:
if labels is not None: labels = labels.float().view(-1).to(logits.device) loss = self.loss_fn(logits, labels)
4. 验证LoRA模块的设备一致性
检查LoRA适配器参数是否跟随编码器分片设备:
for name, param in self.encoder.named_parameters(): if "lora" in name: print(f"{name}: {param.device}")
若存在设备不匹配,需重新初始化LoRA时指定与编码器一致的device_map。
5. 用Accelerator统一设备管理
手动处理设备映射易出错,改用accelerate自动分配:
from accelerate import Accelerator accelerator = Accelerator() # 初始化模型时不手动指定device_map model = ProtT5ForClassification(encoder, device_map=None) # 统一准备模型、优化器和数据加载器 model, optimizer, train_dataloader = accelerator.prepare(model, optimizer, train_dataloader)
6. 确认编码器分片的实际设备
查看编码器的设备映射表,确认最后一层所在设备,再将分类头绑定到该设备:
print(self.encoder.hf_device_map) # 根据输出调整分类头的device_map,比如最后一层在cuda:0则改为 dispatch_model(self.classifier, device_map={"" : "cuda:0"})
内容的提问来源于stack exchange,提问作者Dwi Rezky Fahlan
相关产品推荐
相关产品推荐

