import os import random import warnings from pathlib import Path from PIL import Image import pandas as pd import numpy as np from tqdm import tqdm import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader import torchvision.transforms as T import timm import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix # Подавление предупреждений цветовых профилей warnings.filterwarnings("ignore", message=".*Unknown Adobe color transform code.*") # Настройки окружения DATA_ROOT = Path("./NFS/Thesis/Emoset/EmoSet-118K") # ВАЖНО: Добавили путь для медиа файлов MEDIA_DIR = Path("./src/scripts/media") MEDIA_DIR.mkdir(parents=True, exist_ok=True) BATCH_SIZE = 64 EPOCHS = 30 LR = 5e-5 NUM_WORKERS = 32 PATIENCE = 7 # Маппинг классов CLASS_MAPPING = { "amusement": 0, "anger": 1, "awe": 2, "contentment": 3, "disgust": 4, "excitement": 5, "fear": 6, "sadness": 7 } # Инвертированный маппинг для графиков INV_CLASS_MAPPING = {v: k for k, v in CLASS_MAPPING.items()} CLASS_NAMES = [INV_CLASS_MAPPING[i] for i in range(len(CLASS_MAPPING))] DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Устройство: {DEVICE}") # Фиксация генераторов псевдослучайных чисел def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed() # Инициализация структур данных class EmoSetDataset(Dataset): def __init__(self, root: Path | str, split: str, transform=None): self.root = Path(root) / split self.df = pd.read_csv(self.root / "labels.csv") self.transform = transform # Фильтрация датафрейма self.df = self.df[self.df["label"].isin(CLASS_MAPPING.keys())].reset_index(drop=True) def __len__(self): return len(self.df) def __getitem__(self, idx): row = self.df.iloc[idx] img_path = self.root / "images" / row["filename"] try: img = Image.open(img_path).convert("RGB") except Exception: img = Image.new("RGB", (256, 256), (0, 0, 0)) if self.transform: img_tensor = self.transform(img) else: img_tensor = T.ToTensor()(img) label_idx = CLASS_MAPPING[row["label"]] return img_tensor, label_idx # Трансформации base_tf = [ T.ToTensor(), T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ] train_transform = T.Compose([ T.Resize(256, antialias=True), T.RandomCrop(224), T.RandomHorizontalFlip(), *base_tf ]) val_transform = T.Compose([ T.Resize(256, antialias=True), T.CenterCrop(224), *base_tf ]) train_ds = EmoSetDataset(DATA_ROOT, "train", transform=train_transform) val_ds = EmoSetDataset(DATA_ROOT, "val", transform=val_transform) train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True) # Инициализация модели и оптимизатора model = timm.create_model("resnet50", pretrained=True, num_classes=8, drop_rate=0.3) model.to(DEVICE) criterion = nn.CrossEntropyLoss(label_smoothing=0.1) optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-3) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS) # Функции для отрисовки графиков def plot_learning_curves(history): """Отрисовка графиков функции потерь и точности""" epochs = range(1, len(history['train_loss']) + 1) plt.figure(figsize=(14, 5)) # График Loss plt.subplot(1, 2, 1) plt.plot(epochs, history['train_loss'], 'b-', label='Train Loss') plt.plot(epochs, history['val_loss'], 'r--', label='Validation Loss') plt.title('График функции потерь (Loss)', fontsize=14) plt.xlabel('Эпохи', fontsize=12) plt.ylabel('Loss', fontsize=12) plt.legend() plt.grid(True, linestyle=':', alpha=0.7) # График Accuracy plt.subplot(1, 2, 2) plt.plot(epochs, history['train_acc'], 'b-', label='Train Accuracy') plt.plot(epochs, history['val_acc'], 'r--', label='Validation Accuracy') plt.title('График точности (Accuracy)', fontsize=14) plt.xlabel('Эпохи', fontsize=12) plt.ylabel('Accuracy', fontsize=12) plt.legend() plt.grid(True, linestyle=':', alpha=0.7) plt.tight_layout() plot_path = MEDIA_DIR / "training_history.png" plt.savefig(plot_path, dpi=300, bbox_inches='tight') plt.close() print(f"[INFO] График обучения сохранен в: {plot_path}") def plot_confusion_matrix(y_true, y_pred): """Отрисовка тепловой матрицы ошибок""" cm = confusion_matrix(y_true, y_pred) plt.figure(figsize=(10, 8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES, cbar_kws={'label': 'Количество сэмплов'}) plt.title('Матрица ошибок (Confusion Matrix) - ResNet50', fontsize=16, pad=20) plt.ylabel('Истинные классы (Ground Truth)', fontsize=12) plt.xlabel('Предсказанные классы (Predicted)', fontsize=12) plt.xticks(rotation=45, ha='right') plt.yticks(rotation=0) plt.tight_layout() cm_path = MEDIA_DIR / "confusion_matrix_emoset.png" plt.savefig(cm_path, dpi=300, bbox_inches='tight') plt.close() print(f"[INFO] Матрица ошибок сохранена в: {cm_path}") # Логика эпохи обучения def train_epoch(current_model, loader): current_model.train() total_loss, correct_preds, total_samples = 0.0, 0, 0 for imgs, labels in tqdm(loader, desc="Тренировка", leave=False, smoothing=0): imgs, labels = imgs.to(DEVICE), labels.to(DEVICE) optimizer.zero_grad(set_to_none=True) logits = current_model(imgs) loss = criterion(logits, labels) loss.backward() optimizer.step() total_loss += loss.item() * imgs.size(0) preds = logits.argmax(dim=1) correct_preds += (preds == labels).sum().item() total_samples += labels.size(0) return total_loss / total_samples, correct_preds / total_samples # Логика эпохи валидации с сохранением предсказаний для матрицы ошибок @torch.no_grad() def val_epoch(current_model, loader, return_preds=False): current_model.eval() total_loss, correct_preds, total_samples = 0.0, 0, 0 all_preds, all_labels = [], [] for imgs, labels in tqdm(loader, desc="Валидация", leave=False, smoothing=0): imgs, labels = imgs.to(DEVICE), labels.to(DEVICE) logits = current_model(imgs) loss = criterion(logits, labels) total_loss += loss.item() * imgs.size(0) preds = logits.argmax(dim=1) correct_preds += (preds == labels).sum().item() total_samples += labels.size(0) if return_preds: all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) avg_loss = total_loss / total_samples avg_acc = correct_preds / total_samples if return_preds: return avg_loss, avg_acc, all_labels, all_preds return avg_loss, avg_acc if __name__ == "__main__": best_val_acc = 0.0 best_val_loss = float('inf') epochs_no_improve = 0 checkpoint_path = "./emosetV2_resnet50_best.pth" # Словарь для хранения истории обучения history = { 'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': [] } # Переменные для хранения лучших предсказаний для матрицы best_labels, best_preds = [], [] print("Старт обучения.") for epoch in range(1, EPOCHS + 1): train_loss, train_acc = train_epoch(model, train_loader) # Получаем предсказания только если это может быть лучшая эпоха val_loss, val_acc, val_labels, val_preds = val_epoch(model, val_loader, return_preds=True) scheduler.step() # Запись в историю history['train_loss'].append(train_loss) history['train_acc'].append(train_acc) history['val_loss'].append(val_loss) history['val_acc'].append(val_acc) print(f"[{epoch}/{EPOCHS}] Train Loss: {train_loss:.4f}, Acc: {train_acc:.4f} | Val Loss: {val_loss:.4f}, Acc: {val_acc:.4f}") # Сохранение лучших весов по Accuracy if val_acc > best_val_acc: best_val_acc = val_acc best_labels = val_labels # Сохраняем предсказания лучшей модели best_preds = val_preds torch.save(model.state_dict(), checkpoint_path) print(f"Сохранен чекпоинт (Acc: {best_val_acc:.4f})") # Оценка переобучения по Loss (Early Stopping) if val_loss < best_val_loss: best_val_loss = val_loss epochs_no_improve = 0 else: epochs_no_improve += 1 if epochs_no_improve >= PATIENCE: print(f"Ранняя остановка: метрика валидации не улучшается {PATIENCE} эпох.") break print("Процесс обучения завершен. Генерирую графики для диссертации...") plot_learning_curves(history) plot_confusion_matrix(best_labels, best_preds) print("Все медиафайлы успешно созданы!")