283 lines
10 KiB
Python
283 lines
10 KiB
Python
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("Все медиафайлы успешно созданы!") |