model 的文件夹中,以便其他脚本调用。import torchimport torch.nn as nn# 定义简单的全连接神经网络class SimpleNet(nn.Module):def __init__(self):super(SimpleNet, self).__init__()self.fc1 = nn.Linear(784, 128)# 28*28 = 784 输入节点, 128 输出节点self.fc2 = nn.Linear(128, 64)self.fc3 = nn.Linear(64, 2)# 输出10类,对应MNIST的10个数字def forward(self, x):x = x.view(-1, 784)# flatten the imagex = torch.relu(self.fc1(x))x = torch.relu(self.fc2(x))x = self.fc3(x)return torch.log_softmax(x, dim=1)
# 从model模块导入SimpleNet类from model.simple_nn import SimpleNet# 导入PyTorch库import torch# 设置保存模型参数的文件路径path_state_dict = '3_simple_nn_model.pth'# 创建SimpleNet类的实例,这会初始化模型model = SimpleNet()# 加载保存的模型参数到模型实例model.load_state_dict(torch.load(path_state_dict))# 将模型切换到评估模式,这对于进行预测是必要的,因为它会禁用一些特定于训练阶段的操作,比如Dropoutmodel.eval()# 生成随机输入数据:1个样本,每个样本10个特征input_data = torch.randn(1, 784)# 使用模型进行预测,此时不计算梯度以提高性能和减少内存使用with torch.no_grad(): output = model(input_data)
2. 载入和使用模型
# 从model模块导入SimpleNet类from model.simple_nn import SimpleNet# 导入PyTorch库import torch# 设置保存模型参数的文件路径path_state_dict = '3_simple_nn_model.pth'# 创建SimpleNet类的实例,这会初始化模型model = SimpleNet()# 加载保存的模型参数到模型实例model.load_state_dict(torch.load(path_state_dict))# 将模型切换到评估模式,这对于进行预测是必要的,因为它会禁用一些特定于训练阶段的操作,比如Dropoutmodel.eval()# 生成随机输入数据:1个样本,每个样本10个特征input_data = torch.randn(1, 784)# 使用模型进行预测,此时不计算梯度以提高性能和减少内存使用with torch.no_grad():output = model(input_data)
SimpleNet_Predictor) 来展示如何优雅地在应用中使用模型,并将其存储在python包中(utils_AI)。import torchfrom torchvision import transformsfrom model.simple_nn import SimpleNetclass SimpleNet_Predictor:def __init__(self, model_path):"""初始化模型预测器,加载模型并进行热身。:param model_class: 模型类,用于实例化模型:param model_path: 预训练模型的路径,用于加载模型参数"""# 实例化模型self.model = SimpleNet()# 加载模型参数self.model.load_state_dict(torch.load(model_path))# 切换到评估模式self.model.eval()# 定义数据预处理self.transform = transforms.Compose([# 此处为空,可自定义])# 进行模型热身,以确保模型在首次调用时响应迅速self.warmup()def warmup(self):"""执行一次前向传递以热身模型。"""with torch.no_grad():# 使用一些随机数据进行热身random_input = torch.randn(1, 784)random_input = self.transform(random_input)# 应用转换self.model(random_input)def __call__(self, input_data):"""使得该类实例可以像函数那样被调用,进行模型预测。:param input_data: 输入数据,应为torch.Tensor格式:return: 模型的输出结果"""with torch.no_grad():input_data = self.transform(input_data)# 应用转换return self.model(input_data)
from utils_AI.simplenet_predictor import SimpleNet_Predictorimport torchmodel_path = '3_simple_nn_model.pth'# 假设模型参数已经保存在这个路径predictor = SimpleNet_Predictor(model_path)# 测试模型预测test_input = torch.randn(1, 784)output = predictor(test_input)print("Model output shape:", output.shape)
