PyTorch自定义损失函数权重获取及无传参添加L1正则方法
嘿,我来帮你搞定这两个PyTorch相关的问题:
问题1:如何在PyTorch中获取自定义损失函数所需的权重?
获取模型权重的方式有好几种,根据你的使用场景选就行:
- 直接从模型实例提取:如果你已经初始化了模型,比如
model = AutoEncoder(inp_size=784, hid_size=128),可以通过参数的属性直接获取,比如编码器线性层的权重是model.e1.weight,解码器的是model.d1.weight。如果要遍历所有可训练参数,用model.parameters();要是想拿到参数名方便筛选(比如只挑权重,排除偏置),就用model.named_parameters(),比如:for name, param in model.named_parameters(): if 'weight' in name: # 处理权重参数 pass - 在模型内部访问:如果你的损失计算逻辑是和模型绑定的(比如写成模型的一个方法),那在方法里直接用
self.e1.weight这种方式就能拿到对应权重,完全不用额外传递。 - 注意:
model.parameters()只会返回requires_grad=True的可训练参数,如果你有冻结的层(设置了requires_grad=False),这些参数不会被包含进来,刚好符合正则化只针对可训练参数的需求。
问题2:AutoEncoder添加L1正则化但不传入权重的最优实现
其实最优的方式是把正则化的计算逻辑和模型整合到一起,不用手动传权重,我给你补全并修改你的模型代码,实现这个需求:
import torch import torch.nn as nn class AutoEncoder(nn.Module): def __init__(self, inp_size, hid_size): super(AutoEncoder, self).__init__() self.lambd = 1. # 可以根据需求调整L1正则化的系数 # Encoder self.e1 = nn.Linear(inp_size, hid_size) # Decoder self.d1 = nn.Linear(hid_size, inp_size) self.sigmoid = nn.Sigmoid() def forward(self, x): encoded = torch.relu(self.e1(x)) decoded = self.sigmoid(self.d1(encoded)) return decoded def compute_loss(self, x): # 先得到模型的重构输出 decoded = self.forward(x) # 计算重构损失,这里用MSE,你也可以换成交叉熵等其他损失 recon_loss = nn.MSELoss()(decoded, x) # 计算L1正则化:遍历需要正则化的权重(这里选编码器和解码器的权重) l1_reg = torch.tensor(0., device=x.device) # 确保和输入在同一设备上 for param in [self.e1.weight, self.d1.weight]: l1_reg += torch.norm(param, p=1) # 总损失 = 重构损失 + 正则化项 total_loss = recon_loss + self.lambd * l1_reg return total_loss
训练的时候你只需要这么用:
model = AutoEncoder(inp_size=784, hid_size=128) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) # 假设x是你的输入数据 for epoch in range(epochs): optimizer.zero_grad() loss = model.compute_loss(x) loss.backward() optimizer.step()
这种方式的好处是:
- 不用手动把权重传入损失函数,模型内部直接访问参数,代码更简洁;
- 损失计算逻辑和模型绑定,维护起来更方便,后续调整正则化范围(比如加偏置或者其他层)直接修改
compute_loss里的参数列表就行; - 如果想对所有可训练参数做L1正则化,还可以把遍历部分改成:
这样不管后续模型加多少层,都不用手动更新参数列表。for param in self.parameters(): l1_reg += torch.norm(param, p=1)
另外还有一种方式是自定义损失类,初始化时传入模型实例,但相比上面的方法会多一步实例化损失类的操作,没有整合到模型里那么直观。
内容的提问来源于stack exchange,提问作者oeb
相关产品推荐
相关产品推荐

