ComDesignProject/train.py

71 lines
2.5 KiB
Python

import numpy as np
from model import my_model, judge_loss
from loaders import img_loader
import torch
from torchvision import transforms
from torch.utils.data import DataLoader
device = "cuda" if torch.cuda.is_available() else "cpu"
def train(dataloader: DataLoader, model, loss_fn:judge_loss.Final_loss, optimizer,loss_count,index=1,mini_batches=1,):
size = len(dataloader.dataset)
model.train()
optimizer.zero_grad()
q_list = list(dataloader.dataset.query_dict.keys())
q_dict = dataloader.dataset.query_dict
for batch, (X, index) in enumerate(dataloader):
if batch% 2:
continue
X = X.to(device)
# Compute prediction error
pred = model(X)
for q_index in q_list:
q_tensor: torch.Tensor = q_dict[q_index]
q_id = model(q_tensor.unsqueeze(dim=0)).repeat(len(index),1,1,1).to(device)
loss = loss_fn(pred, q_id, q_index == index)
loss.backward(retain_graph=True)
optimizer.step()
optimizer.zero_grad()
loss, current = loss.item(), (batch+1) * len(X)
loss = loss / 128
loss_count.append(loss)
print(f"[{current:>5d}/{size:>5d}]: loss={loss:>7f} ")
return loss_count
if __name__ == '__main__':
print(torch.__version__)
print(device)
loss_count = []
raw_transformer = transforms.Compose([
transforms.ToPILImage(),
transforms.Resize((128,64)),# hxw
transforms.ToTensor(),
])
model = my_model.ReIDNet()
loss_fn = judge_loss.Final_loss().to(device)
for i in range(1,6+1):
dataset = img_loader.Dataset(gallarys_path="mydataset/gallery",query_path=f"mydataset/query/cam{i}",raw_data_transform=raw_transformer)
my_loader = DataLoader(dataset, batch_size=128,num_workers=0,shuffle=True,pin_memory=True)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4,betas=(0.9,0.99),eps=1e-05)
enpoch=4
for i in range(enpoch):
print(f"--------------------epoch:{i}----------------------")
# model.load_state_dict(torch.load("model.pth"))
model = model.to(device)
train(my_loader,model,loss_fn,optimizer,loss_count)
torch.save(model.state_dict(), f"./model{i}.pth")
print(f"----end----------end-----------end--------end------\n")
Loss = np.array(loss_count)
np.save('./loss/cam_{}_epoch_{}'.format(i,enpoch), Loss)