AI-Project/cnn_train.py

166 lines
5.7 KiB
Python
Raw Permalink Normal View History

2024-07-03 00:42:21 +02:00
from datetime import datetime
import numpy.random
import torch.utils.data
import torch.cuda
2024-07-21 17:36:22 +02:00
from architecture import model
2024-07-03 00:42:21 +02:00
from dataset import ImagesDataset
from AImageDataset import AImagesDataset
2024-07-06 22:27:09 +02:00
from AsyncDataLoader import AsyncDataLoader
2024-07-03 00:42:21 +02:00
def split_data(data: ImagesDataset):
class_nums = (torch.bincount(
torch.Tensor([data.classnames_to_ids[name] for _, name in data.filenames_classnames]).long())
.tolist())
indices = ([], [])
for class_id in range(len(class_nums)):
class_num = class_nums[class_id]
index_perm = torch.randperm(class_num).tolist()
class_indices = [i for i, e in enumerate(data.filenames_classnames) if data.classnames_to_ids[e[1]] == class_id]
indices[0].extend([class_indices[i] for i in index_perm[0::2]])
indices[1].extend([class_indices[i] for i in index_perm[1::2]])
return (torch.utils.data.Subset(data, indices[0]),
torch.utils.data.Subset(data, indices[1]))
2024-07-07 11:56:19 +02:00
2024-07-06 18:20:21 +02:00
def train_model(accuracies,
losses,
progress_epoch,
progress_train_data,
progress_eval_data,
model,
num_epochs,
batch_size,
optimizer,
loss_function,
2024-07-07 11:56:19 +02:00
augment_data,
2024-07-06 18:20:21 +02:00
device,
start_time):
2024-07-03 00:42:21 +02:00
torch.random.manual_seed(42)
numpy.random.seed(42)
2024-07-06 18:20:21 +02:00
torch.multiprocessing.set_start_method('spawn', force=True)
2024-07-03 00:42:21 +02:00
dataset = ImagesDataset("training_data")
2024-07-07 11:56:19 +02:00
train_data, eval_data = split_data(dataset)
2024-07-03 00:42:21 +02:00
2024-07-07 11:56:19 +02:00
augmented_train_data = AImagesDataset(train_data, augment_data)
2024-07-06 22:27:09 +02:00
train_loader = AsyncDataLoader(augmented_train_data,
batch_size=batch_size,
num_workers=3,
pin_memory=True,
shuffle=True)
eval_loader = AsyncDataLoader(eval_data,
batch_size=batch_size,
num_workers=3,
pin_memory=True)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
for epoch in progress_epoch.range(num_epochs):
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
train_positives = torch.tensor(0, device=device)
eval_positives = torch.tensor(0, device=device)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
train_loss = torch.tensor(0.0, device=device)
eval_loss = torch.tensor(0.0, device=device)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
# Start training of model
2024-07-03 00:42:21 +02:00
progress_train_data.reset()
model.train()
2024-07-06 18:20:21 +02:00
for batch_nr, (image_t, transforms, img_index, class_ids, labels, paths) \
in enumerate(progress_train_data.iter(train_loader)):
2024-07-06 22:27:09 +02:00
2024-07-06 18:20:21 +02:00
image_t = image_t.to(device)
class_ids = class_ids.to(device)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
outputs = model(image_t)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
optimizer.zero_grad(set_to_none=True)
loss = loss_function(outputs, class_ids)
2024-07-03 00:42:21 +02:00
loss.backward()
optimizer.step()
2024-07-06 18:20:21 +02:00
train_loss += loss
classes = outputs.argmax(dim=1)
2024-07-06 18:20:21 +02:00
train_positives += torch.sum(torch.eq(classes, class_ids))
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
accuracies.append('train_acc', train_positives.item() / len(augmented_train_data))
losses.append('train_loss', train_loss.item() / len(augmented_train_data))
print("Train: ", train_positives.item(), "/ ", len(augmented_train_data),
" = ", train_positives.item() / len(augmented_train_data))
2024-07-03 00:42:21 +02:00
# evaluation of model
2024-07-06 18:20:21 +02:00
progress_eval_data.reset()
2024-07-03 00:42:21 +02:00
model.eval()
with torch.no_grad():
2024-07-06 18:20:21 +02:00
for (image_t, class_ids, labels, paths) in progress_eval_data.iter(eval_loader):
image_t = image_t.to(device)
class_ids = class_ids.to(device)
outputs = model(image_t)
classes = outputs.argmax(dim=1)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
eval_positives += torch.sum(torch.eq(classes, class_ids))
eval_loss += loss_function(outputs, class_ids)
2024-07-03 00:42:21 +02:00
2024-07-06 18:20:21 +02:00
accuracies.append('eval_acc', eval_positives.item() / len(eval_data))
losses.append('eval_loss', eval_loss.item() / len(eval_data))
print("Eval: ", eval_positives.item(), "/ ", len(eval_data), " = ", eval_positives.item() / len(eval_data))
2024-07-03 00:42:21 +02:00
if eval_positives.item() / len(eval_data) > 0.5:
2024-07-21 17:36:22 +02:00
torch.save(model.state_dict(), f'models/model-{start_time.strftime("%Y%m%d-%H%M%S")}-epoch-{epoch}.pth')
2024-07-06 18:20:21 +02:00
with open(f'models/model-{start_time.strftime("%Y%m%d-%H%M%S")}.csv', 'a') as file:
file.write(f'{epoch};{len(augmented_train_data)};{len(eval_data)};{train_loss.item()};{eval_loss.item()};'
f'{train_positives};{eval_positives}\n')
2024-07-07 11:56:19 +02:00
def train_worker(p_epoch, p_train, p_eval, plotter_accuracies, plotter_loss, start_time):
2024-07-06 18:20:21 +02:00
if not torch.cuda.is_available():
raise RuntimeError("GPU not available")
device = 'cuda'
2024-07-21 17:36:22 +02:00
model.to(device)
2024-07-06 18:20:21 +02:00
num_epochs = 1000000
batch_size = 64
optimizer = torch.optim.Adam(model.parameters(),
2024-07-07 11:56:19 +02:00
lr=0.0001,
2024-07-06 18:20:21 +02:00
fused=True)
loss_function = torch.nn.CrossEntropyLoss()
augment_data = True
2024-07-06 18:20:21 +02:00
file_name = f'models/model-{start_time.strftime("%Y%m%d-%H%M%S")}.csv'
with open(file_name.replace(".csv", ".txt"), 'a') as file:
file.write(f"device: {device}\n")
file.write(f"batch_size: {batch_size}\n")
file.write(f"optimizer: {optimizer}\n")
file.write(f"loss_function: {loss_function}\n")
2024-07-07 11:56:19 +02:00
file.write(f"augment_data: {augment_data}\n")
2024-07-06 18:20:21 +02:00
file.write(f"model: {model}")
train_model(plotter_accuracies, plotter_loss, p_epoch, p_train, p_eval,
model,
num_epochs,
batch_size,
optimizer,
loss_function,
2024-07-07 11:56:19 +02:00
augment_data,
2024-07-06 18:20:21 +02:00
device,
start_time)