GemmaSight / train.py
TouseefAhmad's picture
Upload folder using huggingface_hub
de2e2e5 verified
Raw
History Blame Contribute Delete
7.16 kB
import os
import sys
# Enable importing from parent package
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import json
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import TensorDataset, DataLoader
import numpy as np
from sklearn.metrics import roc_auc_score, accuracy_score, f1_score
from gemmasight.config import (
DATA_DIR, MODELS_DIR, DIM_FUSED, RANDOM_SEED,
EPOCHS, BATCH_SIZE, LEARNING_RATE, WEIGHT_DECAY,
FAISS_INDEX_PATH
)
from gemmasight.models.classifier import MSIClassifier
from gemmasight.models.retriever import CaseRetriever
# Set random seeds for reproducibility
torch.manual_seed(RANDOM_SEED)
np.random.seed(RANDOM_SEED)
def generate_synthetic_data(num_samples=100):
"""
Generates realistic synthetic 1536-dim embeddings and binary labels for simulation training.
"""
print(f"Generating {num_samples} synthetic histopathology samples...")
# Simulate a structure where some features are correlated with label 1 (MSI-High)
embeddings = np.random.randn(num_samples, DIM_FUSED).astype(np.float32)
labels = np.random.randint(0, 2, size=(num_samples,)).astype(np.float32)
# Inject signal correlating with labels to allow classifier learning
for i in range(num_samples):
if labels[i] == 1:
embeddings[i, :200] += 0.5 # Shift some features positively for MSI-High
else:
embeddings[i, :200] -= 0.5 # Shift some features negatively for MSS
# L2 normalize embeddings to match frozen encoder output format
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
norms[norms == 0] = 1e-12
embeddings = embeddings / norms
return embeddings, labels
def main():
print("=== GemmaSight Classifier Training Pipeline ===")
train_emb_path = os.path.join(DATA_DIR, "train_embeddings.npy")
train_lbl_path = os.path.join(DATA_DIR, "train_labels.npy")
val_emb_path = os.path.join(DATA_DIR, "val_embeddings.npy")
val_lbl_path = os.path.join(DATA_DIR, "val_labels.npy")
# 1. Load or Generate Dataset
if os.path.exists(train_emb_path) and os.path.exists(train_lbl_path):
print("Loading pre-extracted training embeddings from disk...")
X_train = np.load(train_emb_path)
y_train = np.load(train_lbl_path)
if os.path.exists(val_emb_path) and os.path.exists(val_lbl_path):
X_val = np.load(val_emb_path)
y_val = np.load(val_lbl_path)
else:
print("Validation files not found. Splitting training data 80/20.")
indices = np.random.permutation(len(X_train))
split_idx = int(0.8 * len(X_train))
X_val = X_train[indices[split_idx:]]
y_val = y_train[indices[split_idx:]]
X_train = X_train[indices[:split_idx]]
y_train = y_train[indices[:split_idx]]
else:
print("No pre-extracted embeddings found. Launching Synthetic Dataset Generator baseline...")
X_train, y_train = generate_synthetic_data(num_samples=160)
X_val, y_val = generate_synthetic_data(num_samples=40)
# Save so they are available for other pipelines
np.save(train_emb_path, X_train)
np.save(train_lbl_path, y_train)
np.save(val_emb_path, X_val)
np.save(val_lbl_path, y_val)
print("Synthetic dataset successfully cached in data/ directory.")
# 2. Build and Save FAISS Index (Once from Training Cohort)
print("\nBuilding CaseRetriever FAISS database from training cohort...")
retriever = CaseRetriever(index_dim=DIM_FUSED)
# Generate metadata for training cohorts
patient_ids = [f"GS-{1000 + i}" for i in range(len(X_train))]
visual_descriptions = []
for lbl in y_train:
if lbl == 1:
visual_descriptions.append("poorly differentiated, medullary architecture, abundant tumor-infiltrating lymphocytes (TILs)")
else:
visual_descriptions.append("well-formed glands, regular tubular structures, preserved crypt morphology")
retriever.build_index(X_train, y_train, visual_descriptions, patient_ids)
# 3. Create PyTorch DataLoaders
train_dataset = TensorDataset(torch.tensor(X_train), torch.tensor(y_train).unsqueeze(1))
val_dataset = TensorDataset(torch.tensor(X_val), torch.tensor(y_val).unsqueeze(1))
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)
# 4. Initialize MLP Classifier
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"\nInitializing 3-layer MSI-Classifier on device: {device}")
model = MSIClassifier(input_dim=DIM_FUSED).to(device)
criterion = nn.BCELoss()
optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)
# 5. Training Loop
best_val_auroc = 0.0
best_weights_path = os.path.join(MODELS_DIR, "best_classifier.pt")
print("\nStarting training loop...")
for epoch in range(1, EPOCHS + 1):
model.train()
train_loss = 0.0
for batch_x, batch_y in train_loader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad()
outputs = model(batch_x)
loss = criterion(outputs, batch_y)
loss.backward()
optimizer.step()
train_loss += loss.item() * batch_x.size(0)
train_loss /= len(train_loader.dataset)
scheduler.step()
# Validation evaluation
model.eval()
val_preds = []
val_targets = []
with torch.no_grad():
for batch_x, batch_y in val_loader:
batch_x = batch_x.to(device)
outputs = model(batch_x)
val_preds.extend(outputs.cpu().numpy())
val_targets.extend(batch_y.numpy())
val_preds = np.array(val_preds)
val_targets = np.array(val_targets)
# Compute metrics
val_auc = roc_auc_score(val_targets, val_preds)
val_preds_bin = (val_preds >= 0.5).astype(int)
val_acc = accuracy_score(val_targets, val_preds_bin)
val_f1 = f1_score(val_targets, val_preds_bin)
print(f"Epoch {epoch}/{EPOCHS} | Train Loss: {train_loss:.4f} | Val AUROC: {val_auc:.4f} | Val Acc: {val_acc:.4f} | Val F1: {val_f1:.4f}")
# Save best checkpoint
if val_auc > best_val_auroc:
best_val_auroc = val_auc
torch.save(model.state_dict(), best_weights_path)
print(f"--> Saved new best checkpoint with Val AUROC: {val_val_auc:.4f}" if False else "--> Saved new best model checkpoint.")
print(f"\nTraining Complete! Best model saved to: {best_weights_path}")
print(f"Best Validation AUROC achieved: {best_val_auroc:.4f}")
if __name__ == "__main__":
main()