feat: add metrics
This commit is contained in:
@@ -0,0 +1,283 @@
|
||||
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("Все медиафайлы успешно созданы!")
|
||||
Reference in New Issue
Block a user