JiRackDeltaNet_27b.py
| 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 | |