simit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,130 @@
1
+ """Fused Triton kernels for the MoT decoder's elementwise ops.
2
+
3
+ Each kernel reproduces the eager PyTorch op sequence including its
4
+ intermediate roundings to the activation dtype (one pass over memory instead
5
+ of 3-7), so outputs match the eager modules up to fp32 reduction order.
6
+ Without Triton (or on CPU) the eager versions are used.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import torch
12
+ import torch.nn.functional as F
13
+
14
+ try:
15
+ import triton
16
+ import triton.language as tl
17
+ _HAVE_TRITON = True
18
+ except Exception: # pragma: no cover
19
+ _HAVE_TRITON = False
20
+
21
+
22
+ def _enabled(x: torch.Tensor) -> bool:
23
+ return _HAVE_TRITON and x.is_cuda and x.dtype == torch.bfloat16
24
+
25
+
26
+ if _HAVE_TRITON:
27
+ @triton.jit
28
+ def _bf16(x):
29
+ """Round fp32 to the nearest-even bfloat16, kept in fp32. Bit-level, because
30
+ the compiler folds a plain fp32->bf16->fp32 cast pair away."""
31
+ b = x.to(tl.uint32, bitcast=True)
32
+ b = b + (0x7FFF + ((b >> 16) & 1))
33
+ return ((b >> 16) << 16).to(tl.float32, bitcast=True)
34
+
35
+ @triton.jit
36
+ def _rms_norm_kernel(X, W, Y, M, N, eps, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
37
+ rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
38
+ cols = tl.arange(0, BLOCK_N)
39
+ mask = (rows[:, None] < M) & (cols[None, :] < N)
40
+ offs = rows[:, None].to(tl.int64) * N + cols[None, :]
41
+ x = tl.load(X + offs, mask=mask, other=0.0).to(tl.float32)
42
+ rstd = 1.0 / tl.sqrt(tl.sum(x * x, axis=1) / N + eps)
43
+ normed = _bf16(x * rstd[:, None]) # x.to(dtype)
44
+ w = tl.load(W + cols, mask=cols < N, other=0.0).to(tl.float32)
45
+ y = w[None, :] * normed
46
+ if Y.dtype.element_ty == tl.bfloat16:
47
+ y = _bf16(y)
48
+ tl.store(Y + offs, y.to(Y.dtype.element_ty), mask=mask)
49
+
50
+ @triton.jit
51
+ def _rope_kernel(X, COS, SIN, Y, M, H, D: tl.constexpr, BLOCK_M: tl.constexpr):
52
+ # X: [T*H, D] rows; COS/SIN: [T, D]; out = x*cos + rotate_half(x)*sin
53
+ rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
54
+ half: tl.constexpr = D // 2
55
+ cols = tl.arange(0, half)
56
+ rmask = rows[:, None] < M
57
+ base = rows[:, None].to(tl.int64) * D
58
+ t = (rows // H)[:, None].to(tl.int64) * D
59
+ x1 = tl.load(X + base + cols[None, :], mask=rmask, other=0.0)
60
+ x2 = tl.load(X + base + half + cols[None, :], mask=rmask, other=0.0)
61
+ c1 = tl.load(COS + t + cols[None, :], mask=rmask, other=0.0)
62
+ c2 = tl.load(COS + t + half + cols[None, :], mask=rmask, other=0.0)
63
+ s1 = tl.load(SIN + t + cols[None, :], mask=rmask, other=0.0)
64
+ s2 = tl.load(SIN + t + half + cols[None, :], mask=rmask, other=0.0)
65
+ a1 = _bf16(x1.to(tl.float32) * c1.to(tl.float32))
66
+ a2 = _bf16(x2.to(tl.float32) * c2.to(tl.float32))
67
+ b1 = _bf16(-x2.to(tl.float32) * s1.to(tl.float32))
68
+ b2 = _bf16(x1.to(tl.float32) * s2.to(tl.float32))
69
+ tl.store(Y + base + cols[None, :], _bf16(a1 + b1).to(tl.bfloat16), mask=rmask)
70
+ tl.store(Y + base + half + cols[None, :], _bf16(a2 + b2).to(tl.bfloat16), mask=rmask)
71
+
72
+ @triton.jit
73
+ def _silu_mul_kernel(G, U, Y, n, BLOCK: tl.constexpr):
74
+ offs = tl.program_id(0).to(tl.int64) * BLOCK + tl.arange(0, BLOCK)
75
+ mask = offs < n
76
+ g = tl.load(G + offs, mask=mask, other=0.0).to(tl.float32)
77
+ u = tl.load(U + offs, mask=mask, other=0.0).to(tl.float32)
78
+ s = _bf16(g / (1.0 + tl.exp(-g))) # F.silu(g)
79
+ tl.store(Y + offs, _bf16(s * u).to(tl.bfloat16), mask=mask)
80
+
81
+
82
+ def rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
83
+ """``weight * (x.float() * rsqrt(mean(x^2) + eps)).to(x.dtype)``"""
84
+ if not _enabled(x) or x.numel() == 0:
85
+ dtype = x.dtype
86
+ xf = x.float()
87
+ xf = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + eps)
88
+ return weight * xf.to(dtype)
89
+ N = x.shape[-1]
90
+ x2 = x.contiguous().view(-1, N)
91
+ y = torch.empty(x.shape, dtype=torch.promote_types(weight.dtype, x.dtype), device=x.device)
92
+ M = x2.shape[0]
93
+ BLOCK_N = triton.next_power_of_2(N)
94
+ BLOCK_M = max(1, min(64, 4096 // BLOCK_N))
95
+ _rms_norm_kernel[(triton.cdiv(M, BLOCK_M),)](x2, weight, y, M, N, eps, BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N)
96
+ return y
97
+
98
+
99
+ def _rotate_half(x):
100
+ x1, x2 = x.chunk(2, dim=-1)
101
+ return torch.cat((-x2, x1), dim=-1)
102
+
103
+
104
+ def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
105
+ """``x * cos + rotate_half(x) * sin`` for x [T, H, D] and cos/sin [T, D]."""
106
+ D = x.shape[-1]
107
+ if (not _enabled(x) or x.numel() == 0 or cos.dtype != x.dtype or D & (D - 1)
108
+ or cos.shape != (x.shape[0], D)):
109
+ c, s = cos[:, None], sin[:, None]
110
+ return x * c + _rotate_half(x) * s
111
+ T, H = x.shape[0], x.shape[1]
112
+ x = x.contiguous()
113
+ cos, sin = cos.contiguous(), sin.contiguous()
114
+ y = torch.empty_like(x)
115
+ M = T * H
116
+ BLOCK_M = 32
117
+ _rope_kernel[(triton.cdiv(M, BLOCK_M),)](x, cos, sin, y, M, H, D=D, BLOCK_M=BLOCK_M)
118
+ return y
119
+
120
+
121
+ def silu_mul(g: torch.Tensor, u: torch.Tensor) -> torch.Tensor:
122
+ """``F.silu(g) * u``"""
123
+ if not _enabled(g) or g.numel() == 0 or g.dtype != u.dtype:
124
+ return F.silu(g) * u
125
+ g, u = g.contiguous(), u.contiguous()
126
+ y = torch.empty_like(g)
127
+ n = g.numel()
128
+ BLOCK = 2048
129
+ _silu_mul_kernel[(triton.cdiv(n, BLOCK),)](g, u, y, n, BLOCK=BLOCK)
130
+ return y
@@ -0,0 +1,514 @@
1
+ """Inference-only modules for BAGEL-style Mixture-of-Transformers UMMs.
2
+
3
+ Re-implemented for batched serving on PyTorch's built-in variable-length
4
+ attention (no flash-attn install needed). Parameter names match the
5
+ released BAGEL checkpoint (``language_model.model.layers.*`` ...), which
6
+ Lance reuses for its language model.
7
+
8
+ The language model processes one flat batch of tokens per step. Tokens are
9
+ grouped into "pseudo-sequences" (a decode token, a prefill chunk, an image
10
+ block) that attend to a contiguous region of a shared KV arena; causal and
11
+ bidirectional pseudo-sequences are attended in two calls per layer.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import math
17
+ from dataclasses import dataclass
18
+ from typing import Optional
19
+
20
+ import numpy as np
21
+ import torch
22
+ import torch.nn.functional as F
23
+ from torch import nn
24
+ from torch.nn.attention.varlen import varlen_attn
25
+
26
+ from .kernels import apply_rope, rms_norm, silu_mul
27
+
28
+
29
+ class DeviceMixin:
30
+ """``.to(device)`` for model containers that hold several modules (load on
31
+ CPU once, move to the GPU per request -- e.g. on Hugging Face ZeroGPU)."""
32
+
33
+ def to(self, device):
34
+ device = torch.device(device)
35
+ for name, value in list(vars(self).items()):
36
+ if isinstance(value, nn.Module):
37
+ value.to(device)
38
+ elif isinstance(value, torch.Tensor):
39
+ setattr(self, name, value.to(device))
40
+ self.device = device
41
+ return self
42
+
43
+
44
+ # ----------------------------------------------------------------- basic layers
45
+ class RMSNorm(nn.Module):
46
+ def __init__(self, dim: int, eps: float = 1e-6):
47
+ super().__init__()
48
+ self.weight = nn.Parameter(torch.ones(dim))
49
+ self.eps = eps
50
+
51
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
52
+ return rms_norm(x, self.weight, self.eps)
53
+
54
+
55
+ class SwiGLU(nn.Module):
56
+ def __init__(self, hidden: int, inter: int):
57
+ super().__init__()
58
+ self.gate_proj = nn.Linear(hidden, inter, bias=False)
59
+ self.up_proj = nn.Linear(hidden, inter, bias=False)
60
+ self.down_proj = nn.Linear(inter, hidden, bias=False)
61
+
62
+ def forward(self, x):
63
+ return self.down_proj(silu_mul(self.gate_proj(x), self.up_proj(x)))
64
+
65
+
66
+ class RotaryEmbedding(nn.Module):
67
+ """1-D RoPE with cached cos/sin tables (positions are integers)."""
68
+
69
+ def __init__(self, head_dim: int, theta: float, max_positions: int = 65536):
70
+ super().__init__()
71
+ inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
72
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
73
+ self.max_positions = 0
74
+ self._cos = self._sin = None
75
+ self._grow(max_positions)
76
+
77
+ def _grow(self, n: int):
78
+ t = torch.arange(n, dtype=torch.float32, device=self.inv_freq.device)
79
+ freqs = torch.outer(t, self.inv_freq)
80
+ emb = torch.cat((freqs, freqs), dim=-1)
81
+ self._cos, self._sin = emb.cos(), emb.sin()
82
+ self.max_positions = n
83
+
84
+ def forward(self, positions: torch.Tensor, dtype):
85
+ if self._cos.device != positions.device:
86
+ self._cos, self._sin = self._cos.to(positions.device), self._sin.to(positions.device)
87
+ if int(positions.max()) >= self.max_positions:
88
+ self._grow(int(positions.max()) * 2)
89
+ self._cos, self._sin = self._cos.to(positions.device), self._sin.to(positions.device)
90
+ return self._cos[positions].to(dtype), self._sin[positions].to(dtype)
91
+
92
+
93
+ class MRotaryEmbedding(RotaryEmbedding):
94
+ """Qwen2-VL multimodal RoPE: positions are (t, h, w) triples and the head
95
+ dimension is split into ``sections`` (e.g. [16, 24, 24]) per axis."""
96
+
97
+ def __init__(self, head_dim: int, theta: float, sections, max_positions: int = 65536):
98
+ super().__init__(head_dim, theta, max_positions)
99
+ axis = []
100
+ for i, n in enumerate(list(sections) * 2):
101
+ axis += [i % 3] * n
102
+ assert len(axis) == head_dim, (sections, head_dim)
103
+ self.register_buffer("axis_of_dim", torch.tensor(axis), persistent=False)
104
+
105
+ def forward(self, positions: torch.Tensor, dtype):
106
+ if positions.dim() == 1:
107
+ return super().forward(positions, dtype)
108
+ cos, sin = super().forward(positions.reshape(-1), dtype) # [3*T, D]
109
+ T, D = positions.shape[1], cos.shape[-1]
110
+ cos, sin = cos.view(3, T, D), sin.view(3, T, D)
111
+ idx = self.axis_of_dim.to(positions.device)[None, None, :].expand(1, T, D)
112
+ return cos.gather(0, idx)[0], sin.gather(0, idx)[0]
113
+
114
+
115
+ # --------------------------------------------------------------- batch metadata
116
+ @dataclass
117
+ class AttnBatch:
118
+ """Where every token of a step lives and what it may attend to.
119
+
120
+ Tokens are ordered causal pseudo-sequences first (``n_causal`` tokens),
121
+ then bidirectional ones. For each group: ``cu_q`` query offsets, ``k_start``
122
+ arena start of the region, ``k_used`` keys visible (prefix + chunk).
123
+ Generation-expert tokens (latents) come last: rows ``[n_und:]``, so the
124
+ experts run on contiguous slices instead of gathered rows."""
125
+
126
+ slots: torch.Tensor # [T] arena index each token's K/V is written to
127
+ n_causal: int
128
+ causal: Optional[tuple] = None # (cu_q, k_start_cu, k_used, max_q, max_k)
129
+ full: Optional[tuple] = None
130
+ n_und: Optional[int] = None # None -> all tokens use the understanding expert
131
+
132
+
133
+ class KVArena:
134
+ """Per-layer contiguous K/V storage: ``k[layer]`` is [capacity, Hkv, D]."""
135
+
136
+ def __init__(self, num_layers: int, capacity: int, num_kv_heads: int, head_dim: int, dtype, device):
137
+ self.k = torch.empty(num_layers, capacity, num_kv_heads, head_dim, dtype=dtype, device=device)
138
+ self.v = torch.empty_like(self.k)
139
+ self.capacity = capacity
140
+
141
+ def copy(self, src: int, dst: int, n: int):
142
+ self.k[:, dst:dst + n].copy_(self.k[:, src:src + n])
143
+ self.v[:, dst:dst + n].copy_(self.v[:, src:src + n])
144
+
145
+
146
+ def _attend(q, arena_k, arena_v, meta, window, scale):
147
+ cu_q, k_start, k_used, max_q, max_k = meta
148
+ return varlen_attn(q, arena_k, arena_v, cu_q, k_start, max_q, max_k, window_size=window,
149
+ enable_gqa=q.shape[1] != arena_k.shape[1], seqused_k=k_used, scale=scale)
150
+
151
+
152
+ # --------------------------------------------------------------- MoT decoder
153
+ def _split(fn_und, fn_gen, x, n_und: Optional[int]):
154
+ """Rows [:n_und] through the understanding expert, the rest through the
155
+ generation expert (functions may return a tensor or a tuple of tensors)."""
156
+ if n_und is None or n_und == x.shape[0]:
157
+ return fn_und(x)
158
+ if n_und == 0:
159
+ return fn_gen(x)
160
+ a, b = fn_und(x[:n_und]), fn_gen(x[n_und:])
161
+ if isinstance(a, tuple):
162
+ return tuple(torch.cat(pair) for pair in zip(a, b))
163
+ return torch.cat([a, b])
164
+
165
+
166
+ class MoTAttention(nn.Module):
167
+ def __init__(self, hidden, n_heads, n_kv, head_dim, eps, layer_idx):
168
+ super().__init__()
169
+ self.n_heads, self.n_kv, self.head_dim, self.layer_idx = n_heads, n_kv, head_dim, layer_idx
170
+ self.scale = head_dim ** -0.5
171
+ for suffix in ("", "_moe_gen"):
172
+ setattr(self, f"q_proj{suffix}", nn.Linear(hidden, n_heads * head_dim, bias=True))
173
+ setattr(self, f"k_proj{suffix}", nn.Linear(hidden, n_kv * head_dim, bias=True))
174
+ setattr(self, f"v_proj{suffix}", nn.Linear(hidden, n_kv * head_dim, bias=True))
175
+ setattr(self, f"o_proj{suffix}", nn.Linear(n_heads * head_dim, hidden, bias=False))
176
+ setattr(self, f"q_norm{suffix}", RMSNorm(head_dim, eps))
177
+ setattr(self, f"k_norm{suffix}", RMSNorm(head_dim, eps))
178
+
179
+ def _qkv(self, x, suffix):
180
+ q = getattr(self, f"q_proj{suffix}")(x).view(-1, self.n_heads, self.head_dim)
181
+ k = getattr(self, f"k_proj{suffix}")(x).view(-1, self.n_kv, self.head_dim)
182
+ v = getattr(self, f"v_proj{suffix}")(x).view(-1, self.n_kv, self.head_dim)
183
+ return getattr(self, f"q_norm{suffix}")(q), getattr(self, f"k_norm{suffix}")(k), v
184
+
185
+ def forward(self, x, cos, sin, batch: AttnBatch, arena: KVArena):
186
+ q, k, v = _split(lambda t: self._qkv(t, ""), lambda t: self._qkv(t, "_moe_gen"), x, batch.n_und)
187
+ q = apply_rope(q, cos, sin)
188
+ k = apply_rope(k, cos, sin)
189
+ ak, av = arena.k[self.layer_idx], arena.v[self.layer_idx]
190
+ ak[batch.slots] = k
191
+ av[batch.slots] = v
192
+ nc = batch.n_causal
193
+ if batch.full is None:
194
+ out = _attend(q, ak, av, batch.causal, (-1, 0), self.scale)
195
+ elif batch.causal is None:
196
+ out = _attend(q, ak, av, batch.full, (-1, -1), self.scale)
197
+ else:
198
+ out = torch.cat([_attend(q[:nc], ak, av, batch.causal, (-1, 0), self.scale),
199
+ _attend(q[nc:], ak, av, batch.full, (-1, -1), self.scale)])
200
+ out = out.view(-1, self.n_heads * self.head_dim)
201
+ return _split(self.o_proj, self.o_proj_moe_gen, out, batch.n_und)
202
+
203
+
204
+ class MoTLayer(nn.Module):
205
+ def __init__(self, cfg, layer_idx):
206
+ super().__init__()
207
+ h = cfg["hidden_size"]
208
+ self.self_attn = MoTAttention(h, cfg["num_attention_heads"], cfg["num_key_value_heads"],
209
+ h // cfg["num_attention_heads"], cfg["rms_norm_eps"], layer_idx)
210
+ self.mlp = SwiGLU(h, cfg["intermediate_size"])
211
+ self.mlp_moe_gen = SwiGLU(h, cfg["intermediate_size"])
212
+ for name in ("input_layernorm", "post_attention_layernorm"):
213
+ setattr(self, name, RMSNorm(h, cfg["rms_norm_eps"]))
214
+ setattr(self, name + "_moe_gen", RMSNorm(h, cfg["rms_norm_eps"]))
215
+
216
+ def forward(self, x, cos, sin, batch, arena):
217
+ h = _split(self.input_layernorm, self.input_layernorm_moe_gen, x, batch.n_und)
218
+ x = x + self.self_attn(h, cos, sin, batch, arena)
219
+ h = _split(lambda t: self.mlp(self.post_attention_layernorm(t)),
220
+ lambda t: self.mlp_moe_gen(self.post_attention_layernorm_moe_gen(t)), x, batch.n_und)
221
+ return x + h
222
+
223
+
224
+ class _Decoder(nn.Module):
225
+ def __init__(self, cfg):
226
+ super().__init__()
227
+ self.embed_tokens = nn.Embedding(cfg["vocab_size"], cfg["hidden_size"])
228
+ self.layers = nn.ModuleList([MoTLayer(cfg, i) for i in range(cfg["num_hidden_layers"])])
229
+ self.norm = RMSNorm(cfg["hidden_size"], cfg["rms_norm_eps"])
230
+ self.norm_moe_gen = RMSNorm(cfg["hidden_size"], cfg["rms_norm_eps"])
231
+
232
+
233
+ class MoTLanguageModel(nn.Module):
234
+ """Qwen2-architecture MoT decoder (understanding + generation experts)."""
235
+
236
+ def __init__(self, cfg: dict):
237
+ super().__init__()
238
+ self.cfg = cfg
239
+ self.model = _Decoder(cfg)
240
+ self.lm_head = nn.Linear(cfg["hidden_size"], cfg["vocab_size"], bias=False)
241
+ head_dim = cfg["hidden_size"] // cfg["num_attention_heads"]
242
+ self.rotary = RotaryEmbedding(head_dim, cfg.get("rope_theta", 1e6))
243
+
244
+ @property
245
+ def num_layers(self):
246
+ return len(self.model.layers)
247
+
248
+ def forward(self, x, positions, batch: AttnBatch, arena: KVArena, cos_sin=None):
249
+ cos, sin = cos_sin if cos_sin is not None else self.rotary(positions, x.dtype)
250
+ for layer in self.model.layers:
251
+ x = layer(x, cos, sin, batch, arena)
252
+ return _split(self.model.norm, self.model.norm_moe_gen, x, batch.n_und)
253
+
254
+
255
+ # --------------------------------------------------------- SigLIP NaViT (BAGEL)
256
+ class FP32LayerNorm(nn.LayerNorm):
257
+ """LayerNorm computed in float32 (as under autocast), returned in input dtype."""
258
+
259
+ def forward(self, x):
260
+ return F.layer_norm(x.float(), self.normalized_shape, self.weight.float(), self.bias.float(),
261
+ self.eps).to(x.dtype)
262
+
263
+
264
+ class _SiglipAttention(nn.Module):
265
+ def __init__(self, dim, heads):
266
+ super().__init__()
267
+ self.heads, self.head_dim = heads, dim // heads
268
+ self.q_proj, self.k_proj, self.v_proj = (nn.Linear(dim, dim) for _ in range(3))
269
+ self.out_proj = nn.Linear(dim, dim)
270
+
271
+ def forward(self, x, cu, max_len):
272
+ T = x.shape[0]
273
+ q = self.q_proj(x).view(T, self.heads, self.head_dim)
274
+ k = self.k_proj(x).view(T, self.heads, self.head_dim)
275
+ v = self.v_proj(x).view(T, self.heads, self.head_dim)
276
+ out = varlen_attn(q, k, v, cu, cu, max_len, max_len, window_size=(-1, -1))
277
+ return self.out_proj(out.reshape(T, -1))
278
+
279
+
280
+ class _SiglipMLP(nn.Module):
281
+ def __init__(self, dim, inter):
282
+ super().__init__()
283
+ self.fc1, self.fc2 = nn.Linear(dim, inter), nn.Linear(inter, dim)
284
+
285
+ def forward(self, x):
286
+ return self.fc2(F.gelu(self.fc1(x), approximate="tanh"))
287
+
288
+
289
+ class _SiglipLayer(nn.Module):
290
+ def __init__(self, dim, heads, inter, eps):
291
+ super().__init__()
292
+ self.layer_norm1, self.layer_norm2 = FP32LayerNorm(dim, eps=eps), FP32LayerNorm(dim, eps=eps)
293
+ self.self_attn = _SiglipAttention(dim, heads)
294
+ self.mlp = _SiglipMLP(dim, inter)
295
+
296
+ def forward(self, x, cu, max_len):
297
+ x = x + self.self_attn(self.layer_norm1(x), cu, max_len)
298
+ return x + self.mlp(self.layer_norm2(x))
299
+
300
+
301
+ class _SiglipEmbeddings(nn.Module):
302
+ def __init__(self, cfg):
303
+ super().__init__()
304
+ p = cfg["patch_size"]
305
+ self.patch_embedding = nn.Linear(cfg.get("num_channels", 3) * p * p, cfg["hidden_size"])
306
+ self.position_embedding = nn.Embedding((cfg["image_size"] // p) ** 2, cfg["hidden_size"])
307
+
308
+
309
+ class _SiglipEncoder(nn.Module):
310
+ def __init__(self, cfg, n_layers):
311
+ super().__init__()
312
+ self.layers = nn.ModuleList([
313
+ _SiglipLayer(cfg["hidden_size"], cfg["num_attention_heads"], cfg["intermediate_size"],
314
+ cfg.get("layer_norm_eps", 1e-6)) for _ in range(n_layers)
315
+ ])
316
+
317
+
318
+ class _SiglipVision(nn.Module):
319
+ def __init__(self, cfg, n_layers):
320
+ super().__init__()
321
+ self.embeddings = _SiglipEmbeddings(cfg)
322
+ self.encoder = _SiglipEncoder(cfg, n_layers)
323
+ self.post_layernorm = FP32LayerNorm(cfg["hidden_size"], eps=cfg.get("layer_norm_eps", 1e-6))
324
+
325
+
326
+ class SiglipNaViT(nn.Module):
327
+ """Variable-resolution SigLIP: packed patches of many images per call."""
328
+
329
+ def __init__(self, cfg: dict, n_layers: int):
330
+ super().__init__()
331
+ self.vision_model = _SiglipVision(cfg, n_layers)
332
+
333
+ def forward(self, patches, pos_ids, cu, max_len):
334
+ vm = self.vision_model
335
+ x = vm.embeddings.patch_embedding(patches) + vm.embeddings.position_embedding(pos_ids)
336
+ for layer in vm.encoder.layers:
337
+ x = layer(x, cu, max_len)
338
+ return vm.post_layernorm(x)
339
+
340
+
341
+ class MLPConnector(nn.Module):
342
+ def __init__(self, in_dim, out_dim):
343
+ super().__init__()
344
+ self.fc1, self.fc2 = nn.Linear(in_dim, out_dim), nn.Linear(out_dim, out_dim)
345
+
346
+ def forward(self, x):
347
+ return self.fc2(F.gelu(self.fc1(x), approximate="tanh"))
348
+
349
+
350
+ def sincos_2d(dim: int, grid: int) -> torch.Tensor:
351
+ """Fixed 2-D sine-cosine position table of shape [grid*grid, dim]."""
352
+ def one_d(d, pos):
353
+ omega = 1.0 / 10000 ** (np.arange(d // 2, dtype=np.float64) / (d / 2.0))
354
+ out = np.einsum("m,d->md", pos.reshape(-1), omega)
355
+ return np.concatenate([np.sin(out), np.cos(out)], axis=1)
356
+
357
+ gw, gh = np.meshgrid(np.arange(grid, dtype=np.float32), np.arange(grid, dtype=np.float32))
358
+ emb = np.concatenate([one_d(dim // 2, gh), one_d(dim // 2, gw)], axis=1)
359
+ return torch.from_numpy(emb).float()
360
+
361
+
362
+ class PositionTable(nn.Module):
363
+ def __init__(self, grid: int, dim: int):
364
+ super().__init__()
365
+ self.pos_embed = nn.Parameter(sincos_2d(dim, grid), requires_grad=False)
366
+
367
+ def forward(self, ids):
368
+ return self.pos_embed[ids]
369
+
370
+
371
+ class TimestepEmbedder(nn.Module):
372
+ def __init__(self, hidden: int, freq_dim: int = 256):
373
+ super().__init__()
374
+ self.mlp = nn.Sequential(nn.Linear(freq_dim, hidden), nn.SiLU(), nn.Linear(hidden, hidden))
375
+ self.freq_dim = freq_dim
376
+
377
+ def forward(self, t):
378
+ half = self.freq_dim // 2
379
+ freqs = torch.exp(-math.log(10000) * torch.arange(half, dtype=torch.float32, device=t.device) / half)
380
+ args = t[:, None].float() * freqs[None]
381
+ emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
382
+ return self.mlp(emb.to(self.mlp[0].weight.dtype))
383
+
384
+
385
+ # ----------------------------------------------------------------- FLUX VAE
386
+ class _ResBlock(nn.Module):
387
+ def __init__(self, cin, cout):
388
+ super().__init__()
389
+ self.norm1 = nn.GroupNorm(32, cin, eps=1e-6)
390
+ self.conv1 = nn.Conv2d(cin, cout, 3, padding=1)
391
+ self.norm2 = nn.GroupNorm(32, cout, eps=1e-6)
392
+ self.conv2 = nn.Conv2d(cout, cout, 3, padding=1)
393
+ self.nin_shortcut = nn.Conv2d(cin, cout, 1) if cin != cout else None
394
+
395
+ def forward(self, x):
396
+ h = self.conv2(F.silu(self.norm2(self.conv1(F.silu(self.norm1(x))))))
397
+ return (x if self.nin_shortcut is None else self.nin_shortcut(x)) + h
398
+
399
+
400
+ class _AttnBlock(nn.Module):
401
+ def __init__(self, c):
402
+ super().__init__()
403
+ self.norm = nn.GroupNorm(32, c, eps=1e-6)
404
+ self.q, self.k, self.v, self.proj_out = (nn.Conv2d(c, c, 1) for _ in range(4))
405
+
406
+ def forward(self, x):
407
+ h = self.norm(x)
408
+ b, c, hh, ww = h.shape
409
+ q, k, v = (t(h).reshape(b, 1, c, hh * ww).transpose(-1, -2) for t in (self.q, self.k, self.v))
410
+ out = F.scaled_dot_product_attention(q, k, v).transpose(-1, -2).reshape(b, c, hh, ww)
411
+ return x + self.proj_out(out)
412
+
413
+
414
+ class _Up(nn.Module):
415
+ def __init__(self, c):
416
+ super().__init__()
417
+ self.conv = nn.Conv2d(c, c, 3, padding=1)
418
+
419
+ def forward(self, x):
420
+ return self.conv(F.interpolate(x, scale_factor=2.0, mode="nearest"))
421
+
422
+
423
+ class _Down(nn.Module):
424
+ def __init__(self, c):
425
+ super().__init__()
426
+ self.conv = nn.Conv2d(c, c, 3, stride=2)
427
+
428
+ def forward(self, x):
429
+ return self.conv(F.pad(x, (0, 1, 0, 1)))
430
+
431
+
432
+ class _Stage(nn.Module):
433
+ pass
434
+
435
+
436
+ class FluxAutoEncoder(nn.Module):
437
+ """The FLUX.1 VAE used by BAGEL (16 latent channels, 8x downsampling)."""
438
+
439
+ def __init__(self, ch=128, ch_mult=(1, 2, 4, 4), num_res_blocks=2, z_channels=16,
440
+ scale_factor=0.3611, shift_factor=0.1159):
441
+ super().__init__()
442
+ self.scale_factor, self.shift_factor = scale_factor, shift_factor
443
+ n = len(ch_mult)
444
+ # encoder
445
+ enc = nn.Module()
446
+ enc.conv_in = nn.Conv2d(3, ch, 3, padding=1)
447
+ enc.down = nn.ModuleList()
448
+ in_mult = (1,) + tuple(ch_mult)
449
+ block_in = ch
450
+ for i in range(n):
451
+ stage = _Stage()
452
+ stage.block = nn.ModuleList()
453
+ stage.attn = nn.ModuleList()
454
+ block_in, block_out = ch * in_mult[i], ch * ch_mult[i]
455
+ for _ in range(num_res_blocks):
456
+ stage.block.append(_ResBlock(block_in, block_out))
457
+ block_in = block_out
458
+ if i != n - 1:
459
+ stage.downsample = _Down(block_in)
460
+ enc.down.append(stage)
461
+ enc.mid = nn.Module()
462
+ enc.mid.block_1, enc.mid.attn_1, enc.mid.block_2 = _ResBlock(block_in, block_in), _AttnBlock(block_in), _ResBlock(block_in, block_in)
463
+ enc.norm_out = nn.GroupNorm(32, block_in, eps=1e-6)
464
+ enc.conv_out = nn.Conv2d(block_in, 2 * z_channels, 3, padding=1)
465
+ self.encoder = enc
466
+ # decoder
467
+ dec = nn.Module()
468
+ block_in = ch * ch_mult[-1]
469
+ dec.conv_in = nn.Conv2d(z_channels, block_in, 3, padding=1)
470
+ dec.mid = nn.Module()
471
+ dec.mid.block_1, dec.mid.attn_1, dec.mid.block_2 = _ResBlock(block_in, block_in), _AttnBlock(block_in), _ResBlock(block_in, block_in)
472
+ ups = [None] * n
473
+ for i in reversed(range(n)):
474
+ stage = _Stage()
475
+ stage.block = nn.ModuleList()
476
+ stage.attn = nn.ModuleList()
477
+ block_out = ch * ch_mult[i]
478
+ for _ in range(num_res_blocks + 1):
479
+ stage.block.append(_ResBlock(block_in, block_out))
480
+ block_in = block_out
481
+ if i != 0:
482
+ stage.upsample = _Up(block_in)
483
+ ups[i] = stage
484
+ dec.up = nn.ModuleList(ups)
485
+ dec.norm_out = nn.GroupNorm(32, block_in, eps=1e-6)
486
+ dec.conv_out = nn.Conv2d(block_in, 3, 3, padding=1)
487
+ self.decoder = dec
488
+ self._n = n
489
+ self._num_res_blocks = num_res_blocks
490
+
491
+ def decode(self, z):
492
+ d = self.decoder
493
+ z = z / self.scale_factor + self.shift_factor
494
+ h = d.conv_in(z)
495
+ h = d.mid.block_2(d.mid.attn_1(d.mid.block_1(h)))
496
+ for i in reversed(range(self._n)):
497
+ for blk in d.up[i].block:
498
+ h = blk(h)
499
+ if i != 0:
500
+ h = d.up[i].upsample(h)
501
+ return d.conv_out(F.silu(d.norm_out(h)))
502
+
503
+ def encode(self, x):
504
+ e = self.encoder
505
+ h = e.conv_in(x)
506
+ for i in range(self._n):
507
+ for blk in e.down[i].block:
508
+ h = blk(h)
509
+ if i != self._n - 1:
510
+ h = e.down[i].downsample(h)
511
+ h = e.mid.block_2(e.mid.attn_1(e.mid.block_1(h)))
512
+ h = e.conv_out(F.silu(e.norm_out(h)))
513
+ mean = h.chunk(2, dim=1)[0]
514
+ return self.scale_factor * (mean - self.shift_factor)