DragonfireCoder/microbiome-ai-classifier
2
1import os2import sys3import json4import torch5import optuna6import gc7from torch.utils.data import DataLoader8from tqdm import tqdm9 10from MainMicrobiome import (11 MicrobiomeBrain, CachedImageDataset, get_train_transform,12 inference_transform, CLASS_NAMES, SAM, FocalLoss13)14 15def objective(trial):16 print(f"\n[โก] Starting Trial #{trial.number}...")17 # 1. Hyperparameter Search Space18 lr = trial.suggest_float("lr", 1e-5, 1e-2, log=True)19 dropout = trial.suggest_float("dropout", 0.2, 0.6)20 weight_decay = trial.suggest_float("weight_decay", 1e-6, 1e-3, log=True)21 22 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")23 24 # 2. Data Setup (Using high-speed cache - Soil Only)25 try:26 from torchvision.datasets import ImageFolder27 full_ds = ImageFolder("data/train")28 soil_indices = [i for i, (p, t) in enumerate(full_ds.samples) if full_ds.classes[t] != "NOT_SOIL"]29 # Use simple label mapping for the 4 ensemble classes30 samples = [full_ds.samples[i] for i in soil_indices]31 train_dataset = CachedImageDataset(samples, transform=get_train_transform(0))32 train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)33 except Exception as e:34 print(f"[!] Data Error: {e}")35 return 0.036 37 # 3. Model Setup (Championship Specs)38 print(f"[*] Moving {115 if trial.number > 5 else 32}M Parameter Brain to {device}...")39 model = MicrobiomeBrain(base_channels=32, d1=dropout).to(device)40 optimizer = SAM(model.parameters(), torch.optim.AdamW, lr=lr, weight_decay=weight_decay)41 criterion = FocalLoss()42 43 # 4. Training Loop (Small burst for tuning)44 model.train()45 total_acc = 046 for epoch in range(3):47 correct, total = 0, 048 for images, labels in train_loader:49 images, labels = images.to(device), labels.to(device)50 51 # SAM Step 152 optimizer.zero_grad()53 outputs = model(images)54 criterion(outputs, labels).backward()55 optimizer.first_step(zero_grad=True)56 57 # SAM Step 258 criterion(model(images), labels).backward()59 optimizer.second_step(zero_grad=True)60 61 _, predicted = torch.max(outputs.data, 1)62 total += labels.size(0)63 correct += (predicted == labels).sum().item()64 65 total_acc = 100 * correct / total66 print(f" - Epoch {epoch+1}/3: Accuracy {total_acc:.2f}%")67 trial.report(total_acc, epoch)68 if trial.should_prune():69 # Clean up before pruning70 del model, optimizer, train_loader71 gc.collect()72 if torch.cuda.is_available(): torch.cuda.empty_cache()73 raise optuna.exceptions.TrialPruned()74 75 # 5. Final Cleanup76 del model, optimizer, train_loader77 gc.collect()78 if torch.cuda.is_available(): torch.cuda.empty_cache()79 80 return total_acc81 82if __name__ == "__main__":83 print("="*50)84 print("๐ SOILSENSE ULTIMATE HYPER-AUTO-TUNER")85 print("="*50)86 87 if not os.path.exists("data/train"):88 print("โ ERROR: No training data found. Run SetupMicrobiomeData.py first!")89 sys.exit(1)90 91 study = optuna.create_study(direction="maximize")92 study.optimize(objective, n_trials=10)93 94 print("\n" + "="*50)95 print("โ
OPTIMIZATION COMPLETE")96 print(f"๐ Best Accuracy: {study.best_value:.2f}%")97 98 # SAVE PARAMS (This is the magic fix for MasterScript!)99 with open("best_params.json", "w") as f:100 json.dump(study.best_params, f, indent=4)101 102 print("โ
Saved: best_params.json")103 print("="*50)