Image-to-Image
Transformers
Safetensors
patchsvae
image-reconstruction
svd
geometric-deep-learning
autoencoder
omega-tokens
geolip
custom_code
Instructions to use AbstractPhil/svae-fresnel-128 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use AbstractPhil/svae-fresnel-128 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-to-image", model="AbstractPhil/svae-fresnel-128", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """PatchSVAE model for HuggingFace AutoModel. | |
| Usage: | |
| from transformers import AutoConfig, AutoModel | |
| config = AutoConfig.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True) | |
| model = AutoModel.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True) | |
| # Full reconstruction | |
| output = model(images) | |
| recon = output["recon"] # (B, 3, 128, 128) | |
| latent = output["latent"] # (B, 16, 8, 8) omega tokens | |
| # Encode to omega tokens | |
| omega = model.encode(images) # (B, 16, 8, 8) | |
| # Full SVD decomposition | |
| svd = model.encode_full(images) # dict with U, S, Vt, M per patch | |
| """ | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from typing import Optional, Dict, Union | |
| from transformers import PreTrainedModel | |
| from .configuration_patchsvae import PatchSVAEConfig | |
| # ── SVD Backend (self-contained, no external deps required) ────── | |
| try: | |
| from geolip_core.linalg.eigh import FLEigh, _FL_MAX_N | |
| _HAS_FL = True | |
| except ImportError: | |
| _HAS_FL = False | |
| def _gram_eigh_svd(A): | |
| """Thin SVD via Gram + eigh in fp64.""" | |
| orig_dtype = A.dtype | |
| with torch.amp.autocast('cuda', enabled=False): | |
| A_d = A.double() | |
| G = torch.bmm(A_d.transpose(1, 2), A_d) | |
| eigenvalues, V = torch.linalg.eigh(G) | |
| eigenvalues = eigenvalues.flip(-1) | |
| V = V.flip(-1) | |
| S = torch.sqrt(eigenvalues.clamp(min=1e-24)) | |
| U = torch.bmm(A_d, V) / S.unsqueeze(1).clamp(min=1e-16) | |
| Vh = V.transpose(-2, -1).contiguous() | |
| return U.to(orig_dtype), S.to(orig_dtype), Vh.to(orig_dtype) | |
| def _svd_fp64(A): | |
| """Auto-dispatch: FL eigh for N<=12, Gram eigh otherwise.""" | |
| B, M, N = A.shape | |
| if _HAS_FL and N <= _FL_MAX_N and A.is_cuda: | |
| orig_dtype = A.dtype | |
| with torch.amp.autocast('cuda', enabled=False): | |
| A_d = A.double() | |
| G = torch.bmm(A_d.transpose(1, 2), A_d) | |
| eigenvalues, V = FLEigh()(G.float()) | |
| eigenvalues = eigenvalues.double().flip(-1) | |
| V = V.double().flip(-1) | |
| S = torch.sqrt(eigenvalues.clamp(min=1e-24)) | |
| U = torch.bmm(A_d, V) / S.unsqueeze(1).clamp(min=1e-16) | |
| Vh = V.transpose(-2, -1).contiguous() | |
| return U.to(orig_dtype), S.to(orig_dtype), Vh.to(orig_dtype) | |
| else: | |
| return _gram_eigh_svd(A) | |
| # ── Patch Utilities ────────────────────────────────────────────── | |
| def _extract_patches(images, patch_size): | |
| B, C, H, W = images.shape | |
| gh, gw = H // patch_size, W // patch_size | |
| x = images.reshape(B, C, gh, patch_size, gw, patch_size) | |
| x = x.permute(0, 2, 4, 1, 3, 5) | |
| return x.reshape(B, gh * gw, C * patch_size * patch_size), gh, gw | |
| def _stitch_patches(patches, gh, gw, patch_size): | |
| B = patches.shape[0] | |
| x = patches.reshape(B, gh, gw, 3, patch_size, patch_size) | |
| x = x.permute(0, 3, 1, 4, 2, 5) | |
| return x.reshape(B, 3, gh * patch_size, gw * patch_size) | |
| # ── Components ─────────────────────────────────────────────────── | |
| class _BoundarySmooth(nn.Module): | |
| def __init__(self, channels=3, mid=16): | |
| super().__init__() | |
| self.net = nn.Sequential( | |
| nn.Conv2d(channels, mid, 3, padding=1), | |
| nn.GELU(), | |
| nn.Conv2d(mid, channels, 3, padding=1), | |
| ) | |
| nn.init.zeros_(self.net[-1].weight) | |
| nn.init.zeros_(self.net[-1].bias) | |
| def forward(self, x): | |
| return x + self.net(x) | |
| class _SpectralCrossAttention(nn.Module): | |
| def __init__(self, D, n_heads=4, max_alpha=0.2, alpha_init=-2.0): | |
| super().__init__() | |
| self.n_heads = n_heads | |
| self.head_dim = D // n_heads | |
| self.max_alpha = max_alpha | |
| assert D % n_heads == 0 | |
| self.qkv = nn.Linear(D, 3 * D) | |
| self.out_proj = nn.Linear(D, D) | |
| self.norm = nn.LayerNorm(D) | |
| self.scale = self.head_dim ** -0.5 | |
| self.alpha_logits = nn.Parameter(torch.full((D,), alpha_init)) | |
| def alpha(self): | |
| return self.max_alpha * torch.sigmoid(self.alpha_logits) | |
| def forward(self, S): | |
| B, N, D = S.shape | |
| S_normed = self.norm(S) | |
| qkv = self.qkv(S_normed).reshape(B, N, 3, self.n_heads, self.head_dim) | |
| qkv = qkv.permute(2, 0, 3, 1, 4) | |
| q, k, v = qkv[0], qkv[1], qkv[2] | |
| attn = (q @ k.transpose(-2, -1)) * self.scale | |
| attn = attn.softmax(dim=-1) | |
| out = (attn @ v).transpose(1, 2).reshape(B, N, D) | |
| gate = torch.tanh(self.out_proj(out)) | |
| return S * (1.0 + self.alpha.unsqueeze(0).unsqueeze(0) * gate) | |
| # ── Model ──────────────────────────────────────────────────────── | |
| class PatchSVAEModel(PreTrainedModel): | |
| """Patch-based SVD Autoencoder — The Fresnel Geometric Compression Lens. | |
| Decomposes images into patches, encodes each to a sphere-normalized | |
| matrix, performs SVD, coordinates spectra via cross-attention, | |
| and reconstructs with 99.993% fidelity. | |
| The spectral vectors S form omega tokens: modality-agnostic, | |
| geometrically structured, universal representations. | |
| """ | |
| config_class = PatchSVAEConfig | |
| _tied_weights_keys = [] | |
| def __init__(self, config: PatchSVAEConfig): | |
| super().__init__(config) | |
| V = config.matrix_v | |
| D = config.D | |
| hidden = config.hidden | |
| depth = config.depth | |
| ps = config.patch_size | |
| patch_dim = 3 * ps * ps | |
| mat_dim = V * D | |
| # Encoder | |
| self.enc_in = nn.Linear(patch_dim, hidden) | |
| self.enc_blocks = nn.ModuleList([ | |
| nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden), | |
| nn.GELU(), nn.Linear(hidden, hidden)) | |
| for _ in range(depth) | |
| ]) | |
| self.enc_out = nn.Linear(hidden, mat_dim) | |
| nn.init.orthogonal_(self.enc_out.weight) | |
| # Decoder | |
| self.dec_in = nn.Linear(mat_dim, hidden) | |
| self.dec_blocks = nn.ModuleList([ | |
| nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden), | |
| nn.GELU(), nn.Linear(hidden, hidden)) | |
| for _ in range(depth) | |
| ]) | |
| self.dec_out = nn.Linear(hidden, patch_dim) | |
| # Cross-attention | |
| self.cross_attn = nn.ModuleList([ | |
| _SpectralCrossAttention(D, n_heads=min(4, D), | |
| max_alpha=config.max_alpha, | |
| alpha_init=config.alpha_init) | |
| for _ in range(config.n_cross_layers) | |
| ]) | |
| # Boundary smoothing | |
| self.boundary_smooth = _BoundarySmooth(channels=3, mid=16) | |
| self.post_init() | |
| def _encode_patches_to_svd(self, patches): | |
| B, N, _ = patches.shape | |
| V, D = self.config.matrix_v, self.config.D | |
| flat = patches.reshape(B * N, -1) | |
| h = F.gelu(self.enc_in(flat)) | |
| for block in self.enc_blocks: | |
| h = h + block(h) | |
| M = self.enc_out(h).reshape(B * N, V, D) | |
| M = F.normalize(M, dim=-1) | |
| U, S, Vt = _svd_fp64(M) | |
| U = U.reshape(B, N, V, D) | |
| S = S.reshape(B, N, D) | |
| Vt = Vt.reshape(B, N, D, D) | |
| M = M.reshape(B, N, V, D) | |
| S_coord = S | |
| for layer in self.cross_attn: | |
| S_coord = layer(S_coord) | |
| return {"U": U, "S_orig": S, "S": S_coord, "Vt": Vt, "M": M} | |
| def _decode_from_svd(self, U, S, Vt): | |
| B, N, V, D = U.shape | |
| U_flat = U.reshape(B * N, V, D) | |
| S_flat = S.reshape(B * N, D) | |
| Vt_flat = Vt.reshape(B * N, D, D) | |
| M_hat = torch.bmm(U_flat * S_flat.unsqueeze(1), Vt_flat) | |
| h = F.gelu(self.dec_in(M_hat.reshape(B * N, -1))) | |
| for block in self.dec_blocks: | |
| h = h + block(h) | |
| return self.dec_out(h).reshape(B, N, -1) | |
| def encode(self, pixel_values: torch.Tensor) -> torch.Tensor: | |
| """Encode images to omega tokens (spatial latent). | |
| Args: | |
| pixel_values: (B, 3, H, W) normalized images | |
| Returns: | |
| (B, D, gh, gw) spectral latent — omega tokens | |
| For 128×128: (B, 16, 8, 8) = 1024 values, 48:1 compression | |
| """ | |
| ps = self.config.patch_size | |
| patches, gh, gw = _extract_patches(pixel_values, ps) | |
| svd = self._encode_patches_to_svd(patches) | |
| S = svd["S"] # (B, N, D) | |
| return S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw) | |
| def encode_full(self, pixel_values: torch.Tensor) -> Dict: | |
| """Encode to full SVD decomposition per patch. | |
| Returns dict with U, S_orig, S, Vt, M, gh, gw. | |
| """ | |
| ps = self.config.patch_size | |
| patches, gh, gw = _extract_patches(pixel_values, ps) | |
| svd = self._encode_patches_to_svd(patches) | |
| svd["gh"] = gh | |
| svd["gw"] = gw | |
| return svd | |
| def decode(self, latent: torch.Tensor, | |
| U: Optional[torch.Tensor] = None, | |
| Vt: Optional[torch.Tensor] = None) -> torch.Tensor: | |
| """Decode from omega tokens to images. | |
| Args: | |
| latent: (B, D, gh, gw) spectral latent | |
| U: optional (B, N, V, D) for lossless reconstruction | |
| Vt: optional (B, N, D, D) for lossless reconstruction | |
| Returns: | |
| (B, 3, H, W) reconstructed image | |
| """ | |
| B, D, gh, gw = latent.shape | |
| N = gh * gw | |
| S = latent.reshape(B, D, N).permute(0, 2, 1) | |
| if U is None or Vt is None: | |
| V = self.config.matrix_v | |
| U = torch.eye(V, D, device=latent.device, dtype=latent.dtype) | |
| U = U.unsqueeze(0).unsqueeze(0).expand(B, N, -1, -1) | |
| Vt = torch.eye(D, device=latent.device, dtype=latent.dtype) | |
| Vt = Vt.unsqueeze(0).unsqueeze(0).expand(B, N, -1, -1) | |
| decoded = self._decode_from_svd(U, S, Vt) | |
| recon = _stitch_patches(decoded, gh, gw, self.config.patch_size) | |
| return self.boundary_smooth(recon) | |
| def forward( | |
| self, | |
| pixel_values: torch.Tensor, | |
| **kwargs, | |
| ) -> Dict[str, torch.Tensor]: | |
| """Full encode → SVD → coordinate → decode pipeline. | |
| Args: | |
| pixel_values: (B, 3, H, W) normalized images | |
| Returns: | |
| dict with "recon", "latent", "svd" keys | |
| """ | |
| ps = self.config.patch_size | |
| patches, gh, gw = _extract_patches(pixel_values, ps) | |
| svd = self._encode_patches_to_svd(patches) | |
| decoded = self._decode_from_svd(svd["U"], svd["S"], svd["Vt"]) | |
| recon = _stitch_patches(decoded, gh, gw, ps) | |
| recon = self.boundary_smooth(recon) | |
| S = svd["S"] | |
| latent = S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw) | |
| return {"recon": recon, "latent": latent, "svd": svd} | |
| def effective_rank(S): | |
| p = S / (S.sum(-1, keepdim=True) + 1e-8) | |
| p = p.clamp(min=1e-8) | |
| return (-(p * p.log()).sum(-1)).exp() | |
| # Register for AutoClass — this is what makes AutoModel.from_pretrained work | |
| PatchSVAEConfig.register_for_auto_class() | |
| PatchSVAEModel.register_for_auto_class("AutoModel") |