博客
关于我
PyTorch-Tutorials【pytorch官方教程中英文详解】- 1 Quickstart
阅读量:796 次
发布时间:2023-03-04

本文共 2935 字,大约阅读时间需要 9 分钟。

PyTorch基础教程:从数据到模型训练

PyTorch是一款强大的机器学习框架,支持深度学习模型的开发与训练。本节将介绍PyTorch中处理数据、构建模型、优化模型参数以及模型的保存与加载流程。


1. 数据处理

在PyTorch中,数据处理是机器学习模型的基础。PyTorch提供了torch.utils.data.DataLoadertorch.utils.data.Dataset两个核心工具。

  • Dataset:用于存储样本数据及对应标签,例如对于FashionMNIST数据集,Dataset会包含每个训练样本的图像和对应的类别标签。
  • DataLoader:将Dataset包装成一个可迭代的对象,支持批处理、随机采样和多进程加载。这里定义了一个批处理大小为64,每个批次将包含64个样本。
batch_size = 64
train_dataloader = DataLoader(training_data, batch_size=batch_size)
test_dataloader = DataLoader(test_data, batch_size=batch_size)

2. 模型构建

在PyTorch中,模型通常由多个神经网络层组成。可以通过继承nn.Module类来定义模型。

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using {device} device")
class NeuralNetwork(nn.Module):
def __init__(self):
super(NeuralNetwork, self).__init__()
self.flatten = nn.Flatten()
self.linear_relu_stack = nn.Sequential(
nn.Linear(28*28, 512),
nn.ReLU(),
nn.Linear(512, 512),
nn.ReLU(),
nn.Linear(512, 10)
)
def forward(self, x):
x = self.flatten(x)
logits = self.linear_relu_stack(x)
return logits

3. 模型优化

模型训练需要损失函数和优化器。

  • 损失函数:用于衡量预测结果与真实标签之间的差异。常用的有交叉熵损失函数。
  • 优化器:用于调整模型参数,SGD是常用的优化算法。
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=1e-3)

4. 训练过程

在训练循环中,模型会逐步优化参数以减少损失。

def train(dataloader, model, loss_fn, optimizer):
model.train()
size = len(dataloader.dataset)
for batch, (X, y) in enumerate(dataloader):
X, y = X.to(device), y.to(device)
pred = model(X)
loss = loss_fn(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if batch % 100 == 0:
print(f"loss: {loss.item():>7f} [{batch*len(X):>5d}/{size:>5d}]")

5. 模型测试

训练结束后,需要在测试集上评估模型性能。

def test(dataloader, model, loss_fn):
model.eval()
test_loss, correct = 0, 0
with torch.no_grad():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model(X)
test_loss += loss_fn(pred, y).item()
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
test_loss /= len(dataloader)
correct /= len(dataloader.dataset)
print(f"Test Error: Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f}")

6. 模型保存

将模型的状态字典保存为文件,便于后续加载和使用。

torch.save(model.state_dict(), "model.pth")
print("Saved PyTorch Model State to model.pth")

7. 模型加载

加载保存的模型状态字典,重新获得模型结构和参数。

model = NeuralNetwork()
model.load_state_dict(torch.load("model.pth"))
model.eval()

模型示例使用

classes = [
"T-shirt/top",
"Trouser",
"Pullover",
"Dress",
"Coat",
"Sandal",
"Shirt",
"Sneaker",
"Bag",
"Ankle boot"
]
x, y = test_data[0][0], test_data[0][1]
with torch.no_grad():
pred = model(x)
predicted, actual = classes[pred[0].argmax(0)], classes[y]
print(f'Predicted: "{predicted}", Actual: "{actual}"')

结果分析

训练过程中,训练损失逐渐降低,测试准确率逐步提升,表明模型正在有效学习数据特征。

转载地址:http://ywafk.baihongyu.com/

你可能感兴趣的文章
PowerShell攻击工具Nishang实战
查看>>
PowerShell攻击工具PowerSploit实战
查看>>
Powershell管理系列(四)Lync server 2013 批量启用语音及分配分机号
查看>>
PowerShell脚本运行完 不要马上关闭用什么命令可以停留窗口窗口
查看>>
PowerShell远程连接到Windows
查看>>
power(8) identity
查看>>
POW的重力之美
查看>>
PO、VO、DAO、BO、DTO、POJO能分清吗?
查看>>
pytorch介绍-ChatGPT4o作答
查看>>
PP-PLL:基于概率传播的部分标签学习
查看>>
pytorch介绍
查看>>
pprint 排序字典但不是集合?
查看>>
pptp拨号上网
查看>>
ppt上的倒计时小工具_PPT中有哪些「看似很 LOW,实则惊艳」的小工具
查看>>
PPT添加视频的路径问题
查看>>
PPT美化插件 islide 安装过程问题“加载com加载项时运行出现错误”
查看>>
Prefix Tuning:详细解读Optimizing Continuous Prompts for Generation
查看>>
PreparedStatement 与Statement 的区别,以及为什么推荐使用 PreparedStatement ?
查看>>
PreparedStatement 查询 In 语句 setArray 等介绍。
查看>>
presentModalViewController显示半透明的一个view
查看>>