GEMS-Generalized-Raman-Classifier / main /hyperpara_optim_downstream.py
JunhanCai's picture
Initial commit with GEMS model and Dockerfile
6918c6b
Raw
History Blame
6.44 kB
import gc
import os
import numpy as np
import optuna
import pandas as pd
import torch
from sklearn.preprocessing import LabelEncoder
from torch.utils.data import DataLoader
from Ramandataset import RamanDataset
from Raman_Task import (
augment_minority_classes,
load_mae_model_for_classification,
load_real_data,
stratified_split_with_minimum_samples,
train_predictor,
unique,
)
def evaluate_classifier_accuracy(classifier, val_loader, device):
classifier.eval()
correct = 0
total = 0
with torch.no_grad():
for inputs, _, labels in val_loader:
inputs = inputs.to(device)
labels = labels.to(device)
if labels.dim() > 1:
labels = labels.squeeze()
labels = labels.long()
logits, _ = classifier(inputs)
preds = torch.argmax(logits, dim=1)
total += labels.size(0)
correct += (preds == labels).sum().item()
return correct / max(1, total)
def run_downstream_hpo(num_trials=25, epochs_per_trial=30, pretrained_path=None):
base_path = "/home/lion/Desktop/splendid_lion/project/2025/multi_source/Raman/support_file/data/"
data_path = os.path.join(base_path, "rruff/classifier_0_3500_spectra.npy")
labels_path = os.path.join(base_path, "rruff/classifier_0_3500_labels.npy")
wavenumbers_path = os.path.join(base_path, "rruff/classifier_0_3500_wavenumbers.npy")
spectra, labels, _ = load_real_data(
data_path, labels_path=labels_path, wavenumbers_path=wavenumbers_path, normalize=True
)
input_length = spectra.shape[1]
num_classes = len(unique(labels))
y = np.array(labels)
le = LabelEncoder()
y_encoded = le.fit_transform(y)
x_train, x_val, x_test, y_train, y_val, y_test = stratified_split_with_minimum_samples(
spectra, y_encoded, test_size=0.15, val_size=0.15, min_samples_per_class=1, random_state=42
)
x_train_aug, y_train_aug = augment_minority_classes(x_train, y_train, min_samples=30, target_samples=159)
para = {"embedding_dim": 512, "num_heads": 16, "num_layers": 12, "patch_num": 100}
if pretrained_path is None:
pretrained_path = (
"/home/lion/Desktop/splendid_lion/project/2025/multi_source/Raman/support_file/model/"
"0.0001_0.5_512_16_12_100/Fine_tuned_baseonALL.pth"
)
if not os.path.exists(pretrained_path):
raise FileNotFoundError(f"Pretrained model not found: {pretrained_path}")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
def objective(trial):
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
lr = trial.suggest_categorical("lr", [1e-3, 5e-4, 1e-4, 5e-5, 1e-5])
batch_size = trial.suggest_categorical("batch_size", [16, 32])
weight_decay = trial.suggest_categorical("weight_decay", [1e-3, 5e-4, 1e-4, 5e-5, 1e-5])
freeze_encoder = trial.suggest_categorical("freeze_encoder", [True, False])
trial_dir = os.path.join(base_path, "optimization_results", f"downstream_trial_{trial.number}")
os.makedirs(trial_dir, exist_ok=True)
train_ds = RamanDataset(x_train_aug, None, labels=y_train_aug, transform=None, is_train=True)
val_ds = RamanDataset(x_val, None, labels=y_val, transform=None, is_train=False)
test_ds = RamanDataset(x_test, None, labels=y_test, transform=None, is_train=False)
train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, drop_last=True, num_workers=0)
val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0)
test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, num_workers=0)
classifier, _, _, mae_model = load_mae_model_for_classification(
pretrained_path,
input_length,
para["patch_num"],
para["embedding_dim"],
para["num_layers"],
para["num_heads"],
num_classes,
device,
)
train_predictor(
classifier=classifier,
mae_model=mae_model,
train_loader=train_loader,
val_loader=val_loader,
test_loader=test_loader,
device=device,
epochs=epochs_per_trial,
lr=lr,
weight_decay=weight_decay,
patience=10,
save_dir=trial_dir,
model_name=f"downstream_trial_{trial.number}",
freeze_encoder=freeze_encoder,
)
best_ckpt = os.path.join(trial_dir, f"downstream_trial_{trial.number}_best_class.pth")
if os.path.exists(best_ckpt):
ckpt = torch.load(best_ckpt, map_location=device)
classifier.load_state_dict(ckpt["model_state_dict"])
val_acc = evaluate_classifier_accuracy(classifier, val_loader, device)
row = {
"trial_id": trial.number,
"val_accuracy": val_acc,
"lr": lr,
"batch_size": batch_size,
"weight_decay": weight_decay,
"freeze_encoder": freeze_encoder,
}
pd.DataFrame([row]).to_csv(
"optuna_downstream_results.csv",
mode="a",
header=not os.path.exists("optuna_downstream_results.csv"),
index=False,
)
return val_acc
study = optuna.create_study(
direction="maximize",
study_name=f"raman_downstream_optimization_{pd.Timestamp.now().strftime('%Y%m%d_%H%M%S')}",
storage="sqlite:///raman_optimization.db",
load_if_exists=False,
)
study.optimize(objective, n_trials=num_trials, n_jobs=1, gc_after_trial=True)
print("=" * 60)
print("DOWNSTREAM HPO SUMMARY")
print(f"Best val acc: {study.best_value:.4f}")
print(f"Best params: {study.best_params}")
print("=" * 60)
best_trial_id = study.best_trial.number
best_trial_dir = os.path.join(base_path, "optimization_results", f"downstream_trial_{best_trial_id}")
summary = {
"best_trial_id": best_trial_id,
"best_value": float(study.best_value),
"best_params": dict(study.best_params),
"best_trial_dir": best_trial_dir,
"pretrained_path": pretrained_path,
}
return summary
if __name__ == "__main__":
run_downstream_hpo(num_trials=25, epochs_per_trial=30)