inference/model.py
44.1 KB · 962 lines · python Raw
1 import math
2 from dataclasses import dataclass
3 from typing import Tuple, Optional, Literal
4 from functools import lru_cache
5 from contextlib import contextmanager
6
7 import torch
8 from torch import nn
9 import torch.nn.functional as F
10 import torch.distributed as dist
11
12 from kernel import act_quant, fp4_act_quant, fp8_gemm, fp4_gemm, sparse_attn, hc_split_sinkhorn
13
14
15 world_size = 1
16 rank = 0
17 block_size = 128
18 fp4_block_size = 32
19 default_dtype = torch.bfloat16
20 scale_fmt = None
21 scale_dtype = torch.float32
22
23
24 @contextmanager
25 def set_dtype(dtype):
26 """Temporarily override torch default dtype, restoring it on exit (even if an exception occurs)."""
27 prev = torch.get_default_dtype()
28 torch.set_default_dtype(dtype)
29 try:
30 yield
31 finally:
32 torch.set_default_dtype(prev)
33
34 @dataclass
35 class ModelArgs:
36 """Model hyperparameters. Field names match the config JSON keys."""
37 max_batch_size: int = 4
38 max_seq_len: int = 4096
39 temperature: float = 1
40 dtype: Literal["bf16", "fp8"] = "fp8"
41 scale_fmt: Literal[None, "ue8m0"] = "ue8m0"
42 expert_dtype: Literal[None, "fp4"] = None
43 scale_dtype: Literal["fp32", "fp8"] = "fp8"
44 vocab_size: int = 129280
45 dim: int = 4096
46 moe_inter_dim: int = 4096
47 n_layers: int = 7
48 n_hash_layers: int = 0
49 n_mtp_layers: int = 1
50 n_heads: int = 64
51 # moe
52 n_routed_experts: int = 8
53 n_shared_experts: int = 1
54 n_activated_experts: int = 2
55 score_func: Literal["softmax", "sigmoid", "sqrtsoftplus"] = "sqrtsoftplus"
56 route_scale: float = 1.
57 swiglu_limit: float = 0.
58 # mqa
59 q_lora_rank: int = 1024
60 head_dim: int = 512
61 rope_head_dim: int = 64
62 norm_eps: float = 1e-6
63 o_groups: int = 8
64 o_lora_rank: int = 1024
65 window_size: int = 128
66 compress_ratios: Tuple[int] = (0, 0, 4, 128, 4, 128, 4, 0)
67 # yarn
68 compress_rope_theta: float = 40000.0
69 original_seq_len: int = 0
70 rope_theta: float = 10000.0
71 rope_factor: float = 40
72 beta_fast: int = 32
73 beta_slow: int = 1
74 # index
75 index_n_heads: int = 64
76 index_head_dim: int = 128
77 index_topk: int = 512
78 # hc
79 hc_mult: int = 4
80 hc_sinkhorn_iters: int = 20
81 hc_eps: float = 1e-6
82 # dspark
83 dspark_block_size: int = 0
84 dspark_noise_token_id: int = 0
85 dspark_target_layer_ids: Tuple[int] = tuple()
86 dspark_markov_rank: int = 256
87
88
89 class ParallelEmbedding(nn.Module):
90 """Embedding sharded along the vocab dimension. Each rank holds vocab_size // world_size rows.
91 Out-of-range indices are zero-masked before all_reduce to combine partial embeddings."""
92 def __init__(self, vocab_size: int, dim: int):
93 super().__init__()
94 self.vocab_size = vocab_size
95 self.dim = dim
96 assert vocab_size % world_size == 0, f"Vocabulary size must be divisible by world size (world_size={world_size})"
97 self.part_vocab_size = (vocab_size // world_size)
98 self.vocab_start_idx = rank * self.part_vocab_size
99 self.vocab_end_idx = self.vocab_start_idx + self.part_vocab_size
100 self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim))
101
102 def forward(self, x: torch.Tensor) -> torch.Tensor:
103 if world_size > 1:
104 mask = (x < self.vocab_start_idx) | (x >= self.vocab_end_idx)
105 x = x - self.vocab_start_idx
106 x[mask] = 0
107 y = F.embedding(x, self.weight)
108 if world_size > 1:
109 y[mask] = 0
110 dist.all_reduce(y)
111 return y
112
113
114 def linear(x: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor] = None) -> torch.Tensor:
115 """Dispatches to fp4_gemm / fp8_gemm / F.linear based on weight dtype.
116 For quantized weights, x is first quantized to FP8 via act_quant."""
117 assert bias is None
118
119 if weight.dtype == torch.float4_e2m1fn_x2:
120 x, s = act_quant(x, block_size, scale_fmt, scale_dtype)
121 return fp4_gemm(x, s, weight, weight.scale, scale_dtype)
122 elif weight.dtype == torch.float8_e4m3fn:
123 x, s = act_quant(x, block_size, scale_fmt, scale_dtype)
124 return fp8_gemm(x, s, weight, weight.scale, scale_dtype)
125 else:
126 return F.linear(x, weight)
127
128
129 class Linear(nn.Module):
130 """Linear layer supporting BF16, FP8, and FP4 weight formats with per-block scaling."""
131
132 def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
133 super().__init__()
134 self.in_features = in_features
135 self.out_features = out_features
136 dtype = dtype or default_dtype
137 if dtype == torch.float4_e2m1fn_x2:
138 # FP4: weight is [out, in//2] in float4_e2m1fn_x2, logically [out, in] in fp4
139 # Scale is [out, in//32] in float8_e8m0fnu (1 scale per 32 fp4 elements along K)
140 self.weight = nn.Parameter(torch.empty(out_features, in_features // 2, dtype=torch.float4_e2m1fn_x2))
141 scale_out_features = out_features
142 scale_in_features = in_features // fp4_block_size
143 self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu))
144 elif dtype == torch.float8_e4m3fn:
145 self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype))
146 scale_out_features = (out_features + block_size - 1) // block_size
147 scale_in_features = (in_features + block_size - 1) // block_size
148 self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu))
149 else:
150 self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype))
151 self.register_parameter("scale", None)
152 if bias:
153 self.bias = nn.Parameter(torch.empty(out_features))
154 else:
155 self.register_parameter("bias", None)
156
157 def forward(self, x: torch.Tensor) -> torch.Tensor:
158 return linear(x, self.weight, self.bias)
159
160
161 class ColumnParallelLinear(Linear):
162 """Shards output dim across TP ranks. No all-reduce needed on output."""
163 def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
164 assert out_features % world_size == 0, f"Output features must be divisible by world size (world_size={world_size})"
165 self.part_out_features = out_features // world_size
166 super().__init__(in_features, self.part_out_features, bias, dtype)
167
168 def forward(self, x: torch.Tensor) -> torch.Tensor:
169 return linear(x, self.weight, self.bias)
170
171
172 class RowParallelLinear(Linear):
173 """Shards input dim across TP ranks. All-reduce on output to sum partial results."""
174 def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None):
175 assert in_features % world_size == 0, f"Input features must be divisible by world size (world_size={world_size})"
176 self.part_in_features = in_features // world_size
177 super().__init__(self.part_in_features, out_features, bias, dtype)
178
179 def forward(self, x: torch.Tensor) -> torch.Tensor:
180 y = linear(x, self.weight, None)
181 if world_size > 1:
182 y = y.float()
183 dist.all_reduce(y)
184 if self.bias is not None:
185 y += self.bias
186 return y.type_as(x)
187
188
189 class RMSNorm(nn.Module):
190 def __init__(self, dim: int, eps: float = 1e-6):
191 super().__init__()
192 self.dim = dim
193 self.eps = eps
194 # rmsnorm in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient.
195 self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32))
196
197 def forward(self, x: torch.Tensor):
198 dtype = x.dtype
199 x = x.float()
200 var = x.square().mean(-1, keepdim=True)
201 x = x * torch.rsqrt(var + self.eps)
202 return (self.weight * x).to(dtype)
203
204
205 @lru_cache(2)
206 def precompute_freqs_cis(dim, seqlen, original_seq_len, base, factor, beta_fast, beta_slow) -> torch.Tensor:
207 """Precomputes complex exponentials for rotary embeddings with YaRN scaling.
208 When original_seq_len > 0, applies frequency interpolation with a smooth
209 linear ramp between beta_fast and beta_slow correction ranges."""
210
211 def find_correction_dim(num_rotations, dim, base, max_seq_len):
212 return dim * math.log(max_seq_len / (num_rotations * 2 * math.pi)) / (2 * math.log(base))
213
214 def find_correction_range(low_rot, high_rot, dim, base, max_seq_len):
215 low = math.floor(find_correction_dim(low_rot, dim, base, max_seq_len))
216 high = math.ceil(find_correction_dim(high_rot, dim, base, max_seq_len))
217 return max(low, 0), min(high, dim-1)
218
219 def linear_ramp_factor(min, max, dim):
220 if min == max:
221 max += 0.001
222 linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min)
223 ramp_func = torch.clamp(linear_func, 0, 1)
224 return ramp_func
225
226 freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
227 if original_seq_len > 0:
228 low, high = find_correction_range(beta_fast, beta_slow, dim, base, original_seq_len)
229 smooth = 1 - linear_ramp_factor(low, high, dim // 2)
230 freqs = freqs / factor * (1 - smooth) + freqs * smooth
231
232 t = torch.arange(seqlen)
233 freqs = torch.outer(t, freqs)
234 freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
235 return freqs_cis
236
237
238 def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor, inverse: bool = False) -> torch.Tensor:
239 """Applies rotary positional embeddings in-place. Uses conjugate for inverse (de-rotation)."""
240 y = x
241 x = torch.view_as_complex(x.float().unflatten(-1, (-1, 2)))
242 if inverse:
243 freqs_cis = freqs_cis.conj()
244 if x.ndim == 3:
245 freqs_cis = freqs_cis.view(1, x.size(1), x.size(-1))
246 else:
247 freqs_cis = freqs_cis.view(1, x.size(1), 1, x.size(-1))
248 x = torch.view_as_real(x * freqs_cis).flatten(-2)
249 y.copy_(x)
250 return y
251
252
253 def rotate_activation(x: torch.Tensor) -> torch.Tensor:
254 """Applies randomized Hadamard rotation to spread information across dims before FP8 quant."""
255 assert x.dtype == torch.bfloat16
256 from fast_hadamard_transform import hadamard_transform
257 return hadamard_transform(x, scale=x.size(-1) ** -0.5)
258
259
260 @lru_cache(1)
261 def get_window_topk_idxs(window_size: int, bsz: int, seqlen: int, start_pos: int):
262 if start_pos >= window_size - 1:
263 start_pos %= window_size
264 matrix = torch.cat([torch.arange(start_pos + 1, window_size), torch.arange(0, start_pos + 1)], dim=0)
265 elif start_pos > 0:
266 matrix = F.pad(torch.arange(start_pos + 1), (0, window_size - start_pos - 1), value=-1)
267 else:
268 base = torch.arange(seqlen).unsqueeze(1)
269 matrix = (base - window_size + 1).clamp(0) + torch.arange(min(seqlen, window_size))
270 matrix = torch.where(matrix > base, -1, matrix)
271 return matrix.int().unsqueeze(0).expand(bsz, -1, -1).contiguous()
272
273
274 @lru_cache(2)
275 def get_compress_topk_idxs(ratio: int, bsz: int, seqlen: int, start_pos: int, offset: int):
276 if start_pos > 0:
277 matrix = torch.arange(0, (start_pos + 1) // ratio) + offset
278 else:
279 matrix = torch.arange(seqlen // ratio).repeat(seqlen, 1)
280 mask = matrix >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
281 matrix = torch.where(mask, -1, matrix + offset)
282 return matrix.int().unsqueeze(0).expand(bsz, -1, -1).contiguous()
283
284
285 class Compressor(nn.Module):
286 """Compresses KV cache via learned gated pooling over `compress_ratio` consecutive tokens.
287 When overlap=True (ratio==4), uses overlapping windows for smoother compression boundaries."""
288
289 def __init__(self, args: ModelArgs, compress_ratio: int = 4, head_dim: int = 512, rotate: bool = False):
290 super().__init__()
291 self.dim = args.dim
292 self.head_dim = head_dim
293 self.rope_head_dim = args.rope_head_dim
294 self.nope_head_dim = head_dim - args.rope_head_dim
295 self.compress_ratio = compress_ratio
296 self.overlap = compress_ratio == 4
297 self.rotate = rotate
298 coff = 1 + self.overlap
299
300 self.ape = nn.Parameter(torch.empty(compress_ratio, coff * self.head_dim, dtype=torch.float32))
301 # wkv and wgate in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient.
302 # When overlap, the first half of dims is for overlapping compression, second half for normal.
303 self.wkv = Linear(self.dim, coff * self.head_dim, dtype=torch.float32)
304 self.wgate = Linear(self.dim, coff * self.head_dim, dtype=torch.float32)
305 self.norm = RMSNorm(self.head_dim, args.norm_eps)
306 self.kv_cache: torch.Tensor = None # assigned lazily from Attention.kv_cache
307 # State buffers for decode-phase incremental compression.
308 # With overlap: state[:, :ratio] = overlapping window, state[:, ratio:] = current window.
309 self.register_buffer("kv_state", torch.zeros(args.max_batch_size, coff * compress_ratio, coff * self.head_dim, dtype=torch.float32), persistent=False)
310 self.register_buffer("score_state", torch.full((args.max_batch_size, coff * compress_ratio, coff * self.head_dim), float("-inf"), dtype=torch.float32), persistent=False)
311 self.freqs_cis: torch.Tensor = None
312
313 def overlap_transform(self, tensor: torch.Tensor, value=0):
314 # tensor: [b,s,r,2d]
315 b, s, _, _ = tensor.size()
316 ratio, d = self.compress_ratio, self.head_dim
317 new_tensor = tensor.new_full((b, s, 2 * ratio, d), value)
318 new_tensor[:, :, ratio:] = tensor[:, :, :, d:]
319 new_tensor[:, 1:, :ratio] = tensor[:, :-1, :, :d]
320 return new_tensor
321
322 def forward(self, x: torch.Tensor, start_pos: int):
323 assert self.kv_cache is not None
324 bsz, seqlen, _ = x.size()
325 ratio, overlap, d, rd = self.compress_ratio, self.overlap, self.head_dim, self.rope_head_dim
326 dtype = x.dtype
327 # compression need fp32
328 x = x.float()
329 kv = self.wkv(x)
330 score = self.wgate(x)
331 if start_pos == 0:
332 should_compress = seqlen >= ratio
333 remainder = seqlen % ratio
334 cutoff = seqlen - remainder
335 offset = ratio if overlap else 0
336 if overlap and cutoff >= ratio:
337 self.kv_state[:bsz, :ratio] = kv[:, cutoff-ratio : cutoff]
338 self.score_state[:bsz, :ratio] = score[:, cutoff-ratio : cutoff] + self.ape
339 if remainder > 0:
340 kv, self.kv_state[:bsz, offset : offset+remainder] = kv.split([cutoff, remainder], dim=1)
341 self.score_state[:bsz, offset : offset+remainder] = score[:, cutoff:] + self.ape[:remainder]
342 score = score[:, :cutoff]
343 kv = kv.unflatten(1, (-1, ratio))
344 score = score.unflatten(1, (-1, ratio)) + self.ape
345 if overlap:
346 kv = self.overlap_transform(kv, 0)
347 score = self.overlap_transform(score, float("-inf"))
348 kv = (kv * score.softmax(dim=2)).sum(dim=2)
349 else:
350 should_compress = (start_pos + 1) % self.compress_ratio == 0
351 score += self.ape[start_pos % ratio]
352 if overlap:
353 self.kv_state[:bsz, ratio + start_pos % ratio] = kv.squeeze(1)
354 self.score_state[:bsz, ratio + start_pos % ratio] = score.squeeze(1)
355 if should_compress:
356 kv_state = torch.cat([self.kv_state[:bsz, :ratio, :d], self.kv_state[:bsz, ratio:, d:]], dim=1)
357 score_state = torch.cat([self.score_state[:bsz, :ratio, :d], self.score_state[:bsz, ratio:, d:]], dim=1)
358 kv = (kv_state * score_state.softmax(dim=1)).sum(dim=1, keepdim=True)
359 self.kv_state[:bsz, :ratio] = self.kv_state[:bsz, ratio:]
360 self.score_state[:bsz, :ratio] = self.score_state[:bsz, ratio:]
361 else:
362 self.kv_state[:bsz, start_pos % ratio] = kv.squeeze(1)
363 self.score_state[:bsz, start_pos % ratio] = score.squeeze(1)
364 if should_compress:
365 kv = (self.kv_state[:bsz] * self.score_state[:bsz].softmax(dim=1)).sum(dim=1, keepdim=True)
366 if not should_compress:
367 return
368 kv = self.norm(kv.to(dtype))
369 if start_pos == 0:
370 freqs_cis = self.freqs_cis[:cutoff:ratio]
371 else:
372 freqs_cis = self.freqs_cis[start_pos + 1 - self.compress_ratio].unsqueeze(0)
373 apply_rotary_emb(kv[..., -rd:], freqs_cis)
374 if self.rotate:
375 kv = rotate_activation(kv)
376 fp4_act_quant(kv, fp4_block_size, True)
377 else:
378 act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
379 if start_pos == 0:
380 self.kv_cache[:bsz, :seqlen // ratio] = kv
381 else:
382 self.kv_cache[:bsz, start_pos // ratio] = kv.squeeze(1)
383 return kv
384
385
386 class Indexer(torch.nn.Module):
387 """Selects top-k compressed KV positions for sparse attention via learned scoring.
388 Has its own Compressor (with Hadamard rotation) to build compressed KV for scoring."""
389
390 def __init__(self, args: ModelArgs, compress_ratio: int = 4):
391 super().__init__()
392 self.dim = args.dim
393 self.n_heads = args.index_n_heads
394 self.n_local_heads = args.index_n_heads // world_size
395 self.head_dim = args.index_head_dim
396 self.rope_head_dim = args.rope_head_dim
397 self.index_topk = args.index_topk
398 self.q_lora_rank = args.q_lora_rank
399 self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim)
400 self.weights_proj = ColumnParallelLinear(self.dim, self.n_heads, dtype=torch.bfloat16)
401 self.softmax_scale = self.head_dim ** -0.5
402 self.compress_ratio = compress_ratio
403
404 self.compressor = Compressor(args, compress_ratio, self.head_dim, True)
405 self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, args.max_seq_len // compress_ratio, self.head_dim), persistent=False)
406 self.freqs_cis = None
407
408 def forward(self, x: torch.Tensor, qr: torch.Tensor, start_pos: int, offset: int):
409 bsz, seqlen, _ = x.size()
410 freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
411 ratio = self.compress_ratio
412 rd = self.rope_head_dim
413 end_pos = start_pos + seqlen
414 if self.compressor.kv_cache is None:
415 self.compressor.kv_cache = self.kv_cache
416 self.compressor.freqs_cis = self.freqs_cis
417 q = self.wq_b(qr)
418 q = q.unflatten(-1, (self.n_local_heads, self.head_dim))
419 apply_rotary_emb(q[..., -rd:], freqs_cis)
420 q = rotate_activation(q)
421 # use fp4 simulation for q and kv in indexer
422 fp4_act_quant(q, fp4_block_size, True)
423 self.compressor(x, start_pos)
424 weights = self.weights_proj(x) * (self.softmax_scale * self.n_heads ** -0.5)
425 # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16
426 index_score = torch.einsum("bshd,btd->bsht", q, self.kv_cache[:bsz, :end_pos // ratio])
427 index_score = (index_score.relu_() * weights.unsqueeze(-1)).sum(dim=2)
428 if world_size > 1:
429 dist.all_reduce(index_score)
430 if start_pos == 0:
431 mask = torch.arange(seqlen // ratio).repeat(seqlen, 1) >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
432 index_score += torch.where(mask, float("-inf"), 0)
433 topk_idxs = index_score.topk(min(self.index_topk, end_pos // ratio), dim=-1)[1]
434 if start_pos == 0:
435 mask = topk_idxs >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio
436 topk_idxs = torch.where(mask, -1, topk_idxs + offset)
437 else:
438 topk_idxs += offset
439 return topk_idxs
440
441
442 class Attention(nn.Module):
443 """Multi-head Latent Attention (MLA) with sliding window + optional KV compression.
444 Uses low-rank Q projection (wq_a -> q_norm -> wq_b) and grouped low-rank O projection."""
445 def __init__(self, layer_id: int, args: ModelArgs):
446 super().__init__()
447 self.layer_id = layer_id
448 self.dim = args.dim
449 self.n_heads = args.n_heads
450 self.n_local_heads = args.n_heads // world_size
451 self.q_lora_rank = args.q_lora_rank
452 self.o_lora_rank = args.o_lora_rank
453 self.head_dim = args.head_dim
454 self.rope_head_dim = args.rope_head_dim
455 self.nope_head_dim = args.head_dim - args.rope_head_dim
456 self.n_groups = args.o_groups
457 self.n_local_groups = self.n_groups // world_size
458 self.window_size = args.window_size
459 self.compress_ratio = args.compress_ratios[layer_id]
460 self.eps = args.norm_eps
461
462 self.attn_sink = nn.Parameter(torch.empty(self.n_local_heads, dtype=torch.float32))
463 self.wq_a = Linear(self.dim, self.q_lora_rank)
464 self.q_norm = RMSNorm(self.q_lora_rank, self.eps)
465 self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim)
466 self.wkv = Linear(self.dim, self.head_dim)
467 self.kv_norm = RMSNorm(self.head_dim, self.eps)
468 self.wo_a = ColumnParallelLinear(self.n_heads * self.head_dim // self.n_groups, self.n_groups * args.o_lora_rank, dtype=torch.bfloat16)
469 self.wo_b = RowParallelLinear(self.n_groups * args.o_lora_rank, self.dim)
470 self.softmax_scale = self.head_dim ** -0.5
471
472 if self.compress_ratio:
473 self.compressor = Compressor(args, self.compress_ratio, self.head_dim)
474 if self.compress_ratio == 4:
475 self.indexer = Indexer(args, self.compress_ratio)
476 else:
477 self.indexer = None
478
479 kv_cache_size = args.window_size + (args.max_seq_len // self.compress_ratio if self.compress_ratio else 0)
480 self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, kv_cache_size, self.head_dim), persistent=False)
481 if self.compress_ratio:
482 original_seq_len, rope_theta = args.original_seq_len, args.compress_rope_theta
483 else:
484 # disable YaRN and use base rope_theta in pure sliding-window attention
485 original_seq_len, rope_theta = 0, args.rope_theta
486 freqs_cis = precompute_freqs_cis(self.rope_head_dim, args.max_seq_len, original_seq_len,
487 rope_theta, args.rope_factor, args.beta_fast, args.beta_slow)
488 self.register_buffer("freqs_cis", freqs_cis, persistent=False)
489
490 def forward(self, x: torch.Tensor, start_pos: int):
491 bsz, seqlen, _ = x.size()
492 freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
493 win = self.window_size
494 ratio = self.compress_ratio
495 rd = self.rope_head_dim
496 if self.compress_ratio and self.compressor.kv_cache is None:
497 self.compressor.kv_cache = self.kv_cache[:, win:]
498 self.compressor.freqs_cis = self.freqs_cis
499 if self.indexer is not None:
500 self.indexer.freqs_cis = self.freqs_cis
501 # q
502 qr = q = self.q_norm(self.wq_a(x))
503 q = self.wq_b(q).unflatten(-1, (self.n_local_heads, self.head_dim))
504 q *= torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps)
505 apply_rotary_emb(q[..., -rd:], freqs_cis)
506
507 # win kv & topk_idxs
508 kv = self.wkv(x)
509 kv = self.kv_norm(kv)
510 apply_rotary_emb(kv[..., -rd:], freqs_cis)
511 # FP8-simulate non-rope dims to match QAT; rope dims stay bf16 for positional precision
512 act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
513 topk_idxs = get_window_topk_idxs(win, bsz, seqlen, start_pos)
514 if self.compress_ratio:
515 offset = kv.size(1) if start_pos == 0 else win
516 if self.indexer is not None:
517 compress_topk_idxs = self.indexer(x, qr, start_pos, offset).int()
518 else:
519 compress_topk_idxs = get_compress_topk_idxs(ratio, bsz, seqlen, start_pos, offset)
520 topk_idxs = torch.cat([topk_idxs, compress_topk_idxs], dim=-1)
521
522 # compress kv & attn
523 if start_pos == 0:
524 if seqlen <= win:
525 self.kv_cache[:bsz, :seqlen] = kv
526 else:
527 cutoff = seqlen % win
528 self.kv_cache[:bsz, cutoff: win], self.kv_cache[:bsz, :cutoff] = kv[:, -win:].split([win - cutoff, cutoff], dim=1)
529 if self.compress_ratio:
530 if (kv_compress := self.compressor(x, start_pos)) is not None:
531 kv = torch.cat([kv, kv_compress], dim=1)
532 # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16
533 o = sparse_attn(q, kv, self.attn_sink, topk_idxs, self.softmax_scale)
534 else:
535 self.kv_cache[:bsz, start_pos % win] = kv.squeeze(1)
536 if self.compress_ratio:
537 self.compressor(x, start_pos)
538 o = sparse_attn(q, self.kv_cache[:bsz], self.attn_sink, topk_idxs, self.softmax_scale)
539 apply_rotary_emb(o[..., -rd:], freqs_cis, True)
540
541 # o
542 o = o.view(bsz, seqlen, self.n_local_groups, -1)
543 wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1)
544 # NOTE: wo_a is FP8 in checkpoint; could do FP8 einsum here for better perf,
545 # but using BF16 for simplicity.
546 o = torch.einsum("bsgd,grd->bsgr", o, wo_a)
547 x = self.wo_b(o.flatten(2))
548 return x
549
550
551 class Gate(nn.Module):
552 """MoE gating: computes expert routing scores and selects top-k experts.
553 Supports hash-based routing (first n_hash_layers) where expert indices are
554 predetermined per token ID, and score-based routing (remaining layers)."""
555 def __init__(self, layer_id: int, args: ModelArgs):
556 super().__init__()
557 self.dim = args.dim
558 self.topk = args.n_activated_experts
559 self.score_func = args.score_func
560 self.route_scale = args.route_scale
561 self.hash = layer_id < args.n_hash_layers
562 self.weight = nn.Parameter(torch.empty(args.n_routed_experts, args.dim))
563 if self.hash:
564 self.tid2eid = nn.Parameter(torch.empty(args.vocab_size, args.n_activated_experts, dtype=torch.int32), requires_grad=False)
565 self.bias = None
566 else:
567 self.bias = nn.Parameter(torch.empty(args.n_routed_experts, dtype=torch.float32))
568
569 def forward(self, x: torch.Tensor, input_ids: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]:
570 scores = linear(x.float(), self.weight.float())
571 if self.score_func == "softmax":
572 scores = scores.softmax(dim=-1)
573 elif self.score_func == "sigmoid":
574 scores = scores.sigmoid()
575 else:
576 scores = F.softplus(scores).sqrt()
577 original_scores = scores
578 # Bias shifts scores for expert selection (topk) but does not affect routing weights.
579 if self.bias is not None:
580 scores = scores + self.bias
581 if self.hash:
582 indices = self.tid2eid[input_ids]
583 else:
584 indices = scores.topk(self.topk, dim=-1)[1]
585 weights = original_scores.gather(1, indices)
586 if self.score_func != "softmax":
587 weights /= weights.sum(dim=-1, keepdim=True)
588 weights *= self.route_scale
589 return weights, indices
590
591
592 class Expert(nn.Module):
593 """Single MoE expert: SwiGLU FFN (w1, w2, w3). Computation in float32 for stability."""
594 def __init__(self, dim: int, inter_dim: int, dtype=None, swiglu_limit=0):
595 super().__init__()
596 self.w1 = Linear(dim, inter_dim, dtype=dtype)
597 self.w2 = Linear(inter_dim, dim, dtype=dtype)
598 self.w3 = Linear(dim, inter_dim, dtype=dtype)
599 self.swiglu_limit = swiglu_limit
600
601 def forward(self, x: torch.Tensor, weights: Optional[torch.Tensor] = None) -> torch.Tensor:
602 dtype = x.dtype
603 gate = self.w1(x).float()
604 up = self.w3(x).float()
605 if self.swiglu_limit > 0:
606 up = torch.clamp(up, min=-self.swiglu_limit, max=self.swiglu_limit)
607 gate = torch.clamp(gate, max=self.swiglu_limit)
608 x = F.silu(gate) * up
609 if weights is not None:
610 x = weights * x
611 return self.w2(x.to(dtype))
612
613
614 class MoE(nn.Module):
615 """Mixture-of-Experts: gate routes each token to top-k routed experts + 1 shared expert.
616 Experts are sharded across TP ranks; each rank handles n_routed_experts // world_size experts."""
617 def __init__(self, layer_id: int, args: ModelArgs):
618 super().__init__()
619 self.layer_id = layer_id
620 self.dim = args.dim
621 assert args.n_routed_experts % world_size == 0, f"Number of experts must be divisible by world size (world_size={world_size})"
622 self.n_routed_experts = args.n_routed_experts
623 self.n_local_experts = args.n_routed_experts // world_size
624 self.n_activated_experts = args.n_activated_experts
625 self.experts_start_idx = rank * self.n_local_experts
626 self.experts_end_idx = self.experts_start_idx + self.n_local_experts
627 self.gate = Gate(layer_id, args)
628 expert_dtype = torch.float4_e2m1fn_x2 if args.expert_dtype == "fp4" else None
629 self.experts = nn.ModuleList([Expert(args.dim, args.moe_inter_dim, dtype=expert_dtype, swiglu_limit=args.swiglu_limit) if self.experts_start_idx <= i < self.experts_end_idx else None
630 for i in range(self.n_routed_experts)])
631 assert args.n_shared_experts == 1
632 self.shared_experts = Expert(args.dim, args.moe_inter_dim, swiglu_limit=args.swiglu_limit)
633
634 def forward(self, x: torch.Tensor, input_ids: torch.Tensor) -> torch.Tensor:
635 shape = x.size()
636 x = x.view(-1, self.dim)
637 weights, indices = self.gate(x, input_ids.flatten())
638 y = torch.zeros_like(x, dtype=torch.float32)
639 counts = torch.bincount(indices.flatten(), minlength=self.n_routed_experts).tolist()
640 for i in range(self.experts_start_idx, self.experts_end_idx):
641 if counts[i] == 0:
642 continue
643 expert = self.experts[i]
644 idx, top = torch.where(indices == i)
645 y[idx] += expert(x[idx], weights[idx, top, None])
646 if world_size > 1:
647 dist.all_reduce(y)
648 y += self.shared_experts(x)
649 return y.type_as(x).view(shape)
650
651
652 class Block(nn.Module):
653 """Transformer block with Hyper-Connections (HC) mixing.
654 Instead of a simple residual, HC maintains `hc_mult` copies of the hidden state.
655 hc_pre: reduces hc copies -> 1 via learned weighted sum (pre-weights from Sinkhorn).
656 hc_post: expands 1 -> hc copies via learned post-weights + combination matrix."""
657 attention_cls = Attention
658
659 def __init__(self, layer_id: int, args: ModelArgs):
660 super().__init__()
661 self.layer_id = layer_id
662 self.norm_eps = args.norm_eps
663 self.attn = self.attention_cls(layer_id, args)
664 self.ffn = MoE(layer_id, args)
665 self.attn_norm = RMSNorm(args.dim, self.norm_eps)
666 self.ffn_norm = RMSNorm(args.dim, self.norm_eps)
667 self.hc_mult = hc_mult = args.hc_mult
668 self.hc_sinkhorn_iters = args.hc_sinkhorn_iters
669 self.hc_eps = args.hc_eps
670 mix_hc = (2 + hc_mult) * hc_mult
671 hc_dim = hc_mult * args.dim
672 with set_dtype(torch.float32):
673 self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim))
674 self.hc_ffn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim))
675 self.hc_attn_base = nn.Parameter(torch.empty(mix_hc))
676 self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc))
677 self.hc_attn_scale = nn.Parameter(torch.empty(3))
678 self.hc_ffn_scale = nn.Parameter(torch.empty(3))
679
680 def hc_pre(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor):
681 # x: [b,s,hc,d], hc_fn: [mix_hc,hc*d], hc_scale: [3], hc_base: [mix_hc], y: [b,s,hc,d]
682 shape, dtype = x.size(), x.dtype
683 x = x.flatten(2).float()
684 rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps)
685 mixes = F.linear(x, hc_fn) * rsqrt
686 pre, post, comb = hc_split_sinkhorn(mixes, hc_scale, hc_base, self.hc_mult, self.hc_sinkhorn_iters, self.hc_eps)
687 y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2)
688 return y.to(dtype), post, comb
689
690 def hc_post(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor):
691 # x: [b,s,d], residual: [b,s,hc,d], post: [b,s,hc], comb: [b,s,hc,hc], y: [b,s,hc,d]
692 y = post.unsqueeze(-1) * x.unsqueeze(-2) + torch.sum(comb.unsqueeze(-1) * residual.unsqueeze(-2), dim=2)
693 return y.type_as(x)
694
695 def forward(self, x: torch.Tensor, start_pos: int, input_ids: Optional[torch.Tensor], *attn_args) -> torch.Tensor:
696 residual = x
697 x, post, comb = self.hc_pre(x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base)
698 x = self.attn_norm(x)
699 x = self.attn(x, start_pos, *attn_args)
700 x = self.hc_post(x, residual, post, comb)
701
702 residual = x
703 x, post, comb = self.hc_pre(x, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base)
704 x = self.ffn_norm(x)
705 x = self.ffn(x, input_ids)
706 x = self.hc_post(x, residual, post, comb)
707 return x
708
709 def hc_head(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor):
710 shape, dtype = x.size(), x.dtype
711 x = x.flatten(2).float()
712 rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps)
713 mixes = F.linear(x, hc_fn) * rsqrt
714 pre = torch.sigmoid(mixes * hc_scale + hc_base) + self.hc_eps
715 y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2)
716 return y.to(dtype)
717
718
719 class ParallelHead(nn.Module):
720
721 def __init__(self, vocab_size: int, dim: int, norm_eps: float = 1e-6, hc_eps: float = 1e-6):
722 super().__init__()
723 self.vocab_size = vocab_size
724 self.dim = dim
725 self.norm_eps = norm_eps
726 self.hc_eps = hc_eps
727 self.part_vocab_size = (vocab_size // world_size)
728 # lm_head in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for easier computation of logits later.
729 self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim, dtype=torch.float32))
730
731 def forward(self, x: torch.Tensor, full_logits=False):
732 # x: [b,s,hc,d]
733 if not full_logits:
734 x = x[:, -1]
735 logits = F.linear(x.float(), self.weight)
736 if world_size > 1:
737 all_logits = [torch.empty_like(logits) for _ in range(world_size)]
738 dist.all_gather(all_logits, logits)
739 logits = torch.cat(all_logits, dim=-1)
740 return logits
741
742
743 @lru_cache(1)
744 def get_dspark_topk_idxs(window_size: int, bsz: int, block_size: int, start_pos: int):
745 assert start_pos > 0
746 matrix = torch.cat([torch.arange(min(window_size, start_pos + 1)), window_size + torch.arange(block_size)])
747 return matrix.int().view(1, 1, -1).expand(bsz, block_size, -1).contiguous()
748
749
750 class DSparkAttention(Attention):
751
752 def forward(self, x: torch.Tensor, start_pos: int, main_x: torch.Tensor):
753 assert self.compress_ratio == 0
754 bsz, seqlen, _ = main_x.size()
755 win = self.window_size
756 rd = self.rope_head_dim
757
758 main_freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen]
759 main_kv = self.kv_norm(self.wkv(main_x))
760 apply_rotary_emb(main_kv[..., -rd:], main_freqs_cis)
761 act_quant(main_kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
762
763 if start_pos == 0:
764 if seqlen <= win:
765 self.kv_cache[:bsz, :seqlen] = main_kv
766 else:
767 cutoff = seqlen % win
768 self.kv_cache[:bsz, cutoff: win], self.kv_cache[:bsz, :cutoff] = main_kv[:, -win:].split([win - cutoff, cutoff], dim=1)
769 return x
770
771 bsz, block_size, _ = x.size()
772 freqs_cis = self.freqs_cis[start_pos+seqlen:start_pos+seqlen+block_size]
773
774 q = self.q_norm(self.wq_a(x))
775 q = self.wq_b(q).unflatten(-1, (self.n_local_heads, self.head_dim))
776 q *= torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps)
777 apply_rotary_emb(q[..., -rd:], freqs_cis)
778 kv = self.kv_norm(self.wkv(x))
779 apply_rotary_emb(kv[..., -rd:], freqs_cis)
780 act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True)
781
782 topk_idxs = get_dspark_topk_idxs(win, bsz, block_size, start_pos)
783 self.kv_cache[:bsz, start_pos % win] = main_kv.squeeze(1)
784 kv = torch.cat([self.kv_cache[:bsz], kv], dim=1)
785 o = sparse_attn(q, kv, self.attn_sink, topk_idxs, self.softmax_scale)
786 apply_rotary_emb(o[..., -rd:], freqs_cis, True)
787
788 o = o.view(bsz, block_size, self.n_local_groups, -1)
789 wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1)
790 o = torch.einsum("bsgd,grd->bsgr", o, wo_a)
791 x = self.wo_b(o.flatten(2))
792 return x
793
794
795 class DSparkMarkovHead(nn.Module):
796 def __init__(self, vocab_size: int, dspark_markov_rank: int):
797 super().__init__()
798 self.markov_w1 = ParallelEmbedding(vocab_size, dspark_markov_rank)
799 self.markov_w2 = ParallelHead(vocab_size, dspark_markov_rank)
800
801 def forward(self, token_ids: torch.Tensor) -> torch.Tensor:
802 embed = self.markov_w1(token_ids)
803 logits = self.markov_w2(embed, full_logits=True)
804 return logits, embed
805
806
807 class DSparkConfidenceHead(nn.Module):
808 def __init__(self, input_dim: int):
809 super().__init__()
810 # proj in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for fp32 confidence score.
811 self.proj = Linear(input_dim, 1, dtype=torch.float32)
812
813 def forward(self, hidden: torch.Tensor, markov_embed: torch.Tensor):
814 hidden = torch.cat([hidden, markov_embed], dim=-1)
815 return self.proj(hidden.float()).squeeze(-1)
816
817
818 class DSparkBlock(Block):
819 """DSpark stage stored under the mtp.* checkpoint namespace."""
820 attention_cls = DSparkAttention
821
822 def __init__(self, layer_id: int, args: ModelArgs):
823 super().__init__(layer_id, args)
824 self.dim = args.dim
825 stage_id = layer_id - args.n_layers
826 self.block_size = args.dspark_block_size
827 self.noise_token_id = args.dspark_noise_token_id
828 self.temperature = args.temperature
829 hc_dim = self.hc_mult * args.dim
830 if stage_id == 0:
831 assert len(args.dspark_target_layer_ids) > 0, "DSpark needs target layers"
832 self.main_proj = Linear(args.dim * len(args.dspark_target_layer_ids), args.dim)
833 self.main_norm = RMSNorm(args.dim, args.norm_eps)
834 if stage_id == args.n_mtp_layers - 1:
835 self.norm = RMSNorm(args.dim, args.norm_eps)
836 self.markov_head = DSparkMarkovHead(args.vocab_size, args.dspark_markov_rank)
837 self.confidence_head = DSparkConfidenceHead(args.dim + args.dspark_markov_rank)
838 with set_dtype(torch.float32):
839 self.hc_head_fn = nn.Parameter(torch.empty(self.hc_mult, hc_dim))
840 self.hc_head_base = nn.Parameter(torch.empty(self.hc_mult))
841 self.hc_head_scale = nn.Parameter(torch.empty(1))
842 self.embed: ParallelEmbedding = None
843 self.head: ParallelHead = None
844
845 def forward(self, x: torch.Tensor, start_pos: int, input_ids: torch.Tensor, main_x: torch.Tensor) -> torch.Tensor:
846 if start_pos > 0:
847 return super().forward(x, start_pos, input_ids, main_x)
848 # only compute KV cache in prefill stage
849 return self.attn(x, start_pos, main_x)
850
851 def forward_embed(self, main_hidden: torch.Tensor, input_ids: torch.Tensor):
852 assert self.embed is not None
853 main_x = self.main_norm(self.main_proj(main_hidden))
854 draft_input_ids = input_ids.new_full([input_ids.size(0), self.block_size], self.noise_token_id)
855 draft_input_ids[:, 0] = input_ids
856 x = self.embed(draft_input_ids)
857 x = x.unsqueeze(2).repeat(1, 1, self.hc_mult, 1)
858 return x, main_x
859
860 def forward_head(self, x: torch.Tensor, input_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
861 assert self.head is not None
862 x = self.hc_head(x, self.hc_head_fn, self.hc_head_scale, self.hc_head_base)
863 logits = self.head(self.norm(x), full_logits=True)
864 output_ids = input_ids.new_empty(input_ids.size(0), self.block_size + 1)
865 output_ids[:, 0] = input_ids
866 markov_embeds = []
867 for i in range(self.block_size):
868 logits_bias, markov_embed = self.markov_head(output_ids[:, i])
869 logits[:, i].add_(logits_bias)
870 markov_embeds.append(markov_embed)
871 output_ids[:, i + 1] = sample(logits[:, i], self.temperature)
872 markov_embed = torch.stack(markov_embeds, dim=1)
873 confidence = self.confidence_head(x, markov_embed)
874 return output_ids, logits, confidence
875
876
877 class Transformer(nn.Module):
878 """Full DeepSeek-V4 model: embed -> HC-expand -> N blocks -> HC-head -> logits.
879 Sets global state (world_size, rank, default_dtype, scale_fmt, scale_dtype) in __init__."""
880 def __init__(self, args: ModelArgs):
881 global world_size, rank, default_dtype, scale_fmt, scale_dtype
882 world_size = dist.get_world_size() if dist.is_initialized() else 1
883 rank = dist.get_rank() if dist.is_initialized() else 0
884 default_dtype = torch.float8_e4m3fn if args.dtype == "fp8" else torch.bfloat16
885 scale_fmt = "ue8m0" if args.scale_dtype == "fp8" else args.scale_fmt
886 scale_dtype = torch.float8_e8m0fnu if args.scale_dtype == "fp8" else torch.float32
887 super().__init__()
888 self.max_seq_len = args.max_seq_len
889 self.temperature = args.temperature
890 self.norm_eps = args.norm_eps
891 self.hc_eps = args.hc_eps
892 self.embed = ParallelEmbedding(args.vocab_size, args.dim)
893 self.layers = torch.nn.ModuleList()
894 for layer_id in range(args.n_layers):
895 self.layers.append(Block(layer_id, args))
896 self.norm = RMSNorm(args.dim, self.norm_eps)
897 self.head = ParallelHead(args.vocab_size, args.dim, self.norm_eps, self.hc_eps)
898 self.mtp = torch.nn.ModuleList()
899 self.target_layer_ids = args.dspark_target_layer_ids
900 if args.dspark_block_size:
901 for layer_id in range(args.n_mtp_layers):
902 self.mtp.append(DSparkBlock(args.n_layers + layer_id, args))
903 self.mtp[-1].embed = self.embed
904 self.mtp[-1].head = self.head
905 self.hc_mult = hc_mult = args.hc_mult
906 hc_dim = hc_mult * args.dim
907 with set_dtype(torch.float32):
908 self.hc_head_fn = nn.Parameter(torch.empty(hc_mult, hc_dim))
909 self.hc_head_base = nn.Parameter(torch.empty(hc_mult))
910 self.hc_head_scale = nn.Parameter(torch.empty(1))
911
912 @torch.inference_mode()
913 def forward(self, input_ids: torch.Tensor, start_pos: int = 0):
914 h = self.embed(input_ids)
915 # Expand to hc_mult copies for Hyper-Connections
916 h = h.unsqueeze(2).repeat(1, 1, self.hc_mult, 1)
917 main_hiddens = []
918 for i, layer in enumerate(self.layers):
919 h = layer(h, start_pos, input_ids)
920 if i in self.target_layer_ids:
921 main_hiddens.append(h.mean(dim=2))
922 h = layer.hc_head(h, self.hc_head_fn, self.hc_head_scale, self.hc_head_base)
923 logits = self.head(self.norm(h))
924 output_ids = sample(logits, self.temperature)
925 main_hidden = torch.cat(main_hiddens, dim=-1) if main_hiddens else None
926 return output_ids, logits, main_hidden
927
928 @torch.inference_mode()
929 def forward_spec(self, input_ids: torch.Tensor, main_hidden: torch.Tensor, start_pos: int = 0):
930 h, main_x = self.mtp[0].forward_embed(main_hidden, input_ids)
931 for layer in self.mtp:
932 h = layer(h, start_pos, input_ids, main_x)
933 if start_pos == 0:
934 return
935 output_ids, logits, confidence = self.mtp[-1].forward_head(h, input_ids)
936 return output_ids, logits, confidence
937
938
939 def sample(logits, temperature: float = 1.0):
940 """Gumbel-max trick: equivalent to multinomial sampling but faster on GPU,
941 since it avoids the GPU-to-CPU sync in torch.multinomial."""
942 if temperature == 0:
943 return logits.argmax(dim=-1)
944 logits = logits / max(temperature, 1e-5)
945 probs = torch.softmax(logits, dim=-1, dtype=torch.float32)
946 return probs.div_(torch.empty_like(probs).exponential_(1)).argmax(dim=-1)
947
948
949 if __name__ == "__main__":
950 torch.set_default_dtype(torch.bfloat16)
951 torch.set_default_device("cuda")
952 torch.manual_seed(0)
953 args = ModelArgs(n_hash_layers=0, dspark_block_size=6, dspark_target_layer_ids=(5, 6))
954 x = torch.randint(0, args.vocab_size, (2, 150))
955 model = Transformer(args)
956
957 output_ids, logits, main_hidden = model(x[:, :128])
958 model.forward_spec(output_ids, main_hidden)
959 for i in range(128, 150):
960 output_ids, logits, main_hidden = model(x[:, i:i+1], i)
961 output_ids, logits, confidence = model.forward_spec(output_ids, main_hidden, i)
962