Download inference_example.py from NJ50/alpine-fewshot: direct link, hf CLI and curl.
- Browser
- Download file 5.5 kB
-
https://huggingface.co/NJ50/alpine-fewshot/resolve/main/inference_example.py
- Command line
-
hf download hf://NJ50/alpine-fewshot/inference_example.py
-
curl -L -o inference_example.py https://huggingface.co/NJ50/alpine-fewshot/resolve/main/inference_example.py
5.5 kB
| """ | |
| 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() | |