PeptEdgeV2 / train_final.py
devansh0703's picture
Initial release: PeptEdgeV2 (3.43M params) w/ trained weights, config, source, model card
9f16c4c verified
Raw
History Blame Contribute Delete
5.53 kB
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%}')