JiRackDeltaNet_27b.py
36.4 KB · 767 lines · python Raw
1 """
2 JiRackDeltaNet_27b.py
3 =============================
4
5 Собственная (hand-rolled) архитектура "JiRack Qwen3.8-27B" -- в стиле
6 прежнего DS32B-проекта: один файл, свой TransformerBlock, никакого
7 AutoConfig/AutoModelForCausalLM внутри класса архитектуры, никакого
8 model_id в конструкторе. Все размерности зашиты константами ниже --
9 подтверждены живым инспектом реального checkpoint'а (inspect_qwen38_keys.py,
10 851 тензор, сверено).
11
12 Математика DeltaNet-рекуррентности, свёртки, гейтинга и partial-rotary
13 mRoPE-attention ПОРТИРОВАНА из реального modeling_qwen3_5.py (Apache 2.0,
14 HuggingFace/Qwen), а не переизобретена -- чтобы не рисковать тихой
15 числовой ошибкой в тонкой рекуррентной формуле. Портирован смысл, но
16 код переписан под наши имена/стиль и урезан: без vision-башни, без
17 KV-cache классов transformers (для тренировки полный форвард по
18 последовательности не нужен; для генерации -- полный recompute на
19 каждом шаге, без инкрементального кэша; это сознательное упрощение,
20 при желании кэш можно добавить отдельно).
21
22 Имена подмодулей (`linear_attn` / `self_attn`, `in_proj_qkv`, `out_proj`,
23 `mlp.gate_proj` и т.д.) намеренно СОВПАДАЮТ с именами в реальном
24 state_dict Qwen3_5ForCausalLM -- поэтому веса из HF-чекпоинта грузятся
25 напрямую через load_state_dict(), без ручного key-mapping.
26
27 Файл умеет:
28 - собрать модель (JiRackQwen38ForCausalLM) с нуля, случайные веса;
29 - принять state_dict в формате HF Qwen3_5ForCausalLM (после
30 AutoModelForCausalLM.from_pretrained(...).state_dict() в конвертере)
31 и загрузить его напрямую;
32 - пропатчить нужные nn.Linear -> BitLinear (TERNARIZE_PATTERNS);
33 - сохранить/загрузить свой JiRack-чекпоинт (save_checkpoint/load_checkpoint).
34
35 Конвертер (convert_qwen38_to_jirack.py) по-прежнему знает про model_id
36 и HF Hub -- но это знание живёт ТОЛЬКО в конвертере, не в архитектуре.
37 """
38
39 import os
40 import re
41 import torch
42 import torch.nn as nn
43 import torch.nn.functional as F
44
45
46 # =====================================================================
47 # 1. CONFIRMED АРХИТЕКТУРА -- жёстко зашитые константы (не AutoConfig)
48 # =====================================================================
49
50 HIDDEN_SIZE = 5120
51 VOCAB_SIZE = 248320
52 INTERMEDIATE_SIZE = 17408
53 NUM_LAYERS = 64
54 RMS_NORM_EPS = 1e-6
55
56 # Чередование слоёв: full_attention_interval = 4 ->
57 # 3x linear_attention, 1x full_attention, повторяется на все 64 слоя.
58 LAYER_TYPES = ["linear_attention"] * 3 + ["full_attention"]
59 LAYER_TYPES = (LAYER_TYPES * (NUM_LAYERS // len(LAYER_TYPES) + 1))[:NUM_LAYERS]
60
61 # --- Gated DeltaNet (linear_attention) ---
62 LINEAR_NUM_KEY_HEADS = 16
63 LINEAR_KEY_HEAD_DIM = 128
64 LINEAR_NUM_VALUE_HEADS = 48
65 LINEAR_VALUE_HEAD_DIM = 128
66 LINEAR_CONV_KERNEL_DIM = 4
67 DELTA_KEY_DIM = LINEAR_NUM_KEY_HEADS * LINEAR_KEY_HEAD_DIM # 2048
68 DELTA_VALUE_DIM = LINEAR_NUM_VALUE_HEADS * LINEAR_VALUE_HEAD_DIM # 6144
69 DELTA_CONV_DIM = DELTA_KEY_DIM * 2 + DELTA_VALUE_DIM # 10240
70 DELTA_CHUNK_SIZE = 64
71
72 # --- Gated Attention (full_attention) ---
73 NUM_ATTENTION_HEADS = 24
74 NUM_KEY_VALUE_HEADS = 4
75 HEAD_DIM = 256
76 PARTIAL_ROTARY_FACTOR = 0.25
77 ROPE_THETA = 1e7
78 MROPE_SECTION = [11, 11, 10]
79 ATTENTION_BIAS = False
80
81 ARCH_SUMMARY = f"""
82 Qwen3.8-27B (qwen3_5), hand-rolled JiRack TransformerBlock, text-only.
83 hidden_size={HIDDEN_SIZE}, vocab_size={VOCAB_SIZE},
84 intermediate_size={INTERMEDIATE_SIZE}, num_hidden_layers={NUM_LAYERS}
85 Чередование (interval=4): 3x linear_attention (DeltaNet), 1x full_attention
86 """
87
88 KEPT_FULL_PRECISION = """
89 Всегда остаются full precision (не GEMM, либо слишком чувствительны):
90 - dt_bias, A_log (сырые Parameters, decay DeltaNet)
91 - linear_attn.in_proj_b/in_proj_a (по умолчанию, см. GATE_PROJ_PATTERNS)
92 - conv1d (depthwise grouped conv, не GEMM)
93 - все нормы: RMSNormGated, q_norm/k_norm, model.norm
94 - embed_tokens, lm_head
95 - mtp.* (вся MTP/NextN-голова) (draft-слой спекулятивного декода;
96 тернаризация убивает acceptance rate,
97 комьюнити-GGUF держат её в Q8_0)
98 """
99
100
101 # =====================================================================
102 # 2. BitLinear -- наш тернарный слой
103 # =====================================================================
104
105 class BitLinear(nn.Linear):
106 """b1.58-style per-tensor absmean ternary weights, per-token int8
107 activations, непрерывный lambda-warmup (STE)."""
108
109 def __init__(self, in_features, out_features, bias=False):
110 super().__init__(in_features, out_features, bias=bias)
111 self.eps = 1e-5
112 self.register_buffer("lambda_", torch.zeros(()), persistent=True)
113
114 @classmethod
115 def from_linear(cls, lin: nn.Linear) -> "BitLinear":
116 bl = cls(lin.in_features, lin.out_features, bias=lin.bias is not None)
117 bl = bl.to(dtype=lin.weight.dtype, device=lin.weight.device)
118 with torch.no_grad():
119 bl.weight.copy_(lin.weight)
120 if lin.bias is not None:
121 bl.bias.copy_(lin.bias)
122 return bl
123
124 def forward(self, x: torch.Tensor) -> torch.Tensor:
125 if not self.training and float(self.lambda_) < 1e-6:
126 return F.linear(x, self.weight, self.bias)
127
128 lam = self.lambda_.to(x.dtype)
129 w = self.weight
130 gamma = w.float().abs().mean().clamp(min=self.eps).to(w.dtype)
131 w_quant = torch.clamp(torch.round(w / gamma), -1.0, 1.0) * gamma
132 w_effective = w + lam * (w_quant - w).detach()
133
134 x_scale = 127.0 / x.abs().max(dim=-1, keepdim=True).values.clamp(min=self.eps)
135 x_quant = torch.clamp(torch.round(x * x_scale), -128.0, 127.0) / x_scale
136 x_effective = x + lam * (x_quant - x).detach()
137
138 return F.linear(x_effective, w_effective, self.bias)
139
140 def set_lambda(self, v: float):
141 self.lambda_.fill_(v)
142
143
144 # =====================================================================
145 # 3. Карта: что тернаризуем, что оставляем full precision
146 # =====================================================================
147
148 TERNARIZE_PATTERNS = [
149 r"linear_attn\.in_proj_qkv$",
150 r"linear_attn\.in_proj_z$",
151 r"linear_attn\.out_proj$",
152 r"self_attn\.q_proj$",
153 r"self_attn\.k_proj$",
154 r"self_attn\.v_proj$",
155 r"self_attn\.o_proj$",
156 r"mlp\.gate_proj$",
157 r"mlp\.up_proj$",
158 r"mlp\.down_proj$",
159 ]
160
161 GATE_PROJ_PATTERNS = [
162 r"linear_attn\.in_proj_b$",
163 r"linear_attn\.in_proj_a$",
164 ]
165
166
167 def apply_bitlinear_patch(model: nn.Module, include_gate_proj: bool = False):
168 """Точечная замена nn.Linear -> BitLinear по TERNARIZE_PATTERNS."""
169 patterns = list(TERNARIZE_PATTERNS)
170 if include_gate_proj:
171 patterns += GATE_PROJ_PATTERNS
172 compiled = [re.compile(p) for p in patterns]
173
174 patched = []
175 for name, module in list(model.named_modules()):
176 if not isinstance(module, nn.Linear):
177 continue
178 # MTP/NextN-голова НЕ тернаризуется -- всегда full precision.
179 # Суффиксные паттерны (self_attn.q_proj$ и т.п.) иначе зацепили бы
180 # mtp.layers.0.self_attn.q_proj.
181 if name.startswith("mtp"):
182 continue
183 if any(p.search(name) for p in compiled):
184 parent_name, _, child_name = name.rpartition(".")
185 parent = model.get_submodule(parent_name) if parent_name else model
186 setattr(parent, child_name, BitLinear.from_linear(module))
187 patched.append(name)
188 return patched
189
190
191 # =====================================================================
192 # 4. Общие блоки (порт математики из реального modeling_qwen3_5.py)
193 # =====================================================================
194
195 class RMSNorm(nn.Module):
196 """(x * w).to(dtype), w хранится как (weight) с offset +1 -- как в
197 оригинале Qwen3_5RMSNorm (инициализация нулями, эффективный вес 1+w)."""
198
199 def __init__(self, dim: int, eps: float = RMS_NORM_EPS):
200 super().__init__()
201 self.eps = eps
202 self.weight = nn.Parameter(torch.zeros(dim))
203
204 def forward(self, x: torch.Tensor) -> torch.Tensor:
205 dtype = x.dtype
206 x = x.float()
207 x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
208 x = x * (1.0 + self.weight.float())
209 return x.to(dtype)
210
211
212 class RMSNormGated(nn.Module):
213 """RMSNorm с гейтингом через SiLU(gate), применяется ПЕРЕД гейтингом
214 (норма считается по hidden_states, гейт применяется после)."""
215
216 def __init__(self, dim: int, eps: float = RMS_NORM_EPS):
217 super().__init__()
218 self.weight = nn.Parameter(torch.ones(dim))
219 self.eps = eps
220
221 def forward(self, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor:
222 dtype = x.dtype
223 x = x.float()
224 var = x.pow(2).mean(-1, keepdim=True)
225 x = x * torch.rsqrt(var + self.eps)
226 x = self.weight * x.to(dtype)
227 x = x * F.silu(gate.float())
228 return x.to(dtype)
229
230
231 def rotate_half(x: torch.Tensor) -> torch.Tensor:
232 x1 = x[..., : x.shape[-1] // 2]
233 x2 = x[..., x.shape[-1] // 2:]
234 return torch.cat((-x2, x1), dim=-1)
235
236
237 def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):
238 """Partial rotary: только первые cos.shape[-1] каналов вращаются,
239 остальное (q_pass/k_pass) проходит без изменений."""
240 cos = cos.unsqueeze(unsqueeze_dim)
241 sin = sin.unsqueeze(unsqueeze_dim)
242 rotary_dim = cos.shape[-1]
243 q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
244 k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
245 q_embed = torch.cat([(q_rot * cos) + (rotate_half(q_rot) * sin), q_pass], dim=-1)
246 k_embed = torch.cat([(k_rot * cos) + (rotate_half(k_rot) * sin), k_pass], dim=-1)
247 return q_embed, k_embed
248
249
250 class TextRotaryEmbedding(nn.Module):
251 """mRoPE с partial_rotary_factor. Для чистого текста (без vision)
252 3 позиционные "оси" (T,H,W) совпадают -> interleaved mRoPE
253 вырождается в обычный RoPE, но мы всё равно считаем честно тем же
254 путём, что и оригинал, для полной числовой идентичности."""
255
256 def __init__(self, head_dim=HEAD_DIM, partial_rotary_factor=PARTIAL_ROTARY_FACTOR,
257 theta=ROPE_THETA, mrope_section=MROPE_SECTION):
258 super().__init__()
259 self.rotary_dim = int(head_dim * partial_rotary_factor)
260 inv_freq = 1.0 / (theta ** (torch.arange(0, self.rotary_dim, 2, dtype=torch.float32) / self.rotary_dim))
261 self.register_buffer("inv_freq", inv_freq, persistent=False)
262 self.mrope_section = mrope_section
263
264 def _apply_interleaved_mrope(self, freqs: torch.Tensor) -> torch.Tensor:
265 # freqs: (3, bs, seq, rotary_dim // 2)
266 freqs_t = freqs[0].clone()
267 for dim, offset in enumerate((1, 2), start=1):
268 length = self.mrope_section[dim] * 3
269 idx = slice(offset, length, 3)
270 freqs_t[..., idx] = freqs[dim, ..., idx]
271 return freqs_t
272
273 @torch.no_grad()
274 def forward(self, x: torch.Tensor, position_ids: torch.LongTensor):
275 # position_ids: (bs, seq) -> реплицируем в 3 идентичные оси (text-only)
276 if position_ids.ndim == 2:
277 position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
278 inv_freq_exp = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)
279 pos_exp = position_ids[:, :, None, :].float()
280 freqs = (inv_freq_exp @ pos_exp).transpose(2, 3) # (3, bs, seq, rotary_dim//2)
281 freqs = self._apply_interleaved_mrope(freqs) # (bs, seq, rotary_dim//2)
282 emb = torch.cat((freqs, freqs), dim=-1) # (bs, seq, rotary_dim)
283 return emb.cos().to(x.dtype), emb.sin().to(x.dtype)
284
285
286 def l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6):
287 return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
288
289
290 def causal_conv1d(hidden_states: torch.Tensor, weight: torch.Tensor, activation: str = "silu"):
291 """hidden_states: (batch, channels, seq_len). Depthwise causal conv,
292 left-padded by kernel_size - 1, затем активация."""
293 channels, seq_len = hidden_states.shape[1], hidden_states.shape[2]
294 padding = weight.shape[-1] - 1
295 out = F.conv1d(hidden_states, weight=weight.unsqueeze(1), bias=None, padding=padding, groups=channels)
296 out = out[:, :, :seq_len]
297 return F.silu(out) if activation == "silu" else out
298
299
300 def chunked_gated_delta_rule(query, key, value, g, beta, chunk_size=DELTA_CHUNK_SIZE):
301 """Порт torch_chunk_gated_delta_rule из modeling_qwen3_5.py -- без
302 initial_state/output_final_state (тренировка/полный форвард без
303 инкрементального кэша). query/key/value: (b, h, t, d), g/beta: (b, h, t)."""
304 initial_dtype = query.dtype
305 query = l2norm(query, dim=-1, eps=1e-6)
306 key = l2norm(key, dim=-1, eps=1e-6)
307 query, key, value, beta, g = [x.contiguous().to(torch.float32) for x in (query, key, value, beta, g)]
308
309 b, h, t, k_dim = key.shape
310 v_dim = value.shape[-1]
311 pad = (chunk_size - t % chunk_size) % chunk_size
312 query = F.pad(query, (0, 0, 0, pad))
313 key = F.pad(key, (0, 0, 0, pad))
314 value = F.pad(value, (0, 0, 0, pad))
315 beta = F.pad(beta, (0, pad))
316 g = F.pad(g, (0, pad))
317 total_t = t + pad
318 scale = 1 / (query.shape[-1] ** 0.5)
319 query = query * scale
320
321 v_beta = value * beta.unsqueeze(-1)
322 k_beta = key * beta.unsqueeze(-1)
323 query, key, value, k_beta, v_beta = [
324 x.reshape(x.shape[0], x.shape[1], -1, chunk_size, x.shape[-1]) for x in (query, key, value, k_beta, v_beta)
325 ]
326 g = g.reshape(g.shape[0], g.shape[1], -1, chunk_size)
327 mask = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device), diagonal=0)
328
329 g = g.cumsum(dim=-1)
330 decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril()
331 attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask, 0)
332 for i in range(1, chunk_size):
333 row = attn[..., i, :i].clone()
334 sub = attn[..., :i, :i].clone()
335 attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
336 attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
337 value = attn @ v_beta
338 k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
339
340 last_state = torch.zeros(b, h, k_dim, v_dim, dtype=value.dtype, device=value.device)
341 core_out = torch.zeros_like(value)
342
343 for i in range(total_t // chunk_size):
344 q_i, k_i, v_i = query[:, :, i], key[:, :, i], value[:, :, i]
345 attn_i = q_i @ k_i.transpose(-1, -2) * decay_mask[:, :, i]
346 v_prime = k_cumdecay[:, :, i] @ last_state
347 v_new = v_i - v_prime
348 attn_inter = (q_i * g[:, :, i, :, None].exp()) @ last_state
349 core_out[:, :, i] = attn_inter + attn_i @ v_new
350 last_state = (
351 last_state * g[:, :, i, -1, None, None].exp()
352 + (k_i * (g[:, :, i, -1, None] - g[:, :, i]).exp()[..., None]).transpose(-1, -2) @ v_new
353 )
354
355 core_out = core_out.reshape(core_out.shape[0], core_out.shape[1], -1, core_out.shape[-1])
356 core_out = core_out[:, :, :t]
357 core_out = core_out.transpose(1, 2).contiguous().to(initial_dtype)
358 return core_out
359
360
361 # =====================================================================
362 # 5. Наш собственный TransformerBlock (и подмодули linear_attn/self_attn)
363 # =====================================================================
364
365 class GatedDeltaNetBlock(nn.Module):
366 """linear_attention слой (Gated DeltaNet). Имена подмодулей совпадают
367 с HF state_dict, чтобы веса грузились напрямую."""
368
369 def __init__(self):
370 super().__init__()
371 self.conv1d = nn.Conv1d(
372 in_channels=DELTA_CONV_DIM, out_channels=DELTA_CONV_DIM, bias=False,
373 kernel_size=LINEAR_CONV_KERNEL_DIM, groups=DELTA_CONV_DIM,
374 padding=LINEAR_CONV_KERNEL_DIM - 1,
375 )
376 self.dt_bias = nn.Parameter(torch.ones(LINEAR_NUM_VALUE_HEADS))
377 A = torch.empty(LINEAR_NUM_VALUE_HEADS).uniform_(0.01, 16)
378 self.A_log = nn.Parameter(torch.log(A))
379 self.norm = RMSNormGated(LINEAR_VALUE_HEAD_DIM, eps=RMS_NORM_EPS)
380 self.out_proj = nn.Linear(DELTA_VALUE_DIM, HIDDEN_SIZE, bias=False)
381 self.in_proj_qkv = nn.Linear(HIDDEN_SIZE, DELTA_KEY_DIM * 2 + DELTA_VALUE_DIM, bias=False)
382 self.in_proj_z = nn.Linear(HIDDEN_SIZE, DELTA_VALUE_DIM, bias=False)
383 self.in_proj_b = nn.Linear(HIDDEN_SIZE, LINEAR_NUM_VALUE_HEADS, bias=False)
384 self.in_proj_a = nn.Linear(HIDDEN_SIZE, LINEAR_NUM_VALUE_HEADS, bias=False)
385
386 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
387 bsz, seq_len, _ = hidden_states.shape
388
389 mixed_qkv = self.in_proj_qkv(hidden_states).transpose(1, 2) # (b, conv_dim, t)
390 z = self.in_proj_z(hidden_states).reshape(bsz, seq_len, -1, LINEAR_VALUE_HEAD_DIM)
391 b = self.in_proj_b(hidden_states)
392 a = self.in_proj_a(hidden_states)
393
394 mixed_qkv = causal_conv1d(mixed_qkv, self.conv1d.weight.squeeze(1), activation="silu")
395 mixed_qkv = mixed_qkv.transpose(1, 2) # (b, t, conv_dim)
396
397 query, key, value = torch.split(mixed_qkv, [DELTA_KEY_DIM, DELTA_KEY_DIM, DELTA_VALUE_DIM], dim=-1)
398 query = query.reshape(bsz, seq_len, -1, LINEAR_KEY_HEAD_DIM)
399 key = key.reshape(bsz, seq_len, -1, LINEAR_KEY_HEAD_DIM)
400 value = value.reshape(bsz, seq_len, -1, LINEAR_VALUE_HEAD_DIM)
401
402 beta = b.sigmoid()
403 g = -self.A_log.float().exp() * F.softplus(a.float() + self.dt_bias)
404
405 rep = LINEAR_NUM_VALUE_HEADS // LINEAR_NUM_KEY_HEADS
406 if rep > 1:
407 query = query.repeat_interleave(rep, dim=2)
408 key = key.repeat_interleave(rep, dim=2)
409
410 # -> (b, h, t, d) for the chunked kernel
411 query_h = query.transpose(1, 2)
412 key_h = key.transpose(1, 2)
413 value_h = value.transpose(1, 2)
414 beta_h = beta.transpose(1, 2)
415 g_h = g.transpose(1, 2)
416
417 core_attn_out = chunked_gated_delta_rule(query_h, key_h, value_h, g_h, beta_h)
418
419 core_attn_out = core_attn_out.reshape(-1, LINEAR_VALUE_HEAD_DIM)
420 z_flat = z.reshape(-1, LINEAR_VALUE_HEAD_DIM)
421 core_attn_out = self.norm(core_attn_out, z_flat)
422 core_attn_out = core_attn_out.reshape(bsz, seq_len, -1)
423
424 return self.out_proj(core_attn_out)
425
426
427 class AttentionBlock(nn.Module):
428 """full_attention слой (Gated Attention + partial rotary mRoPE)."""
429
430 def __init__(self):
431 super().__init__()
432 self.num_key_value_groups = NUM_ATTENTION_HEADS // NUM_KEY_VALUE_HEADS
433 self.scaling = HEAD_DIM ** -0.5
434 self.q_proj = nn.Linear(HIDDEN_SIZE, NUM_ATTENTION_HEADS * HEAD_DIM * 2, bias=ATTENTION_BIAS)
435 self.k_proj = nn.Linear(HIDDEN_SIZE, NUM_KEY_VALUE_HEADS * HEAD_DIM, bias=ATTENTION_BIAS)
436 self.v_proj = nn.Linear(HIDDEN_SIZE, NUM_KEY_VALUE_HEADS * HEAD_DIM, bias=ATTENTION_BIAS)
437 self.o_proj = nn.Linear(NUM_ATTENTION_HEADS * HEAD_DIM, HIDDEN_SIZE, bias=ATTENTION_BIAS)
438 self.q_norm = RMSNorm(HEAD_DIM, eps=RMS_NORM_EPS)
439 self.k_norm = RMSNorm(HEAD_DIM, eps=RMS_NORM_EPS)
440
441 def forward(self, hidden_states, cos, sin, causal_mask):
442 bsz, seq_len, _ = hidden_states.shape
443 hidden_shape = (bsz, seq_len, -1, HEAD_DIM)
444
445 query_states, gate = torch.chunk(
446 self.q_proj(hidden_states).view(bsz, seq_len, -1, HEAD_DIM * 2), 2, dim=-1
447 )
448 gate = gate.reshape(bsz, seq_len, -1)
449
450 query_states = self.q_norm(query_states.view(hidden_shape)).transpose(1, 2)
451 key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
452 value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
453
454 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
455
456 rep = self.num_key_value_groups
457 key_states = key_states.repeat_interleave(rep, dim=1)
458 value_states = value_states.repeat_interleave(rep, dim=1)
459
460 attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) * self.scaling
461 attn_weights = attn_weights + causal_mask
462 attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
463 attn_output = torch.matmul(attn_weights, value_states)
464 attn_output = attn_output.transpose(1, 2).contiguous().reshape(bsz, seq_len, -1)
465 attn_output = attn_output * torch.sigmoid(gate)
466
467 return self.o_proj(attn_output)
468
469
470 class MLPBlock(nn.Module):
471 def __init__(self):
472 super().__init__()
473 self.gate_proj = nn.Linear(HIDDEN_SIZE, INTERMEDIATE_SIZE, bias=False)
474 self.up_proj = nn.Linear(HIDDEN_SIZE, INTERMEDIATE_SIZE, bias=False)
475 self.down_proj = nn.Linear(INTERMEDIATE_SIZE, HIDDEN_SIZE, bias=False)
476
477 def forward(self, x):
478 return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
479
480
481 class TransformerBlock(nn.Module):
482 """Наш decoder layer -- linear_attn ИЛИ self_attn, в зависимости от
483 layer_type, + общий MLP. Имена атрибутов совпадают с HF state_dict."""
484
485 def __init__(self, layer_idx: int, layer_type: str = None):
486 super().__init__()
487 # layer_type задан явно -> используем его (нужно для MTP-слоя,
488 # который всегда full_attention и лежит ВНЕ LAYER_TYPES).
489 self.block_type = layer_type if layer_type is not None else LAYER_TYPES[layer_idx]
490 if self.block_type == "linear_attention":
491 self.linear_attn = GatedDeltaNetBlock()
492 else:
493 self.self_attn = AttentionBlock()
494 self.mlp = MLPBlock()
495 self.input_layernorm = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
496 self.post_attention_layernorm = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
497
498 def forward(self, hidden_states, cos, sin, causal_mask):
499 residual = hidden_states
500 hidden_states = self.input_layernorm(hidden_states)
501 if self.block_type == "linear_attention":
502 hidden_states = self.linear_attn(hidden_states)
503 else:
504 hidden_states = self.self_attn(hidden_states, cos, sin, causal_mask)
505 hidden_states = residual + hidden_states
506
507 residual = hidden_states
508 hidden_states = self.post_attention_layernorm(hidden_states)
509 hidden_states = self.mlp(hidden_states)
510 hidden_states = residual + hidden_states
511 return hidden_states
512
513
514 # =====================================================================
515 # 5b. MTP / NextN-голова (draft-слой спекулятивного декода)
516 # =====================================================================
517
518 class MTPBlock(nn.Module):
519 """Multi-Token Prediction голова -- ровно как в реальном чекпоинте
520 Qwen3.5/3.8 (mtp_num_hidden_layers=1, mtp_use_dedicated_embeddings=false).
521
522 Имена подмодулей СОВПАДАЮТ с ключами HF state_dict (top-level `mtp.`):
523 mtp.pre_fc_norm_embedding.weight
524 mtp.pre_fc_norm_hidden.weight
525 mtp.fc.weight [HIDDEN, 2*HIDDEN]
526 mtp.layers.0.* (обычный full_attention decoder)
527 mtp.norm.weight (final norm перед общим lm_head)
528
529 embed_tokens и lm_head ОБЩИЕ со стволом (dedicated embeddings = false),
530 здесь их нет. В GGUF это blk.{NUM_LAYERS}.nextn.eh_proj /
531 shared_head_norm + тензоры draft-слоя.
532
533 Механика (DeepSeek-style, порядок как в vLLM Qwen3NextMTP):
534 h = fc(cat(pre_fc_norm_embedding(emb_next), pre_fc_norm_hidden(h_prev)))
535 h = layers[0](h) # full attention
536 logits2 = lm_head(norm(h))
537 """
538
539 def __init__(self):
540 super().__init__()
541 self.pre_fc_norm_embedding = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
542 self.pre_fc_norm_hidden = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
543 self.fc = nn.Linear(2 * HIDDEN_SIZE, HIDDEN_SIZE, bias=False)
544 self.layers = nn.ModuleList([
545 TransformerBlock(layer_idx=NUM_LAYERS, layer_type="full_attention"),
546 ])
547 self.norm = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
548
549 def forward(self, inputs_embeds, previous_hidden, cos, sin, causal_mask):
550 """inputs_embeds -- embed_tokens(токены, сдвинутые на +1)
551 previous_hidden -- hidden ствола ДО финального model.norm
552 (см. контракт: MTP сам применяет pre_fc_norm_hidden)."""
553 h = self.fc(torch.cat(
554 (self.pre_fc_norm_embedding(inputs_embeds),
555 self.pre_fc_norm_hidden(previous_hidden)),
556 dim=-1,
557 ))
558 h = self.layers[0](h, cos, sin, causal_mask)
559 return self.norm(h)
560
561
562 # =====================================================================
563 # 6. Полная модель
564 # =====================================================================
565
566 class JiRackQwen38Model(nn.Module):
567 def __init__(self):
568 super().__init__()
569 self.embed_tokens = nn.Embedding(VOCAB_SIZE, HIDDEN_SIZE)
570 self.layers = nn.ModuleList([TransformerBlock(i) for i in range(NUM_LAYERS)])
571 self.norm = RMSNorm(HIDDEN_SIZE, eps=RMS_NORM_EPS)
572 self.rotary_emb = TextRotaryEmbedding()
573
574 def forward(self, input_ids: torch.LongTensor, return_prenorm: bool = False):
575 bsz, seq_len = input_ids.shape
576 hidden_states = self.embed_tokens(input_ids)
577
578 position_ids = torch.arange(seq_len, device=input_ids.device).unsqueeze(0).expand(bsz, -1)
579 cos, sin = self.rotary_emb(hidden_states, position_ids)
580
581 causal_mask = torch.full((seq_len, seq_len), float("-inf"), device=input_ids.device, dtype=torch.float32)
582 causal_mask = torch.triu(causal_mask, diagonal=1)[None, None, :, :].to(hidden_states.dtype)
583
584 for layer in self.layers:
585 hidden_states = layer(hidden_states, cos, sin, causal_mask)
586
587 if return_prenorm:
588 # Для MTP: голова сама применяет pre_fc_norm_hidden, финальный
589 # model.norm сюда НЕ входит (иначе двойная нормализация).
590 return self.norm(hidden_states), hidden_states, (cos, sin, causal_mask)
591 return self.norm(hidden_states)
592
593
594 class JiRackQwen38ForCausalLM(nn.Module):
595 """Верхнеуровневая модель. Имена атрибутов (`model`, `lm_head`)
596 совпадают с HF Qwen3_5ForCausalLM -- state_dict грузится напрямую."""
597
598 def __init__(self):
599 super().__init__()
600 self.model = JiRackQwen38Model()
601 self.lm_head = nn.Linear(HIDDEN_SIZE, VOCAB_SIZE, bias=False)
602 # MTP/NextN-голова: top-level атрибут `mtp` -- ключи state_dict
603 # получаются ровно mtp.*, как в реальном HF-чекпоинте.
604 self.mtp = MTPBlock()
605
606 def forward(self, input_ids: torch.LongTensor) -> torch.Tensor:
607 hidden_states = self.model(input_ids)
608 return self.lm_head(hidden_states)
609
610 def mtp_forward(self, input_ids: torch.LongTensor):
611 """Возвращает (logits_t1, logits_t2): обычные логиты ствола и
612 draft-логиты MTP-головы (предсказание токена t+2 из позиции t).
613 Для инференса/валидации MTP; в GGUF-конвертации не участвует."""
614 _, prenorm_hidden, (cos, sin, causal_mask) = self.model(
615 input_ids, return_prenorm=True)
616 logits_t1 = self.lm_head(self.model.norm(prenorm_hidden))
617
618 # Вход MTP -- эмбеддинги токенов, сдвинутых на +1 (следующий токен
619 # уже известен на шаге верификации). Последнюю позицию дублируем,
620 # чтобы сохранить форму (draft для неё не используется).
621 shifted = torch.cat([input_ids[:, 1:], input_ids[:, -1:]], dim=1)
622 emb_next = self.model.embed_tokens(shifted)
623 h2 = self.mtp(emb_next, prenorm_hidden, cos, sin, causal_mask)
624 logits_t2 = self.lm_head(h2)
625 return logits_t1, logits_t2
626
627 # -----------------------------------------------------------------
628 # Загрузка реальных весов (state_dict уже скачан конвертером через
629 # HF Hub -- сюда просто передаётся словарь тензоров)
630 # -----------------------------------------------------------------
631 def load_hf_state_dict(self, state_dict: dict, strict: bool = False):
632 """Грузит state_dict в формате оригинального Qwen3_5ForCausalLM.
633 mtp.* ключи теперь ГРУЗЯТСЯ в self.mtp (MTPBlock) -- фильтра
634 больше нет. strict=False оставлен из-за служебных lambda_-буферов
635 BitLinear, которых нет в реальном чекпоинте."""
636 missing, unexpected = self.load_state_dict(state_dict, strict=strict)
637 real_missing = [k for k in missing if not k.endswith("lambda_")]
638 if real_missing:
639 print(f"⚠️ {len(real_missing)} missing keys (первые 10): {real_missing[:10]}")
640 if unexpected:
641 print(f"⚠️ {len(unexpected)} unexpected keys (первые 10): {unexpected[:10]}")
642 return missing, unexpected
643
644 # -----------------------------------------------------------------
645 # Подсадка РЕАЛЬНЫХ mtp.* весов из оригинального HF-чекпоинта.
646 # Нужна, потому что старые JiRack .pt сохранены ДО появления MTPBlock
647 # и mtp-весов не содержат (там была бы случайная инициализация).
648 # -----------------------------------------------------------------
649 def load_mtp_from_hf_dir(self, hf_dir: str):
650 """Читает из папки оригинального HF-чекпоинта ТОЛЬКО mtp.* тензоры
651 (по index.json или перебором *.safetensors) и грузит их в self.mtp.
652 Возвращает число загруженных тензоров."""
653 import glob
654 import json as _json
655 from safetensors import safe_open
656
657 index_path = os.path.join(hf_dir, "model.safetensors.index.json")
658 mtp_map = {} # key -> shard file
659 if os.path.exists(index_path):
660 with open(index_path) as f:
661 weight_map = _json.load(f)["weight_map"]
662 for k, shard in weight_map.items():
663 if k.startswith("mtp."):
664 mtp_map[k] = os.path.join(hf_dir, shard)
665 else:
666 for shard in glob.glob(os.path.join(hf_dir, "*.safetensors")):
667 with safe_open(shard, framework="pt", device="cpu") as f:
668 for k in f.keys():
669 if k.startswith("mtp."):
670 mtp_map[k] = shard
671
672 if not mtp_map:
673 raise FileNotFoundError(
674 f"В {hf_dir} не найдено ни одного mtp.* тензора. Нужен "
675 f"ОРИГИНАЛЬНЫЙ чекпоинт Qwen3.8-27B (не JiRack-экспорт!). "
676 f"Если оригинальные шарды удалены -- перекачай с HF Hub "
677 f"только шарды с mtp.* по model.safetensors.index.json.")
678
679 own_dtype = next(self.mtp.parameters()).dtype
680 loaded = {}
681 by_shard = {}
682 for k, shard in mtp_map.items():
683 by_shard.setdefault(shard, []).append(k)
684 for shard, keys in by_shard.items():
685 with safe_open(shard, framework="pt", device="cpu") as f:
686 for k in keys:
687 loaded[k] = f.get_tensor(k).to(own_dtype)
688
689 missing, unexpected = self.load_state_dict(loaded, strict=False)
690 n_ok = len(loaded)
691 print(f"✅ MTP: загружено {n_ok} тензоров из {len(by_shard)} шардов ({hf_dir})")
692 if unexpected:
693 print(f"⚠️ MTP: unexpected keys: {unexpected[:10]}")
694 return n_ok
695
696 # -----------------------------------------------------------------
697 # JiRack-чекпоинт (свой формат, с lambda/patched_linear_names)
698 # -----------------------------------------------------------------
699 def save_checkpoint(self, path: str, lam: float, global_step: int = 0, patched_linear_names=None):
700 torch.save({
701 "model": self.state_dict(),
702 "lambda": lam,
703 "global_step": global_step,
704 "patched_linear_names": patched_linear_names or [],
705 }, path)
706
707 @classmethod
708 def load_checkpoint(cls, path: str, dtype=torch.bfloat16, include_gate_proj: bool = False):
709 ckpt = torch.load(path, map_location="cpu", weights_only=False)
710 model = cls()
711 patched_names = ckpt.get("patched_linear_names", [])
712 if patched_names:
713 for name in patched_names:
714 *parent_path, leaf = name.split(".")
715 parent = model
716 for p in parent_path:
717 parent = getattr(parent, p)
718 old = getattr(parent, leaf)
719 if not isinstance(old, BitLinear):
720 setattr(parent, leaf, BitLinear.from_linear(old))
721 else:
722 apply_bitlinear_patch(model, include_gate_proj=include_gate_proj)
723
724 has_mtp_in_ckpt = any(k.startswith("mtp.") for k in ckpt["model"])
725 model.load_state_dict(ckpt["model"], strict=False)
726 model = model.to(dtype)
727
728 if not has_mtp_in_ckpt:
729 print(
730 "⚠️ В этом .pt НЕТ mtp.* весов (чекпоинт сохранён до "
731 "появления MTPBlock). Голова сейчас со СЛУЧАЙНОЙ "
732 "инициализацией! Перед экспортом обязательно вызови "
733 "model.load_mtp_from_hf_dir('<папка оригинального HF-чекпоинта>')."
734 )
735
736 lam = ckpt.get("lambda", 0.0)
737 for m in model.modules():
738 if isinstance(m, BitLinear):
739 m.set_lambda(lam)
740
741 print(f"✅ Загружено. global_step={ckpt.get('global_step', 0)}, lambda={lam}")
742 return model
743
744
745 # =====================================================================
746 # 7. Демонстрация (без скачивания весов -- только структура)
747 # =====================================================================
748
749 if __name__ == "__main__":
750 print(ARCH_SUMMARY)
751 print(KEPT_FULL_PRECISION)
752
753 model = JiRackQwen38ForCausalLM()
754 n_params = sum(p.numel() for p in model.parameters())
755 print(f"✅ Модель собрана (случайные веса). Параметров: {n_params / 1e9:.2f}B")
756
757 patched = apply_bitlinear_patch(model)
758 print(f"✅ Заменено подмодулей на BitLinear: {len(patched)}")
759 print("Примеры первых 10:")
760 for n in patched[:10]:
761 print(f" {n}")
762
763 ids = torch.randint(0, VOCAB_SIZE, (1, 8))
764 with torch.no_grad():
765 logits = model(ids)
766 print(f"✅ Пробный forward (случайные веса) -> logits shape: {tuple(logits.shape)}")
767