PyTorch相同模型实例输出为何不同?如何实现一致输出?
如何让PyTorch中多次实例化的同结构模型输出一致
核心原因
你设置的随机种子仅在代码启动时生效一次,而nn.Linear这类层的权重、偏置依赖随机数生成器初始化。第一次创建模型后,随机数生成器的状态会被消耗(前进),后续实例化模型时,生成器已处于新状态,导致初始化参数不同,最终输出存在差异。
解决方案
方案一:每次实例化模型前重置随机种子
通过封装种子设置函数,每次创建新模型前重新初始化所有随机种子,确保每次模型初始化都使用相同的随机状态。
修改后的代码示例:
import os import random import numpy as np import torch from torch import nn from torch.backends import cudnn seed = 0 def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed(seed) cudnn.deterministic = True cudnn.benchmark = False class SimpleNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) def forward(self, x): out = self.fc1(x) return out input_size = 10 hidden_size = 20 output_size = 5 device = 'cpu' input_data = torch.randn(32, input_size) input_data = input_data.to(device) # 第一次实例化模型 set_seed(seed) model = SimpleNet(input_size, hidden_size, output_size).to(device) output1 = model(input_data) print(output1) # 第二次实例化前重置种子 set_seed(seed) model = SimpleNet(input_size, hidden_size, output_size).to(device) output2 = model(input_data) print(output2) # 第三次实例化前重置种子 set_seed(seed) model = SimpleNet(input_size, hidden_size, output_size).to(device) output3 = model(input_data) print(output3)
方案二:保存初始模型的状态字典,后续加载复用
先创建一个基准模型并保存其权重参数,之后每次实例化新模型后,直接加载该状态字典,确保所有模型使用完全相同的初始化参数。
代码示例:
import os import random import numpy as np import torch from torch import nn from torch.backends import cudnn seed = 0 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed(seed) cudnn.deterministic = True cudnn.benchmark = False class SimpleNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) def forward(self, x): out = self.fc1(x) return out input_size = 10 hidden_size = 20 output_size = 5 device = 'cpu' input_data = torch.randn(32, input_size) input_data = input_data.to(device) # 创建基准模型并保存状态字典 base_model = SimpleNet(input_size, hidden_size, output_size).to(device) base_state_dict = base_model.state_dict() # 第一次使用 model = SimpleNet(input_size, hidden_size, output_size).to(device) model.load_state_dict(base_state_dict) output1 = model(input_data) print(output1) # 第二次使用 model = SimpleNet(input_size, hidden_size, output_size).to(device) model.load_state_dict(base_state_dict) output2 = model(input_data) print(output2) # 第三次使用 model = SimpleNet(input_size, hidden_size, output_size).to(device) model.load_state_dict(base_state_dict) output3 = model(input_data) print(output3)
内容的提问来源于stack exchange,提问作者user22477480
相关产品推荐
相关产品推荐

