import os, sys, json, time, itertools, math import numpy as np 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 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_epoch(model, loader, criterion, optimizer, scaler, device): model.train() total_loss = 0 for x, y in 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() return total_loss / len(loader) def train(config): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f'\nDevice: {device}') print(f'Config: {json.dumps(config, indent=2)}') sequences, labels = load_genpept_data() train_loader, val_loader, test_loader = get_dataloaders( sequences, labels, batch_size=config['batch_size'], max_len=config['max_len'] ) model = PeptEdgeV2( vocab_size=21, max_len=config['max_len'], d_model=config['d_model'], n_heads=config['n_heads'], num_layers=config['num_layers'], ff_dim=config['ff_dim'], num_classes=2, dropout=config['dropout'], sd_prob=config['sd_prob'], ).to(device) total_params = count_parameters(model) print(f'Params: {total_params:,}') criterion = nn.CrossEntropyLoss(label_smoothing=config['label_smoothing']) optimizer = optim.AdamW(model.parameters(), lr=config['lr'], weight_decay=config['weight_decay']) warmup_steps = config.get('warmup', 10) total_steps = config['epochs'] def lr_lambda(step): if step < warmup_steps: return step / warmup_steps return 0.5 * (1 + math.cos(math.pi * (step - warmup_steps) / (total_steps - warmup_steps))) scheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) scaler = GradScaler('cuda') best_val_f1 = 0 best_state = None patience_counter = 0 history = [] print(f'\n{"Ep":>3} | {"Loss":>7} | {"Acc":>6} | {"F1":>6} | {"AUC":>6} | {"Sp":>6} | {"MCC":>6} | {"Best":>4} | {"LR":>8}') print('-' * 65) for epoch in range(config['epochs']): train_loss = train_epoch(model, train_loader, criterion, optimizer, scaler, device) 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() patience_counter = 0 else: patience_counter += 1 history.append({'epoch': epoch+1, 'train_loss': train_loss, **val_metrics}) if epoch == 0 or (epoch+1) % 2 == 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["specificity"]:>6.4f} | {val_metrics["mcc"]:>6.4f} | ' f'{"*BEST" if is_best else "":>4} | {scheduler.get_last_lr()[0]:>8.2e}') if patience_counter >= config['patience']: 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:,}') os.makedirs('results', exist_ok=True) torch.save({ 'model_state_dict': best_state, 'config': config, 'test_metrics': test_metrics, 'total_params': total_params, }, 'results/peptedge_v2.pt') with open('results/training_v2_results.json', 'w') as f: json.dump({'config': config, 'total_params': total_params, 'best_val_f1': best_val_f1, 'test_metrics': test_metrics, 'history': history}, f, indent=2, default=str) return test_metrics, total_params def run_grid_search(): base_config = { 'batch_size': 64, 'epochs': 100, 'patience': 25, 'max_len': 200, 'label_smoothing': 0.1, 'sd_prob': 0.05, 'warmup': 10, } grid = { 'd_model': [192], 'n_heads': [6], 'num_layers': [4, 5], 'ff_dim': [384], 'dropout': [0.15, 0.25], 'lr': [3e-4, 5e-4], 'weight_decay': [1e-4, 5e-5], } keys = list(grid.keys()) best_f1 = 0 best_result = None for values in itertools.product(*grid.values()): config = {**base_config, **dict(zip(keys, values))} print(f'\n{"#"*60}') print(f'RUNNING CONFIG: {dict(zip(keys, values))}') print(f'{"#"*60}') try: metrics, params = train(config) f1 = metrics['f1'] print(f'RESULT: F1={f1:.4f}, AUC={metrics["auc"]:.4f}') if f1 > best_f1: best_f1 = f1 best_result = (config, metrics, params) except Exception as e: print(f'ERROR: {e}') continue print(f'\n{"="*60}') print(f'BEST RESULT: F1={best_f1:.4f}') print(f'Config: {json.dumps(best_result[0], indent=2)}') print(f'Metrics: {json.dumps(best_result[1], indent=2)}') print(f'Params: {best_result[2]:,}') print(f'{"="*60}') return best_result if __name__ == '__main__': run_grid_search()