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

如何创建权重和为1的简单PyTorch神经网络

实现所有权重总和为1的PyTorch神经网络

下面分不同场景给出实现方案:

方案1:强制每一步权重总和严格等于1

适合对权重和有硬性要求的场景,每次前向传播前都会对权重做归一化,保证总和始终为1:

import torch
import torch.nn as nn

class StrictSumOneNet(nn.Module):
    def __init__(self, input_dim, output_dim, use_bias=False):
        super().__init__()
        self.fc = nn.Linear(input_dim, output_dim, bias=use_bias)
        # 初始化权重后先做一次归一化
        nn.init.normal_(self.fc.weight)
        with torch.no_grad():
            self.fc.weight.div_(self.fc.weight.sum())
    
    def forward(self, x):
        # 前向传播前重新归一化权重,避免训练过程中权重和偏离1
        with torch.no_grad():
            self.fc.weight.div_(self.fc.weight.sum())
        return self.fc(x)

方案2:通过损失函数约束权重和趋近于1

如果不需要严格保证每一步权重和都为1,只需要训练收敛后权重和接近1,可在损失中加入正则项,不会额外增加前向传播的计算开销:

import torch.optim as optim

# 初始化模型、优化器、损失函数
model = StrictSumOneNet(10, 2)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.MSELoss()
reg_lambda = 1e-2 # 正则项系数,可根据实际效果调整

# 训练循环示例
for inputs, labels in train_dataloader:
    optimizer.zero_grad()
    outputs = model(inputs)
    task_loss = criterion(outputs, labels)
    # 加入权重和约束正则项
    weight_sum_loss = reg_lambda * torch.abs(model.fc.weight.sum() - 1)
    total_loss = task_loss + weight_sum_loss
    total_loss.backward()
    optimizer.step()

方案3:非负权重且总和为1

如果要求所有权重需要同时满足非负、总和为1,直接用softmax处理原始权重即可,不需要额外做归一化:

class PositiveSumOneNet(nn.Module):
    def __init__(self, input_dim, output_dim):
        super().__init__()
        # 定义可训练的原始权重参数
        self.weight_raw = nn.Parameter(torch.randn(output_dim, input_dim))
    
    def forward(self, x):
        # softmax输出天然满足所有元素非负、总和为1
        norm_weight = torch.softmax(self.weight_raw.flatten(), dim=0).reshape_as(self.weight_raw)
        return nn.functional.linear(x, norm_weight, bias=None)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 18:15:03