Files
SortMeow/nn/train/train_resnet18.py
T
2026-07-09 17:48:39 +08:00

51 lines
1.4 KiB
Python

import os
import torch
from nn.nets.resnet18 import resnet18
from nn.dataloader.dataloader import train_dataloader
from nn.const import epoch, lr, batch_size
def train():
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device: ", device)
net = resnet18().to(device)
loss_func = torch.nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6) # 余弦退火
for e in range(epoch):
print("epoch: ", e)
net.train()
for i, data in enumerate(train_dataloader):
inputs, labels = data
inputs, labels = inputs.to(device), labels.to(device)
outputs = net(inputs)
loss = loss_func(outputs, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
_, pred = torch.max(outputs, dim=1)
correct = pred.eq(labels.data).cpu().sum()
print("step: ", i, "loss: ", loss.item(), "correct: ", 1.0 * correct / batch_size)
scheduler.step()
print("lr: ", optimizer.state_dict()['param_groups'][0]['lr'])
if not os.path.exists("./models"):
os.makedirs("./models")
torch.save(net.state_dict(), "./models/resnet18_epoch_{}.pth".format(e + 1))
if __name__ == "__main__":
train()