PyTorch替换加权MSE损失遇设备迁移错误的解决方案求助
解决PyTorch中替换加权MSE损失的报错问题
错误原因
你自定义的weighted_mse_loss是普通Python函数,而非PyTorch的nn.Module子类。原代码中调用self.mse_criterion.to(device)时,会触发Module类的内部逻辑,而普通函数没有_modules等Module专属属性,因此报错。
两种可行解决方案
方案一:将加权MSE封装为nn.Module子类
这种方式可以完全兼容原有代码结构,包括保留to(device)调用:
import torch class WeightedMSELoss(torch.nn.Module): def __init__(self): super().__init__() # 注册权重为buffer(不可训练,随模块自动迁移设备) # 你的权重列表是24个2 + 18个1,简化写法 self.register_buffer('weight', torch.FloatTensor([2]*24 + [1]*18)) def forward(self, input, target): return self.weight * (input - target) ** 2
替换原有代码:
# 初始化损失模块 self.mse_criterion = WeightedMSELoss() # 损失计算逻辑不变 loss_mot_rec = self.mse_criterion(self.fake_noise, self.real_noise).mean(dim=-1)
to方法无需修改,因为WeightedMSELoss是Module子类,调用to(device)时会自动将buffer中的权重迁移到目标设备。
方案二:使用普通函数,修改设备适配逻辑
如果不想封装成Module,需手动处理权重的设备匹配,并修改to方法:
首先调整损失函数,确保权重与输入张量同设备:
def weighted_mse_loss(input, target): # 从输入张量的设备和 dtype 创建权重,避免设备不匹配 weight = torch.tensor([2]*24 + [1]*18, dtype=input.dtype, device=input.device) return weight * (input - target) ** 2
替换原有代码:
# 无需初始化mse_criterion模块,直接调用函数 loss_mot_rec = weighted_mse_loss(self.fake_noise, self.real_noise).mean(dim=-1)
修改to方法,移除对普通函数的to调用:
def to(self, device): if self.opt.is_train: # 普通函数无需迁移设备,此处留空或删除原行 pass self.encoder = self.encoder.to(device)
关键注意事项
- 无论哪种方案,都要保证权重张量与输入/目标张量在同一设备上(CPU/GPU),否则会触发设备不匹配报错。
- 若需要权重可训练,方案一中可将
register_buffer改为self.weight = torch.nn.Parameter(...),并设置requires_grad=True。
内容的提问来源于stack exchange,提问作者Nikita Belooussov
相关产品推荐
相关产品推荐

