Image Classification
PyTorch
few-shot-learning
parameter-efficient
alpine-fewshot / inference_example.py
NJ50's picture
Initial release: ALPINE checkpoints (canonical EXP-F3 & variant-35k), model definitions, and inference example
6edcaab verified
Raw History Blame Contribute Delete
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()