""" ALPINE: Adaptive Localization for Parameter- and Sample-Efficient Few-Shot Learning Inference Example & 5-Way Episode Demonstration This script demonstrates: 1. Loading the Canonical EXP-F3 checkpoint (22,249 params). 2. Loading the Optional Variant EXP-F3-35k checkpoint (34,917 params). 3. Running a 5-Way 5-Shot classification episode (Prototypical classification). 4. Inspecting the Gabor-guided adaptive patch centers. Requirements: torch >= 2.0.0 numpy """ import os import torch from src.models import load_alpine_model, ALPINE_CIFAR, ALPINE_Native def run_5way_episode_demo(model, img_size=(3, 32, 32), n_way=5, k_shot=5, q_query=15, device="cpu"): """ Simulates a standard n-way k-shot episodic classification task. """ print(f"\n--- Running {n_way}-Way {k_shot}-Shot Episodic Classification Demo ---") model.eval() # 1. Generate synthetic support and query images for demonstration # In practice, replace these with real images from CIFAR-FS or MiniImageNet total_support = n_way * k_shot total_query = n_way * q_query # Class-specific mean color offsets to simulate distinct visual classes support_images = [] support_labels = [] for c in range(n_way): # Class-distinct pattern class_pattern = torch.randn(1, *img_size) * 0.2 + (c * 0.5 - 1.0) class_samples = class_pattern + torch.randn(k_shot, *img_size) * 0.1 support_images.append(class_samples) support_labels.extend([c] * k_shot) support_x = torch.cat(support_images, dim=0).to(device) support_y = torch.tensor(support_labels, dtype=torch.long, device=device) query_images = [] query_labels = [] for c in range(n_way): class_pattern = torch.randn(1, *img_size) * 0.2 + (c * 0.5 - 1.0) class_samples = class_pattern + torch.randn(q_query, *img_size) * 0.1 query_images.append(class_samples) query_labels.extend([c] * q_query) query_x = torch.cat(query_images, dim=0).to(device) query_y = torch.tensor(query_labels, dtype=torch.long, device=device) # 2. Extract prototypes from support set with torch.no_grad(): prototypes = model.compute_prototypes(support_x, support_y, n=n_way) print(f"Computed {n_way} class prototypes. Shape: {prototypes.shape}") # 3. Predict query set logits (negative squared Euclidean distance) logits = model.predict_proto(query_x, prototypes) predictions = logits.argmax(dim=-1) acc = (predictions == query_y).float().mean().item() * 100.0 print(f"Query set size: {total_query} images ({q_query} per class)") print(f"Sample predictions: {predictions[:10].tolist()}") print(f"Ground truth labels: {query_y[:10].tolist()}") print(f"Episode accuracy: {acc:.2f}%") # 4. Inspect adaptive patch locator centers on a sample image with torch.no_grad(): features, centers, rel_tokens = model.extract_with_rel_tokens(query_x[:1]) print(f"\nAdaptive Patch Locator Inspection:") print(f"Extracted feature representation shape: {features.shape}") print(f"Gabor-guided patch sampling centers (x, y in [-1, 1]):") for idx, (cx, cy) in enumerate(centers[0].tolist()): print(f" Patch {idx + 1}: cx = {cx:+.4f}, cy = {cy:+.4f}") def main(): device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using device: {device}") # ========================================================================= # 1. PRIMARY MODEL: Canonical ALPINE / EXP-F3 (22,249 parameters) # Used for all primary benchmarks, robustness tests, and cross-domain eval # ========================================================================= cifar_canonical_path = "checkpoints/canonical/exp_f3_cifar_5shot_seed1.pt" if os.path.exists(cifar_canonical_path): print(f"\nLoading Primary Model: Canonical EXP-F3 (22,249 params) from {cifar_canonical_path}") model_canonical = load_alpine_model( cifar_canonical_path, model_type="canonical", dataset="cifar", device=device ) total_params = sum(p.numel() for p in model_canonical.parameters() if p.requires_grad) print(f"Loaded successfully! Trainable parameters: {total_params:,}") # Run 5-way 5-shot episode run_5way_episode_demo(model_canonical, img_size=(3, 32, 32), k_shot=5, device=device) # ========================================================================= # 2. OPTIONAL VARIANT: ALPINE / EXP-F3-35k (34,917 parameters) # Wider configuration, parameter-capacity sweet spot (Table 1 of paper) # ========================================================================= cifar_35k_path = "checkpoints/variant-35k/exp_f3_35k_cifar_5shot_seed1.pt" if os.path.exists(cifar_35k_path): print(f"\nLoading Optional Variant: EXP-F3-35k (34,917 params) from {cifar_35k_path}") model_35k = load_alpine_model( cifar_35k_path, model_type="variant-35k", dataset="cifar", device=device ) total_params_35k = sum(p.numel() for p in model_35k.parameters() if p.requires_grad) print(f"Loaded successfully! Trainable parameters: {total_params_35k:,}") # Run 5-way 5-shot episode run_5way_episode_demo(model_35k, img_size=(3, 32, 32), k_shot=5, device=device) if __name__ == "__main__": main()