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 # Warning: Pipeline type "image-to-image" is no longer supported in transformers v5. # You must load the model directly (see below) or downgrade to v4.x with: # pip install "transformers<5.0.0" from transformers import pipeline pipe = pipeline("image-to-image", model="AbstractPhil/svae-fresnel-128", trust_remote_code=True)# pip install -U transformers accelerate # 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
Update modeling_patchsvae.py
Browse files- modeling_patchsvae.py +98 -114
modeling_patchsvae.py
CHANGED
|
@@ -1,19 +1,20 @@
|
|
| 1 |
"""PatchSVAE model for HuggingFace AutoModel.
|
| 2 |
|
| 3 |
Usage:
|
| 4 |
-
from transformers import AutoModel
|
|
|
|
|
|
|
| 5 |
model = AutoModel.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True)
|
| 6 |
|
| 7 |
# Full reconstruction
|
| 8 |
-
output = model(images)
|
| 9 |
-
|
| 10 |
-
#
|
| 11 |
-
latent = model.encode(images) # (B, D, gh, gw) = (B, 16, 8, 8)
|
| 12 |
|
| 13 |
-
#
|
| 14 |
-
|
| 15 |
|
| 16 |
-
#
|
| 17 |
svd = model.encode_full(images) # dict with U, S, Vt, M per patch
|
| 18 |
"""
|
| 19 |
|
|
@@ -21,22 +22,22 @@ import math
|
|
| 21 |
import torch
|
| 22 |
import torch.nn as nn
|
| 23 |
import torch.nn.functional as F
|
| 24 |
-
from
|
| 25 |
-
from typing import Optional, Dict
|
| 26 |
from transformers import PreTrainedModel
|
| 27 |
from .configuration_patchsvae import PatchSVAEConfig
|
| 28 |
|
| 29 |
|
| 30 |
-
# ── SVD Backend ──────
|
| 31 |
|
| 32 |
try:
|
| 33 |
from geolip_core.linalg.eigh import FLEigh, _FL_MAX_N
|
| 34 |
-
|
| 35 |
except ImportError:
|
| 36 |
-
|
| 37 |
|
| 38 |
|
| 39 |
-
def
|
|
|
|
| 40 |
orig_dtype = A.dtype
|
| 41 |
with torch.amp.autocast('cuda', enabled=False):
|
| 42 |
A_d = A.double()
|
|
@@ -51,8 +52,9 @@ def _gram_eigh_svd_fp64(A):
|
|
| 51 |
|
| 52 |
|
| 53 |
def _svd_fp64(A):
|
|
|
|
| 54 |
B, M, N = A.shape
|
| 55 |
-
if
|
| 56 |
orig_dtype = A.dtype
|
| 57 |
with torch.amp.autocast('cuda', enabled=False):
|
| 58 |
A_d = A.double()
|
|
@@ -65,30 +67,29 @@ def _svd_fp64(A):
|
|
| 65 |
Vh = V.transpose(-2, -1).contiguous()
|
| 66 |
return U.to(orig_dtype), S.to(orig_dtype), Vh.to(orig_dtype)
|
| 67 |
else:
|
| 68 |
-
return
|
| 69 |
|
| 70 |
|
| 71 |
# ── Patch Utilities ──────────────────────────────────────────────
|
| 72 |
|
| 73 |
-
def
|
| 74 |
B, C, H, W = images.shape
|
| 75 |
gh, gw = H // patch_size, W // patch_size
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
return patches, gh, gw
|
| 80 |
|
| 81 |
|
| 82 |
-
def
|
| 83 |
B = patches.shape[0]
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
return
|
| 87 |
|
| 88 |
|
| 89 |
# ── Components ───────────────────────────────────────────────────
|
| 90 |
|
| 91 |
-
class
|
| 92 |
def __init__(self, channels=3, mid=16):
|
| 93 |
super().__init__()
|
| 94 |
self.net = nn.Sequential(
|
|
@@ -103,7 +104,7 @@ class BoundarySmooth(nn.Module):
|
|
| 103 |
return x + self.net(x)
|
| 104 |
|
| 105 |
|
| 106 |
-
class
|
| 107 |
def __init__(self, D, n_heads=4, max_alpha=0.2, alpha_init=-2.0):
|
| 108 |
super().__init__()
|
| 109 |
self.n_heads = n_heads
|
|
@@ -130,80 +131,68 @@ class SpectralCrossAttention(nn.Module):
|
|
| 130 |
attn = attn.softmax(dim=-1)
|
| 131 |
out = (attn @ v).transpose(1, 2).reshape(B, N, D)
|
| 132 |
gate = torch.tanh(self.out_proj(out))
|
| 133 |
-
|
| 134 |
-
return S * (1.0 + alpha.unsqueeze(0).unsqueeze(0) * gate)
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
# ── Output ───────────────────────────────────────────────────────
|
| 138 |
-
|
| 139 |
-
@dataclass
|
| 140 |
-
class PatchSVAEOutput:
|
| 141 |
-
"""Output from PatchSVAE forward pass."""
|
| 142 |
-
recon: torch.Tensor # (B, 3, H, W) reconstructed image
|
| 143 |
-
latent: torch.Tensor # (B, D, gh, gw) omega tokens
|
| 144 |
-
svd: Optional[Dict] = None # full SVD components if requested
|
| 145 |
|
| 146 |
|
| 147 |
# ── Model ────────────────────────────────────────────────────────
|
| 148 |
|
| 149 |
class PatchSVAEModel(PreTrainedModel):
|
| 150 |
-
"""Patch-based SVD Autoencoder — The Geometric
|
| 151 |
|
| 152 |
Decomposes images into patches, encodes each to a sphere-normalized
|
| 153 |
matrix, performs SVD, coordinates spectra via cross-attention,
|
| 154 |
-
and reconstructs.
|
| 155 |
|
| 156 |
The spectral vectors S form omega tokens: modality-agnostic,
|
| 157 |
geometrically structured, universal representations.
|
| 158 |
"""
|
| 159 |
config_class = PatchSVAEConfig
|
| 160 |
-
|
| 161 |
|
| 162 |
def __init__(self, config: PatchSVAEConfig):
|
| 163 |
super().__init__(config)
|
| 164 |
-
self.config = config
|
| 165 |
|
| 166 |
V = config.matrix_v
|
| 167 |
D = config.D
|
| 168 |
hidden = config.hidden
|
| 169 |
depth = config.depth
|
| 170 |
ps = config.patch_size
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
self.mat_dim = V * D
|
| 174 |
|
| 175 |
# Encoder
|
| 176 |
-
self.enc_in = nn.Linear(
|
| 177 |
self.enc_blocks = nn.ModuleList([
|
| 178 |
nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden),
|
| 179 |
nn.GELU(), nn.Linear(hidden, hidden))
|
| 180 |
for _ in range(depth)
|
| 181 |
])
|
| 182 |
-
self.enc_out = nn.Linear(hidden,
|
| 183 |
nn.init.orthogonal_(self.enc_out.weight)
|
| 184 |
|
| 185 |
# Decoder
|
| 186 |
-
self.dec_in = nn.Linear(
|
| 187 |
self.dec_blocks = nn.ModuleList([
|
| 188 |
nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden),
|
| 189 |
nn.GELU(), nn.Linear(hidden, hidden))
|
| 190 |
for _ in range(depth)
|
| 191 |
])
|
| 192 |
-
self.dec_out = nn.Linear(hidden,
|
| 193 |
|
| 194 |
# Cross-attention
|
| 195 |
self.cross_attn = nn.ModuleList([
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
for _ in range(config.n_cross_layers)
|
| 200 |
])
|
| 201 |
|
| 202 |
# Boundary smoothing
|
| 203 |
-
self.boundary_smooth =
|
|
|
|
|
|
|
| 204 |
|
| 205 |
-
def
|
| 206 |
-
"""Encode flat patches to SVD components."""
|
| 207 |
B, N, _ = patches.shape
|
| 208 |
V, D = self.config.matrix_v, self.config.D
|
| 209 |
|
|
@@ -212,7 +201,7 @@ class PatchSVAEModel(PreTrainedModel):
|
|
| 212 |
for block in self.enc_blocks:
|
| 213 |
h = h + block(h)
|
| 214 |
M = self.enc_out(h).reshape(B * N, V, D)
|
| 215 |
-
M = F.normalize(M, dim=-1)
|
| 216 |
|
| 217 |
U, S, Vt = _svd_fp64(M)
|
| 218 |
|
|
@@ -221,17 +210,14 @@ class PatchSVAEModel(PreTrainedModel):
|
|
| 221 |
Vt = Vt.reshape(B, N, D, D)
|
| 222 |
M = M.reshape(B, N, V, D)
|
| 223 |
|
| 224 |
-
# Cross-patch spectral coordination
|
| 225 |
S_coord = S
|
| 226 |
for layer in self.cross_attn:
|
| 227 |
S_coord = layer(S_coord)
|
| 228 |
|
| 229 |
-
return {
|
| 230 |
|
| 231 |
-
def
|
| 232 |
-
"""Decode from SVD components to flat patches."""
|
| 233 |
B, N, V, D = U.shape
|
| 234 |
-
|
| 235 |
U_flat = U.reshape(B * N, V, D)
|
| 236 |
S_flat = S.reshape(B * N, D)
|
| 237 |
Vt_flat = Vt.reshape(B * N, D, D)
|
|
@@ -240,98 +226,96 @@ class PatchSVAEModel(PreTrainedModel):
|
|
| 240 |
h = F.gelu(self.dec_in(M_hat.reshape(B * N, -1)))
|
| 241 |
for block in self.dec_blocks:
|
| 242 |
h = h + block(h)
|
| 243 |
-
|
| 244 |
-
return patches.reshape(B, N, -1)
|
| 245 |
|
| 246 |
-
def encode(self,
|
| 247 |
"""Encode images to omega tokens (spatial latent).
|
| 248 |
|
| 249 |
Args:
|
| 250 |
-
|
| 251 |
|
| 252 |
Returns:
|
| 253 |
-
|
| 254 |
-
|
| 255 |
"""
|
| 256 |
ps = self.config.patch_size
|
| 257 |
-
patches, gh, gw =
|
| 258 |
-
svd = self.
|
| 259 |
-
S = svd[
|
| 260 |
return S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
|
| 261 |
|
| 262 |
-
def encode_full(self,
|
| 263 |
-
"""Encode
|
| 264 |
|
| 265 |
-
|
| 266 |
-
images: (B, 3, H, W) normalized images
|
| 267 |
-
|
| 268 |
-
Returns:
|
| 269 |
-
dict with U, S_orig, S (coordinated), Vt, M per patch,
|
| 270 |
-
plus gh, gw grid dimensions.
|
| 271 |
"""
|
| 272 |
ps = self.config.patch_size
|
| 273 |
-
patches, gh, gw =
|
| 274 |
-
svd = self.
|
| 275 |
-
svd[
|
| 276 |
-
svd[
|
| 277 |
return svd
|
| 278 |
|
| 279 |
-
def decode(self, latent: torch.Tensor,
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
If U and Vt are provided, uses them for full reconstruction.
|
| 284 |
-
If only latent is provided, uses identity directions (lossy).
|
| 285 |
|
| 286 |
Args:
|
| 287 |
latent: (B, D, gh, gw) spectral latent
|
| 288 |
-
U: (B, N, V, D)
|
| 289 |
-
Vt: (B, N, D, D)
|
| 290 |
|
| 291 |
Returns:
|
| 292 |
-
|
| 293 |
"""
|
| 294 |
B, D, gh, gw = latent.shape
|
| 295 |
N = gh * gw
|
| 296 |
-
S = latent.reshape(B, D, N).permute(0, 2, 1)
|
| 297 |
|
| 298 |
if U is None or Vt is None:
|
| 299 |
-
# Lossy decode: construct approximate U, Vt from identity
|
| 300 |
V = self.config.matrix_v
|
| 301 |
-
U = torch.eye(V, D, device=latent.device
|
| 302 |
-
U = U.expand(B, N, -1, -1)
|
| 303 |
-
Vt = torch.eye(D, device=latent.device
|
| 304 |
-
Vt = Vt.expand(B, N, -1, -1)
|
| 305 |
-
|
| 306 |
-
decoded = self.
|
| 307 |
-
recon =
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
|
|
|
|
|
|
|
|
|
| 312 |
"""Full encode → SVD → coordinate → decode pipeline.
|
| 313 |
|
| 314 |
Args:
|
| 315 |
pixel_values: (B, 3, H, W) normalized images
|
| 316 |
|
| 317 |
Returns:
|
| 318 |
-
|
| 319 |
"""
|
| 320 |
ps = self.config.patch_size
|
| 321 |
-
patches, gh, gw =
|
| 322 |
-
svd = self.
|
| 323 |
-
decoded = self.
|
| 324 |
-
recon =
|
| 325 |
recon = self.boundary_smooth(recon)
|
| 326 |
|
| 327 |
-
|
| 328 |
-
S = svd['S']
|
| 329 |
latent = S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
|
| 330 |
|
| 331 |
-
return
|
| 332 |
|
| 333 |
@staticmethod
|
| 334 |
def effective_rank(S):
|
| 335 |
p = S / (S.sum(-1, keepdim=True) + 1e-8)
|
| 336 |
p = p.clamp(min=1e-8)
|
| 337 |
-
return (-(p * p.log()).sum(-1)).exp()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
"""PatchSVAE model for HuggingFace AutoModel.
|
| 2 |
|
| 3 |
Usage:
|
| 4 |
+
from transformers import AutoConfig, AutoModel
|
| 5 |
+
|
| 6 |
+
config = AutoConfig.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True)
|
| 7 |
model = AutoModel.from_pretrained("AbstractPhil/svae-fresnel-128", trust_remote_code=True)
|
| 8 |
|
| 9 |
# Full reconstruction
|
| 10 |
+
output = model(images)
|
| 11 |
+
recon = output["recon"] # (B, 3, 128, 128)
|
| 12 |
+
latent = output["latent"] # (B, 16, 8, 8) omega tokens
|
|
|
|
| 13 |
|
| 14 |
+
# Encode to omega tokens
|
| 15 |
+
omega = model.encode(images) # (B, 16, 8, 8)
|
| 16 |
|
| 17 |
+
# Full SVD decomposition
|
| 18 |
svd = model.encode_full(images) # dict with U, S, Vt, M per patch
|
| 19 |
"""
|
| 20 |
|
|
|
|
| 22 |
import torch
|
| 23 |
import torch.nn as nn
|
| 24 |
import torch.nn.functional as F
|
| 25 |
+
from typing import Optional, Dict, Union
|
|
|
|
| 26 |
from transformers import PreTrainedModel
|
| 27 |
from .configuration_patchsvae import PatchSVAEConfig
|
| 28 |
|
| 29 |
|
| 30 |
+
# ── SVD Backend (self-contained, no external deps required) ──────
|
| 31 |
|
| 32 |
try:
|
| 33 |
from geolip_core.linalg.eigh import FLEigh, _FL_MAX_N
|
| 34 |
+
_HAS_FL = True
|
| 35 |
except ImportError:
|
| 36 |
+
_HAS_FL = False
|
| 37 |
|
| 38 |
|
| 39 |
+
def _gram_eigh_svd(A):
|
| 40 |
+
"""Thin SVD via Gram + eigh in fp64."""
|
| 41 |
orig_dtype = A.dtype
|
| 42 |
with torch.amp.autocast('cuda', enabled=False):
|
| 43 |
A_d = A.double()
|
|
|
|
| 52 |
|
| 53 |
|
| 54 |
def _svd_fp64(A):
|
| 55 |
+
"""Auto-dispatch: FL eigh for N<=12, Gram eigh otherwise."""
|
| 56 |
B, M, N = A.shape
|
| 57 |
+
if _HAS_FL and N <= _FL_MAX_N and A.is_cuda:
|
| 58 |
orig_dtype = A.dtype
|
| 59 |
with torch.amp.autocast('cuda', enabled=False):
|
| 60 |
A_d = A.double()
|
|
|
|
| 67 |
Vh = V.transpose(-2, -1).contiguous()
|
| 68 |
return U.to(orig_dtype), S.to(orig_dtype), Vh.to(orig_dtype)
|
| 69 |
else:
|
| 70 |
+
return _gram_eigh_svd(A)
|
| 71 |
|
| 72 |
|
| 73 |
# ── Patch Utilities ──────────────────────────────────────────────
|
| 74 |
|
| 75 |
+
def _extract_patches(images, patch_size):
|
| 76 |
B, C, H, W = images.shape
|
| 77 |
gh, gw = H // patch_size, W // patch_size
|
| 78 |
+
x = images.reshape(B, C, gh, patch_size, gw, patch_size)
|
| 79 |
+
x = x.permute(0, 2, 4, 1, 3, 5)
|
| 80 |
+
return x.reshape(B, gh * gw, C * patch_size * patch_size), gh, gw
|
|
|
|
| 81 |
|
| 82 |
|
| 83 |
+
def _stitch_patches(patches, gh, gw, patch_size):
|
| 84 |
B = patches.shape[0]
|
| 85 |
+
x = patches.reshape(B, gh, gw, 3, patch_size, patch_size)
|
| 86 |
+
x = x.permute(0, 3, 1, 4, 2, 5)
|
| 87 |
+
return x.reshape(B, 3, gh * patch_size, gw * patch_size)
|
| 88 |
|
| 89 |
|
| 90 |
# ── Components ───────────────────────────────────────────────────
|
| 91 |
|
| 92 |
+
class _BoundarySmooth(nn.Module):
|
| 93 |
def __init__(self, channels=3, mid=16):
|
| 94 |
super().__init__()
|
| 95 |
self.net = nn.Sequential(
|
|
|
|
| 104 |
return x + self.net(x)
|
| 105 |
|
| 106 |
|
| 107 |
+
class _SpectralCrossAttention(nn.Module):
|
| 108 |
def __init__(self, D, n_heads=4, max_alpha=0.2, alpha_init=-2.0):
|
| 109 |
super().__init__()
|
| 110 |
self.n_heads = n_heads
|
|
|
|
| 131 |
attn = attn.softmax(dim=-1)
|
| 132 |
out = (attn @ v).transpose(1, 2).reshape(B, N, D)
|
| 133 |
gate = torch.tanh(self.out_proj(out))
|
| 134 |
+
return S * (1.0 + self.alpha.unsqueeze(0).unsqueeze(0) * gate)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 135 |
|
| 136 |
|
| 137 |
# ── Model ────────────────────────────────────────────────────────
|
| 138 |
|
| 139 |
class PatchSVAEModel(PreTrainedModel):
|
| 140 |
+
"""Patch-based SVD Autoencoder — The Fresnel Geometric Compression Lens.
|
| 141 |
|
| 142 |
Decomposes images into patches, encodes each to a sphere-normalized
|
| 143 |
matrix, performs SVD, coordinates spectra via cross-attention,
|
| 144 |
+
and reconstructs with 99.993% fidelity.
|
| 145 |
|
| 146 |
The spectral vectors S form omega tokens: modality-agnostic,
|
| 147 |
geometrically structured, universal representations.
|
| 148 |
"""
|
| 149 |
config_class = PatchSVAEConfig
|
| 150 |
+
_tied_weights_keys = []
|
| 151 |
|
| 152 |
def __init__(self, config: PatchSVAEConfig):
|
| 153 |
super().__init__(config)
|
|
|
|
| 154 |
|
| 155 |
V = config.matrix_v
|
| 156 |
D = config.D
|
| 157 |
hidden = config.hidden
|
| 158 |
depth = config.depth
|
| 159 |
ps = config.patch_size
|
| 160 |
+
patch_dim = 3 * ps * ps
|
| 161 |
+
mat_dim = V * D
|
|
|
|
| 162 |
|
| 163 |
# Encoder
|
| 164 |
+
self.enc_in = nn.Linear(patch_dim, hidden)
|
| 165 |
self.enc_blocks = nn.ModuleList([
|
| 166 |
nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden),
|
| 167 |
nn.GELU(), nn.Linear(hidden, hidden))
|
| 168 |
for _ in range(depth)
|
| 169 |
])
|
| 170 |
+
self.enc_out = nn.Linear(hidden, mat_dim)
|
| 171 |
nn.init.orthogonal_(self.enc_out.weight)
|
| 172 |
|
| 173 |
# Decoder
|
| 174 |
+
self.dec_in = nn.Linear(mat_dim, hidden)
|
| 175 |
self.dec_blocks = nn.ModuleList([
|
| 176 |
nn.Sequential(nn.LayerNorm(hidden), nn.Linear(hidden, hidden),
|
| 177 |
nn.GELU(), nn.Linear(hidden, hidden))
|
| 178 |
for _ in range(depth)
|
| 179 |
])
|
| 180 |
+
self.dec_out = nn.Linear(hidden, patch_dim)
|
| 181 |
|
| 182 |
# Cross-attention
|
| 183 |
self.cross_attn = nn.ModuleList([
|
| 184 |
+
_SpectralCrossAttention(D, n_heads=min(4, D),
|
| 185 |
+
max_alpha=config.max_alpha,
|
| 186 |
+
alpha_init=config.alpha_init)
|
| 187 |
for _ in range(config.n_cross_layers)
|
| 188 |
])
|
| 189 |
|
| 190 |
# Boundary smoothing
|
| 191 |
+
self.boundary_smooth = _BoundarySmooth(channels=3, mid=16)
|
| 192 |
+
|
| 193 |
+
self.post_init()
|
| 194 |
|
| 195 |
+
def _encode_patches_to_svd(self, patches):
|
|
|
|
| 196 |
B, N, _ = patches.shape
|
| 197 |
V, D = self.config.matrix_v, self.config.D
|
| 198 |
|
|
|
|
| 201 |
for block in self.enc_blocks:
|
| 202 |
h = h + block(h)
|
| 203 |
M = self.enc_out(h).reshape(B * N, V, D)
|
| 204 |
+
M = F.normalize(M, dim=-1)
|
| 205 |
|
| 206 |
U, S, Vt = _svd_fp64(M)
|
| 207 |
|
|
|
|
| 210 |
Vt = Vt.reshape(B, N, D, D)
|
| 211 |
M = M.reshape(B, N, V, D)
|
| 212 |
|
|
|
|
| 213 |
S_coord = S
|
| 214 |
for layer in self.cross_attn:
|
| 215 |
S_coord = layer(S_coord)
|
| 216 |
|
| 217 |
+
return {"U": U, "S_orig": S, "S": S_coord, "Vt": Vt, "M": M}
|
| 218 |
|
| 219 |
+
def _decode_from_svd(self, U, S, Vt):
|
|
|
|
| 220 |
B, N, V, D = U.shape
|
|
|
|
| 221 |
U_flat = U.reshape(B * N, V, D)
|
| 222 |
S_flat = S.reshape(B * N, D)
|
| 223 |
Vt_flat = Vt.reshape(B * N, D, D)
|
|
|
|
| 226 |
h = F.gelu(self.dec_in(M_hat.reshape(B * N, -1)))
|
| 227 |
for block in self.dec_blocks:
|
| 228 |
h = h + block(h)
|
| 229 |
+
return self.dec_out(h).reshape(B, N, -1)
|
|
|
|
| 230 |
|
| 231 |
+
def encode(self, pixel_values: torch.Tensor) -> torch.Tensor:
|
| 232 |
"""Encode images to omega tokens (spatial latent).
|
| 233 |
|
| 234 |
Args:
|
| 235 |
+
pixel_values: (B, 3, H, W) normalized images
|
| 236 |
|
| 237 |
Returns:
|
| 238 |
+
(B, D, gh, gw) spectral latent — omega tokens
|
| 239 |
+
For 128×128: (B, 16, 8, 8) = 1024 values, 48:1 compression
|
| 240 |
"""
|
| 241 |
ps = self.config.patch_size
|
| 242 |
+
patches, gh, gw = _extract_patches(pixel_values, ps)
|
| 243 |
+
svd = self._encode_patches_to_svd(patches)
|
| 244 |
+
S = svd["S"] # (B, N, D)
|
| 245 |
return S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
|
| 246 |
|
| 247 |
+
def encode_full(self, pixel_values: torch.Tensor) -> Dict:
|
| 248 |
+
"""Encode to full SVD decomposition per patch.
|
| 249 |
|
| 250 |
+
Returns dict with U, S_orig, S, Vt, M, gh, gw.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 251 |
"""
|
| 252 |
ps = self.config.patch_size
|
| 253 |
+
patches, gh, gw = _extract_patches(pixel_values, ps)
|
| 254 |
+
svd = self._encode_patches_to_svd(patches)
|
| 255 |
+
svd["gh"] = gh
|
| 256 |
+
svd["gw"] = gw
|
| 257 |
return svd
|
| 258 |
|
| 259 |
+
def decode(self, latent: torch.Tensor,
|
| 260 |
+
U: Optional[torch.Tensor] = None,
|
| 261 |
+
Vt: Optional[torch.Tensor] = None) -> torch.Tensor:
|
| 262 |
+
"""Decode from omega tokens to images.
|
|
|
|
|
|
|
| 263 |
|
| 264 |
Args:
|
| 265 |
latent: (B, D, gh, gw) spectral latent
|
| 266 |
+
U: optional (B, N, V, D) for lossless reconstruction
|
| 267 |
+
Vt: optional (B, N, D, D) for lossless reconstruction
|
| 268 |
|
| 269 |
Returns:
|
| 270 |
+
(B, 3, H, W) reconstructed image
|
| 271 |
"""
|
| 272 |
B, D, gh, gw = latent.shape
|
| 273 |
N = gh * gw
|
| 274 |
+
S = latent.reshape(B, D, N).permute(0, 2, 1)
|
| 275 |
|
| 276 |
if U is None or Vt is None:
|
|
|
|
| 277 |
V = self.config.matrix_v
|
| 278 |
+
U = torch.eye(V, D, device=latent.device, dtype=latent.dtype)
|
| 279 |
+
U = U.unsqueeze(0).unsqueeze(0).expand(B, N, -1, -1)
|
| 280 |
+
Vt = torch.eye(D, device=latent.device, dtype=latent.dtype)
|
| 281 |
+
Vt = Vt.unsqueeze(0).unsqueeze(0).expand(B, N, -1, -1)
|
| 282 |
+
|
| 283 |
+
decoded = self._decode_from_svd(U, S, Vt)
|
| 284 |
+
recon = _stitch_patches(decoded, gh, gw, self.config.patch_size)
|
| 285 |
+
return self.boundary_smooth(recon)
|
| 286 |
+
|
| 287 |
+
def forward(
|
| 288 |
+
self,
|
| 289 |
+
pixel_values: torch.Tensor,
|
| 290 |
+
**kwargs,
|
| 291 |
+
) -> Dict[str, torch.Tensor]:
|
| 292 |
"""Full encode → SVD → coordinate → decode pipeline.
|
| 293 |
|
| 294 |
Args:
|
| 295 |
pixel_values: (B, 3, H, W) normalized images
|
| 296 |
|
| 297 |
Returns:
|
| 298 |
+
dict with "recon", "latent", "svd" keys
|
| 299 |
"""
|
| 300 |
ps = self.config.patch_size
|
| 301 |
+
patches, gh, gw = _extract_patches(pixel_values, ps)
|
| 302 |
+
svd = self._encode_patches_to_svd(patches)
|
| 303 |
+
decoded = self._decode_from_svd(svd["U"], svd["S"], svd["Vt"])
|
| 304 |
+
recon = _stitch_patches(decoded, gh, gw, ps)
|
| 305 |
recon = self.boundary_smooth(recon)
|
| 306 |
|
| 307 |
+
S = svd["S"]
|
|
|
|
| 308 |
latent = S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
|
| 309 |
|
| 310 |
+
return {"recon": recon, "latent": latent, "svd": svd}
|
| 311 |
|
| 312 |
@staticmethod
|
| 313 |
def effective_rank(S):
|
| 314 |
p = S / (S.sum(-1, keepdim=True) + 1e-8)
|
| 315 |
p = p.clamp(min=1e-8)
|
| 316 |
+
return (-(p * p.log()).sum(-1)).exp()
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
# Register for AutoClass — this is what makes AutoModel.from_pretrained work
|
| 320 |
+
PatchSVAEConfig.register_for_auto_class()
|
| 321 |
+
PatchSVAEModel.register_for_auto_class("AutoModel")
|