Download src/models.py from NJ50/alpine-fewshot: direct link, hf CLI and curl.
- Browser
- Download file 22.8 kB
-
https://huggingface.co/NJ50/alpine-fewshot/resolve/main/src/models.py
- Command line
-
hf download hf://NJ50/alpine-fewshot/src/models.py
-
curl -L -o models.py https://huggingface.co/NJ50/alpine-fewshot/resolve/main/src/models.py
22.8 kB
| """ | |
| ALPINE: Adaptive Localization for Parameter- and Sample-Efficient Few-Shot Learning | |
| Official Model Architecture & Checkpoint Loader | |
| This module defines: | |
| - Canonical EXP-F3 (22,249 parameters): ALPINE_CIFAR, ALPINE_Native | |
| - Optional Variant EXP-F3-35k (34,917 parameters): ALPINE_35k_CIFAR, ALPINE_35k_Native | |
| - Gabor edge-energy guided windowed patch locator | |
| - Convenient `load_alpine_model` checkpoint loader | |
| """ | |
| import math | |
| import os | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| try: | |
| from .irfe_p1_canonical import ( | |
| GaborPreprocess, | |
| PatchEncoderCIFARP1, | |
| SubPatchEncoderCIFARP1, | |
| BoundaryEncoderCIFARP1, | |
| PatchEncoderNativeP1, | |
| SubPatchEncoderNativeP1, | |
| BoundaryEncoderNativeP1 | |
| ) | |
| except ImportError: | |
| from irfe_p1_canonical import ( | |
| GaborPreprocess, | |
| PatchEncoderCIFARP1, | |
| SubPatchEncoderCIFARP1, | |
| BoundaryEncoderCIFARP1, | |
| PatchEncoderNativeP1, | |
| SubPatchEncoderNativeP1, | |
| BoundaryEncoderNativeP1 | |
| ) | |
| def get_exp_p1_base_centers(): | |
| """Returns canonical base 5-patch center coordinates in normalized [-1, 1] range.""" | |
| return torch.tensor([ | |
| [-0.5161, -0.5161], | |
| [-0.5161, 0.5161], | |
| [ 0.0000, 0.0000], | |
| [ 0.5161, -0.5161], | |
| [ 0.5161, 0.5161] | |
| ], dtype=torch.float32) | |
| class WideWindowAdaptivePatchLocator(nn.Module): | |
| """ | |
| Windowed Adaptive Patch Locator (EXP-F3). | |
| Dynamically displaces patch sampling centers toward salient features | |
| within a Gaussian spatial window around canonical base centers, | |
| guided by Gabor edge-energy response E(y, x). | |
| """ | |
| def __init__(self, base_centers, window_frac=0.50, temperature=1.0): | |
| super().__init__() | |
| self.register_buffer("base_centers", base_centers) | |
| self.num_patches = base_centers.shape[0] | |
| self.window_frac = window_frac | |
| self.temperature = nn.Parameter(torch.tensor(temperature, dtype=torch.float32)) | |
| def forward(self, energy_map): | |
| B, _, H, W = energy_map.shape | |
| ys = torch.linspace(-1, 1, H, device=energy_map.device) | |
| xs = torch.linspace(-1, 1, W, device=energy_map.device) | |
| gy, gx = torch.meshgrid(ys, xs, indexing='ij') | |
| flat_energy = energy_map.squeeze(1) | |
| centers = [] | |
| for k in range(self.num_patches): | |
| cy0 = self.base_centers[k, 0] | |
| cx0 = self.base_centers[k, 1] | |
| dist_sq = (gy - cy0)**2 + (gx - cx0)**2 | |
| log_window = -dist_sq / (2 * (self.window_frac**2) + 1e-6) | |
| weighted_energy = flat_energy / (self.temperature.abs() + 1e-4) + log_window | |
| w = F.softmax(weighted_energy.view(B, -1), dim=-1) | |
| cy = (w * gy.reshape(-1)).sum(-1) | |
| cx = (w * gx.reshape(-1)).sum(-1) | |
| centers.append(torch.stack([cx, cy], dim=-1)) | |
| return torch.stack(centers, dim=1) # Shape: (B, 5, 2) | |
| def extract_patches_grid_sample(x, centers, patch_scale=0.5): | |
| """Bilinearly extracts image patches around adaptive centers using grid_sample.""" | |
| B, C, H, W = x.shape | |
| patch_h = int(round(H * patch_scale)) | |
| patch_w = int(round(W * patch_scale)) | |
| patches = [] | |
| for k in range(centers.size(1)): | |
| cx = centers[:, k, 0] | |
| cy = centers[:, k, 1] | |
| theta = torch.zeros(B, 2, 3, device=x.device, dtype=x.dtype) | |
| theta[:, 0, 0] = patch_scale | |
| theta[:, 1, 1] = patch_scale | |
| theta[:, 0, 2] = cx | |
| theta[:, 1, 2] = cy | |
| grid = F.affine_grid(theta, torch.Size([B, C, patch_h, patch_w]), align_corners=False) | |
| p = F.grid_sample(x, grid, align_corners=False, mode='bilinear', padding_mode='reflection') | |
| patches.append(p) | |
| return patches | |
| # ===================================================================== | |
| # CANONICAL EXP-F3 (22,249 PARAMETERS) | |
| # ===================================================================== | |
| class ALPINE_CIFAR(nn.Module): | |
| """ | |
| Canonical ALPINE / EXP-F3 architecture for CIFAR-FS (32x32). | |
| Total Trainable Parameters: 22,249. | |
| """ | |
| def __init__(self, window_frac=0.50, embed_dim=16, num_heads=4): | |
| super().__init__() | |
| self.embed_dim = embed_dim | |
| base_centers = get_exp_p1_base_centers() | |
| self.locator = WideWindowAdaptivePatchLocator(base_centers, window_frac=window_frac) | |
| self.encoder = PatchEncoderCIFARP1(out_dim=embed_dim) | |
| self.sub_encoder = SubPatchEncoderCIFARP1(out_dim=embed_dim) | |
| self.boundary_encoder = BoundaryEncoderCIFARP1(out_dim=embed_dim) | |
| self.sub_fusion = nn.Sequential( | |
| nn.Linear(embed_dim * 4, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.rel_proj = nn.Sequential( | |
| nn.Linear(embed_dim * 2, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.mha = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) | |
| self.ln_attn = nn.LayerNorm(embed_dim) | |
| self.ws_query = nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02) | |
| self.gate_net = nn.Sequential( | |
| nn.Linear(embed_dim * 2, 16), nn.ReLU(), nn.Linear(16, embed_dim * 2) | |
| ) | |
| def _split_boundaries(self, x): | |
| return (x[:,:, 8:16, 8:16], x[:,:, 8:16, 16:24], | |
| x[:,:, 16:24, 8:16], x[:,:, 16:24, 16:24]) | |
| def _encode_patch(self, p): | |
| v_c = self.encoder(p) | |
| sp = [p[:,:, r:r+8, c:c+8] for r in [0, 8] for c in [0, 8]] | |
| v_f = self.sub_fusion(torch.cat([self.sub_encoder(s) for s in sp], dim=1)) | |
| return v_c + v_f | |
| def _rel(self, va, vb): | |
| return self.rel_proj(torch.cat([va - vb, va * vb], dim=1)) | |
| def extract_with_rel_tokens(self, x): | |
| B = x.size(0) | |
| gabor_out = self.encoder.gabor(x) | |
| energy_map = gabor_out[:, 3:7].abs().sum(dim=1, keepdim=True) | |
| centers = self.locator(energy_map) | |
| raw_patches = extract_patches_grid_sample(x, centers, patch_scale=0.5) | |
| patches = [self._encode_patch(p) for p in raw_patches] | |
| edges = [self._rel(patches[i], patches[j]) for i in range(5) for j in range(i+1, 5)] | |
| boundaries = [self.boundary_encoder(b) for b in self._split_boundaries(x)] | |
| tokens = torch.stack(patches + edges + boundaries, dim=1) | |
| attn, _ = self.mha(tokens, tokens, tokens) | |
| refined = self.ln_attn(tokens + attn) | |
| ws = self.ws_query.expand(B, -1, -1) | |
| scores = torch.bmm(ws, refined.transpose(1, 2)) / math.sqrt(self.embed_dim) | |
| W = torch.bmm(F.softmax(scores, dim=-1), refined).squeeze(1) | |
| patch_pool = refined[:, :5, :].mean(dim=1) | |
| g = torch.cat([patch_pool, W], dim=1) | |
| gate = 2.0 * torch.sigmoid(self.gate_net(g)) - 1.0 | |
| out_feat= g + g * gate | |
| rel_tokens = torch.stack(edges, dim=1) | |
| return out_feat, centers, rel_tokens | |
| def extract(self, x): | |
| """Extracts fixed-dimensional feature representations for few-shot metric classification.""" | |
| feat, _, _ = self.extract_with_rel_tokens(x) | |
| return feat | |
| def compute_prototypes(self, sx, sy, n=5): | |
| """Computes support prototypes for n-way classification.""" | |
| f = self.extract(sx) | |
| return torch.stack([f[sy == c].mean(0) for c in range(n)]) | |
| def predict_proto(self, qx, protos): | |
| """Computes negative squared Euclidean distance between query representations and prototypes.""" | |
| return -(torch.cdist(self.extract(qx), protos) ** 2) | |
| class ALPINE_Native(nn.Module): | |
| """ | |
| Canonical ALPINE / EXP-F3 architecture for MiniImageNet Native (84x84). | |
| Total Trainable Parameters: 22,249. | |
| """ | |
| def __init__(self, window_frac=0.50, embed_dim=16, num_heads=4): | |
| super().__init__() | |
| self.embed_dim = embed_dim | |
| base_centers = get_exp_p1_base_centers() | |
| self.locator = WideWindowAdaptivePatchLocator(base_centers, window_frac=window_frac) | |
| self.encoder = PatchEncoderNativeP1(out_dim=embed_dim) | |
| self.sub_encoder = SubPatchEncoderNativeP1(out_dim=embed_dim) | |
| self.boundary_encoder = BoundaryEncoderNativeP1(out_dim=embed_dim) | |
| self.sub_fusion = nn.Sequential( | |
| nn.Linear(embed_dim * 4, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.rel_proj = nn.Sequential( | |
| nn.Linear(embed_dim * 2, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.mha = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) | |
| self.ln_attn = nn.LayerNorm(embed_dim) | |
| self.ws_query = nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02) | |
| self.gate_net = nn.Sequential( | |
| nn.Linear(embed_dim * 2, 16), nn.ReLU(), nn.Linear(16, embed_dim * 2) | |
| ) | |
| def _split_boundaries(self, x): | |
| return (x[:,:, 18:48, 18:48], x[:,:, 18:48, 36:66], | |
| x[:,:, 36:66, 18:48], x[:,:, 36:66, 36:66]) | |
| def _encode_patch(self, p): | |
| v_c = self.encoder(p) | |
| sp = [p[:,:, r:r+24, c:c+24] for r in [0, 24] for c in [0, 24]] | |
| v_f = self.sub_fusion(torch.cat([self.sub_encoder(s) for s in sp], dim=1)) | |
| return v_c + v_f | |
| def _rel(self, va, vb): | |
| return self.rel_proj(torch.cat([va - vb, va * vb], dim=1)) | |
| def extract_with_rel_tokens(self, x): | |
| B = x.size(0) | |
| gabor_out = self.encoder.gabor(x) | |
| energy_map = gabor_out[:, 3:7].abs().sum(dim=1, keepdim=True) | |
| centers = self.locator(energy_map) | |
| raw_patches = extract_patches_grid_sample(x, centers, patch_scale=48.0/84.0) | |
| patches = [self._encode_patch(p) for p in raw_patches] | |
| edges = [self._rel(patches[i], patches[j]) for i in range(5) for j in range(i+1, 5)] | |
| boundaries = [self.boundary_encoder(b) for b in self._split_boundaries(x)] | |
| tokens = torch.stack(patches + edges + boundaries, dim=1) | |
| attn, _ = self.mha(tokens, tokens, tokens) | |
| refined = self.ln_attn(tokens + attn) | |
| ws = self.ws_query.expand(B, -1, -1) | |
| scores = torch.bmm(ws, refined.transpose(1, 2)) / math.sqrt(self.embed_dim) | |
| W = torch.bmm(F.softmax(scores, dim=-1), refined).squeeze(1) | |
| patch_pool = refined[:, :5, :].mean(dim=1) | |
| g = torch.cat([patch_pool, W], dim=1) | |
| gate = 2.0 * torch.sigmoid(self.gate_net(g)) - 1.0 | |
| out_feat= g + g * gate | |
| rel_tokens = torch.stack(edges, dim=1) | |
| return out_feat, centers, rel_tokens | |
| def extract(self, x): | |
| feat, _, _ = self.extract_with_rel_tokens(x) | |
| return feat | |
| def compute_prototypes(self, sx, sy, n=5): | |
| f = self.extract(sx) | |
| return torch.stack([f[sy == c].mean(0) for c in range(n)]) | |
| def predict_proto(self, qx, protos): | |
| return -(torch.cdist(self.extract(qx), protos) ** 2) | |
| # Aliases for backward compatibility | |
| IRFEExpF3_CIFAR = ALPINE_CIFAR | |
| IRFEExpF3_Native = ALPINE_Native | |
| # ===================================================================== | |
| # OPTIONAL VARIANT EXP-F3-35k (34,917 PARAMETERS) | |
| # ===================================================================== | |
| class PatchEncoderCIFARGeneric(nn.Module): | |
| def __init__(self, c1, c2, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv1 = nn.Conv2d(7, c1, 3, padding=1) | |
| self.conv2 = nn.Conv2d(c1, c2, 3, padding=1) | |
| self.pool = nn.MaxPool2d(2, 2) | |
| self.fc = nn.Linear(c2 * 4 * 4, out_dim) | |
| self.ln = nn.LayerNorm(out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| x = self.pool(F.relu(self.conv1(x))) | |
| x = self.pool(F.relu(self.conv2(x))) | |
| return self.ln(self.fc(x.view(x.size(0), -1))) | |
| class SubPatchEncoderCIFARGeneric(nn.Module): | |
| def __init__(self, sub_c, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv = nn.Conv2d(7, sub_c, 3, padding=1) | |
| self.pool = nn.AdaptiveAvgPool2d((4, 4)) | |
| self.fc = nn.Linear(sub_c * 4 * 4, out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| return F.relu(self.fc(self.pool(F.relu(self.conv(x))).view(x.size(0), -1))) | |
| class BoundaryEncoderCIFARGeneric(nn.Module): | |
| def __init__(self, sub_c, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv = nn.Conv2d(7, sub_c, 3, padding=1) | |
| self.pool = nn.AdaptiveAvgPool2d((4, 4)) | |
| self.fc = nn.Linear(sub_c * 4 * 4, out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| return F.relu(self.fc(self.pool(F.relu(self.conv(x))).view(x.size(0), -1))) | |
| class PatchEncoderNativeGeneric(nn.Module): | |
| def __init__(self, c1, c2, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv1 = nn.Conv2d(7, c1, 3, padding=1) | |
| self.conv2 = nn.Conv2d(c1, c2, 3, padding=1) | |
| self.pool = nn.MaxPool2d(2, 2) | |
| self.adap = nn.AdaptiveAvgPool2d((4, 4)) | |
| self.fc = nn.Linear(c2 * 4 * 4, out_dim) | |
| self.ln = nn.LayerNorm(out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| x = self.pool(F.relu(self.conv1(x))) | |
| x = self.pool(F.relu(self.conv2(x))) | |
| x = self.adap(x) | |
| return self.ln(self.fc(x.view(x.size(0), -1))) | |
| class SubPatchEncoderNativeGeneric(nn.Module): | |
| def __init__(self, sub_c, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv = nn.Conv2d(7, sub_c, 3, padding=1) | |
| self.pool = nn.AdaptiveAvgPool2d((4, 4)) | |
| self.fc = nn.Linear(sub_c * 4 * 4, out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| return F.relu(self.fc(self.pool(F.relu(self.conv(x))).view(x.size(0), -1))) | |
| class BoundaryEncoderNativeGeneric(nn.Module): | |
| def __init__(self, sub_c, out_dim): | |
| super().__init__() | |
| self.gabor = GaborPreprocess(ksize=5) | |
| self.conv = nn.Conv2d(7, sub_c, 3, padding=1) | |
| self.pool = nn.AdaptiveAvgPool2d((4, 4)) | |
| self.fc = nn.Linear(sub_c * 4 * 4, out_dim) | |
| def forward(self, x): | |
| x = self.gabor(x) | |
| return F.relu(self.fc(self.pool(F.relu(self.conv(x))).view(x.size(0), -1))) | |
| class ALPINE_35k_CIFAR(nn.Module): | |
| """ | |
| Optional Variant: ALPINE / EXP-F3-35k for CIFAR-FS (32x32). | |
| Total Trainable Parameters: 34,917. | |
| """ | |
| def __init__(self, c1=15, c2=19, sub_c=8, gate_h=12, embed_dim=32, window_frac=0.50, num_heads=4): | |
| super().__init__() | |
| self.embed_dim = embed_dim | |
| base_centers = get_exp_p1_base_centers() | |
| self.locator = WideWindowAdaptivePatchLocator(base_centers, window_frac=window_frac) | |
| self.encoder = PatchEncoderCIFARGeneric(c1, c2, out_dim=embed_dim) | |
| self.sub_encoder = SubPatchEncoderCIFARGeneric(sub_c, out_dim=embed_dim) | |
| self.boundary_encoder = BoundaryEncoderCIFARGeneric(sub_c, out_dim=embed_dim) | |
| self.sub_fusion = nn.Sequential( | |
| nn.Linear(embed_dim * 4, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.rel_proj = nn.Sequential( | |
| nn.Linear(embed_dim * 2, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.mha = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) | |
| self.ln_attn = nn.LayerNorm(embed_dim) | |
| self.ws_query = nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02) | |
| self.gate_net = nn.Sequential( | |
| nn.Linear(embed_dim * 2, gate_h), nn.ReLU(), nn.Linear(gate_h, embed_dim * 2) | |
| ) | |
| def _split_boundaries(self, x): | |
| return (x[:,:, 8:16, 8:16], x[:,:, 8:16, 16:24], | |
| x[:,:, 16:24, 8:16], x[:,:, 16:24, 16:24]) | |
| def _encode_patch(self, p): | |
| v_c = self.encoder(p) | |
| sp = [p[:,:, r:r+8, c:c+8] for r in [0, 8] for c in [0, 8]] | |
| v_f = self.sub_fusion(torch.cat([self.sub_encoder(s) for s in sp], dim=1)) | |
| return v_c + v_f | |
| def _rel(self, va, vb): | |
| return self.rel_proj(torch.cat([va - vb, va * vb], dim=1)) | |
| def extract_with_rel_tokens(self, x): | |
| B = x.size(0) | |
| gabor_out = self.encoder.gabor(x) | |
| energy_map = gabor_out[:, 3:7].abs().sum(dim=1, keepdim=True) | |
| centers = self.locator(energy_map) | |
| raw_patches = extract_patches_grid_sample(x, centers, patch_scale=0.5) | |
| patches = [self._encode_patch(p) for p in raw_patches] | |
| edges = [self._rel(patches[i], patches[j]) for i in range(5) for j in range(i+1, 5)] | |
| boundaries = [self.boundary_encoder(b) for b in self._split_boundaries(x)] | |
| tokens = torch.stack(patches + edges + boundaries, dim=1) | |
| attn, _ = self.mha(tokens, tokens, tokens) | |
| refined = self.ln_attn(tokens + attn) | |
| ws = self.ws_query.expand(B, -1, -1) | |
| scores = torch.bmm(ws, refined.transpose(1, 2)) / math.sqrt(self.embed_dim) | |
| W = torch.bmm(F.softmax(scores, dim=-1), refined).squeeze(1) | |
| patch_pool = refined[:, :5, :].mean(dim=1) | |
| g = torch.cat([patch_pool, W], dim=1) | |
| gate = 2.0 * torch.sigmoid(self.gate_net(g)) - 1.0 | |
| out_feat= g + g * gate | |
| rel_tokens = torch.stack(edges, dim=1) | |
| return out_feat, centers, rel_tokens | |
| def extract(self, x): | |
| feat, _, _ = self.extract_with_rel_tokens(x) | |
| return feat | |
| def compute_prototypes(self, sx, sy, n=5): | |
| f = self.extract(sx) | |
| return torch.stack([f[sy == c].mean(0) for c in range(n)]) | |
| def predict_proto(self, qx, protos): | |
| return -(torch.cdist(self.extract(qx), protos) ** 2) | |
| class ALPINE_35k_Native(nn.Module): | |
| """ | |
| Optional Variant: ALPINE / EXP-F3-35k for MiniImageNet Native (84x84). | |
| Total Trainable Parameters: 34,917. | |
| """ | |
| def __init__(self, c1=15, c2=19, sub_c=8, gate_h=12, embed_dim=32, window_frac=0.50, num_heads=4): | |
| super().__init__() | |
| self.embed_dim = embed_dim | |
| base_centers = get_exp_p1_base_centers() | |
| self.locator = WideWindowAdaptivePatchLocator(base_centers, window_frac=window_frac) | |
| self.encoder = PatchEncoderNativeGeneric(c1, c2, out_dim=embed_dim) | |
| self.sub_encoder = SubPatchEncoderNativeGeneric(sub_c, out_dim=embed_dim) | |
| self.boundary_encoder = BoundaryEncoderNativeGeneric(sub_c, out_dim=embed_dim) | |
| self.sub_fusion = nn.Sequential( | |
| nn.Linear(embed_dim * 4, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.rel_proj = nn.Sequential( | |
| nn.Linear(embed_dim * 2, embed_dim), nn.LayerNorm(embed_dim), nn.ReLU() | |
| ) | |
| self.mha = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) | |
| self.ln_attn = nn.LayerNorm(embed_dim) | |
| self.ws_query = nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02) | |
| self.gate_net = nn.Sequential( | |
| nn.Linear(embed_dim * 2, gate_h), nn.ReLU(), nn.Linear(gate_h, embed_dim * 2) | |
| ) | |
| def _split_boundaries(self, x): | |
| return (x[:,:, 18:48, 18:48], x[:,:, 18:48, 36:66], | |
| x[:,:, 36:66, 18:48], x[:,:, 36:66, 36:66]) | |
| def _encode_patch(self, p): | |
| v_c = self.encoder(p) | |
| sp = [p[:,:, r:r+24, c:c+24] for r in [0, 24] for c in [0, 24]] | |
| v_f = self.sub_fusion(torch.cat([self.sub_encoder(s) for s in sp], dim=1)) | |
| return v_c + v_f | |
| def _rel(self, va, vb): | |
| return self.rel_proj(torch.cat([va - vb, va * vb], dim=1)) | |
| def extract_with_rel_tokens(self, x): | |
| B = x.size(0) | |
| gabor_out = self.encoder.gabor(x) | |
| energy_map = gabor_out[:, 3:7].abs().sum(dim=1, keepdim=True) | |
| centers = self.locator(energy_map) | |
| raw_patches = extract_patches_grid_sample(x, centers, patch_scale=48.0/84.0) | |
| patches = [self._encode_patch(p) for p in raw_patches] | |
| edges = [self._rel(patches[i], patches[j]) for i in range(5) for j in range(i+1, 5)] | |
| boundaries = [self.boundary_encoder(b) for b in self._split_boundaries(x)] | |
| tokens = torch.stack(patches + edges + boundaries, dim=1) | |
| attn, _ = self.mha(tokens, tokens, tokens) | |
| refined = self.ln_attn(tokens + attn) | |
| ws = self.ws_query.expand(B, -1, -1) | |
| scores = torch.bmm(ws, refined.transpose(1, 2)) / math.sqrt(self.embed_dim) | |
| W = torch.bmm(F.softmax(scores, dim=-1), refined).squeeze(1) | |
| patch_pool = refined[:, :5, :].mean(dim=1) | |
| g = torch.cat([patch_pool, W], dim=1) | |
| gate = 2.0 * torch.sigmoid(self.gate_net(g)) - 1.0 | |
| out_feat= g + g * gate | |
| rel_tokens = torch.stack(edges, dim=1) | |
| return out_feat, centers, rel_tokens | |
| def extract(self, x): | |
| feat, _, _ = self.extract_with_rel_tokens(x) | |
| return feat | |
| def compute_prototypes(self, sx, sy, n=5): | |
| f = self.extract(sx) | |
| return torch.stack([f[sy == c].mean(0) for c in range(n)]) | |
| def predict_proto(self, qx, protos): | |
| return -(torch.cdist(self.extract(qx), protos) ** 2) | |
| # Aliases for 35k variant | |
| IRFEExpF3_35k_CIFAR = ALPINE_35k_CIFAR | |
| IRFEExpF3_35k_Native = ALPINE_35k_Native | |
| def load_alpine_model(checkpoint_path, model_type="canonical", dataset="cifar", device="cpu"): | |
| """ | |
| Convenience loader for ALPINE model checkpoints. | |
| Args: | |
| checkpoint_path (str): Path to .pt checkpoint file. | |
| model_type (str): 'canonical' (22,249 params) or 'variant-35k' (34,917 params). | |
| dataset (str): 'cifar' (32x32 images) or 'mini' (84x84 images). | |
| device (str or torch.device): Device to load model onto. | |
| Returns: | |
| nn.Module: Loaded ALPINE model ready in eval mode. | |
| """ | |
| device = torch.device(device) | |
| model_type = model_type.lower() | |
| dataset = dataset.lower() | |
| if "35k" in model_type: | |
| if dataset in ["cifar", "cifar-fs", "cifar_fs"]: | |
| model = ALPINE_35k_CIFAR() | |
| else: | |
| model = ALPINE_35k_Native() | |
| else: | |
| if dataset in ["cifar", "cifar-fs", "cifar_fs"]: | |
| model = ALPINE_CIFAR() | |
| else: | |
| model = ALPINE_Native() | |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) | |
| state_dict = checkpoint["model_state_dict"] if isinstance(checkpoint, dict) and "model_state_dict" in checkpoint else checkpoint | |
| model.load_state_dict(state_dict) | |
| model.to(device) | |
| model.eval() | |
| return model | |