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.
- simit/__init__.py +39 -0
- simit/api.py +258 -0
- simit/backends/__init__.py +112 -0
- simit/backends/bagel.py +359 -0
- simit/backends/base.py +114 -0
- simit/backends/hf.py +336 -0
- simit/backends/imagegen.py +136 -0
- simit/backends/lance.py +338 -0
- simit/backends/mot/engine.py +844 -0
- simit/backends/mot/kernels.py +130 -0
- simit/backends/mot/modeling.py +514 -0
- simit/backends/mot/wan_vae.py +872 -0
- simit/backends/vllm.py +204 -0
- simit/config.py +112 -0
- simit/metrics.py +113 -0
- simit/pipeline/core.py +586 -0
- simit/pipeline/prompts.py +167 -0
- simit/skills/__init__.py +6 -0
- simit/skills/assets/MERMAID_LICENSE +21 -0
- simit/skills/assets/mermaid.min.js +3587 -0
- simit/skills/base.py +172 -0
- simit/skills/builtin.py +614 -0
- simit/skills/helpers.py +197 -0
- simit/skills/pool.py +262 -0
- simit/skills/prompts.py +2374 -0
- simit/skills/renderers.py +1292 -0
- simit/skills/routing.py +89 -0
- simit/skills/sandbox.py +188 -0
- simit/skills/web.py +253 -0
- simit/tune.py +323 -0
- simit/types.py +96 -0
- simit/utils.py +67 -0
- simit-0.1.0.dist-info/METADATA +388 -0
- simit-0.1.0.dist-info/RECORD +38 -0
- simit-0.1.0.dist-info/WHEEL +5 -0
- simit-0.1.0.dist-info/licenses/LICENSE +202 -0
- simit-0.1.0.dist-info/licenses/src/simit/skills/assets/MERMAID_LICENSE +21 -0
- simit-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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)
|