/
grenki70
/
siames
Обзор
Документация
Войти
/
grenki70
/
siames
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
trainer.py
209 строк
7 KB
grenki70
first_commit
28 май 2026, 06:18
28 май 2026, 06:18
f3a78a4
Код
Авторство
О чём код?
import torch import torch.nn as nn import torch.optim as optim import matplotlib.pyplot as plt import mlflow import mlflow.pytorch from sklearn.metrics import f1_score from torch.utils.data import DataLoader, random_split from config import Config from dataset import SiameseAudioDataset from model import SiameseNetwork class Trainer: def __init__(self, cfg: Config): self.cfg = cfg self.device = cfg.device self.model = SiameseNetwork( cfg.audio.n_mfcc, cfg.max_audio_len, cfg.training.dropout ).to(self.device) self.criterion = nn.CrossEntropyLoss() self.optimizer = optim.Adam( self.model.parameters(), lr=cfg.training.learning_rate, weight_decay=cfg.training.weight_decay, ) self.train_losses = [] self.train_accuracies = [] self.val_losses = [] self.val_accuracies = [] self.val_f1_scores = [] def prepare_data(self): full_dataset = SiameseAudioDataset(self.cfg.paths.dataset, self.cfg) if len(full_dataset) == 0: raise ValueError("Набор данных пустой") train_size = int(self.cfg.training.train_ratio * len(full_dataset)) val_size = len(full_dataset) - train_size train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size]) self.train_loader = DataLoader( train_dataset, batch_size=self.cfg.training.batch_size, shuffle=True, ) self.val_loader = DataLoader( val_dataset, batch_size=self.cfg.training.batch_size, shuffle=False, ) print("\n" + "=" * 30) print(f"Пары: {len(full_dataset)}") print(f"Обуч. выборка: {len(train_dataset)} пар ({self.cfg.training.train_ratio * 100:.0f}%)") print(f"Вал. Выборка: {len(val_dataset)} пар ({self.cfg.training.val_ratio * 100:.0f}%)") print("=" * 30) def _train_epoch(self): self.model.train() running_loss = 0.0 train_correct = 0 train_total = 0 for audio1, audio2, labels in self.train_loader: audio1 = audio1.to(self.device) audio2 = audio2.to(self.device) labels = labels.to(self.device) self.optimizer.zero_grad() outputs = self.model(audio1, audio2) loss = self.criterion(outputs, labels) loss.backward() self.optimizer.step() running_loss += loss.item() _, predicted = torch.max(outputs.data, 1) train_total += labels.size(0) train_correct += (predicted == labels).sum().item() avg_loss = running_loss / len(self.train_loader) avg_acc = train_correct / train_total return avg_loss, avg_acc def _validate_epoch(self): self.model.eval() running_loss = 0.0 correct = 0 total = 0 all_preds = [] all_labels = [] with torch.no_grad(): for audio1, audio2, labels in self.val_loader: audio1 = audio1.to(self.device) audio2 = audio2.to(self.device) labels = labels.to(self.device) outputs = self.model(audio1, audio2) loss = self.criterion(outputs, labels) running_loss += loss.item() _, predicted = torch.max(outputs.data, 1) total += labels.size(0) correct += (predicted == labels).sum().item() all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) avg_loss = running_loss / len(self.val_loader) acc = correct / total if total > 0 else 0 f1 = f1_score(all_labels, all_preds, average='macro') return avg_loss, acc, f1 def train(self): self.prepare_data() mlflow.set_tracking_uri(self.cfg.paths.mlflow_tracking_uri) mlflow.set_experiment(self.cfg.paths.mlflow_experiment) with mlflow.start_run() as run: # Логируем гиперпараметры mlflow.log_params({ "epochs": self.cfg.training.epochs, "learning_rate": self.cfg.training.learning_rate, "batch_size": self.cfg.training.batch_size, "sample_rate": self.cfg.audio.sample_rate, "n_mfcc": self.cfg.audio.n_mfcc, "max_audio_len": self.cfg.max_audio_len, "train_ratio": self.cfg.training.train_ratio, }) for epoch in range(self.cfg.training.epochs): train_loss, train_acc = self._train_epoch() val_loss, val_acc, val_f1 = self._validate_epoch() self.train_losses.append(train_loss) self.train_accuracies.append(train_acc) self.val_losses.append(val_loss) self.val_accuracies.append(val_acc) self.val_f1_scores.append(val_f1) # Логируем метрики в MLflow mlflow.log_metrics({ "train_loss": train_loss, "train_acc": train_acc, "val_loss": val_loss, "val_acc": val_acc, "val_f1": val_f1, }, step=epoch) print( f'Epoch [{epoch + 1}/{self.cfg.training.epochs}] | ' f'Loss: {train_loss:.4f} | ' f'Val Loss: {val_loss:.4f} | ' f'Train Acc: {train_acc:.4f} | ' f'Val Acc: {val_acc:.4f} | ' f'F1: {val_f1:.4f}' ) # Сохраняем модель в MLflow mlflow.pytorch.log_model(self.model, "siamese_model") # Строим и логируем графики self._plot_metrics() mlflow.log_artifact("training_plots.png") print(f"\nMLflow Run ID: {run.info.run_id}") def _plot_metrics(self): plt.figure(figsize=(15, 15)) plt.subplot(3, 2, 1) plt.plot(self.train_accuracies, label='Train Accuracy', color='cyan') plt.title('Train acc') plt.xlabel('Epochs') plt.ylabel('Accuracy') plt.grid(True) plt.legend() plt.subplot(3, 2, 2) plt.plot(self.val_accuracies, label='Val Accuracy', color='green') plt.title('Val Accuracy') plt.xlabel('Epochs') plt.ylabel('Accuracy') plt.grid(True) plt.legend() plt.subplot(3, 2, 3) plt.plot(self.train_losses, label='Train Loss', color='blue') plt.title('Train Loss') plt.xlabel('Epochs') plt.ylabel('Loss') plt.grid(True) plt.legend() plt.subplot(3, 2, 4) plt.plot(self.val_losses, label='Val Loss', color='orange') plt.title('Val Loss') plt.xlabel('Epochs') plt.ylabel('Loss') plt.grid(True) plt.legend() plt.tight_layout() plt.savefig("training_plots.png", dpi=150) plt.show()