import os, sys, json, math import torch import torch.nn as nn import torch.nn.functional as F import torch.optim as optim from torch.amp import autocast, GradScaler from sklearn.metrics import (accuracy_score, precision_score, recall_score, f1_score, roc_auc_score, matthews_corrcoef) from collections import Counter import numpy as np from model_v2 import PeptEdgeV2, count_parameters from data_utils import load_genpept_data, get_dataloaders def evaluate(model, loader, device): model.eval() all_preds, all_labels, all_probs = [], [], [] with torch.no_grad(): for x, y in loader: x, y = x.to(device), y.to(device) with autocast(device_type='cuda'): logits = model(x) probs = F.softmax(logits, dim=1) preds = logits.argmax(dim=1) all_preds.append(preds.cpu()) all_labels.append(y.cpu()) all_probs.append(probs.cpu()) preds = torch.cat(all_preds).numpy() labels = torch.cat(all_labels).numpy() probs = torch.cat(all_probs).numpy() return { 'accuracy': float(accuracy_score(labels, preds)), 'precision': float(precision_score(labels, preds, zero_division=0)), 'recall': float(recall_score(labels, preds, zero_division=0)), 'specificity': float(recall_score(labels, 1 - preds, zero_division=0)), 'f1': float(f1_score(labels, preds, zero_division=0)), 'auc': float(roc_auc_score(labels, probs[:, 1])), 'mcc': float(matthews_corrcoef(labels, preds)), } def train(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f'Device: {device}') sequences, labels = load_genpept_data() train_loader, val_loader, test_loader = get_dataloaders( sequences, labels, batch_size=64, max_len=200 ) model = PeptEdgeV2( vocab_size=21, max_len=200, d_model=192, n_heads=6, num_layers=5, ff_dim=384, num_classes=2, dropout=0.25, sd_prob=0.05, ).to(device) total_params = count_parameters(model) print(f'Params: {total_params:,}') criterion = nn.CrossEntropyLoss(label_smoothing=0.1) optimizer = optim.AdamW(model.parameters(), lr=3e-4, weight_decay=5e-5) warmup = 15 total_epochs = 100 def lr_lambda(step): if step < warmup: return step / warmup return 0.5 * (1 + math.cos(math.pi * (step - warmup) / (total_epochs - warmup))) scheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) scaler = GradScaler('cuda') best_val_f1 = 0 best_state = None patience_counter = 0 history = [] ckpt_dir = 'checkpoints' os.makedirs(ckpt_dir, exist_ok=True) print(f'\n{"Ep":>3} | {"Loss":>7} | {"Acc":>6} | {"F1":>6} | {"AUC":>6} | {"MCC":>6} | Best | {"LR":>8}') print('-' * 55) for epoch in range(total_epochs): model.train() total_loss = 0 for x, y in train_loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() with autocast(device_type='cuda'): logits = model(x) loss = criterion(logits, y) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() total_loss += loss.item() train_loss = total_loss / len(train_loader) val_metrics = evaluate(model, val_loader, device) scheduler.step() is_best = val_metrics['f1'] > best_val_f1 if is_best: best_val_f1 = val_metrics['f1'] best_state = model.state_dict().copy() torch.save({ 'epoch': epoch, 'model_state_dict': best_state, 'val_metrics': val_metrics, 'config': {'d_model': 192, 'n_heads': 6, 'num_layers': 5, 'ff_dim': 384, 'dropout': 0.25}, 'total_params': total_params, }, f'{ckpt_dir}/best_model.pt') patience_counter = 0 else: patience_counter += 1 history.append({'epoch': epoch+1, 'train_loss': train_loss, **val_metrics}) if epoch < 5 or (epoch+1) % 3 == 0 or is_best: print(f'{epoch+1:>3} | {train_loss:>7.4f} | {val_metrics["accuracy"]:>6.4f} | ' f'{val_metrics["f1"]:>6.4f} | {val_metrics["auc"]:>6.4f} | ' f'{val_metrics["mcc"]:>6.4f} | {"*" if is_best else " "} | {scheduler.get_last_lr()[0]:>8.2e}') if patience_counter >= 30: print(f'Early stop at epoch {epoch+1}') break model.load_state_dict(best_state) test_metrics = evaluate(model, test_loader, device) print('\n' + '='*55) print('TEST SET RESULTS') print('='*55) for k, v in test_metrics.items(): print(f' {k}: {v:.4f}') print(f' params: {total_params:,}') results = { 'test_metrics': test_metrics, 'total_params': total_params, 'best_val_f1': best_val_f1, 'history': history, } with open('results/final_results.json', 'w') as f: json.dump(results, f, indent=2, default=str) return test_metrics, total_params if __name__ == '__main__': metrics, params = train() sota_f1 = 0.883 our_f1 = metrics['f1'] print(f'\nSOTA (ESM-2 LoRA 650M): {sota_f1:.2%} F1') print(f'PeptEdgeV2 ({params:,} params): {our_f1:.2%} F1') print(f'Δ: {our_f1 - sota_f1:+.2%}')