PeptEdgeV2 / train.py
devansh0703's picture
Initial release: PeptEdgeV2 (3.43M params) w/ trained weights, config, source, model card
9f16c4c verified
Raw
History Blame Contribute Delete
7.52 kB
import os
import sys
import time
import json
import numpy as np
import torch
import torch.nn as nn
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,
confusion_matrix)
from collections import Counter
from model import PeptEdge, count_parameters
from data_utils import load_genpept_data, get_dataloaders
def evaluate(model, loader, device, multilabel=False):
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)
logits = model(x)
if multilabel:
probs = torch.sigmoid(logits)
preds = (probs > 0.5).long()
else:
probs = torch.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()
if multilabel:
from sklearn.metrics import average_precision_score
ap_per_class = []
for i in range(probs.shape[1]):
ap_per_class.append(average_precision_score(labels[:, i], probs[:, i]))
mAP = np.mean(ap_per_class)
f1 = f1_score(labels, preds, average='macro', zero_division=0)
acc = accuracy_score(labels.flatten(), preds.flatten())
return {
'accuracy': float(acc),
'f1_macro': float(f1),
'mAP': float(mAP),
}
else:
if probs.shape[1] == 2:
auc = roc_auc_score(labels, probs[:, 1])
else:
auc = roc_auc_score(labels, probs, multi_class='ovr')
mcc = matthews_corrcoef(labels, preds)
sens = recall_score(labels, preds, pos_label=1, zero_division=0)
spec = recall_score(labels, preds, pos_label=0, zero_division=0)
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(spec),
'f1': float(f1_score(labels, preds, zero_division=0)),
'auc': float(auc),
'mcc': float(mcc),
}
def train_epoch(model, loader, criterion, optimizer, scaler, device, use_amp=True):
model.train()
total_loss = 0
for x, y in loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
if use_amp:
with autocast(device_type='cuda'):
logits = model(x)
loss = criterion(logits, y)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
else:
logits = model(x)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
total_loss += loss.item()
return total_loss / len(loader)
def train(config=None):
if config is None:
config = {
'batch_size': 64,
'lr': 5e-4,
'weight_decay': 1e-4,
'epochs': 150,
'd_model': 128,
'n_heads': 4,
'num_layers': 3,
'ff_dim': 256,
'dropout': 0.2,
'max_len': 200,
'use_amp': True,
'patience': 20,
'label_smoothing': 0.05,
}
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f'Device: {device}')
print(f'Config: {json.dumps(config, indent=2)}')
sequences, labels = load_genpept_data()
print(f'Dataset: {len(sequences)} sequences, {Counter(labels)}')
train_loader, val_loader, test_loader = get_dataloaders(
sequences, labels, batch_size=config['batch_size'], max_len=config['max_len']
)
model = PeptEdge(
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'],
).to(device)
total_params = count_parameters(model)
print(f'Model params: {total_params:,}')
criterion = nn.CrossEntropyLoss(label_smoothing=config['label_smoothing'])
optimizer = optim.AdamW(model.parameters(), lr=config['lr'],
weight_decay=config['weight_decay'])
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=config['epochs'])
scaler = GradScaler('cuda') if device.type == 'cuda' else None
best_val_f1 = 0
best_state = None
patience_counter = 0
history = []
print(f'\n{"Epoch":>5} | {"Train Loss":>10} | {"Val Acc":>8} | {"Val F1":>8} | {"Val AUC":>8} | {"Best":>5} | {"LR":>10}')
print('-' * 60)
for epoch in range(config['epochs']):
train_loss = train_epoch(model, train_loader, criterion, optimizer,
scaler, device, config['use_amp'])
val_metrics = evaluate(model, val_loader, device)
scheduler.step()
current_lr = scheduler.get_last_lr()[0]
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 + 1) % 2 == 0 or epoch == 0:
print(f'{epoch + 1:>5} | {train_loss:>10.4f} | {val_metrics["accuracy"]:>8.4f} | '
f'{val_metrics["f1"]:>8.4f} | {val_metrics["auc"]:>8.4f} | '
f'{"*" if is_best else "":>5} | {current_lr:>10.2e}')
if patience_counter >= config['patience']:
print(f'Early stopping at epoch {epoch + 1}')
break
model.load_state_dict(best_state)
test_metrics = evaluate(model, test_loader, device)
print('\n' + '=' * 50)
print('TEST SET RESULTS')
print('=' * 50)
for k, v in test_metrics.items():
print(f' {k}: {v:.4f}')
print(f' params: {total_params:,}')
results = {
'config': config,
'total_params': total_params,
'best_val_f1': best_val_f1,
'test_metrics': test_metrics,
'history': history,
}
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_genpept.pt')
with open('results/training_results.json', 'w') as f:
json.dump(results, f, indent=2, default=str)
print(f'\nModel saved to results/peptedge_genpept.pt')
print(f'Results saved to results/training_results.json')
return results
if __name__ == '__main__':
results = train()
print(f"\nSOTA Comparison:")
print(f" ESM-2 LoRA (650M params): 88.3% F1")
print(f" PeptEdge ({results['total_params']:,} params): {results['test_metrics']['f1']:.2%} F1")
improvement = results['test_metrics']['f1'] - 0.883
print(f" Δ: {improvement:+.2%}")