1. 模型定义

在PyTorch中,我们通过继承nn.Module类来定义神经网络模型。以下是三个不同的网络结构示例:

1.1 基础卷积网络(Net类)

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(3, 16, 5)  # 输入通道3,输出通道16,卷积核5x5
        self.pool1 = nn.MaxPool2d(2, 2)   # 2x2最大池化
        self.conv2 = nn.Conv2d(16, 30, 5) # 输入通道16,输出通道30,卷积核5x5
        self.pool2 = nn.MaxPool2d(2, 2)   # 2x2最大池化
        self.aap = nn.AdaptiveAvgPool2d(1) # 自适应平均池化到1x1
        self.fc3 = nn.Linear(30, 10)      # 全连接层,输出10个类别
        
    def forward(self, x):
        x = self.pool1(F.relu(self.conv1(x)))
        x = self.pool2(F.relu(self.conv2(x)))
        x = self.aap(x)                   # 全局平均池化
        x = x.view(x.shape[0], -1)        # 展平
        x = self.fc3(x)
        return x

1.2 类LeNet网络

class LeNet(nn.Module):
    def __init__(self):
        super(LeNet, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)    # 输入通道3,输出通道6,卷积核5x5
        self.conv2 = nn.Conv2d(6, 16, 5)   # 输入通道6,输出通道16,卷积核5x5
        self.fc1 = nn.Linear(16*5*5, 120)  # 全连接层
        self.fc2 = nn.Linear(120, 84)      # 全连接层
        self.fc3 = nn.Linear(84, 10)       # 输出层,10个类别
        
    def forward(self, x):
        out = F.relu(self.conv1(x))
        out = F.max_pool2d(out, 2)         # 2x2最大池化
        out = F.relu(self.conv2(out))
        out = F.max_pool2d(out, 2)         # 2x2最大池化
        out = out.view(out.size(0), -1)    # 展平
        out = F.relu(self.fc1(out))
        out = F.relu(self.fc2(out))
        out = self.fc3(out)
        return out

2. 数据准备与预处理

数据预处理是深度学习中的重要环节,以下是CIFAR-10数据集的处理示例:

# 设备配置
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# 训练数据变换
transform_train = transforms.Compose([
    transforms.RandomCrop(32, padding=4),      # 随机裁剪
    transforms.RandomHorizontalFlip(),         # 随机水平翻转
    transforms.ToTensor(),                     # 转换为张量
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), # 标准化
])

# 测试数据变换
transform_test = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
])

# 加载数据集
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, 
                                       download=False, transform=transform_train)
trainloader = DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)

testset = torchvision.datasets.CIFAR10(root='./data', train=False, 
                                      download=False, transform=transform_test)
testloader = DataLoader(testset, batch_size=100, shuffle=False, num_workers=2)

# CIFAR-10类别
classes = ('plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck')

3. 模型训练流程

训练循环是深度学习的核心部分,以下是标准的训练代码:

# 定义损失函数和优化器
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)

# 训练循环
for epoch in range(10):
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):
        # 获取训练数据
        inputs, labels = data
        inputs, labels = inputs.to(device), labels.to(device)
        
        # 梯度清零
        optimizer.zero_grad()
        
        # 前向传播 + 反向传播 + 优化
        outputs = net(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        # 统计损失
        running_loss += loss.item()
        if i % 2000 == 1999:    # 每2000个小批次打印一次
            print('[%d, %5d] loss: %.3f' % (epoch + 1, i + 1, running_loss / 2000))
            running_loss = 0.0

print('Finished Training')

4. 实用工具函数

4.1 模型参数统计

类似于Keras的model.summary()功能,可以显示各层参数信息:

import collections
import torch

def paras_summary(input_size, model):
    # 实现各层参数统计的函数
    # 可以显示每层的输入输出形状、参数数量等信息
    # 具体实现见原始代码
    pass

5. 关键知识点总结

5.1 模型设计要点

  1. 卷积层配置:合理设置输入输出通道数、卷积核大小和步长

  2. 池化层选择:最大池化、平均池化或全局池化的应用场景

  3. 全连接层设计:注意输入维度与前一层的输出维度匹配

  4. 激活函数:ReLU是最常用的激活函数

5.2 训练技巧

  1. 学习率设置:初始学习率不宜过大,可使用学习率调度器

  2. 动量优化:momentum参数可以加速收敛并减少震荡

  3. 批次大小:根据GPU内存选择合适的大小

  4. 数据增强:提升模型泛化能力的重要手段

5.3 注意事项

  1. 训练前务必进行梯度清零(optimizer.zero_grad()

  2. 合理设置验证频率,避免过拟合

  3. 使用GPU加速训练时,确保数据和模型都在同一设备上

  4. 保存最佳模型权重,便于后续使用

   通过以上步骤,我们可以构建一个完整的图像分类流程。实际应用中,还需要根据具体任务调整网络结构、超参数和数据预处理方法。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐