Image Classification
PyTorch
few-shot-learning
parameter-efficient
File size: 5,502 Bytes
6edcaab
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""
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()