Spaces:
Runtime error
Runtime error
| import os | |
| import time | |
| import torch | |
| import torch.nn as nn | |
| from PIL import Image | |
| import pandas as pd | |
| import gradio as gr | |
| from torchvision import models, transforms | |
| # 1. Dynamic Class Identification from your training CSV | |
| CSV_FILE_PATH = "Training_set.csv" | |
| LABEL_COLUMN_NAME = "label" | |
| if os.path.exists(CSV_FILE_PATH): | |
| metadata_df = pd.read_csv(CSV_FILE_PATH) | |
| CLASS_NAMES = sorted(metadata_df[LABEL_COLUMN_NAME].dropna().unique().tolist()) | |
| print(f"Dynamically discovered {len(CLASS_NAMES)} classes from CSV.") | |
| else: | |
| CLASS_NAMES = ['Butterfly_Species_1', 'Butterfly_Species_2', 'Butterfly_Species_3'] | |
| print(f"WARNING: '{CSV_FILE_PATH}' not found. Defaulting to placeholders.") | |
| NUM_CLASSES = len(CLASS_NAMES) | |
| # 2. Reconstruct Models to Match Your Training Notebook Exactly | |
| # --- Model A: MobileNetV2 (Aug) --- | |
| mobilenet = models.mobilenet_v2(weights=None) | |
| num_features_mono = mobilenet.classifier[1].in_features | |
| mobilenet.classifier = nn.Sequential( | |
| nn.Dropout(p=0.2), | |
| nn.Linear(num_features_mono, NUM_CLASSES) | |
| ) | |
| # --- Model B: Unmodified VGG16 (Stock Head) --- | |
| vgg16_stock = models.vgg16(weights=None) | |
| num_features_stock = vgg16_stock.classifier[6].in_features | |
| vgg16_stock.classifier[6] = nn.Linear(num_features_stock, NUM_CLASSES) | |
| # --- Model C: Modified VGG16 with GAP --- | |
| vgg16_gap = models.vgg16(weights=None) | |
| vgg16_gap.avgpool = nn.AdaptiveAvgPool2d((1, 1)) # Replaced avgpool step | |
| vgg16_gap.classifier = nn.Sequential( | |
| nn.Flatten(), | |
| nn.Dropout(p=0.3), | |
| nn.Linear(512, NUM_CLASSES) | |
| ) | |
| # 3. Safe Model Checkpoint Loading Strategy | |
| models_dict = { | |
| "MobileNetV2": mobilenet, | |
| "VGG16 (Stock Head)": vgg16_stock, | |
| "VGG16 (GAP Modified Head)": vgg16_gap | |
| } | |
| files_dict = { | |
| "MobileNetV2": "mobilenetv2_butterfly.pth", | |
| "VGG16 (Stock Head)": "unmodified_vgg16_butterfly.pth", | |
| "VGG16 (GAP Modified Head)": "modified_vgg16_gap_butterfly.pth" | |
| } | |
| for name, model in models_dict.items(): | |
| file_path = files_dict[name] | |
| if os.path.exists(file_path): | |
| model.load_state_dict(torch.load(file_path, map_location=torch.device('cpu'))) | |
| model.eval() | |
| print(f"Successfully loaded {name} weights structure without mismatch.") | |
| else: | |
| print(f"WARNING: Checkpoint file '{file_path}' missing!") | |
| # 4. Processing Pipelines | |
| img_transforms = transforms.Compose([ | |
| transforms.Resize((224, 224)), | |
| transforms.ToTensor(), | |
| transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) | |
| ]) | |
| # 5. High Precision Benchmark Handler | |
| def benchmark_inference(input_image): | |
| if input_image is None: | |
| return None | |
| img_tensor = img_transforms(input_image).unsqueeze(0) | |
| results_data = [] | |
| with torch.no_grad(): | |
| for name, model in models_dict.items(): | |
| # Warm-up loop iteration | |
| _ = model(img_tensor) | |
| # Start tracking precision loop speed | |
| start_time = time.perf_counter() | |
| outputs = model(img_tensor) | |
| probabilities = torch.nn.functional.softmax(outputs[0], dim=0) | |
| best_prob, best_class_idx = torch.max(probabilities, dim=0) | |
| end_time = time.perf_counter() | |
| latency_ms = (end_time - start_time) * 1000 | |
| predicted_label = CLASS_NAMES[best_class_idx.item()] | |
| confidence_pct = f"{best_prob.item() * 100:.2f}%" | |
| results_data.append({ | |
| "Model Architecture Blueprint": name, | |
| "Predicted Butterfly Label": predicted_label, | |
| "Confidence Score": confidence_pct, | |
| "Inference Latency (Milliseconds)": f"{latency_ms:.2f} ms" | |
| }) | |
| return pd.DataFrame(results_data) | |
| # 6. Gradio Render Interface | |
| with gr.Blocks(theme=gr.themes.Soft()) as demo: | |
| gr.Markdown("# 🦋 Scenario 1: Butterfly Classifier & Performance Benchmarking Space") | |
| gr.Markdown("Upload an image of a butterfly to compare deployment metrics, classification targets, and computational speed tracking across your models.") | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| input_img_slot = gr.Image(type="pil", label="Input Target Field Photo File") | |
| submit_btn = gr.Button("Execute Multi-Model Run", variant="primary") | |
| with gr.Column(scale=2): | |
| output_table_slot = gr.Dataframe( | |
| headers=["Model Architecture Blueprint", "Predicted Butterfly Label", "Confidence Score", "Inference Latency (Milliseconds)"], | |
| datatype=["str", "str", "str", "str"], | |
| label="Benchmark Results Table" | |
| ) | |
| submit_btn.click( | |
| fn=benchmark_inference, | |
| inputs=input_img_slot, | |
| outputs=output_table_slot | |
| ) | |
| demo.launch() |