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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 02:05:24