AbstractPhil commited on
Commit
a3d42fb
·
verified ·
1 Parent(s): 3e1d693

Update modeling_patchsvae.py

Browse files
Files changed (1) hide show
  1. 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) # output.recon, output.svd
9
-
10
- # Encode to omega tokens (latent)
11
- latent = model.encode(images) # (B, D, gh, gw) = (B, 16, 8, 8)
12
 
13
- # Decode from omega tokens
14
- recon = model.decode(latent) # (B, 3, 128, 128)
15
 
16
- # Encode to full SVD components
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 dataclasses import dataclass
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
- HAS_FL = True
35
  except ImportError:
36
- HAS_FL = False
37
 
38
 
39
- def _gram_eigh_svd_fp64(A):
 
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 HAS_FL and N <= _FL_MAX_N and A.is_cuda:
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 _gram_eigh_svd_fp64(A)
69
 
70
 
71
  # ── Patch Utilities ──────────────────────────────────────────────
72
 
73
- def extract_patches(images, patch_size):
74
  B, C, H, W = images.shape
75
  gh, gw = H // patch_size, W // patch_size
76
- patches = images.reshape(B, C, gh, patch_size, gw, patch_size)
77
- patches = patches.permute(0, 2, 4, 1, 3, 5)
78
- patches = patches.reshape(B, gh * gw, C * patch_size * patch_size)
79
- return patches, gh, gw
80
 
81
 
82
- def stitch_patches(patches, gh, gw, patch_size):
83
  B = patches.shape[0]
84
- patches = patches.reshape(B, gh, gw, 3, patch_size, patch_size)
85
- patches = patches.permute(0, 3, 1, 4, 2, 5)
86
- return patches.reshape(B, 3, gh * patch_size, gw * patch_size)
87
 
88
 
89
  # ── Components ───────────────────────────────────────────────────
90
 
91
- class BoundarySmooth(nn.Module):
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 SpectralCrossAttention(nn.Module):
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
- alpha = self.alpha
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 Engine.
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
- supports_gradient_checkpointing = False
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
- self.patch_dim = 3 * ps * ps
173
- self.mat_dim = V * D
174
 
175
  # Encoder
176
- self.enc_in = nn.Linear(self.patch_dim, hidden)
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, self.mat_dim)
183
  nn.init.orthogonal_(self.enc_out.weight)
184
 
185
  # Decoder
186
- self.dec_in = nn.Linear(self.mat_dim, hidden)
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, self.patch_dim)
193
 
194
  # Cross-attention
195
  self.cross_attn = nn.ModuleList([
196
- SpectralCrossAttention(D, n_heads=min(4, D),
197
- max_alpha=config.max_alpha,
198
- alpha_init=config.alpha_init)
199
  for _ in range(config.n_cross_layers)
200
  ])
201
 
202
  # Boundary smoothing
203
- self.boundary_smooth = BoundarySmooth(channels=3, mid=16)
 
 
204
 
205
- def _encode_patches(self, patches):
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) # rows to S^(D-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 {'U': U, 'S_orig': S, 'S': S_coord, 'Vt': Vt, 'M': M}
230
 
231
- def _decode_patches(self, U, S, Vt):
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
- patches = self.dec_out(h)
244
- return patches.reshape(B, N, -1)
245
 
246
- def encode(self, images: torch.Tensor) -> torch.Tensor:
247
  """Encode images to omega tokens (spatial latent).
248
 
249
  Args:
250
- images: (B, 3, H, W) normalized images
251
 
252
  Returns:
253
- latent: (B, D, gh, gw) spectral latent — the omega tokens
254
- For 128×128 with 16×16 patches: (B, 16, 8, 8)
255
  """
256
  ps = self.config.patch_size
257
- patches, gh, gw = extract_patches(images, ps)
258
- svd = self._encode_patches(patches)
259
- S = svd['S'] # (B, N, D)
260
  return S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
261
 
262
- def encode_full(self, images: torch.Tensor) -> Dict:
263
- """Encode images to full SVD decomposition per patch.
264
 
265
- Args:
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 = extract_patches(images, ps)
274
- svd = self._encode_patches(patches)
275
- svd['gh'] = gh
276
- svd['gw'] = gw
277
  return svd
278
 
279
- def decode(self, latent: torch.Tensor, U: torch.Tensor = None,
280
- Vt: torch.Tensor = None) -> torch.Tensor:
281
- """Decode from omega tokens back to images.
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) optional left singular vectors
289
- Vt: (B, N, D, D) optional right singular vectors
290
 
291
  Returns:
292
- recon: (B, 3, H, W) reconstructed image
293
  """
294
  B, D, gh, gw = latent.shape
295
  N = gh * gw
296
- S = latent.reshape(B, D, N).permute(0, 2, 1) # (B, N, D)
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).unsqueeze(0).unsqueeze(0)
302
- U = U.expand(B, N, -1, -1)
303
- Vt = torch.eye(D, device=latent.device).unsqueeze(0).unsqueeze(0)
304
- Vt = Vt.expand(B, N, -1, -1)
305
-
306
- decoded = self._decode_patches(U, S, Vt)
307
- recon = stitch_patches(decoded, gh, gw, self.config.patch_size)
308
- recon = self.boundary_smooth(recon)
309
- return recon
310
-
311
- def forward(self, pixel_values: torch.Tensor, **kwargs) -> PatchSVAEOutput:
 
 
 
312
  """Full encode → SVD → coordinate → decode pipeline.
313
 
314
  Args:
315
  pixel_values: (B, 3, H, W) normalized images
316
 
317
  Returns:
318
- PatchSVAEOutput with recon, latent, and svd components
319
  """
320
  ps = self.config.patch_size
321
- patches, gh, gw = extract_patches(pixel_values, ps)
322
- svd = self._encode_patches(patches)
323
- decoded = self._decode_patches(svd['U'], svd['S'], svd['Vt'])
324
- recon = stitch_patches(decoded, gh, gw, ps)
325
  recon = self.boundary_smooth(recon)
326
 
327
- # Omega tokens as spatial latent
328
- S = svd['S']
329
  latent = S.permute(0, 2, 1).reshape(S.shape[0], self.config.D, gh, gw)
330
 
331
- return PatchSVAEOutput(recon=recon, latent=latent, svd=svd)
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")