/
akp1n
/
ADVML
Обзор
Документация
Войти
/
akp1n
/
ADVML
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
task1/functions.py
132 строки
4 KB
Artemiy
fix
24 фев 2025, 20:45
24 фев 2025, 20:45
57cd6cd
Код
Авторство
О чём код?
import torch import torch.nn as nn import torch.nn.functional as F import matplotlib.pyplot as plt import numpy as np import os from torch.utils.data import DataLoader def test_model(model, test_loader, device): model.eval() all_preds = [] all_labels = [] all_images = [] with torch.no_grad(): for images, labels in test_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) preds = torch.argmax(F.softmax(outputs, dim=1), dim=1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) all_images.extend(images.cpu().numpy()) return all_preds, all_labels, all_images def plot_results(images, labels, preds, n=16, save_path=None): plt.figure(figsize=(12, 12)) for i in range(n): plt.subplot(int(np.sqrt(n)), int(np.sqrt(n)), i + 1) image = images[i].squeeze() plt.imshow(image, cmap='gray') plt.title(f"True: {labels[i]}, Pred: {preds[i]}") plt.axis('off') plt.tight_layout() if save_path: plt.savefig(save_path) plt.close() def calculate_accuracy(labels, preds): correct = np.sum(np.array(labels) == np.array(preds)) total = len(labels) accuracy = correct / total * 100 return accuracy def train_model(model: nn.Module, train_loader: DataLoader, val_loader: DataLoader, device: torch.device, epochs: int = 10, activation_name: str = 'activation'): optimizer = torch.optim.Adam(model.parameters(), lr=0.001) criterion = nn.CrossEntropyLoss() train_losses = [] val_losses = [] batch_losses = [] batch_steps = [] steps_per_epoch = len(train_loader) step_global = 0 for epoch in range(epochs): model.train() running_loss = 0.0 for images, labels in train_loader: images = images.to(device) labels = labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() batch_losses.append(loss.item()) batch_steps.append(step_global) step_global += 1 avg_loss = running_loss / len(train_loader) train_losses.append(avg_loss) model.eval() val_loss = 0.0 correct = 0 total = 0 with torch.no_grad(): for images, labels in val_loader: images = images.to(device) labels = labels.to(device) outputs = model(images) loss = criterion(outputs, labels) val_loss += loss.item() _, preds = torch.max(outputs, 1) correct += (preds == labels).sum().item() total += labels.size(0) avg_val_loss = val_loss / len(val_loader) val_losses.append(avg_val_loss) val_accuracy = correct / total * 100 print(f"Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}, " f"Val Loss: {avg_val_loss:.4f}, Val Acc: {val_accuracy:.2f}%") fig, ax1 = plt.subplots(figsize=(8, 5)) ax1.plot(batch_steps, batch_losses, label='Train Loss per Batch') ax1.set_xlabel('Global Step') ax1.set_ylabel('Loss') ax1.legend(loc='upper right') ax2 = ax1.twiny() ax2.set_xlim(ax1.get_xlim()) epoch_ticks = [i * steps_per_epoch for i in range(epochs + 1)] epoch_labels = [str(i) for i in range(epochs + 1)] ax2.set_xticks(epoch_ticks) ax2.set_xticklabels(epoch_labels) ax2.set_xlabel('Epoch') os.makedirs('graphics', exist_ok=True) plt.title(f"Loss Curve — {activation_name}") plt.savefig(f"graphics/loss_{activation_name}.png") plt.close(fig)