import json import os from datetime import datetime import torch import torch.optim as optim import os import json from datetime import datetime import torch from sklearn.metrics import ( accuracy_score, precision_score, recall_score, f1_score, roc_auc_score ) def evaluate_metrics(model, dataloader, device, threshold=0.5): """Run model on a dataloader and compute classification metrics.""" model.eval() all_logits = [] all_targets = [] with torch.no_grad(): for batch in dataloader: x_batch, y_batch = batch[0].to(device), batch[1].to(device) logits = model(x_batch) all_logits.append(logits.cpu()) all_targets.append(y_batch.cpu()) logits = torch.cat(all_logits) targets = torch.cat(all_targets) probs = torch.sigmoid(logits).numpy() preds = (probs > threshold).astype(int) targets_np = targets.numpy().astype(int) # average='macro' works for both binary and multi-label; # switch to 'binary' if you have a single output column and want binary-specific behavior avg = "macro" if targets_np.ndim > 1 and targets_np.shape[1] > 1 else "binary" metrics = { "accuracy": float(accuracy_score(targets_np, preds)), "precision": float(precision_score(targets_np, preds, average=avg, zero_division=0)), "recall": float(recall_score(targets_np, preds, average=avg, zero_division=0)), "f1": float(f1_score(targets_np, preds, average=avg, zero_division=0)), } return metrics def train_mlp(model, train_loader, optimizer, criterion, num_epochs=300, save_dir="checkpoints", device="cpu", patience=20, val_loader=None): os.makedirs(save_dir, exist_ok=True) hyperparams = { "num_epochs": num_epochs, "learning_rate": optimizer.param_groups[0]["lr"], "optimizer": optimizer.__class__.__name__, "criterion": criterion.__class__.__name__, "batch_size": train_loader.batch_size, "model_class": model.__class__.__name__, "device": str(device), "patience": patience, "timestamp": datetime.now().isoformat(), } train_loss_history = [] val_loss_history = [] best_loss = float("inf") epochs_without_improvement = 0 model.to(device) for epoch in range(num_epochs): # ---------- Training phase ---------- model.train() total_loss = 0.0 for batch in train_loader: x_batch, y_batch = batch[0].to(device), batch[1].to(device) optimizer.zero_grad() logits = model(x_batch) loss = criterion(logits, y_batch) loss.backward() optimizer.step() total_loss += loss.item() avg_train_loss = total_loss / len(train_loader) train_loss_history.append(avg_train_loss) # ---------- Validation phase ---------- avg_val_loss = None if val_loader is not None: model.eval() val_loss = 0.0 with torch.no_grad(): for batch in val_loader: x_batch, y_batch = batch[0].to(device), batch[1].to(device) logits = model(x_batch) val_loss += criterion(logits, y_batch).item() avg_val_loss = val_loss / len(val_loader) val_loss_history.append(avg_val_loss) # ---------- Checkpointing ---------- monitor_loss = avg_val_loss if avg_val_loss is not None else avg_train_loss if monitor_loss < best_loss: best_loss = monitor_loss epochs_without_improvement = 0 torch.save({ "epoch": epoch + 1, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "loss": best_loss, "architecture": { "input_size": model.input_size, "output_size": model.output_size, "n_neurons": model.n_neurons, "dropout_rates": model.dropout_rates, }, "hyperparams": hyperparams, }, os.path.join(save_dir, "best_model.pt")) else: epochs_without_improvement += 1 # ---------- Logging ---------- log_msg = f"Epoch [{epoch+1}/{num_epochs}] | Train Loss: {avg_train_loss:.6f}" if avg_val_loss is not None: log_msg += f" | Val Loss: {avg_val_loss:.6f}" log_msg += f" | Best: {best_loss:.6f}" print(log_msg) # ---------- Early stopping ---------- if epochs_without_improvement >= patience: print(f"Early stopping at epoch {epoch+1} " f"(no improvement for {patience} epochs)") break # ---------- Final save ---------- torch.save({ "epoch": epoch + 1, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "loss": avg_train_loss if avg_val_loss is None else avg_val_loss, "architecture": { "input_size": model.input_size, "output_size": model.output_size, "n_neurons": model.n_neurons, "dropout_rates": model.dropout_rates, }, "hyperparams": hyperparams, }, os.path.join(save_dir, "final_model.pt")) # ---------- Final evaluation: load best model, compute all metrics ---------- print("\nLoading best model for final evaluation...") checkpoint = torch.load(os.path.join(save_dir, "best_model.pt"), map_location=device) model.load_state_dict(checkpoint["model_state_dict"]) print("\nComputing metrics on training set...") train_metrics = evaluate_metrics(model, train_loader, device) print(f" Train: {train_metrics}") val_metrics = None if val_loader is not None: print("\nComputing metrics on validation set...") val_metrics = evaluate_metrics(model, val_loader, device) print(f" Val: {val_metrics}") # ---------- JSON log ---------- log_data = { "hyperparams": hyperparams, "train_loss_history": train_loss_history, "val_loss_history": val_loss_history, "final_metrics": { "train": train_metrics, "val": val_metrics, }, } with open(os.path.join(save_dir, "training_log.json"), "w") as f: json.dump(log_data, f, indent=2) print(f"\nTraining complete. Best loss: {best_loss:.6f}") print(f"Checkpoints saved to {save_dir}/") return { "train_loss_history": train_loss_history, "val_loss_history": val_loss_history, "final_metrics": {"train": train_metrics, "val": val_metrics}, }