Spaces:
Runtime error
Runtime error
| 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}, | |
| } |