CoolFace
Apppublic

DragonfireCoder/microbiome-ai-classifier

sourceHugging Facemitupdated 5mo agoView on Hugging Face
2likes
AutoTune.py103 linesDownload Raw Back to root
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)