diff --git a/README.md b/README.md index 2a56a0865..057b02655 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,7 @@ FreeToken is an edge-native Mixture-of-Experts (MoE) serving engine designed for - **Fast Edge-Native Runtime**: Provides efficient MoE serving with bandwidth-adaptive CPU–GPU co-execution ($q^\star$ policy), full-layer double-buffered prefill streaming, global LRU expert caching, graph-compatible execution, and the FTW fast weight format. - **Semantic-Aware Caching**: Features semantic anchor checkpoints for recurrent state and KV caches, allowing agentic context edits (e.g., tool calls, thinking blocks) to avoid redundant context recomputation. - **Elastic Memory Management**: Supports dynamic, runtime VRAM re-allocation between expert caches and KV memory without engine restarts or weight reloading. -- **Broad MoE & Ecosystem Support**: Supports frontier open-weight MoE models (e.g., DeepSeek-V4-Flash, Qwen3.6-35B-A3B, GLM-5.2) across various parameter scales and quantization formats (e.g., MXFP4, NVFP4, FP8, BF16), with Anthropic/OpenAI-compatible APIs for seamless integration with real-world coding and tool-calling agents (e.g., Codex, Claude Code, OpenCode, OpenClaw, DeepSeek Harness). +- **Broad MoE & Ecosystem Support**: Supports frontier open-weight MoE models (e.g., DeepSeek-V4-Flash, Qwen3.6-35B-A3B, GLM-5.3-Flash) across various parameter scales and quantization formats (e.g., MXFP4, NVFP4, FP8, BF16), with Anthropic/OpenAI-compatible APIs for seamless integration with real-world coding and tool-calling agents (e.g., Codex, Claude Code, OpenCode, OpenClaw, DeepSeek Harness). - **Diverse Consumer Hardware**: Scales across consumer laptops, gaming desktops, and workstation GPUs, with native support for NVIDIA RTX 30, RTX 40, and RTX 50 series GPUs. ## Getting Started diff --git a/benchmarks/bench_decode_moe.py b/benchmarks/bench_decode_moe.py index 566217927..c22d882ac 100644 --- a/benchmarks/bench_decode_moe.py +++ b/benchmarks/bench_decode_moe.py @@ -99,10 +99,27 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: default=-1, help="hybrid: max PCIe fetches/layer; -1 = auto (benched pcie/cpu bandwidth fraction)", ) + p.add_argument( + "--cpu-threads", + type=int, + default=0, + help="CPU MoE worker threads; 0 = runtime auto-selection", + ) + p.add_argument( + "--kv-reserve-tokens", + type=int, + default=0, + help="KV token floor reserved before auto-sizing the expert cache; 0 = server default", + ) p.add_argument("--mem-ratio", type=float, default=0.9, help="target VRAM utilization") p.add_argument("--gpu", default=None, help="GPU for the serve: a UUID or nvidia-smi index (as ft serve --gpu)") p.add_argument("--no-graph", action="store_true", help="eager decode instead of CUDA graph") + p.add_argument( + "--disable-prefill-overlap", + action="store_true", + help="disable MoE prefill overlap (permits expert caches below 2 * num_experts)", + ) p.add_argument( "--greedy", action="store_true", @@ -187,6 +204,12 @@ def serve_cmd(args: argparse.Namespace, backend: str, port: int) -> list[str]: ] if args.gpu: cmd += ["--gpu", args.gpu] + if args.cpu_threads > 0: + cmd += ["--moe-cpu-threads", str(args.cpu_threads)] + if args.kv_reserve_tokens > 0: + cmd += ["--kv-reserve-tokens", str(args.kv_reserve_tokens)] + if args.disable_prefill_overlap: + cmd.append("--disable-moe-prefill-overlap") if args.cache > 0: cmd += ["--moe-cache-size", str(args.cache)] elif args.cache_rate is not None: diff --git a/docs/models.md b/docs/models.md index 41b79cd68..7c4f4f3c3 100644 --- a/docs/models.md +++ b/docs/models.md @@ -7,6 +7,7 @@ for them; other checkpoints of the same architectures work too. | Model | HF checkpoints | |---|---| | DeepSeek-V4 | [deepseek-ai/DeepSeek-V4-Flash-0731](https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash-0731) | +| GLM-5.3-Flash | [LibertAIDAI/GLM-5.3-Flash-NVFP4](https://huggingface.co/LibertAIDAI/GLM-5.3-Flash-NVFP4) | | GLM-5.2 | [nvidia/GLM-5.2-NVFP4](https://huggingface.co/nvidia/GLM-5.2-NVFP4) | | GLM-4.7 | [nvidia/GLM-4.7-NVFP4](https://huggingface.co/nvidia/GLM-4.7-NVFP4) | | Qwen3.8-Flash-Next | [Qwen/Qwen3.8-Flash-Next-FP8](https://huggingface.co/Qwen/Qwen3.8-Flash-Next-FP8), [RadixArk/Qwen3.8-Flash-Next-NVFP4](https://huggingface.co/RadixArk/Qwen3.8-Flash-Next-NVFP4) | @@ -38,5 +39,7 @@ for them; other checkpoints of the same architectures work too. FreeToken's fast-load format, and `ft serve --model` auto-detects the result. - DeepSeek-V4 checkpoints must keep the `inference/config.json` subdir — the authoritative model args are read from there. +- GLM-5.3-Flash currently supports TP=1 with its NVFP4 routed experts served + from host RAM by the offload backend. - Qwen3.8-Flash-Next keeps a 47.7 GiB PLE n-gram table pinned in host RAM. - Multimodal checkpoints are served text-only. diff --git a/python/freetoken/attention/dsa.py b/python/freetoken/attention/dsa.py index ff4fa7ca6..9425017bc 100644 --- a/python/freetoken/attention/dsa.py +++ b/python/freetoken/attention/dsa.py @@ -31,9 +31,10 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Dict, List, Tuple +from typing import TYPE_CHECKING import torch + from freetoken.core import Batch, get_global_ctx from .base import AttentionSpec, BaseAttnBackend, BaseAttnMetadata @@ -73,8 +74,8 @@ class DSAAttnBackend(DSAIndexerMixin, BaseAttnBackend): def __init__(self, config: ModelConfig) -> None: from freetoken.kvcache.dsa_pool import DSAKVCache, MLAKVCache - args = config.glm_dsa_args - assert args is not None, "dsa backend needs ModelConfig.glm_dsa_args (MLA dims)" + args = config.glm_dsa_args or config.glm5_args + assert args is not None, "dsa backend needs GLM MLA dimensions" self.config = config self.num_heads = config.num_qo_heads self.kv_lora_rank = args.kv_lora_rank @@ -96,13 +97,29 @@ def __init__(self, config: ModelConfig) -> None: # layer -> group leader (most recent "full" layer); leader -> pool slot. # Only built when DSA serves: the dense ablation never consults indexer_types, # so a checkpoint with a malformed list cannot crash the ablation. - self._leader: Dict[int, int] = {} - self._idx_slot: Dict[int, int] = {} + self._leader: dict[int, int] = {} + self._idx_slot: dict[int, int] = {} if self.dsa_enabled: lead = None + # Hybrid models (GLM-5.3-Flash) carry an ``indexer_types`` entry for + # every decoder layer, including KDA layers that never enter this + # backend. The index slab is deliberately sized only for MLA/DSA + # layers, so enumerate that group's layer ids rather than treating a + # ``full`` marker on a linear-attention layer as an allocated slot. + if getattr(config, "glm5_args", None) is not None: + full_group = next( + group for group in config.attention_groups if group.name == "full" + ) + served_layers = set(full_group.layer_ids) + else: + # GLM-5.2 is all-MLA and several backend-level tests intentionally + # pass a minimal config namespace without generic group metadata. + served_layers = set(range(config.num_layers)) # Capped to the SERVED layer count (dev num_layers overrides must not # index slots past the pool the factory sized from the same cap). for lid, kind in enumerate(args.indexer_types[: config.num_layers]): + if lid not in served_layers: + continue if kind == "full": lead = lid self._idx_slot[lid] = len(self._idx_slot) @@ -113,7 +130,7 @@ def __init__(self, config: ModelConfig) -> None: self._rows_buf: torch.Tensor | None = None self._kvlen_buf: torch.Tensor | None = None self.max_seq_len = 0 - self.capture_bs: List[int] = [] + self.capture_bs: list[int] = [] def forward(self, q, k, v, layer_id, batch, attn_spec: AttentionSpec | None = None): raise NotImplementedError("MLA models use mla_forward(), not forward().") @@ -128,7 +145,9 @@ def prepare_metadata(self, batch: Batch) -> None: # active_table_idx (which the decode path's addressing requires) for # phase == "decode". The prefill path handles extend_len == 1 fine. is_decode = getattr(batch, "phase", None) == "decode" - qo_indptr = torch.tensor([0] + seqlens_q, **_CPU_PINNED).cumsum_(0).to(torch.int32) + qo_indptr = ( + torch.tensor([0] + seqlens_q, **_CPU_PINNED).cumsum_(0).to(torch.int32) + ) kv_len = torch.tensor(seqlens_k, **_CPU_PINNED) last = (qo_indptr[1:].to(torch.int32) - 1).to(self.device, non_blocking=True) md = DSAMetadata( @@ -151,8 +170,12 @@ def _attend( from freetoken.kernel.triton.glm_dsa_sparse import glm_dsa_sparse_attn return glm_dsa_sparse_attn( - q_cat, self.kvcache.latent_rows(layer_id), sel, self.sm_scale, - counts=cnt, d_v=self.kv_lora_rank, + q_cat, + self.kvcache.latent_rows(layer_id), + sel, + self.sm_scale, + counts=cnt, + d_v=self.kv_lora_rank, ) def mla_forward( @@ -177,7 +200,9 @@ def mla_forward( if self.dsa_enabled and indexer_qkw is not None: # Scatter index keys unconditionally: short prefills serve through the # identity path TODAY, but their keys must exist once decode passes topk. - self.kvcache.store_index_k(indexer_qkw[1], batch.out_loc, self._idx_slot[layer_id]) + self.kvcache.store_index_k( + indexer_qkw[1], batch.out_loc, self._idx_slot[layer_id] + ) if md.is_decode: return self._decode(md, layer_id, q_nope, q_pe, indexer_qkw) @@ -194,7 +219,9 @@ def _decode(self, md, layer_id, q_nope, q_pe, indexer_qkw) -> torch.Tensor: else: if indexer_qkw is not None: q_idx, _, w = indexer_qkw - s = self.dsa_decode_scores(q_idx, w, self._idx_slot[layer_id], rows, kvlen) + s = self.dsa_decode_scores( + q_idx, w, self._idx_slot[layer_id], rows, kvlen + ) k_sel = min(self.index_topk, s.shape[-1]) picks = self.indexer_select_decode( s.view(bs, 1, -1), valid=kvlen, topk=k_sel, offset=0 @@ -205,15 +232,21 @@ def _decode(self, md, layer_id, q_nope, q_pe, indexer_qkw) -> torch.Tensor: md.sel.clear() md.sel[layer_id] = (sel, cnt) sel, cnt = md.sel[self._leader[layer_id]] - q_cat = torch.cat([q_nope, q_pe], dim=-1).view(bs, 1, self.num_heads, self.latent_dim) + q_cat = torch.cat([q_nope, q_pe], dim=-1).view( + bs, 1, self.num_heads, self.latent_dim + ) o = self._attend(q_cat, layer_id, sel, cnt) return o.view(bs, self.num_heads, self.kv_lora_rank) # ----- prefill / extend (eager) ------------------------------------------------------ def _select_prefill( - self, slot: int, q_idx: torch.Tensor, w: torch.Tensor, - rows: torch.Tensor, positions: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor]: + self, + slot: int, + q_idx: torch.Tensor, + w: torch.Tensor, + rows: torch.Tensor, + positions: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: """Per-request causal top-k: ([1, m, K] physical rows, [1, m] counts).""" kv_len = rows.numel() k_all = self.kvcache.index_k_cache(slot).index_select(0, rows.long()) @@ -223,14 +256,20 @@ def _select_prefill( start_pos = int(positions[0]) # Bound the fp32 [chunk, kv_len] logits transient (worst case is capped by the # model's max_position: floor 16 x 1M x 4 B = 64 MB, see _PREFILL_SCORE_BYTES). - chunk = max(16, min(_PREFILL_SCORE_CHUNK, _PREFILL_SCORE_BYTES // max(kv_len * 4, 1))) + chunk = max( + 16, min(_PREFILL_SCORE_CHUNK, _PREFILL_SCORE_BYTES // max(kv_len * 4, 1)) + ) for s0 in range(0, m, chunk): s1 = min(s0 + chunk, m) scores = self.dsa_prefill_logits(q_idx[s0:s1], k_all, w[s0:s1]) # Shared selection semantics (dsv4_indexer): token-granular == ratio 1. picks = self.indexer_select_prefill( - scores.unsqueeze(0), start_pos=start_pos + s0, seqlen=s1 - s0, - ratio=1, topk=k_sel, offset=0, + scores.unsqueeze(0), + start_pos=start_pos + s0, + seqlen=s1 - s0, + ratio=1, + topk=k_sel, + offset=0, )[0] sel[s0:s1] = self.dsa_map_rows(picks, rows.view(1, -1).expand(s1 - s0, -1)) cnt = torch.clamp(positions + 1, max=k_sel).to(torch.int32) @@ -249,7 +288,8 @@ def _prefill(self, md, layer_id, q_nope, q_pe, batch, indexer_qkw) -> torch.Tens md.sel[layer_id] = [ self._select_prefill( self._idx_slot[layer_id], - q_idx[qo[i] : qo[i + 1]], w[qo[i] : qo[i + 1]], + q_idx[qo[i] : qo[i + 1]], + w[qo[i] : qo[i + 1]], page_table[r.table_idx, : r.device_len], batch.positions[qo[i] : qo[i + 1]], ) @@ -267,34 +307,50 @@ def _prefill(self, md, layer_id, q_nope, q_pe, batch, indexer_qkw) -> torch.Tens # token at kv <= index_topk, and the ablation attends everything). # One shared row list broadcast across queries (stride 0), causality # through per-query counts. - sel = page_table[r.table_idx, : r.device_len].view(1, 1, -1).to(torch.int32) - cnt = (batch.positions[qo[i] : qo[i + 1]] + 1).to(torch.int32).view(1, m) + sel = ( + page_table[r.table_idx, : r.device_len] + .view(1, 1, -1) + .to(torch.int32) + ) + cnt = ( + (batch.positions[qo[i] : qo[i + 1]] + 1).to(torch.int32).view(1, m) + ) o[qo[i] : qo[i + 1]] = self._attend( q_cat[qo[i] : qo[i + 1]].view(1, m, self.num_heads, self.latent_dim), - layer_id, sel, cnt, + layer_id, + sel, + cnt, ).view(m, self.num_heads, self.kv_lora_rank) return o # ----- CUDA graph (decode) ---------------------------------------------------------- - def init_capture_graph(self, max_seq_len: int, bs_list: List[int]) -> None: + def init_capture_graph(self, max_seq_len: int, bs_list: list[int]) -> None: self.max_seq_len = max_seq_len self.capture_bs = sorted(bs_list) max_bs = max(bs_list) width = get_global_ctx().page_table.shape[1] - self._rows_buf = torch.full((max_bs, width), -1, dtype=torch.int32, device=self.device) + self._rows_buf = torch.full( + (max_bs, width), -1, dtype=torch.int32, device=self.device + ) self._kvlen_buf = torch.zeros(max_bs, dtype=torch.int32, device=self.device) def _decode_rows(self, batch: Batch) -> torch.Tensor: """This decode step's per-request page-table rows [bs, W], gathered off the scheduler-staged ``active_table_idx`` (a device tensor -- no host loop).""" - assert batch.active_table_idx is not None, "decode batch is missing its page-table rows" - return get_global_ctx().page_table.index_select(0, batch.active_table_idx.to(torch.int64)) + assert batch.active_table_idx is not None, ( + "decode batch is missing its page-table rows" + ) + return get_global_ctx().page_table.index_select( + 0, batch.active_table_idx.to(torch.int64) + ) def _stage_decode(self, batch: Batch, bs: int, table_idx: torch.Tensor) -> None: """Copy this step's addressing into the static graph buffers and point the metadata at them (restage-per-replay, same shape as the generic backends).""" md = batch.attn_metadata - self._rows_buf[:bs].copy_(get_global_ctx().page_table.index_select(0, table_idx)) + self._rows_buf[:bs].copy_( + get_global_ctx().page_table.index_select(0, table_idx) + ) self._kvlen_buf[:bs].copy_(md.kv_len_cpu.to(self.device, non_blocking=True)) md.rows = self._rows_buf[:bs] md.kvlen = self._kvlen_buf[:bs] @@ -311,7 +367,9 @@ def prepare_for_capture(self, batch: Batch) -> None: self._stage_decode(batch, bs, dummy) def prepare_for_replay(self, batch: Batch) -> None: - assert batch.active_table_idx is not None, "decode batch is missing its page-table rows" + assert batch.active_table_idx is not None, ( + "decode batch is missing its page-table rows" + ) self._stage_decode( batch, batch.padded_size, batch.active_table_idx.to(torch.int64) ) diff --git a/python/freetoken/kernel/aot_models.py b/python/freetoken/kernel/aot_models.py index c9c2fb98e..4d94f99c1 100644 --- a/python/freetoken/kernel/aot_models.py +++ b/python/freetoken/kernel/aot_models.py @@ -269,6 +269,17 @@ def expert_bank_row_bytes(fmt: str, hidden_size: int, moe_intermediate_size: int moe_intermediate_size=2048, expert_formats=_NVFP4_FORMATS, ), + AotModel( + # Hybrid KDA + MLA sparse attention. KDA state and the MLA latent/index + # slabs bypass store_cache, so this family contributes no paged-KV row. + name="zai-org/GLM-5.3-Flash", + architecture="Glm5NextForConditionalGeneration", + hidden_size=4096, + kv_groups=(), + top_k=8, + moe_intermediate_size=2048, + expert_formats=("fp8_block",), + ), AotModel( # MiniMaxAI/MiniMax-M2.5 ships block-fp8, which has no expert-bank # provider for this arch on main -- the NVFP4 release is the servable @@ -317,7 +328,13 @@ def expert_bank_row_bytes(fmt: str, hidden_size: int, moe_intermediate_size: int architecture="Qwen3_5ForConditionalGeneration", hidden_size=5120, kv_groups=((4, 256),), - aliases=("Qwen/Qwen3.6-27B-FP8", "nvidia/Qwen3.6-27B-NVFP4"), + aliases=( + "Qwen/Qwen3.6-27B-FP8", + "nvidia/Qwen3.6-27B-NVFP4", + "Qwen/Qwen3.8-27B", + "Qwen/Qwen3.8-27B-FP8", + "RadixArk/Qwen3.8-27B-NVFP4", + ), ), AotModel( name="google/gemma-4-12B-it", diff --git a/python/freetoken/kernel/fla/fused_sigmoid_gating_recurrent.py b/python/freetoken/kernel/fla/fused_sigmoid_gating_recurrent.py index 339005c0b..f64b63134 100644 --- a/python/freetoken/kernel/fla/fused_sigmoid_gating_recurrent.py +++ b/python/freetoken/kernel/fla/fused_sigmoid_gating_recurrent.py @@ -53,6 +53,8 @@ def fused_sigmoid_gating_delta_rule_update_kernel( USE_QK_L2NORM_IN_KERNEL: tl.constexpr, IS_VARLEN: tl.constexpr, IS_KDA: tl.constexpr, + HAS_GATE_LOWER_BOUND: tl.constexpr, + GATE_LOWER_BOUND: tl.constexpr, # Optional flags for target_verify support (default False for decode) DISABLE_STATE_UPDATE: tl.constexpr = False, CACHE_INTERMEDIATE_STATES: tl.constexpr = False, @@ -165,16 +167,21 @@ def fused_sigmoid_gating_delta_rule_update_kernel( b_a = tl.load(p_a).to(tl.float32) b_dt_bias = tl.load(p_dt_bias).to(tl.float32) - # Compute g = -exp(A_log) * softplus(a + dt_bias) + # Compute the log-decay. GLM-5.3's KDA uses its safe bounded gate + # ``lower_bound * sigmoid(exp(A_log) * x)``; GDN and unbounded KDA use + # ``-exp(A_log) * softplus(x)``. x = b_a + b_dt_bias - beta_x = softplus_beta * x - # Apply softplus with numerical stability - softplus_x = tl.where( - beta_x <= softplus_threshold, - (1.0 / softplus_beta) * tl.log(1.0 + tl.exp(beta_x)), - x, - ) - b_g = -tl.exp(b_A_log) * softplus_x + if IS_KDA and HAS_GATE_LOWER_BOUND: + b_g = GATE_LOWER_BOUND * tl.sigmoid(tl.exp(b_A_log) * x) + else: + beta_x = softplus_beta * x + # Apply softplus with numerical stability + softplus_x = tl.where( + beta_x <= softplus_threshold, + (1.0 / softplus_beta) * tl.log(1.0 + tl.exp(beta_x)), + x, + ) + b_g = -tl.exp(b_A_log) * softplus_x # Compute beta = sigmoid(b) b_beta = 1.0 / (1.0 + tl.exp(-b_b)) @@ -260,6 +267,7 @@ def fused_sigmoid_gating_delta_rule_update( use_qk_l2norm_in_kernel: bool = False, cu_seqlens: Optional[torch.Tensor] = None, is_kda: bool = False, + gate_lower_bound: Optional[float] = None, # Optional parameters for target_verify support disable_state_update: bool = False, intermediate_states_buffer: Optional[torch.Tensor] = None, @@ -362,6 +370,8 @@ def fused_sigmoid_gating_delta_rule_update( USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, IS_VARLEN=cu_seqlens is not None, IS_KDA=is_kda, + HAS_GATE_LOWER_BOUND=gate_lower_bound is not None, + GATE_LOWER_BOUND=gate_lower_bound if gate_lower_bound is not None else 0.0, DISABLE_STATE_UPDATE=disable_state_update, CACHE_INTERMEDIATE_STATES=intermediate_states_buffer is not None, HAS_EAGLE_TREE_CUSTOM_ATTN_MASK=retrieve_parent_token is not None, diff --git a/python/freetoken/kernel/triton/glm_dsa_sparse.py b/python/freetoken/kernel/triton/glm_dsa_sparse.py index d3219f950..50660cd9b 100644 --- a/python/freetoken/kernel/triton/glm_dsa_sparse.py +++ b/python/freetoken/kernel/triton/glm_dsa_sparse.py @@ -36,14 +36,29 @@ @triton.jit def _glm_dsa_sparse_kernel( - q_ptr, pool_ptr, o_ptr, idx_ptr, cnt_ptr, + q_ptr, + pool_ptr, + o_ptr, + idx_ptr, + cnt_ptr, scale, - H, TOPK, - stride_qb, stride_qm, stride_qh, stride_qd, - stride_pn, stride_pd, - stride_ob, stride_om, stride_oh, stride_od, - stride_ib, stride_im, stride_it, - stride_nb, stride_nm, + H, + TOPK, + stride_qb, + stride_qm, + stride_qh, + stride_qd, + stride_pn, + stride_pd, + stride_ob, + stride_om, + stride_oh, + stride_od, + stride_ib, + stride_im, + stride_it, + stride_nb, + stride_nm, D_V: tl.constexpr, D_R: tl.constexpr, BLOCK_H: tl.constexpr, @@ -57,11 +72,18 @@ def _glm_dsa_sparse_kernel( offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H) h_mask = offs_h < H offs_v = tl.arange(0, D_V) - offs_r = tl.arange(0, D_R) q_base = q_ptr + pid_b * stride_qb + pid_m * stride_qm + offs_h[:, None] * stride_qh - q_v = tl.load(q_base + offs_v[None, :] * stride_qd, mask=h_mask[:, None], other=0.0).to(tl.float32) - q_r = tl.load(q_base + (D_V + offs_r[None, :]) * stride_qd, mask=h_mask[:, None], other=0.0).to(tl.float32) + q_v = tl.load( + q_base + offs_v[None, :] * stride_qd, mask=h_mask[:, None], other=0.0 + ).to(tl.float32) + if D_R > 0: + offs_r = tl.arange(0, D_R) + q_r = tl.load( + q_base + (D_V + offs_r[None, :]) * stride_qd, + mask=h_mask[:, None], + other=0.0, + ).to(tl.float32) m_i = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) l_i = tl.zeros((BLOCK_H,), dtype=tl.float32) @@ -72,16 +94,24 @@ def _glm_dsa_sparse_kernel( n_active = tl.load(cnt_ptr + pid_b * stride_nb + pid_m * stride_nm) idx_base = idx_ptr + pid_b * stride_ib + pid_m * stride_im - for t in range(0, tl.cdiv(n_active, BLOCK_T)): + for t in range(tl.cdiv(n_active, BLOCK_T)): offs_t = t * BLOCK_T + tl.arange(0, BLOCK_T) t_mask = offs_t < n_active idxs = tl.load(idx_base + offs_t * stride_it, mask=t_mask, other=-1) valid = idxs >= 0 kv_base = pool_ptr + idxs[:, None] * stride_pn - kv_v = tl.load(kv_base + offs_v[None, :] * stride_pd, mask=valid[:, None], other=0.0).to(tl.float32) - kv_r = tl.load(kv_base + (D_V + offs_r[None, :]) * stride_pd, mask=valid[:, None], other=0.0).to(tl.float32) - - scores = (tl.dot(q_v, tl.trans(kv_v)) + tl.dot(q_r, tl.trans(kv_r))) * scale + kv_v = tl.load( + kv_base + offs_v[None, :] * stride_pd, mask=valid[:, None], other=0.0 + ).to(tl.float32) + scores = tl.dot(q_v, tl.trans(kv_v)) + if D_R > 0: + kv_r = tl.load( + kv_base + (D_V + offs_r[None, :]) * stride_pd, + mask=valid[:, None], + other=0.0, + ).to(tl.float32) + scores += tl.dot(q_r, tl.trans(kv_r)) + scores *= scale scores = tl.where(valid[None, :], scores, -float("inf")) m_new = tl.maximum(m_i, tl.max(scores, axis=1)) @@ -93,19 +123,36 @@ def _glm_dsa_sparse_kernel( m_i = m_new o = acc / l_i[:, None] - o_ptrs = o_ptr + pid_b * stride_ob + pid_m * stride_om + offs_h[:, None] * stride_oh + offs_v[None, :] * stride_od + o_ptrs = ( + o_ptr + + pid_b * stride_ob + + pid_m * stride_om + + offs_h[:, None] * stride_oh + + offs_v[None, :] * stride_od + ) tl.store(o_ptrs, o.to(o_ptr.dtype.element_ty), mask=h_mask[:, None]) @triton.jit def _glm_dsa_decode_logits_kernel( - q_ptr, w_ptr, pool_ptr, rows_ptr, valid_ptr, out_ptr, + q_ptr, + w_ptr, + pool_ptr, + rows_ptr, + valid_ptr, + out_ptr, N_STAGE, - stride_qb, stride_qh, stride_qd, - stride_wb, stride_wh, - stride_pr, stride_pd, - stride_rb, stride_rw, - stride_ob, stride_ot, + stride_qb, + stride_qh, + stride_qd, + stride_wb, + stride_wh, + stride_pr, + stride_pd, + stride_rb, + stride_rw, + stride_ob, + stride_ot, H: tl.constexpr, D: tl.constexpr, BLOCK_H: tl.constexpr, @@ -131,7 +178,9 @@ def _glm_dsa_decode_logits_kernel( n_valid = tl.load(valid_ptr + pid_b) if pid_t * BLOCK_T >= n_valid: - tl.store(out_ptrs, tl.full((BLOCK_T,), float("-inf"), tl.float32), mask=store_mask) + tl.store( + out_ptrs, tl.full((BLOCK_T,), float("-inf"), tl.float32), mask=store_mask + ) return t_mask = offs_t < n_valid @@ -139,13 +188,26 @@ def _glm_dsa_decode_logits_kernel( offs_h = tl.arange(0, BLOCK_H) h_mask = offs_h < H - rows = tl.load(rows_ptr + pid_b * stride_rb + offs_t * stride_rw, mask=t_mask, other=0) + rows = tl.load( + rows_ptr + pid_b * stride_rb + offs_t * stride_rw, mask=t_mask, other=0 + ) rows = tl.maximum(rows, 0) - k = tl.load(pool_ptr + rows[:, None] * stride_pr + offs_d[None, :] * stride_pd, - mask=t_mask[:, None], other=0.0) - q = tl.load(q_ptr + pid_b * stride_qb + offs_h[:, None] * stride_qh + offs_d[None, :] * stride_qd, - mask=h_mask[:, None], other=0.0) - w = tl.load(w_ptr + pid_b * stride_wb + offs_h * stride_wh, mask=h_mask, other=0.0).to(tl.float32) + k = tl.load( + pool_ptr + rows[:, None] * stride_pr + offs_d[None, :] * stride_pd, + mask=t_mask[:, None], + other=0.0, + ) + q = tl.load( + q_ptr + + pid_b * stride_qb + + offs_h[:, None] * stride_qh + + offs_d[None, :] * stride_qd, + mask=h_mask[:, None], + other=0.0, + ) + w = tl.load( + w_ptr + pid_b * stride_wb + offs_h * stride_wh, mask=h_mask, other=0.0 + ).to(tl.float32) score = tl.dot(q, tl.trans(k)) # [BLOCK_H, BLOCK_T] fp32 score = tl.maximum(score, 0.0) * w[:, None] @@ -155,11 +217,11 @@ def _glm_dsa_decode_logits_kernel( def glm_dsa_decode_logits( - q: torch.Tensor, # [B, H, D] bf16, one index query per request - weights: torch.Tensor, # [B, H] fp32 (head gate; softmax scale folded in by caller) - idx_pool: torch.Tensor,# [R, D] bf16, this slot's paged index keys - rows: torch.Tensor, # [B, W] int, physical-row snapshot in position order - valid: torch.Tensor, # [B] int32, live length per request (device-read) + q: torch.Tensor, # [B, H, D] bf16, one index query per request + weights: torch.Tensor, # [B, H] fp32 (head gate; softmax scale folded in by caller) + idx_pool: torch.Tensor, # [R, D] bf16, this slot's paged index keys + rows: torch.Tensor, # [B, W] int, physical-row snapshot in position order + valid: torch.Tensor, # [B] int32, live length per request (device-read) out: torch.Tensor | None = None, ) -> torch.Tensor: """Head-reduced indexer logits ``[B, W]`` fp32 for a decode step. @@ -177,31 +239,65 @@ def glm_dsa_decode_logits( out = torch.empty(b, w_stage, dtype=torch.float32, device=q.device) BLOCK_T = 64 _glm_dsa_decode_logits_kernel[(b, triton.cdiv(w_stage, BLOCK_T))]( - q, weights, idx_pool, rows, valid, out, + q, + weights, + idx_pool, + rows, + valid, + out, w_stage, - q.stride(0), q.stride(1), q.stride(2), - weights.stride(0), weights.stride(1), - idx_pool.stride(0), idx_pool.stride(1), - rows.stride(0), rows.stride(1), - out.stride(0), out.stride(1), - H=h, D=d, - BLOCK_H=triton.next_power_of_2(h), BLOCK_T=BLOCK_T, - num_warps=4, num_stages=2, + q.stride(0), + q.stride(1), + q.stride(2), + weights.stride(0), + weights.stride(1), + idx_pool.stride(0), + idx_pool.stride(1), + rows.stride(0), + rows.stride(1), + out.stride(0), + out.stride(1), + H=h, + D=d, + BLOCK_H=triton.next_power_of_2(h), + BLOCK_T=BLOCK_T, + num_warps=4, + num_stages=2, ) return out @triton.jit def _glm_dsa_splitk_kernel( - q_ptr, pool_ptr, mid_o_ptr, mid_lse_ptr, idx_ptr, cnt_ptr, + q_ptr, + pool_ptr, + mid_o_ptr, + mid_lse_ptr, + idx_ptr, + cnt_ptr, scale, - H, TOPK, - stride_qb, stride_qm, stride_qh, stride_qd, - stride_pn, stride_pd, - stride_mb, stride_mm, stride_mh, stride_ms, stride_md, - stride_lb, stride_lm, stride_lh, stride_ls, - stride_ib, stride_im, stride_it, - stride_nb, stride_nm, + H, + TOPK, + stride_qb, + stride_qm, + stride_qh, + stride_qd, + stride_pn, + stride_pd, + stride_mb, + stride_mm, + stride_mh, + stride_ms, + stride_md, + stride_lb, + stride_lm, + stride_lh, + stride_ls, + stride_ib, + stride_im, + stride_it, + stride_nb, + stride_nm, D_V: tl.constexpr, D_R: tl.constexpr, BLOCK_H: tl.constexpr, @@ -221,7 +317,8 @@ def _glm_dsa_splitk_kernel( offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H) h_mask = offs_h < H offs_v = tl.arange(0, D_V) - offs_r = tl.arange(0, D_R) + if D_R > 0: + offs_r = tl.arange(0, D_R) n_active = TOPK if HAS_COUNTS: @@ -236,9 +333,18 @@ def _glm_dsa_splitk_kernel( acc = tl.zeros((BLOCK_H, D_V), dtype=tl.float32) if split_end > split_start: - q_base = q_ptr + pid_b * stride_qb + pid_m * stride_qm + offs_h[:, None] * stride_qh - q_v = tl.load(q_base + offs_v[None, :] * stride_qd, mask=h_mask[:, None], other=0.0).to(tl.float32) - q_r = tl.load(q_base + (D_V + offs_r[None, :]) * stride_qd, mask=h_mask[:, None], other=0.0).to(tl.float32) + q_base = ( + q_ptr + pid_b * stride_qb + pid_m * stride_qm + offs_h[:, None] * stride_qh + ) + q_v = tl.load( + q_base + offs_v[None, :] * stride_qd, mask=h_mask[:, None], other=0.0 + ).to(tl.float32) + if D_R > 0: + q_r = tl.load( + q_base + (D_V + offs_r[None, :]) * stride_qd, + mask=h_mask[:, None], + other=0.0, + ).to(tl.float32) idx_base = idx_ptr + pid_b * stride_ib + pid_m * stride_im for start in range(split_start, split_end, BLOCK_T): @@ -247,10 +353,18 @@ def _glm_dsa_splitk_kernel( idxs = tl.load(idx_base + offs_t * stride_it, mask=t_mask, other=-1) valid = idxs >= 0 kv_base = pool_ptr + idxs[:, None] * stride_pn - kv_v = tl.load(kv_base + offs_v[None, :] * stride_pd, mask=valid[:, None], other=0.0).to(tl.float32) - kv_r = tl.load(kv_base + (D_V + offs_r[None, :]) * stride_pd, mask=valid[:, None], other=0.0).to(tl.float32) - - scores = (tl.dot(q_v, tl.trans(kv_v)) + tl.dot(q_r, tl.trans(kv_r))) * scale + kv_v = tl.load( + kv_base + offs_v[None, :] * stride_pd, mask=valid[:, None], other=0.0 + ).to(tl.float32) + scores = tl.dot(q_v, tl.trans(kv_v)) + if D_R > 0: + kv_r = tl.load( + kv_base + (D_V + offs_r[None, :]) * stride_pd, + mask=valid[:, None], + other=0.0, + ).to(tl.float32) + scores += tl.dot(q_r, tl.trans(kv_r)) + scores *= scale scores = tl.where(valid[None, :], scores, -float("inf")) m_new = tl.maximum(m_i, tl.max(scores, axis=1)) @@ -264,23 +378,42 @@ def _glm_dsa_splitk_kernel( lse = tl.where(l_i == 0.0, -float("inf"), m_i + tl.log(l_i)) mid_base = ( - mid_o_ptr + pid_b * stride_mb + pid_m * stride_mm - + offs_h[:, None] * stride_mh + split_id * stride_ms + offs_v[None, :] * stride_md + mid_o_ptr + + pid_b * stride_mb + + pid_m * stride_mm + + offs_h[:, None] * stride_mh + + split_id * stride_ms + + offs_v[None, :] * stride_md ) tl.store(mid_base, out, mask=h_mask[:, None]) lse_base = ( - mid_lse_ptr + pid_b * stride_lb + pid_m * stride_lm - + offs_h * stride_lh + split_id * stride_ls + mid_lse_ptr + + pid_b * stride_lb + + pid_m * stride_lm + + offs_h * stride_lh + + split_id * stride_ls ) tl.store(lse_base, lse, mask=h_mask) @triton.jit def _glm_dsa_merge_kernel( - mid_o_ptr, mid_lse_ptr, o_ptr, - stride_mb, stride_mm, stride_mh, stride_ms, stride_md, - stride_lb, stride_lm, stride_lh, stride_ls, - stride_ob, stride_om, stride_oh, stride_od, + mid_o_ptr, + mid_lse_ptr, + o_ptr, + stride_mb, + stride_mm, + stride_mh, + stride_ms, + stride_md, + stride_lb, + stride_lm, + stride_lh, + stride_ls, + stride_ob, + stride_om, + stride_oh, + stride_od, D_V: tl.constexpr, NUM_SPLITS: tl.constexpr, ): @@ -296,7 +429,10 @@ def _glm_dsa_merge_kernel( acc = tl.zeros((D_V,), dtype=tl.float32) mid_base = ( - mid_o_ptr + pid_b * stride_mb + pid_m * stride_mm + pid_h * stride_mh + mid_o_ptr + + pid_b * stride_mb + + pid_m * stride_mm + + pid_h * stride_mh + offs_v * stride_md ) lse_base = mid_lse_ptr + pid_b * stride_lb + pid_m * stride_lm + pid_h * stride_lh @@ -313,7 +449,10 @@ def _glm_dsa_merge_kernel( o = acc / l_i o_ptrs = ( - o_ptr + pid_b * stride_ob + pid_m * stride_om + pid_h * stride_oh + o_ptr + + pid_b * stride_ob + + pid_m * stride_om + + pid_h * stride_oh + offs_v * stride_od ) tl.store(o_ptrs, o.to(o_ptr.dtype.element_ty)) @@ -333,13 +472,14 @@ def _split_count(b: int, m: int, h: int, topk: int, device) -> int: def glm_dsa_sparse_attn( - q: torch.Tensor, # [b, m, h, d_v + d_r] (ckv-absorbed | rope) - pool: torch.Tensor, # [rows, d_v + d_r] GLOBAL latent pool slab (this layer) + q: torch.Tensor, # [b, m, h, d_v + d_r] (ckv-absorbed | rope) + pool: torch.Tensor, # [rows, d_v + d_r] GLOBAL latent pool slab (this layer) topk_idxs: torch.Tensor, # [b, m|1, topk] int32 global rows, -1 masked softmax_scale: float, - counts: torch.Tensor | None = None, # [b, m] int32 live columns per query (device-read) + counts: torch.Tensor + | None = None, # [b, m] int32 live columns per query (device-read) d_v: int = 512, - force_splits: int | None = None, # tests only: 0 = single-program, N = split-k N + force_splits: int | None = None, # tests only: 0 = single-program, N = split-k N ) -> torch.Tensor: """Sparse MLA attention over gathered latent rows; returns ``[b, m, h, d_v]``. @@ -367,53 +507,110 @@ def glm_dsa_sparse_attn( else: cnt, stride_nb, stride_nm = idx, 0, 0 - n_splits = _split_count(b, m, h, topk, q.device) if force_splits is None else force_splits + n_splits = ( + _split_count(b, m, h, topk, q.device) if force_splits is None else force_splits + ) if n_splits: mid_o = q.new_empty(b, m, h, n_splits, d_v, dtype=torch.float32) mid_lse = q.new_empty(b, m, h, n_splits, dtype=torch.float32) grid1 = (m * n_splits, b, triton.cdiv(h, BLOCK_H)) _glm_dsa_splitk_kernel[grid1]( - q, pool_2d, mid_o, mid_lse, idx, cnt, + q, + pool_2d, + mid_o, + mid_lse, + idx, + cnt, float(softmax_scale), - h, topk, - q.stride(0), q.stride(1), q.stride(2), q.stride(3), - pool_2d.stride(0), pool_2d.stride(1), - mid_o.stride(0), mid_o.stride(1), mid_o.stride(2), mid_o.stride(3), mid_o.stride(4), - mid_lse.stride(0), mid_lse.stride(1), mid_lse.stride(2), mid_lse.stride(3), - idx.stride(0), 0 if broadcast_m else idx.stride(1), idx.stride(2), - stride_nb, stride_nm, - D_V=d_v, D_R=d_r, - BLOCK_H=BLOCK_H, BLOCK_T=BLOCK_T, - HAS_COUNTS=has_counts, NUM_SPLITS=n_splits, - num_warps=4, num_stages=2, + h, + topk, + q.stride(0), + q.stride(1), + q.stride(2), + q.stride(3), + pool_2d.stride(0), + pool_2d.stride(1), + mid_o.stride(0), + mid_o.stride(1), + mid_o.stride(2), + mid_o.stride(3), + mid_o.stride(4), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + mid_lse.stride(3), + idx.stride(0), + 0 if broadcast_m else idx.stride(1), + idx.stride(2), + stride_nb, + stride_nm, + D_V=d_v, + D_R=d_r, + BLOCK_H=BLOCK_H, + BLOCK_T=BLOCK_T, + HAS_COUNTS=has_counts, + NUM_SPLITS=n_splits, + num_warps=4, + num_stages=2, ) grid2 = (m, b, h) _glm_dsa_merge_kernel[grid2]( - mid_o, mid_lse, o, - mid_o.stride(0), mid_o.stride(1), mid_o.stride(2), mid_o.stride(3), mid_o.stride(4), - mid_lse.stride(0), mid_lse.stride(1), mid_lse.stride(2), mid_lse.stride(3), - o.stride(0), o.stride(1), o.stride(2), o.stride(3), - D_V=d_v, NUM_SPLITS=n_splits, + mid_o, + mid_lse, + o, + mid_o.stride(0), + mid_o.stride(1), + mid_o.stride(2), + mid_o.stride(3), + mid_o.stride(4), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + mid_lse.stride(3), + o.stride(0), + o.stride(1), + o.stride(2), + o.stride(3), + D_V=d_v, + NUM_SPLITS=n_splits, num_warps=4, ) return o grid = (m, b, triton.cdiv(h, BLOCK_H)) _glm_dsa_sparse_kernel[grid]( - q, pool_2d, o, idx, cnt, + q, + pool_2d, + o, + idx, + cnt, float(softmax_scale), - h, topk, - q.stride(0), q.stride(1), q.stride(2), q.stride(3), - pool_2d.stride(0), pool_2d.stride(1), - o.stride(0), o.stride(1), o.stride(2), o.stride(3), - idx.stride(0), 0 if broadcast_m else idx.stride(1), idx.stride(2), - stride_nb, stride_nm, - D_V=d_v, D_R=d_r, - BLOCK_H=BLOCK_H, BLOCK_T=BLOCK_T, + h, + topk, + q.stride(0), + q.stride(1), + q.stride(2), + q.stride(3), + pool_2d.stride(0), + pool_2d.stride(1), + o.stride(0), + o.stride(1), + o.stride(2), + o.stride(3), + idx.stride(0), + 0 if broadcast_m else idx.stride(1), + idx.stride(2), + stride_nb, + stride_nm, + D_V=d_v, + D_R=d_r, + BLOCK_H=BLOCK_H, + BLOCK_T=BLOCK_T, HAS_COUNTS=has_counts, - num_warps=4, num_stages=2, + num_warps=4, + num_stages=2, ) return o -__all__ = ["glm_dsa_sparse_attn", "glm_dsa_decode_logits"] +__all__ = ["glm_dsa_decode_logits", "glm_dsa_sparse_attn"] diff --git a/python/freetoken/kvcache/__init__.py b/python/freetoken/kvcache/__init__.py index c6c0f1bb9..d9e942f94 100644 --- a/python/freetoken/kvcache/__init__.py +++ b/python/freetoken/kvcache/__init__.py @@ -6,6 +6,7 @@ if TYPE_CHECKING: import torch + from freetoken.models import ModelConfig from .base import ( @@ -75,8 +76,8 @@ def create_kv_pool(config, num_pages: int, device: torch.device, dtype: torch.dt secondary tier -- window pool, index slab, state rings -- are derived here or inside the pool). Single factory entry for all pool families, DSV4 included.""" from .dsv4_cost_model import _dsv4_pool_sizes - from .hybrid_swa_pool import _naive_swa_num_tokens, _swa_paged_num_tokens from .dsv4_paged_pool import DSV4PagedKVCache + from .hybrid_swa_pool import _naive_swa_num_tokens, _swa_paged_num_tokens model_config = config.model_config if resolve_pool_class(model_config) is DSV4PagedKVCache: @@ -147,7 +148,9 @@ def create_kvcache_pool( head_dim = model_config.head_dim if model_config.has_linear_attention: specs = [s for s in model_config.kv_cache_group_specs() if s.num_layers > 0] - assert len(specs) == 1, f"expected one paged-KV group, got {[s.name for s in specs]}" + assert len(specs) == 1, ( + f"expected one paged-KV group, got {[s.name for s in specs]}" + ) spec = specs[0] layer_ids = spec.layer_ids num_kv_heads = spec.num_kv_heads @@ -219,6 +222,7 @@ def create_kvcache_pool( device=device, index_head_dim=spec.index_head_dim, num_index_layers=spec.num_index_layers, + layer_ids=spec.layer_ids, ) return MLAKVCache( latent_dim=spec.head_dim, @@ -227,6 +231,7 @@ def create_kvcache_pool( page_size=page_size, dtype=dtype, device=device, + layer_ids=spec.layer_ids, ) return MHAKVCache( @@ -268,14 +273,14 @@ def create_prefix_cache( __all__ = [ + "SUPPORTED_CACHE_MANAGER", + "BaseCacheHandle", + "BaseKVCachePool", + "BasePrefixCache", + "MatchResult", + "SizeInfo", "create_kv_pool", "create_kvcache_pool", "create_prefix_cache", "resolve_pool_class", - "BaseKVCachePool", - "BaseCacheHandle", - "BasePrefixCache", - "SizeInfo", - "MatchResult", - "SUPPORTED_CACHE_MANAGER", ] diff --git a/python/freetoken/kvcache/dsa_pool.py b/python/freetoken/kvcache/dsa_pool.py index e6a51ac93..91ad3d6b7 100644 --- a/python/freetoken/kvcache/dsa_pool.py +++ b/python/freetoken/kvcache/dsa_pool.py @@ -17,6 +17,8 @@ from __future__ import annotations +from collections.abc import Sequence + import torch from .base import BaseKVCachePool @@ -37,9 +39,23 @@ def __init__( page_size: int, dtype: torch.dtype, device: torch.device, + layer_ids: Sequence[int] | None = None, ) -> None: self._latent_dim = latent_dim self._num_layers = num_layers + if layer_ids is None: + self._num_storage_layers = num_layers + self._layer_map: list[int] | None = None + else: + self._num_storage_layers = len(layer_ids) + layer_map = [-1] * num_layers + for dense, global_id in enumerate(layer_ids): + if global_id < 0 or global_id >= num_layers: + raise ValueError( + f"KV layer id {global_id} outside [0, {num_layers})" + ) + layer_map[global_id] = dense + self._layer_map = layer_map self._page_size = page_size self._dtype = dtype self._device = device @@ -48,15 +64,32 @@ def __init__( def _alloc(self, num_pages: int) -> None: self._num_pages = num_pages self._kv_buffer = torch.empty( - (1, self._num_layers, num_pages, self._page_size, 1, self._latent_dim), + ( + 1, + self._num_storage_layers, + num_pages, + self._page_size, + 1, + self._latent_dim, + ), device=self._device, dtype=self._dtype, ) # -- views ------------------------------------------------------------------ + def _dense(self, layer_id: int) -> int: + if self._layer_map is None: + return layer_id + dense = self._layer_map[layer_id] + if dense < 0: + raise KeyError(f"layer {layer_id} has no paged KV storage") + return dense + def k_cache(self, layer_id: int) -> torch.Tensor: """Paged latent view ``[num_pages, page_size, latent_dim]``.""" - return self._kv_buffer[0, layer_id].view(self._num_pages, self._page_size, -1) + return self._kv_buffer[0, self._dense(layer_id)].view( + self._num_pages, self._page_size, -1 + ) def v_cache(self, layer_id: int) -> torch.Tensor: # MLA: K == V (single latent); same buffer, dsv4_paged_pool precedent. @@ -64,7 +97,7 @@ def v_cache(self, layer_id: int) -> torch.Tensor: def latent_rows(self, layer_id: int) -> torch.Tensor: """Row-flat latent view ``[num_pages * page_size, latent_dim]``.""" - return self._kv_buffer[0, layer_id].view(-1, self._latent_dim) + return self._kv_buffer[0, self._dense(layer_id)].view(-1, self._latent_dim) # -- writes ----------------------------------------------------------------- def store_kv( @@ -107,11 +140,15 @@ def kv_cost(cls, config) -> tuple[int, int, int, int]: def rebuild_from_config( self, config, num_pages: int, *, num_swa_pages: int | None = None ) -> None: - self.rebuild(num_pages + 1) # +1 for the dummy page (matches create_kvcache_pool) + self.rebuild( + num_pages + 1 + ) # +1 for the dummy page (matches create_kvcache_pool) def unit_bytes(self) -> tuple[int, int]: buf = self._kv_buffer - return int(buf.numel() * buf.element_size()) // (self._num_pages * self._page_size), 0 + return int(buf.numel() * buf.element_size()) // ( + self._num_pages * self._page_size + ), 0 # -- pool properties ---------------------------------------------------------- @property @@ -141,10 +178,19 @@ def __init__( device: torch.device, index_head_dim: int, num_index_layers: int, + layer_ids: Sequence[int] | None = None, ) -> None: self._index_head_dim = index_head_dim self._num_index_layers = num_index_layers - super().__init__(latent_dim, num_layers, num_pages, page_size, dtype, device) + super().__init__( + latent_dim, + num_layers, + num_pages, + page_size, + dtype, + device, + layer_ids=layer_ids, + ) def _alloc(self, num_pages: int) -> None: # Both slabs in one allocation step: rebuild can never leave the pool with a @@ -180,4 +226,4 @@ def store_index_k(self, k: torch.Tensor, out_loc: torch.Tensor, slot: int) -> No self._index_k_buffer[slot][out_loc] = k -__all__ = ["MLAKVCache", "DSAKVCache"] +__all__ = ["DSAKVCache", "MLAKVCache"] diff --git a/python/freetoken/models/config.py b/python/freetoken/models/config.py index 229cce812..0fa6575e4 100644 --- a/python/freetoken/models/config.py +++ b/python/freetoken/models/config.py @@ -307,6 +307,10 @@ class ModelConfig: # DSA indexer geometry the model module needs. Opaque to model-agnostic engine code; # None for every other model. glm_dsa_args: Any | None = None + # GLM-5.3-Flash (glm5_next) payload: KDA recurrent geometry, mHC controls and the + # compressed DSA indexer fields. Kept separate from glm_dsa_args because GLM-5.3 is + # a hybrid KDA/MLA model and its checkpoint layout is not GLM-5.2-compatible. + glm5_args: Any | None = None # MiniMax-M3 (minimax_m3) payload (MiniMaxM3Args): the block-sparse indexer geometry # (index heads/dim, top-k blocks, init/local blocks, sparse layer set) plus the # swigluoai/dense-MLP scalars the model module needs. Opaque to model-agnostic engine diff --git a/python/freetoken/models/glm5_next/__init__.py b/python/freetoken/models/glm5_next/__init__.py new file mode 100644 index 000000000..ecf514dcc --- /dev/null +++ b/python/freetoken/models/glm5_next/__init__.py @@ -0,0 +1,16 @@ +from .config import Glm5NextArgs, parse_config +from .model import Glm5NextForCausalLM +from .weight import ( + iter_weights, + load_nvfp4_expert_sources, + load_nvfp4_expert_sources_parallel, +) + +__all__ = [ + "Glm5NextArgs", + "Glm5NextForCausalLM", + "iter_weights", + "load_nvfp4_expert_sources", + "load_nvfp4_expert_sources_parallel", + "parse_config", +] diff --git a/python/freetoken/models/glm5_next/attention.py b/python/freetoken/models/glm5_next/attention.py new file mode 100644 index 000000000..0b4fe332c --- /dev/null +++ b/python/freetoken/models/glm5_next/attention.py @@ -0,0 +1,129 @@ +"""NoPE MLA attention and the GLM-5.3 compressed-indexer parameter layout.""" + +from __future__ import annotations + +import torch + +from freetoken.core import get_global_ctx +from freetoken.layers import BaseOP, LinearReplicated, RMSNorm + + +class _LayerNorm(BaseOP): + def __init__(self, size: int, eps: float = 1e-6) -> None: + self.weight = torch.empty(size) + self.bias = torch.empty(size) + self.eps = eps + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return torch.nn.functional.layer_norm( + x, (x.shape[-1],), self.weight, self.bias, self.eps + ) + + +class Glm5NextIndexer(BaseOP): + """Indexer projections. The current DSA backend uses token keys for short-context + identity and compatibility testing; the pool-compression tensors are retained so + the checkpoint contract is exact while compressed selection is added upstream.""" + + def __init__(self, config) -> None: + args = config.glm5_args + self.n_heads = args.index_n_heads + self.head_dim = args.index_head_dim + self.wq_b = LinearReplicated( + args.q_lora_rank, self.n_heads * self.head_dim, False + ) + self.wk = LinearReplicated(config.hidden_size, self.head_dim, False) + self.k_norm = _LayerNorm(self.head_dim) + self.weights_proj = LinearReplicated(config.hidden_size, self.n_heads, False) + self.index_kpool_compress_ape = torch.empty(args.index_kpool, self.head_dim) + self.index_kpool_compress_gate = torch.empty(self.head_dim, config.hidden_size) + + def compute(self, x: torch.Tensor, q_resid: torch.Tensor): + t = x.shape[0] + q = self.wq_b.forward(q_resid).view(t, self.n_heads, self.head_dim) + k = self.k_norm.forward(self.wk.forward(x)) + w = self.weights_proj.forward(x).float() * self.n_heads**-0.5 + return q, k, w + + +class Glm5NextAttention(BaseOP): + def __init__(self, config, layer_id: int) -> None: + args = config.glm5_args + self.layer_id = layer_id + self.num_heads = config.num_qo_heads + self.qk_nope_head_dim = args.qk_nope_head_dim + self.qk_rope_head_dim = args.qk_rope_head_dim + self.qk_head_dim = args.qk_head_dim + self.v_head_dim = args.v_head_dim + self.kv_lora_rank = args.kv_lora_rank + self.q_a_proj = LinearReplicated(config.hidden_size, args.q_lora_rank, False) + self.q_a_layernorm = RMSNorm(args.q_lora_rank, eps=config.rms_norm_eps) + self.q_b_proj = LinearReplicated( + args.q_lora_rank, self.num_heads * self.qk_head_dim, False + ) + self.kv_a_proj_with_mqa = LinearReplicated( + config.hidden_size, self.kv_lora_rank + self.qk_rope_head_dim, False + ) + self.kv_a_layernorm = RMSNorm(self.kv_lora_rank, eps=config.rms_norm_eps) + self.kv_b_proj = LinearReplicated( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + False, + ) + self.o_proj = LinearReplicated( + self.num_heads * self.v_head_dim, config.hidden_size, False + ) + self.indexer = ( + Glm5NextIndexer(config) if args.indexer_types[layer_id] == "full" else None + ) + self._w_uk = None + self._w_uv = None + + def _kv_b(self): + if self._w_uk is None: + w = self.kv_b_proj.weight.view( + self.num_heads, + self.qk_nope_head_dim + self.v_head_dim, + self.kv_lora_rank, + ) + self._w_uk = w[:, : self.qk_nope_head_dim].contiguous() + self._w_uv = w[:, self.qk_nope_head_dim :].transpose(1, 2).contiguous() + return self._w_uk, self._w_uv + + def prepare_for_runtime(self) -> None: + self._kv_b() + self.kv_b_proj.weight = None + + def forward(self, x: torch.Tensor) -> torch.Tensor: + ctx = get_global_ctx() + t = x.shape[0] + w_uk, w_uv = self._kv_b() + q_resid = self.q_a_layernorm.forward(self.q_a_proj.forward(x)) + q = self.q_b_proj.forward(q_resid).view(t, self.num_heads, self.qk_head_dim) + q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + kv = self.kv_a_proj_with_mqa.forward(x) + c_kv, k_pe = kv.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + c_kv = self.kv_a_layernorm.forward(c_kv) + q_absorbed = torch.bmm(q_nope.transpose(0, 1).contiguous(), w_uk).transpose( + 0, 1 + ) + indexer_qkw = ( + self.indexer.compute(x, q_resid) + if self.indexer is not None + and getattr(ctx.attn_backend, "dsa_enabled", False) + else None + ) + o_latent = ctx.attn_backend.mla_forward( + q_absorbed.contiguous(), + q_pe.contiguous(), + c_kv.contiguous(), + k_pe.contiguous(), + self.layer_id, + ctx.batch, + indexer_qkw=indexer_qkw, + ) + o = torch.bmm(o_latent.transpose(0, 1).contiguous(), w_uv).transpose(0, 1) + return self.o_proj.forward(o.reshape(t, self.num_heads * self.v_head_dim)) + + +__all__ = ["Glm5NextAttention"] diff --git a/python/freetoken/models/glm5_next/config.py b/python/freetoken/models/glm5_next/config.py new file mode 100644 index 000000000..f9b263d48 --- /dev/null +++ b/python/freetoken/models/glm5_next/config.py @@ -0,0 +1,222 @@ +"""Engine-facing configuration for GLM-5.3-Flash (``glm5_next``). + +GLM-5.3 alternates three Kimi Delta Attention layers with one compressed DSA/MLA +layer and carries four manifold-constrained residual streams through every block. +This parser describes that hybrid layout for the runnable FreeToken implementation. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from fnmatch import fnmatch +from typing import Any + +from freetoken.models.config import ( + FullAttentionGroupConfig, + LinearGatedDeltaGroupConfig, + ModelConfig, + RotaryConfig, +) + + +@dataclass(frozen=True) +class Glm5NextArgs: + hc_mult: int + hc_eps: float + hc_sinkhorn_iters: int + linear_num_heads: int + linear_head_dim: int + linear_conv_kernel_dim: int + linear_lower_bound: float | None + kv_lora_rank: int + q_lora_rank: int + qk_head_dim: int + qk_nope_head_dim: int + qk_rope_head_dim: int + v_head_dim: int + index_n_heads: int + index_head_dim: int + index_topk: int + index_kpool: int + indexer_types: tuple[str, ...] + layer_types: tuple[str, ...] + mlp_layer_types: tuple[str, ...] + + +def _quant_get(hf_config: Any): + quant = getattr(hf_config, "quantization_config", None) + if quant is None: + return None + return ( + quant.get + if isinstance(quant, dict) + else (lambda k, d=None: getattr(quant, k, d)) + ) + + +def _ignored(patterns: list[str], module_name: str) -> bool: + return any(fnmatch(module_name, pattern) for pattern in patterns) + + +def _quant_modes(hf_config: Any) -> tuple[str, str, str, str, tuple[int, int] | None]: + """Return expert/attention/dense/head formats plus optional block geometry.""" + get = _quant_get(hf_config) + if get is None: + return "none", "none", "none", "none", None + + algo = str(get("quant_algo") or get("quant_method") or "").lower() + block = get("weight_block_size") + if algo == "fp8" and block: + block_size = tuple(int(x) for x in block) + if block_size != (128, 128): + raise ValueError(f"only 128x128 block-fp8 is supported, got {block_size}") + # The official checkpoint's modules_to_not_convert keeps attention, KDA, + # mHC, norms and the head resident in bf16; routed experts are block-fp8. + return "fp8_block", "none", "none", "none", block_size + + if "fp4" not in algo: + return algo or "none", "none", "none", "none", None + + ignore = list(get("ignore") or ()) + prefix = "model.language_model.layers.3" + + def mode(probe: str) -> str: + return "none" if _ignored(ignore, probe) else "nvfp4" + + return ( + mode(f"{prefix}.mlp.experts.0.gate_proj"), + mode(f"{prefix}.self_attn.q_proj"), + mode("model.language_model.layers.0.mlp.gate_proj"), + mode("lm_head"), + None, + ) + + +def parse_config(hf_config: Any) -> ModelConfig: + text = getattr(hf_config, "text_config", hf_config) + layer_types = tuple(str(x) for x in text.layer_types) + mlp_layer_types = tuple(str(x) for x in text.mlp_layer_types) + if len(layer_types) != int(text.num_hidden_layers): + raise ValueError("layer_types must contain one entry per decoder layer") + if len(mlp_layer_types) != int(text.num_hidden_layers): + raise ValueError("mlp_layer_types must contain one entry per decoder layer") + + linear_ids = tuple( + i for i, kind in enumerate(layer_types) if kind == "linear_attention" + ) + full_ids = tuple( + i for i, kind in enumerate(layer_types) if kind == "deepseek_sparse_attention" + ) + unknown = set(layer_types) - {"linear_attention", "deepseek_sparse_attention"} + if unknown: + raise ValueError( + f"unsupported GLM-5.3 attention layer types: {sorted(unknown)}" + ) + + linear = text.linear_attn_config + linear_get = ( + linear.get + if isinstance(linear, dict) + else lambda k, d=None: getattr(linear, k, d) + ) + qk_rope = int(getattr(text, "qk_rope_head_dim", 0) or 0) + qk_head = int(text.qk_head_dim) + rotary = RotaryConfig( + head_dim=qk_head, + rotary_dim=qk_rope, + max_position=int(text.max_position_embeddings), + # The release sets rope_theta=null because qk_rope_head_dim is zero. Keep a + # harmless finite default so the generic cache schema remains well-formed. + base=float(getattr(text, "rope_theta", None) or 10000.0), + scaling=None, + ) + indexer_types = tuple(str(x) for x in text.indexer_types) + full_indexers = sum(1 for i in full_ids if indexer_types[i] == "full") + + groups = ( + LinearGatedDeltaGroupConfig( + name="linear", + layer_ids=linear_ids, + num_key_heads=int(linear_get("num_heads")), + num_value_heads=int(linear_get("num_heads")), + key_head_dim=int(linear_get("head_dim")), + value_head_dim=int(linear_get("head_dim")), + conv_kernel_dim=int(linear_get("short_conv_kernel_size")), + output_gate="sigmoid", + ), + FullAttentionGroupConfig( + name="full", + layer_ids=full_ids, + num_kv_heads=1, + head_dim=int(text.kv_lora_rank) + qk_rope, + rotary_config=rotary, + mla=True, + index_head_dim=int(text.index_head_dim), + num_index_layers=full_indexers, + index_ratio=int(text.index_kpool), + ), + ) + expert_quant, attn_quant, dense_quant, lm_head_quant, block_size = _quant_modes( + hf_config + ) + args = Glm5NextArgs( + hc_mult=int(text.hc_mult), + hc_eps=float(text.hc_eps), + hc_sinkhorn_iters=int(text.hc_sinkhorn_iters), + linear_num_heads=int(linear_get("num_heads")), + linear_head_dim=int(linear_get("head_dim")), + linear_conv_kernel_dim=int(linear_get("short_conv_kernel_size")), + linear_lower_bound=linear_get("gate_lower_bound"), + kv_lora_rank=int(text.kv_lora_rank), + q_lora_rank=int(text.q_lora_rank), + qk_head_dim=qk_head, + qk_nope_head_dim=int(text.qk_nope_head_dim), + qk_rope_head_dim=qk_rope, + v_head_dim=int(text.v_head_dim), + index_n_heads=int(text.index_n_heads), + index_head_dim=int(text.index_head_dim), + index_topk=int(text.index_topk), + index_kpool=int(text.index_kpool), + indexer_types=indexer_types, + layer_types=layer_types, + mlp_layer_types=mlp_layer_types, + ) + return ModelConfig( + num_layers=int(text.num_hidden_layers), + num_qo_heads=int(text.num_attention_heads), + num_kv_heads=1, + head_dim=int(text.kv_lora_rank) + qk_rope, + hidden_size=int(text.hidden_size), + vocab_size=int(text.vocab_size), + intermediate_size=int(text.intermediate_size), + hidden_act=str(text.hidden_act), + rms_norm_eps=float(text.rms_norm_eps), + tie_word_embeddings=bool(getattr(text, "tie_word_embeddings", False)), + rotary_config=rotary, + attention_groups=groups, + num_experts=int(text.n_routed_experts), + num_experts_per_tok=int(text.num_experts_per_tok), + moe_intermediate_size=int(text.moe_intermediate_size), + norm_topk_prob=bool(text.norm_topk_prob), + model_type=str(getattr(hf_config, "model_type", "glm5_next")), + architectures=list( + getattr(hf_config, "architectures", ["Glm5NextForConditionalGeneration"]) + ), + moe_enabled=True, + first_k_dense_replace=int(text.first_k_dense_replace), + n_shared_experts=int(text.n_shared_experts), + routed_scaling_factor=float(text.routed_scaling_factor), + n_group=int(text.n_group), + topk_group=int(text.topk_group), + swiglu_limit=float(text.swiglu_limit), + attn_sm_scale=qk_head**-0.5, + expert_quant=expert_quant, + attn_quant=attn_quant, + dense_quant=dense_quant, + lm_head_quant=lm_head_quant, + weight_block_size=block_size, + glm5_args=args, + ) + + +__all__ = ["Glm5NextArgs", "parse_config"] diff --git a/python/freetoken/models/glm5_next/hc.py b/python/freetoken/models/glm5_next/hc.py new file mode 100644 index 000000000..20a9b1ba7 --- /dev/null +++ b/python/freetoken/models/glm5_next/hc.py @@ -0,0 +1,110 @@ +"""GLM-5.3 manifold-constrained Hyper-Connections. + +The released ``glm5_next`` equations match FreeToken's DeepSeek-V4 mHC kernels. +This module supplies the GLM parameter layout and deliberately reuses those fused +Sinkhorn/collapse/expand kernels instead of maintaining a second implementation. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from freetoken.layers import BaseOP + + +class HyperConnection(BaseOP): + """One attention or FFN mHC site over ``hc_mult`` residual streams.""" + + def __init__( + self, + hidden_size: int, + hc_mult: int = 4, + norm_eps: float = 1e-5, + sinkhorn_iters: int = 20, + hc_eps: float = 1e-6, + ) -> None: + self.hidden_size = hidden_size + self.hc_mult = hc_mult + self.norm_eps = norm_eps + self.sinkhorn_iters = sinkhorn_iters + self.hc_eps = hc_eps + mix_size = (2 + hc_mult) * hc_mult + self.fn = torch.empty(mix_size, hc_mult * hidden_size, dtype=torch.float32) + self.base = torch.empty(mix_size, dtype=torch.float32) + self.scale = torch.empty(3, dtype=torch.float32) + + def to(self, device) -> HyperConnection: + """Small compatibility helper for direct module tests; engine loading replaces + BaseOP tensors through its state-dict path.""" + self.fn = self.fn.to(device) + self.base = self.base.to(device) + self.scale = self.scale.to(device) + return self + + def mix(self, hidden_streams: torch.Tensor): + """Return collapsed sublayer input plus the placement/mixing coefficients.""" + if hidden_streams.shape[-2:] != (self.hc_mult, self.hidden_size): + raise ValueError( + "expected trailing hidden-stream shape " + f"({self.hc_mult}, {self.hidden_size}), got {tuple(hidden_streams.shape[-2:])}" + ) + shape = hidden_streams.shape + flat = hidden_streams.flatten(-2).float() + inv_rms = torch.rsqrt(flat.square().mean(-1, keepdim=True) + self.norm_eps) + mixes = F.linear(flat, self.fn) * inv_rms + pre_w, post_w, comb_w = mixes.split( + [self.hc_mult, self.hc_mult, self.hc_mult * self.hc_mult], dim=-1 + ) + pre_b, post_b, comb_b = self.base.split( + [self.hc_mult, self.hc_mult, self.hc_mult * self.hc_mult] + ) + pre = torch.sigmoid(pre_w * self.scale[0] + pre_b) + self.hc_eps + post = 2 * torch.sigmoid(post_w * self.scale[1] + post_b) + comb = torch.softmax( + comb_w.reshape(-1, self.hc_mult, self.hc_mult) * self.scale[2] + + comb_b.view(self.hc_mult, self.hc_mult), + dim=-1, + ) + comb = comb + self.hc_eps + comb = comb / (comb.sum(-2, keepdim=True) + self.hc_eps) + for _ in range(self.sinkhorn_iters - 1): + comb = comb / (comb.sum(-1, keepdim=True) + self.hc_eps) + comb = comb / (comb.sum(-2, keepdim=True) + self.hc_eps) + collapsed = ( + (pre.unsqueeze(-1) * hidden_streams.float()) + .sum(-2) + .to(hidden_streams.dtype) + ) + return ( + collapsed.reshape(*shape[:-2], self.hidden_size), + post.reshape(*shape[:-2], self.hc_mult), + comb.reshape(*shape[:-2], self.hc_mult, self.hc_mult), + ) + + def combine( + self, + residual: torch.Tensor, + block_output: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + ) -> torch.Tensor: + """Place a sublayer output back into and mix the residual streams.""" + shape = residual.shape + tokens = residual.numel() // (self.hc_mult * self.hidden_size) + block = block_output.reshape(tokens, self.hidden_size).float() + streams = residual.reshape(tokens, self.hc_mult, self.hidden_size).float() + post = post.reshape(tokens, self.hc_mult).float() + comb = comb.reshape(tokens, self.hc_mult, self.hc_mult).float() + mixed = post.unsqueeze(-1) * block.unsqueeze(-2) + mixed = mixed + torch.matmul(comb.transpose(-1, -2), streams) + mixed = mixed.to(residual.dtype) + return mixed.reshape(shape) + + +def collapse_head(hidden_streams: torch.Tensor) -> torch.Tensor: + """GLM-5.3's final stream collapse is an unweighted mean.""" + return hidden_streams.mean(dim=-2) + + +__all__ = ["HyperConnection", "collapse_head"] diff --git a/python/freetoken/models/glm5_next/kda.py b/python/freetoken/models/glm5_next/kda.py new file mode 100644 index 000000000..a7ec0504e --- /dev/null +++ b/python/freetoken/models/glm5_next/kda.py @@ -0,0 +1,49 @@ +"""GLM-5.3 Kimi Delta Attention recurrence adapters.""" + +from __future__ import annotations + +import torch + + +def kda_decode( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + gate_a: torch.Tensor, + beta_logits: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + state_source: torch.Tensor, + indices: torch.Tensor, + cu_seqlens: torch.Tensor | None = None, + gate_lower_bound: float | None = -5.0, +) -> torch.Tensor: + """Fused vector-decay KDA update with indexed in-place recurrent state. + + ``gate_a`` and ``dt_bias`` are flattened ``[..., heads * key_dim]`` tensors, + matching GLM's ``f_b_proj`` output and checkpoint parameter respectively. + """ + from freetoken.kernel.fla import fused_sigmoid_gating_delta_rule_update + + return fused_sigmoid_gating_delta_rule_update( + A_log=A_log, + a=gate_a, + dt_bias=dt_bias, + softplus_beta=1.0, + softplus_threshold=20.0, + q=q, + k=k, + v=v, + b=beta_logits, + initial_state_source=state_source, + initial_state_indices=indices, + scale=q.shape[-1] ** -0.5, + use_qk_l2norm_in_kernel=True, + cu_seqlens=cu_seqlens, + is_kda=True, + gate_lower_bound=gate_lower_bound, + ) + + +__all__ = ["kda_decode"] diff --git a/python/freetoken/models/glm5_next/linear_attention.py b/python/freetoken/models/glm5_next/linear_attention.py new file mode 100644 index 000000000..db4edf48f --- /dev/null +++ b/python/freetoken/models/glm5_next/linear_attention.py @@ -0,0 +1,141 @@ +"""Kimi Delta Attention used by GLM-5.3-Flash's linear layers.""" + +from __future__ import annotations + +import torch + +from freetoken.core import get_global_ctx +from freetoken.kernel.causal_conv1d import causal_conv1d_decode, causal_conv1d_varlen +from freetoken.layers import BaseOP, LinearReplicated + +from .kda import kda_decode + + +class _DepthwiseConv1d(BaseOP): + def __init__(self, dim: int, kernel: int) -> None: + self.weight = torch.empty(dim, 1, kernel) + + +class _SigmoidGatedRMSNorm(BaseOP): + def __init__(self, dim: int, eps: float) -> None: + self.weight = torch.empty(dim) + self.eps = eps + + def forward(self, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: + from freetoken.kernel.fla import rms_norm_gated + + return rms_norm_gated( + x=x, + weight=self.weight, + bias=None, + z=gate, + eps=self.eps, + is_rms_norm=True, + norm_before_gate=True, + activation="sigmoid", + ) + + +class Glm5NextLinearAttention(BaseOP): + """FreeToken stateful KDA op, preserving the released checkpoint key layout.""" + + def __init__(self, config, layer_id: int) -> None: + args = config.glm5_args + self.layer_id = layer_id + self.num_heads = args.linear_num_heads + self.head_dim = args.linear_head_dim + self.qkv_dim = self.num_heads * self.head_dim + self.conv_dim = 3 * self.qkv_dim + self.conv_kernel_size = args.linear_conv_kernel_dim + self.gate_lower_bound = args.linear_lower_bound + + self.q_proj = LinearReplicated(config.hidden_size, self.qkv_dim, has_bias=False) + self.k_proj = LinearReplicated(config.hidden_size, self.qkv_dim, has_bias=False) + self.v_proj = LinearReplicated(config.hidden_size, self.qkv_dim, has_bias=False) + self.conv1d = _DepthwiseConv1d(self.conv_dim, self.conv_kernel_size) + self.b_proj = LinearReplicated( + config.hidden_size, self.num_heads, has_bias=False + ) + self.f_a_proj = LinearReplicated( + config.hidden_size, self.head_dim, has_bias=False + ) + self.f_b_proj = LinearReplicated(self.head_dim, self.qkv_dim, has_bias=False) + self.g_a_proj = LinearReplicated( + config.hidden_size, self.head_dim, has_bias=False + ) + self.g_b_proj = LinearReplicated(self.head_dim, self.qkv_dim, has_bias=False) + self.A_log = torch.empty(self.num_heads, dtype=torch.float32) + self.dt_bias = torch.empty(self.qkv_dim, dtype=torch.float32) + self.o_norm = _SigmoidGatedRMSNorm(self.head_dim, config.rms_norm_eps) + self.o_proj = LinearReplicated(self.qkv_dim, config.hidden_size, has_bias=False) + + def _conv_weight(self) -> torch.Tensor: + return self.conv1d.weight.squeeze(1) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + ctx = get_global_ctx() + batch = ctx.batch + pool = ctx.linear_state_pool + fla = batch.fla_metadata + if fla is None: + from freetoken.attention.linear import build_fla_metadata + + fla = build_fla_metadata(batch, hidden_states.device) + batch.fla_metadata = fla + + raw = torch.cat( + [ + self.q_proj.forward(hidden_states), + self.k_proj.forward(hidden_states), + self.v_proj.forward(hidden_states), + ], + dim=-1, + ) + li = pool.local_index(self.layer_id) + if batch.is_decode: + mixed = causal_conv1d_decode( + raw, pool.conv_states[li], self._conv_weight(), fla.cache_indices + ) + else: + mixed = ( + causal_conv1d_varlen( + raw.transpose(0, 1).contiguous(), + self._conv_weight(), + pool.conv_states[li], + fla.cu_seqlens, + fla.cache_indices, + fla.has_initial_state, + ) + .transpose(0, 1) + .contiguous() + ) + + total = hidden_states.shape[0] + q, k, v = torch.split(mixed, [self.qkv_dim] * 3, dim=-1) + shape = (1, total, self.num_heads, self.head_dim) + q, k, v = q.view(shape), k.view(shape), v.view(shape) + gate_a = self.f_b_proj.forward(self.f_a_proj.forward(hidden_states)) + beta = self.b_proj.forward(hidden_states) + if fla.fresh_state_indices is not None: + pool.recurrent_states[li].index_fill_(0, fla.fresh_state_indices, 0.0) + core = kda_decode( + q, + k, + v, + gate_a, + beta, + A_log=self.A_log, + dt_bias=self.dt_bias, + state_source=pool.recurrent_states[li], + indices=fla.cache_indices, + cu_seqlens=fla.cu_seqlens, + gate_lower_bound=self.gate_lower_bound, + ) + gate = self.g_b_proj.forward(self.g_a_proj.forward(hidden_states)) + out = self.o_norm.forward( + core.reshape(-1, self.head_dim), gate.reshape(-1, self.head_dim) + ) + return self.o_proj.forward(out.reshape(total, self.qkv_dim)) + + +__all__ = ["Glm5NextLinearAttention"] diff --git a/python/freetoken/models/glm5_next/mlp.py b/python/freetoken/models/glm5_next/mlp.py new file mode 100644 index 000000000..8a6bf242d --- /dev/null +++ b/python/freetoken/models/glm5_next/mlp.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from freetoken.layers import BaseOP, LinearReplicated + + +class Glm5NextMLP(BaseOP): + def __init__( + self, hidden_size: int, intermediate_size: int, limit: float | None + ) -> None: + self.gate_proj = LinearReplicated( + hidden_size, intermediate_size, has_bias=False + ) + self.up_proj = LinearReplicated(hidden_size, intermediate_size, has_bias=False) + self.down_proj = LinearReplicated( + intermediate_size, hidden_size, has_bias=False + ) + self.limit = limit + + def forward(self, x: torch.Tensor) -> torch.Tensor: + gate = self.gate_proj.forward(x) + up = self.up_proj.forward(x) + if self.limit is not None: + gate = gate.clamp(max=self.limit) + up = up.clamp(min=-self.limit, max=self.limit) + return self.down_proj.forward(F.silu(gate) * up) + + +__all__ = ["Glm5NextMLP"] diff --git a/python/freetoken/models/glm5_next/model.py b/python/freetoken/models/glm5_next/model.py new file mode 100644 index 000000000..937d7c977 --- /dev/null +++ b/python/freetoken/models/glm5_next/model.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import torch + +from freetoken.core import get_global_ctx +from freetoken.layers import ( + BaseOP, + OPList, + ParallelLMHead, + RMSNorm, + VocabParallelEmbedding, +) +from freetoken.models.blocks import BaseLLMModel +from freetoken.utils import nvtx_annotate + +from .attention import Glm5NextAttention +from .hc import HyperConnection, collapse_head +from .linear_attention import Glm5NextLinearAttention +from .mlp import Glm5NextMLP +from .moe import Glm5NextSparseBlock + + +class Glm5NextDecoderLayer(BaseOP): + def __init__(self, config, layer_id: int) -> None: + args = config.glm5_args + self._layer_id = layer_id + self.self_attn = ( + Glm5NextLinearAttention(config, layer_id) + if config.is_linear_layer(layer_id) + else Glm5NextAttention(config, layer_id) + ) + self.mlp = ( + Glm5NextMLP( + config.hidden_size, config.intermediate_size, config.swiglu_limit + ) + if layer_id < config.first_k_dense_replace + else Glm5NextSparseBlock(config, layer_id) + ) + hc_kw = { + "hidden_size": config.hidden_size, + "hc_mult": args.hc_mult, + "norm_eps": config.rms_norm_eps, + "sinkhorn_iters": args.hc_sinkhorn_iters, + "hc_eps": args.hc_eps, + } + self.attn_hc = HyperConnection(**hc_kw) + self.ffn_hc = HyperConnection(**hc_kw) + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + @nvtx_annotate("Layer_{}", layer_id_field="_layer_id") + def forward(self, hidden: torch.Tensor) -> torch.Tensor: + x, post, comb = self.attn_hc.mix(hidden) + y = self.self_attn.forward(self.input_layernorm.forward(x)) + hidden = self.attn_hc.combine(hidden, y, post, comb) + x, post, comb = self.ffn_hc.mix(hidden) + y = self.mlp.forward(self.post_attention_layernorm.forward(x)) + return self.ffn_hc.combine(hidden, y, post, comb) + + +class Glm5NextModel(BaseOP): + def __init__(self, config) -> None: + self.hc_mult = config.glm5_args.hc_mult + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, config.hidden_size + ) + self.layers = OPList( + [Glm5NextDecoderLayer(config, i) for i in range(config.num_layers)] + ) + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + + def forward(self, input_ids: torch.Tensor) -> torch.Tensor: + hidden = ( + self.embed_tokens.forward(input_ids) + .unsqueeze(-2) + .expand(-1, self.hc_mult, -1) + .contiguous() + ) + for layer in self.layers.op_list: + hidden = layer.forward(hidden) + return self.norm.forward(collapse_head(hidden)) + + +class Glm5NextForCausalLM(BaseLLMModel): + def __init__(self, config) -> None: + self.model = Glm5NextModel(config) + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + tie_word_embeddings=config.tie_word_embeddings, + tied_embedding=self.model.embed_tokens + if config.tie_word_embeddings + else None, + ) + super().__init__() + + def prepare_for_runtime(self) -> None: + for layer in self.model.layers.op_list: + if isinstance(layer.self_attn, Glm5NextAttention): + layer.self_attn.prepare_for_runtime() + torch.cuda.empty_cache() + + def forward(self) -> torch.Tensor: + return self.lm_head.forward( + self.model.forward(get_global_ctx().batch.input_ids) + ) + + +__all__ = ["Glm5NextForCausalLM"] diff --git a/python/freetoken/models/glm5_next/moe.py b/python/freetoken/models/glm5_next/moe.py new file mode 100644 index 000000000..f447ffcd9 --- /dev/null +++ b/python/freetoken/models/glm5_next/moe.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from freetoken.layers import BaseOP, LinearReplicated, make_moe_layer + +from .mlp import Glm5NextMLP + + +class Glm5NextSparseBlock(BaseOP): + def __init__(self, config, layer_id: int) -> None: + self.top_k = config.num_experts_per_tok + self.num_experts = config.num_experts + self.norm_topk_prob = config.norm_topk_prob + self.routed_scaling_factor = config.routed_scaling_factor + self.n_group = config.n_group + self.topk_group = config.topk_group + self.gate = LinearReplicated( + config.hidden_size, config.num_experts, has_bias=False + ) + self.e_score_correction_bias = torch.empty(config.num_experts) + self.experts = make_moe_layer( + config, + layer_id=layer_id - config.first_k_dense_replace, + renormalize=config.norm_topk_prob, + extra_attrs={"swiglu_limit": config.swiglu_limit}, + ) + self.shared_experts = Glm5NextMLP( + config.hidden_size, + config.moe_intermediate_size * max(1, config.n_shared_experts), + config.swiglu_limit, + ) + + def _route(self, x: torch.Tensor): + scores = F.linear(x.float(), self.gate.weight.float()).sigmoid() + choice = scores + self.e_score_correction_bias.float() + if self.n_group > 1: + m, e, g = x.shape[0], self.num_experts, self.n_group + group_scores = choice.view(m, g, e // g).topk(2, dim=-1)[0].sum(-1) + group_idx = group_scores.topk(self.topk_group, dim=-1, sorted=False)[1] + mask = torch.zeros_like(group_scores).scatter_(1, group_idx, 1).bool() + choice = choice.masked_fill( + ~mask.unsqueeze(-1).expand(m, g, e // g).reshape(m, e), float("-inf") + ) + ids = choice.topk(self.top_k, dim=-1)[1] + weights = scores.gather(-1, ids) + if self.norm_topk_prob: + weights = weights / (weights.sum(-1, keepdim=True) + 1e-20) + return ( + weights * self.routed_scaling_factor + ).float().contiguous(), ids.int().contiguous() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + shape = x.shape + x = x.view(-1, shape[-1]) + weights, ids = self._route(x) + routed = self.experts.routed_forward(x, weights, ids) + return (routed + self.shared_experts.forward(x)).view(shape) + + +__all__ = ["Glm5NextSparseBlock"] diff --git a/python/freetoken/models/glm5_next/weight.py b/python/freetoken/models/glm5_next/weight.py new file mode 100644 index 000000000..d878d19a0 --- /dev/null +++ b/python/freetoken/models/glm5_next/weight.py @@ -0,0 +1,183 @@ +"""GLM-5.3-Flash ModelOpt checkpoint reader.""" + +from __future__ import annotations + +import re +from collections.abc import Iterator + +import safetensors +import torch +from tqdm import tqdm + +from freetoken.distributed import get_tp_info +from freetoken.models.loader import drop_page_cache, iter_weight_files +from freetoken.models.nvfp4_banks import ( + Nvfp4ExpertSourceSpec, + load_nvfp4_expert_source_banks, + load_nvfp4_expert_source_banks_parallel, +) + +_EXPERT_RE = re.compile(r"\.mlp\.experts\.\d+\.") +_BASE_LAYER_RE = re.compile(r"^(?:model\.)?language_model\.layers\.(?P\d+)\.") +# The released checkpoint appends one next-token-prediction block as language +# layer 45 instead of placing it below ``mtp.*``. The serving graph contains +# decoder layers 0..44 only. +_NUM_BASE_LAYERS = 45 +_EXPERT_KEY_RE = re.compile( + r"^model\.language_model\.layers\.(?P\d+)\.mlp\.experts\.(?P\d+)\." + r"(?Pgate_proj|up_proj|down_proj)\.(?Pweight|weight_scale|weight_scale_2)$" +) +_SOURCE_SPEC = Nvfp4ExpertSourceSpec( + key_pattern=_EXPERT_KEY_RE, + proj_to_role={"gate_proj": "gate", "up_proj": "up", "down_proj": "down"}, + layer_to_bank=lambda layer, config: ( + layer - config.first_k_dense_replace + if config.first_k_dense_replace <= layer < config.num_layers + else None + ), + desc="GLM-5.3-Flash NVFP4 experts", +) + +# The release stores one depthwise short-convolution kernel beside each of the +# q/k/v projections. The runtime executes a single concatenated q|k|v +# convolution, so combine those three small tensors while streaming the shards. +_KDA_CONV_PARTS = ( + ".self_attn.q_conv1d.weight", + ".self_attn.k_conv1d.weight", + ".self_attn.v_conv1d.weight", +) +_KDA_CONV_FUSED = ".self_attn.conv1d.weight" + + +def _rename(raw: str) -> str | None: + layer_match = _BASE_LAYER_RE.match(raw) + if ( + raw.startswith(("model.visual.", "visual.", "mtp.")) + or _EXPERT_RE.search(raw) + or ( + layer_match is not None + and int(layer_match.group("layer")) >= _NUM_BASE_LAYERS + ) + ): + return None + if raw.endswith((".weight_scale", ".weight_scale_2", ".input_scale")): + return None + if raw.startswith("model.language_model."): + name = "model." + raw[len("model.language_model.") :] + elif raw.startswith("language_model."): + name = "model." + raw[len("language_model.") :] + else: + name = raw + for checkpoint, runtime in ( + (".hc_attn_fn", ".attn_hc.fn"), + (".hc_attn_base", ".attn_hc.base"), + (".hc_attn_scale", ".attn_hc.scale"), + (".hc_ffn_fn", ".ffn_hc.fn"), + (".hc_ffn_base", ".ffn_hc.base"), + (".hc_ffn_scale", ".ffn_hc.scale"), + (".mlp.gate.e_score_correction_bias", ".mlp.e_score_correction_bias"), + ): + if checkpoint in name: + name = name.replace(checkpoint, runtime) + return name + + +def _try_fuse_kda_conv( + name: str, + tensor: torch.Tensor, + buf: dict[str, dict[int, torch.Tensor]], +) -> tuple[str, torch.Tensor] | tuple[()] | None: + """Fuse checkpoint q/k/v depthwise kernels in runtime channel order. + + Returns ``None`` for an unrelated tensor, ``()`` while a layer is still + incomplete, and ``(runtime_name, concatenated_tensor)`` after all three + pieces have arrived. The buffer intentionally lives across shard files. + """ + for idx, suffix in enumerate(_KDA_CONV_PARTS): + if not name.endswith(suffix): + continue + key = name[: -len(suffix)] + _KDA_CONV_FUSED + slots = buf.setdefault(key, {}) + if idx in slots: + raise ValueError(f"duplicate GLM-5.3 KDA convolution part: {name}") + slots[idx] = tensor + if len(slots) < len(_KDA_CONV_PARTS): + return () + del buf[key] + return key, torch.cat([slots[i] for i in range(len(_KDA_CONV_PARTS))], dim=0) + return None + + +def iter_weights( + model_path: str, + device: torch.device, + *, + include_moe_experts: bool, + include_non_moe: bool, +) -> Iterator[tuple[str, torch.Tensor]]: + if get_tp_info().size > 1: + raise NotImplementedError("GLM-5.3 weight loading supports TP=1 only") + if include_moe_experts: + raise ValueError("GLM-5.3 NVFP4 experts must use the offload source banks") + if not include_non_moe: + return + conv_buf: dict[str, dict[int, torch.Tensor]] = {} + for file in tqdm( + iter_weight_files(model_path), desc="Loading GLM-5.3 dense weights" + ): + with safetensors.safe_open(file, framework="pt", device=str(device)) as f: + # ``safe_open`` exposes keys but is not itself iterable. + for raw in f.keys(): # noqa: SIM118 + name = _rename(raw) + if name is not None: + tensor = f.get_tensor(raw) + fused = _try_fuse_kda_conv(name, tensor, conv_buf) + if fused is not None: + if fused != (): + yield fused + continue + yield name, tensor + drop_page_cache(file) + if conv_buf: + missing = { + key: [ + _KDA_CONV_PARTS[i] + for i in range(len(_KDA_CONV_PARTS)) + if i not in parts + ] + for key, parts in conv_buf.items() + } + raise ValueError(f"incomplete GLM-5.3 KDA convolution fusions: {missing}") + + +def load_nvfp4_expert_sources(model_path: str, config, *, layer_sink=None): + return load_nvfp4_expert_source_banks( + model_path, + config, + _SOURCE_SPEC, + drop_page_cache=drop_page_cache, + primary=get_tp_info().is_primary(), + layer_sink=layer_sink, + ) + + +def load_nvfp4_expert_sources_parallel( + model_path: str, config, *, workers: int = 8, chunk: int = 8 << 20, layer_sink=None +): + return load_nvfp4_expert_source_banks_parallel( + model_path, + config, + _SOURCE_SPEC, + drop_page_cache=drop_page_cache, + primary=get_tp_info().is_primary(), + workers=workers, + chunk=chunk, + layer_sink=layer_sink, + ) + + +__all__ = [ + "iter_weights", + "load_nvfp4_expert_sources", + "load_nvfp4_expert_sources_parallel", +] diff --git a/python/freetoken/models/nvfp4_banks.py b/python/freetoken/models/nvfp4_banks.py index 0e3ab6a51..f4d486948 100644 --- a/python/freetoken/models/nvfp4_banks.py +++ b/python/freetoken/models/nvfp4_banks.py @@ -4,18 +4,31 @@ import json import os import re +from collections.abc import Callable +from contextlib import contextmanager from dataclasses import dataclass -from typing import Callable import safetensors import torch -from freetoken.utils import download_hf_weight from tqdm import tqdm +from freetoken.utils import download_hf_weight + LayerToBank = Callable[[int, object], int | None] DropPageCache = Callable[[str], None] +@contextmanager +def _single_threaded_torch_copies(): + """Avoid intra-op fanout for the loader's many small host tensor copies.""" + previous = torch.get_num_threads() + try: + torch.set_num_threads(1) + yield + finally: + torch.set_num_threads(previous) + + @dataclass(frozen=True) class Nvfp4ExpertSourceSpec: key_pattern: re.Pattern[str] @@ -52,14 +65,17 @@ def _alloc_nvfp4_host_banks(num_layers: int, E: int, H: int, I: int): from freetoken.moe.host_banks import alloc_layer_banks fp8 = torch.float8_e4m3fn - return alloc_layer_banks({ - "gate_up_packed": ((E, 2 * I, H // 2), torch.uint8), - "gate_up_scale": ((E, 2 * I, H // 16), fp8), - "gate_up_global": ((E, 2 * I), torch.float16), - "down_packed": ((E, H, I // 2), torch.uint8), - "down_scale": ((E, H, I // 16), fp8), - "down_global": ((E, H), torch.float16), - }, num_layers) + return alloc_layer_banks( + { + "gate_up_packed": ((E, 2 * I, H // 2), torch.uint8), + "gate_up_scale": ((E, 2 * I, H // 16), fp8), + "gate_up_global": ((E, 2 * I), torch.float16), + "down_packed": ((E, H, I // 2), torch.uint8), + "down_scale": ((E, H, I // 16), fp8), + "down_global": ((E, H), torch.float16), + }, + num_layers, + ) def load_nvfp4_expert_source_banks( @@ -88,8 +104,20 @@ def load_nvfp4_expert_source_banks( """ folder = download_hf_weight(model_path) index_path = os.path.join(folder, "model.safetensors.index.json") - with open(index_path, encoding="utf-8") as f: - weight_map = json.load(f)["weight_map"] + if os.path.exists(index_path): + with open(index_path, encoding="utf-8") as f: + weight_map = json.load(f)["weight_map"] + else: + # Tiny/reference checkpoints are commonly emitted as one safetensors file. + # Build the same name->shard map the indexed path supplies. + weight_map = {} + for shard_path in sorted( + os.path.join(folder, name) + for name in os.listdir(folder) + if name.endswith(".safetensors") + ): + with safetensors.safe_open(shard_path, framework="pt", device="cpu") as f: + weight_map.update({name: os.path.basename(shard_path) for name in f}) E = config.num_experts H = config.hidden_size @@ -99,8 +127,12 @@ def load_nvfp4_expert_source_banks( for shard in sorted(set(weight_map.values())): drop_page_cache(os.path.join(folder, shard)) - weight_shards: dict[str, list[tuple[str, re.Match[str], int]]] = collections.defaultdict(list) - global_shards: dict[str, list[tuple[str, re.Match[str], int]]] = collections.defaultdict(list) + weight_shards: dict[str, list[tuple[str, re.Match[str], int]]] = ( + collections.defaultdict(list) + ) + global_shards: dict[str, list[tuple[str, re.Match[str], int]]] = ( + collections.defaultdict(list) + ) for name, shard in weight_map.items(): match = spec.key_pattern.match(name) if match is None: @@ -146,7 +178,9 @@ def load_nvfp4_expert_source_banks( def _load(sink) -> int: tracker = LayerCompletionTracker(E * 6, _hb, sink) placed = 0 - for shard in tqdm(sorted(weight_shards), desc=f"Loading {spec.desc}", disable=not primary): + for shard in tqdm( + sorted(weight_shards), desc=f"Loading {spec.desc}", disable=not primary + ): path = os.path.join(folder, shard) with safetensors.safe_open(path, framework="pt", device="cpu") as f: for name, match, bank_layer_id in weight_shards[shard]: @@ -164,7 +198,9 @@ def _load(sink) -> int: elif role == "down": down_packed[bank_layer_id][expert] = tensor else: - raise ValueError(f"{spec.desc}: unknown projection role {role!r}") + raise ValueError( + f"{spec.desc}: unknown projection role {role!r}" + ) else: global_scale = globals_map[(layer, expert, proj)] if role == "gate": @@ -177,20 +213,29 @@ def _load(sink) -> int: down_scale[bank_layer_id][expert] = tensor down_global[bank_layer_id][expert] = global_scale else: - raise ValueError(f"{spec.desc}: unknown projection role {role!r}") + raise ValueError( + f"{spec.desc}: unknown projection role {role!r}" + ) tracker.note(bank_layer_id) placed += 1 drop_page_cache(path) return placed - if layer_sink is not None: - placed = _load(layer_sink) - else: - with PinPipeline() as pins: - placed = _load(pins) + # These are tens of thousands of small, disjoint host copies. Letting PyTorch fan + # each assignment across a large intra-op pool is dramatically slower on high-core + # hosts (and competes with the O_DIRECT reader workers). Keep placement single- + # threaded, then restore the process setting before the runtime is constructed. + with _single_threaded_torch_copies(): + if layer_sink is not None: + placed = _load(layer_sink) + else: + with PinPipeline() as pins: + placed = _load(pins) expected = num_layers * E * 6 - assert placed == expected, f"{spec.desc}: loaded {placed} expert tensors, expected {expected}" + assert placed == expected, ( + f"{spec.desc}: loaded {placed} expert tensors, expected {expected}" + ) return { "gate_up_packed": gate_up_packed, "gate_up_scale": gate_up_scale, @@ -219,7 +264,9 @@ def load_nvfp4_expert_source_banks_parallel( from freetoken.models.weight import iter_expert_tensors_parallel folder = download_hf_weight(model_path) - with open(os.path.join(folder, "model.safetensors.index.json"), encoding="utf-8") as f: + with open( + os.path.join(folder, "model.safetensors.index.json"), encoding="utf-8" + ) as f: weight_map = json.load(f)["weight_map"] E = config.num_experts @@ -227,7 +274,9 @@ def load_nvfp4_expert_source_banks_parallel( I = config.moe_intermediate_size num_layers = _num_moe_layers(config) - weight_info: dict[str, tuple[re.Match[str], int]] = {} # name -> (match, bank_layer) + weight_info: dict[ + str, tuple[re.Match[str], int] + ] = {} # name -> (match, bank_layer) global_names_by_shard: dict[str, list[str]] = collections.defaultdict(list) for name, shard in weight_map.items(): match = spec.key_pattern.match(name) @@ -252,9 +301,9 @@ def load_nvfp4_expert_source_banks_parallel( with safetensors.safe_open(path, framework="pt", device="cpu") as f: for name in global_names_by_shard[shard]: m = spec.key_pattern.match(name) - globals_map[(int(m.group("layer")), int(m.group("expert")), m.group("proj"))] = ( - f.get_tensor(name).to(torch.float16) - ) + globals_map[ + (int(m.group("layer")), int(m.group("expert")), m.group("proj")) + ] = f.get_tensor(name).to(torch.float16) drop_page_cache(path) _hb = _alloc_nvfp4_host_banks(num_layers, E, H, I) # unpinned; pinned after fill @@ -302,14 +351,19 @@ def _load(sink) -> int: placed += 1 return placed - if layer_sink is not None: - placed = _load(layer_sink) - else: - with PinPipeline() as pins: - placed = _load(pins) + # See the serial loader above: intra-op fanout makes these small host copies + # dramatically slower on high-core machines and competes with reader workers. + with _single_threaded_torch_copies(): + if layer_sink is not None: + placed = _load(layer_sink) + else: + with PinPipeline() as pins: + placed = _load(pins) expected = num_layers * E * 6 - assert placed == expected, f"{spec.desc}: loaded {placed} expert tensors, expected {expected}" + assert placed == expected, ( + f"{spec.desc}: loaded {placed} expert tensors, expected {expected}" + ) return { "gate_up_packed": gate_up_packed, "gate_up_scale": gate_up_scale, diff --git a/python/freetoken/models/register.py b/python/freetoken/models/register.py index b94d8291b..248d5e5e5 100644 --- a/python/freetoken/models/register.py +++ b/python/freetoken/models/register.py @@ -130,6 +130,10 @@ class ModelSpec: "freetoken.models.glm_moe_dsa", "GlmMoeDsaForCausalLM", ), + "Glm5NextForConditionalGeneration": ModelSpec( + "freetoken.models.glm5_next", + "Glm5NextForCausalLM", + ), } @@ -137,7 +141,9 @@ def get_model_spec(model_architecture: str) -> ModelSpec: try: return _MODEL_REGISTRY[model_architecture] except KeyError as exc: - raise ValueError(f"Model architecture {model_architecture} not supported") from exc + raise ValueError( + f"Model architecture {model_architecture} not supported" + ) from exc def _load_attr(module_path: str, attr_name: str) -> Any: @@ -151,4 +157,4 @@ def get_model_class(model_architecture: str, model_config: ModelConfig): return model_cls(model_config) -__all__ = ["ModelSpec", "get_model_spec", "get_model_class"] +__all__ = ["ModelSpec", "get_model_class", "get_model_spec"] diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 000000000..ede574799 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +# Vast's pinned PyWorker bootstrap installs the Vast SDK separately. diff --git a/scripts/vast_deepseek_v4_provision.sh b/scripts/vast_deepseek_v4_provision.sh new file mode 100755 index 000000000..2392a80f2 --- /dev/null +++ b/scripts/vast_deepseek_v4_provision.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# DeepSeek-V4-Flash uses the same cacheable Vast/PyWorker bootstrap as GLM. +# Keep the model-specific contract in this public wrapper so Vast workers can +# cold-start without credentials for a separate deployment repository. +script_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +shared_provisioner="$script_dir/vast_glm53_provision.sh" +freetoken_ref="${TEKIZAI_FREETOKEN_REF:-feat/glm53-flash}" + +export TEKIZAI_MODEL_SOURCE_PATH="${TEKIZAI_MODEL_SOURCE_PATH:-/workspace/models/DeepSeek-V4-Flash-0731-hf}" +export TEKIZAI_FREETOKEN_MODEL_PATH="${TEKIZAI_FREETOKEN_MODEL_PATH:-/workspace/models/DeepSeek-V4-Flash-0731-ftw}" +export TEKIZAI_CONVERT_FTW="${TEKIZAI_CONVERT_FTW:-1}" +export TEKIZAI_MODEL_REPO="${TEKIZAI_MODEL_REPO:-deepseek-ai/DeepSeek-V4-Flash-0731}" +export TEKIZAI_MODEL_BENCH_DTYPE="${TEKIZAI_MODEL_BENCH_DTYPE:-ds_fp4}" +export TEKIZAI_SERVED_MODEL="${TEKIZAI_SERVED_MODEL:-deepseek-v4-flash}" +export TEKIZAI_FREETOKEN_LOG="${TEKIZAI_FREETOKEN_LOG:-/workspace/logs/freetoken-deepseek-v4-flash.log}" + +if [[ ! -x "$shared_provisioner" ]]; then + curl -fsSL \ + "https://raw.githubusercontent.com/earlvanze/FreeToken/${freetoken_ref}/scripts/vast_glm53_provision.sh" \ + -o "$shared_provisioner" + chmod 0755 "$shared_provisioner" +fi + +exec "$shared_provisioner" diff --git a/scripts/vast_glm53_persistent_provision.sh b/scripts/vast_glm53_persistent_provision.sh new file mode 100755 index 000000000..ab8c1aa14 --- /dev/null +++ b/scripts/vast_glm53_persistent_provision.sh @@ -0,0 +1,90 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +workspace="${WORKSPACE_DIR:-/workspace}" +freetoken_dir="${TEKIZAI_FREETOKEN_DIR:-${workspace}/freetoken}" +model_dir="${TEKIZAI_FREETOKEN_MODEL_PATH:-${workspace}/models/GLM-5.3-Flash-NVFP4}" +repo="${TEKIZAI_FREETOKEN_REPO:-https://github.com/earlvanze/FreeToken.git}" +ref="${TEKIZAI_FREETOKEN_REF:-feat/glm53-flash}" +expected_commit="${TEKIZAI_FREETOKEN_EXPECTED_COMMIT:-}" +model_repo="${TEKIZAI_MODEL_REPO:-LibertAIDAI/GLM-5.3-Flash-NVFP4}" + +export DEBIAN_FRONTEND=noninteractive +export PATH="${HOME}/.local/bin:${PATH}" +export HF_XET_HIGH_PERFORMANCE="${HF_XET_HIGH_PERFORMANCE:-1}" + +mkdir -p "$workspace" "${workspace}/logs" + +if ! command -v uv >/dev/null 2>&1; then + curl -LsSf https://astral.sh/uv/install.sh | sh + export PATH="${HOME}/.local/bin:${PATH}" +fi + +download_model() { + local attempt status + for attempt in 1 2 3 4 5; do + echo "FREETOKEN_PROVISION_STAGE=model_download attempt=${attempt}" + if uvx --from huggingface-hub==1.29.0 hf download "$model_repo" --local-dir "$model_dir"; then + echo "FREETOKEN_PROVISION_STAGE=model_download_complete" + return 0 + else + status=$? + fi + echo "FREETOKEN_PROVISION_STAGE=model_download_retry status=${status}" + sleep "$((attempt * 10))" + done + return "$status" +} + +model_download_pid="" +cleanup() { + if [[ -n "$model_download_pid" ]] && kill -0 "$model_download_pid" 2>/dev/null; then + kill "$model_download_pid" + wait "$model_download_pid" || true + fi +} +trap cleanup EXIT INT TERM + +download_model & +model_download_pid=$! + +apt-get update +apt-get install -y --no-install-recommends \ + build-essential ca-certificates cuda-compiler-13-0 cuda-cudart-dev-13-0 \ + libcurand-dev-13-0 \ + curl git ninja-build numactl python3.12-dev util-linux +apt-get clean + +export CUDA_HOME="${CUDA_HOME:-/usr/local/cuda-13.0}" +export PATH="${CUDA_HOME}/bin:${PATH}" +export LD_LIBRARY_PATH="${CUDA_HOME}/lib64:${LD_LIBRARY_PATH:-}" + +if [[ ! -d "$freetoken_dir/.git" ]]; then + mkdir -p "$freetoken_dir" + git -C "$freetoken_dir" init + git -C "$freetoken_dir" remote add origin "$repo" +fi +git -C "$freetoken_dir" fetch --depth 1 origin "$ref" +git -C "$freetoken_dir" checkout --detach --force FETCH_HEAD +if [[ -n "$expected_commit" ]]; then + resolved_commit="$(git -C "$freetoken_dir" rev-parse HEAD)" + if [[ "$resolved_commit" != "$expected_commit" && "$resolved_commit" != "$expected_commit"* ]]; then + echo "FreeToken ref resolved to unexpected commit: ${resolved_commit}" >&2 + exit 1 + fi +fi + +echo "FREETOKEN_PROVISION_STAGE=dependencies" +if [[ ! -x "$freetoken_dir/.venv/bin/python" ]]; then + uv venv --python 3.12 "$freetoken_dir/.venv" +fi +uv pip install --python "$freetoken_dir/.venv/bin/python" -e "$freetoken_dir[accel]" + +wait "$model_download_pid" +model_download_pid="" + +echo "FREETOKEN_PROVISION_STAGE=bandwidth_check" +"$freetoken_dir/.venv/bin/ft" bench bw --dtype nvfp4 + +echo "FREETOKEN_PROVISION_STAGE=model_start" +exec "$freetoken_dir/scripts/vast_glm53_start.sh" diff --git a/scripts/vast_glm53_provision.sh b/scripts/vast_glm53_provision.sh new file mode 100755 index 000000000..910fe2d2a --- /dev/null +++ b/scripts/vast_glm53_provision.sh @@ -0,0 +1,180 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +workspace="${WORKSPACE_DIR:-/workspace}" +freetoken_dir="${TEKIZAI_FREETOKEN_DIR:-${workspace}/freetoken}" +model_dir="${TEKIZAI_FREETOKEN_MODEL_PATH:-${workspace}/models/GLM-5.3-Flash-NVFP4}" +model_source_dir="${TEKIZAI_MODEL_SOURCE_PATH:-$model_dir}" +convert_ftw="${TEKIZAI_CONVERT_FTW:-0}" +worker_source_dir="${TEKIZAI_WORKER_SOURCE_DIR:-${workspace}/vast-pyworker}" +pyworker_uv_cache="${TEKIZAI_PYWORKER_UV_CACHE:-${workspace}/pyworker-uv-cache}" +repo="${TEKIZAI_FREETOKEN_REPO:-https://github.com/earlvanze/FreeToken.git}" +ref="${TEKIZAI_FREETOKEN_REF:-feat/glm53-flash}" +expected_commit="${TEKIZAI_FREETOKEN_EXPECTED_COMMIT:-}" +model_repo="${TEKIZAI_MODEL_REPO:-LibertAIDAI/GLM-5.3-Flash-NVFP4}" +bootstrap_ref="${TEKIZAI_PYWORKER_BOOTSTRAP_REF:-2207a3f94b55a0921c1641520eeb83de5a0c1611}" +bootstrap="${workspace}/vast-pyworker-bootstrap.sh" +provision_marker="${TEKIZAI_PROVISION_MARKER:-${workspace}/.tekizai-glm53-provisioned}" +provision_marker_value="${expected_commit:-$ref}" +model_bench_dtype="${TEKIZAI_MODEL_BENCH_DTYPE:-nvfp4}" + +export DEBIAN_FRONTEND=noninteractive +export PATH="${HOME}/.local/bin:${PATH}" +export MODEL_LOG="${TEKIZAI_FREETOKEN_LOG:-${workspace}/logs/freetoken-glm53.log}" +export PYWORKER_REPO="${PYWORKER_REPO:-$repo}" +export PYWORKER_REF="${PYWORKER_REF:-$ref}" +export HF_XET_HIGH_PERFORMANCE="${HF_XET_HIGH_PERFORMANCE:-1}" + +mkdir -p "$workspace" "$(dirname "$MODEL_LOG")" + +model_download_pid="" +pyworker_pid="" +model_launcher_pid="" +cleanup() { + local pid + for pid in "$model_launcher_pid" "$pyworker_pid" "$model_download_pid"; do + if [[ -n "$pid" ]] && kill -0 "$pid" 2>/dev/null; then + kill "$pid" + wait "$pid" || true + fi + done +} +trap cleanup EXIT INT TERM + +start_runtime() { + echo "FREETOKEN_PROVISION_STAGE=pyworker_start" + UV_CACHE_DIR="$pyworker_uv_cache" \ + USE_SYSTEM_PYTHON=true \ + ROTATE_MODEL_LOG=true \ + "$bootstrap" & + pyworker_pid=$! + + echo "FREETOKEN_PROVISION_STAGE=model_start" + "$freetoken_dir/scripts/vast_glm53_start.sh" & + model_launcher_pid=$! + + wait "$pyworker_pid" +} + +checkout_matches_expected_commit() { + local checkout + [[ -n "$expected_commit" ]] || return 0 + for checkout in "$freetoken_dir" "$worker_source_dir"; do + [[ -d "$checkout/.git" ]] || return 1 + [[ "$(git -C "$checkout" rev-parse HEAD 2>/dev/null)" == "$expected_commit" ]] || return 1 + done +} + +if [[ -f "$provision_marker" ]] \ + && [[ "$(<"$provision_marker")" == "$provision_marker_value" ]] \ + && [[ -x "$bootstrap" ]] \ + && [[ -x "$freetoken_dir/.venv/bin/ft" ]] \ + && [[ -x "$freetoken_dir/scripts/vast_glm53_start.sh" ]] \ + && [[ -s "$model_dir/config.json" ]] \ + && [[ -x /usr/local/cuda-13.0/bin/nvcc ]] \ + && checkout_matches_expected_commit; then + echo "FREETOKEN_PROVISION_STAGE=fast_resume" + export CUDA_HOME="${CUDA_HOME:-/usr/local/cuda-13.0}" + export PATH="${CUDA_HOME}/bin:${PATH}" + export LD_LIBRARY_PATH="${CUDA_HOME}/lib64:${LD_LIBRARY_PATH:-}" + start_runtime + exit $? +fi + +if ! command -v uv >/dev/null 2>&1; then + curl -LsSf https://astral.sh/uv/install.sh | sh + export PATH="${HOME}/.local/bin:${PATH}" +fi + +download_model() { + local attempt status + for attempt in 1 2 3 4 5; do + echo "FREETOKEN_PROVISION_STAGE=model_download attempt=${attempt}" + if uvx --from huggingface-hub==1.29.0 hf download "$model_repo" --local-dir "$model_source_dir"; then + echo "FREETOKEN_PROVISION_STAGE=model_download_complete" + return 0 + else + status=$? + fi + echo "FREETOKEN_PROVISION_STAGE=model_download_retry status=${status}" + sleep "$((attempt * 10))" + done + return "$status" +} + +download_model & +model_download_pid=$! + +apt-get update +apt-get install -y --no-install-recommends \ + build-essential ca-certificates cuda-compiler-13-0 cuda-cudart-dev-13-0 \ + curl git libcurand-dev-13-0 ninja-build numactl python3.12-dev util-linux +apt-get clean + +export CUDA_HOME="${CUDA_HOME:-/usr/local/cuda-13.0}" +export PATH="${CUDA_HOME}/bin:${PATH}" +export LD_LIBRARY_PATH="${CUDA_HOME}/lib64:${LD_LIBRARY_PATH:-}" + +checkout_ref() { + local destination="$1" + if [[ ! -d "$destination/.git" ]]; then + mkdir -p "$destination" + git -C "$destination" init + git -C "$destination" remote add origin "$repo" + fi + git -C "$destination" fetch --depth 1 origin "$ref" + git -C "$destination" checkout --detach --force FETCH_HEAD +} + +checkout_ref "$freetoken_dir" +checkout_ref "$worker_source_dir" + +if [[ -n "$expected_commit" ]]; then + for checkout in "$freetoken_dir" "$worker_source_dir"; do + actual_commit="$(git -C "$checkout" rev-parse HEAD)" + if [[ "$actual_commit" != "$expected_commit" ]]; then + printf 'Expected FreeToken commit %s at %s, got %s\n' \ + "$expected_commit" "$checkout" "$actual_commit" >&2 + exit 1 + fi + done +fi + +echo "FREETOKEN_PROVISION_STAGE=pyworker_bootstrap" +curl -fsSL \ + "https://raw.githubusercontent.com/vast-ai/pyworker/${bootstrap_ref}/start_server.sh" \ + -o "$bootstrap" +chmod 0755 "$bootstrap" +UV_CACHE_DIR="$pyworker_uv_cache" \ + USE_SYSTEM_PYTHON=true \ + ROTATE_MODEL_LOG=true \ + "$bootstrap" & +pyworker_pid=$! + +echo "FREETOKEN_PROVISION_STAGE=dependencies" +if [[ ! -x "$freetoken_dir/.venv/bin/python" ]]; then + uv venv --python 3.12 "$freetoken_dir/.venv" +fi +uv pip install --python "$freetoken_dir/.venv/bin/python" -e "$freetoken_dir[accel]" + +wait "$model_download_pid" +if [[ "$convert_ftw" == "1" ]] && [[ ! -s "$model_dir/freetoken_weight.json" ]]; then + echo "FREETOKEN_PROVISION_STAGE=ftw_conversion" + rm -rf "$model_dir" + FREETOKEN_SKIP_BANK_PIN=1 \ + "$freetoken_dir/.venv/bin/ft" checkpoint \ + --model "$model_source_dir" \ + --out "$model_dir" \ + --moe-backend offload + echo "FREETOKEN_PROVISION_STAGE=ftw_conversion_complete" +fi +echo "FREETOKEN_PROVISION_STAGE=bandwidth_check" +"$freetoken_dir/.venv/bin/ft" bench bw --dtype "$model_bench_dtype" + +printf '%s\n' "$provision_marker_value" >"$provision_marker" + +echo "FREETOKEN_PROVISION_STAGE=model_start" +"$freetoken_dir/scripts/vast_glm53_start.sh" & +model_launcher_pid=$! + +wait "$pyworker_pid" diff --git a/scripts/vast_glm53_start.sh b/scripts/vast_glm53_start.sh new file mode 100755 index 000000000..0e4e4ceb3 --- /dev/null +++ b/scripts/vast_glm53_start.sh @@ -0,0 +1,76 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +model_path="${TEKIZAI_FREETOKEN_MODEL_PATH:-/workspace/models/GLM-5.3-Flash-NVFP4}" +served_model="${TEKIZAI_SERVED_MODEL:-glm-5.3-flash-nvfp4}" +ft_executable="${TEKIZAI_FREETOKEN_EXECUTABLE:-/workspace/freetoken/.venv/bin/ft}" +port="${TEKIZAI_FREETOKEN_PORT:-1919}" +log_file="${TEKIZAI_FREETOKEN_LOG:-/workspace/logs/freetoken-glm53.log}" +max_running_requests="${TEKIZAI_MAX_RUNNING_REQUESTS:-1}" +moe_backend="${TEKIZAI_MOE_BACKEND:-auto}" + +if [[ ! "$max_running_requests" =~ ^[1-9][0-9]*$ ]]; then + printf 'TEKIZAI_MAX_RUNNING_REQUESTS must be a positive integer, got %q\n' \ + "$max_running_requests" >&2 + exit 2 +fi + +mkdir -p "$(dirname "$log_file")" +: >"$log_file" + +serve_cmd=("$ft_executable" serve \ + --model "$model_path" \ + --served-model-name "$served_model" \ + --host 127.0.0.1 \ + --port "$port" \ + --moe-backend "$moe_backend" \ + --moe-cpu-threads "${TEKIZAI_CPU_THREADS:-48}" \ + --memory-ratio "${TEKIZAI_MEMORY_RATIO:-0.95}" \ + --max-seq-len-override "${TEKIZAI_MAX_SEQ_LEN:-32768}" \ + --max-running-requests "$max_running_requests" \ + --disable-moe-prefill-overlap) + +gpu_cpu_affinity="${TEKIZAI_GPU_CPU_AFFINITY:-}" +if [[ -z "$gpu_cpu_affinity" ]] && command -v nvidia-smi >/dev/null; then + gpu_cpu_affinity="$(nvidia-smi topo -m 2>/dev/null | awk '$1 == "GPU0" { print $3; exit }')" +fi +if [[ "$gpu_cpu_affinity" =~ ^[0-9,-]+$ ]] && command -v taskset >/dev/null; then + serve_cmd=(taskset -c "$gpu_cpu_affinity" "${serve_cmd[@]}") +fi + +"${serve_cmd[@]}" >>"$log_file" 2>&1 & +backend_pid=$! +warmup_request="${TEKIZAI_WARMUP_REQUEST:-1}" +warmup_timeout="${TEKIZAI_WARMUP_TIMEOUT:-600}" + +on_exit() { + if kill -0 "$backend_pid" 2>/dev/null; then + kill "$backend_pid" + wait "$backend_pid" || true + fi +} +trap on_exit EXIT INT TERM + +for _ in $(seq 1 "${TEKIZAI_READY_POLLS:-900}"); do + if ! kill -0 "$backend_pid" 2>/dev/null; then + printf '%s\n' 'FREETOKEN_SERVER_EXITED' >>"$log_file" + wait "$backend_pid" + fi + health="$(curl --fail --silent --max-time 2 "http://127.0.0.1:${port}/health" || true)" + if [[ "$health" == *'"status":"ok"'* ]]; then + if [[ "$warmup_request" == "1" ]]; then + curl --fail --silent --show-error --max-time "$warmup_timeout" \ + "http://127.0.0.1:${port}/v1/chat/completions" \ + -H 'Content-Type: application/json' \ + --data-binary "{\"model\":\"${served_model}\",\"messages\":[{\"role\":\"user\",\"content\":\"Reply OK.\"}],\"max_tokens\":1,\"temperature\":0,\"reasoning_effort\":\"none\",\"stream\":false}" \ + >/dev/null + fi + printf '%s\n' 'FREETOKEN_SERVER_READY' >>"$log_file" + wait "$backend_pid" + exit $? + fi + sleep 2 +done + +printf '%s\n' 'FREETOKEN_SERVER_EXITED readiness_timeout' >>"$log_file" +exit 1 diff --git a/scripts/vast_qwen38_27b_provision.sh b/scripts/vast_qwen38_27b_provision.sh new file mode 100755 index 000000000..cddf8b5f0 --- /dev/null +++ b/scripts/vast_qwen38_27b_provision.sh @@ -0,0 +1,30 @@ +#!/usr/bin/env bash +set -Eeuo pipefail + +# Qwen3.8-27B NVFP4 uses the cacheable Vast/PyWorker bootstrap shared with +# GLM and DeepSeek. Keep source and FTW paths distinct so conversion is +# resumable and a stopped worker can restart directly from the FTW checkpoint. +script_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +shared_provisioner="$script_dir/vast_glm53_provision.sh" +freetoken_ref="${TEKIZAI_FREETOKEN_REF:-feat/glm53-flash}" + +export TEKIZAI_MODEL_SOURCE_PATH="${TEKIZAI_MODEL_SOURCE_PATH:-/workspace/models/Qwen3.8-27B-NVFP4-hf}" +export TEKIZAI_FREETOKEN_MODEL_PATH="${TEKIZAI_FREETOKEN_MODEL_PATH:-/workspace/models/Qwen3.8-27B-NVFP4-ftw}" +export TEKIZAI_CONVERT_FTW="${TEKIZAI_CONVERT_FTW:-1}" +export TEKIZAI_MODEL_REPO="${TEKIZAI_MODEL_REPO:-RadixArk/Qwen3.8-27B-NVFP4}" +export TEKIZAI_MODEL_BENCH_DTYPE="${TEKIZAI_MODEL_BENCH_DTYPE:-nvfp4}" +export TEKIZAI_SERVED_MODEL="${TEKIZAI_SERVED_MODEL:-qwen3.8:27b}" +export TEKIZAI_FREETOKEN_LOG="${TEKIZAI_FREETOKEN_LOG:-/workspace/logs/freetoken-qwen38-27b.log}" +export TEKIZAI_PROVISION_MARKER="${TEKIZAI_PROVISION_MARKER:-/workspace/.tekizai-qwen38-27b-provisioned}" +export TEKIZAI_MEMORY_RATIO="${TEKIZAI_MEMORY_RATIO:-0.90}" +export TEKIZAI_MAX_SEQ_LEN="${TEKIZAI_MAX_SEQ_LEN:-8192}" +export TEKIZAI_MAX_RUNNING_REQUESTS="${TEKIZAI_MAX_RUNNING_REQUESTS:-1}" + +if [[ ! -x "$shared_provisioner" ]]; then + curl -fsSL \ + "https://raw.githubusercontent.com/earlvanze/FreeToken/${freetoken_ref}/scripts/vast_glm53_provision.sh" \ + -o "$shared_provisioner" + chmod 0755 "$shared_provisioner" +fi + +exec "$shared_provisioner" diff --git a/tests/kernels/test_glm5_kda.py b/tests/kernels/test_glm5_kda.py new file mode 100644 index 000000000..db6be8fbd --- /dev/null +++ b/tests/kernels/test_glm5_kda.py @@ -0,0 +1,104 @@ +"""Numerical parity for GLM-5.3's bounded Kimi Delta Attention recurrence.""" + +from __future__ import annotations + +import pytest +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") + + +def _reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + beta_logits: torch.Tensor, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + cu_seqlens: torch.Tensor, + state_source: torch.Tensor, + state_indices: torch.Tensor, + lower_bound: float, +) -> tuple[torch.Tensor, torch.Tensor]: + total, heads, key_dim = q.shape[1:] + q = q.float() + k = k.float() + v = v.float() + a = a.float().view(total, heads, key_dim) + dt_bias = dt_bias.float().view(heads, key_dim) + outputs = torch.empty_like(v) + expected_states = state_source.float().clone() + scale = key_dim**-0.5 + + for request in range(len(cu_seqlens) - 1): + begin = int(cu_seqlens[request]) + end = int(cu_seqlens[request + 1]) + slot = int(state_indices[request]) + # Runtime storage is [head, value, key]; recurrence math is [head, key, value]. + state = expected_states[slot].transpose(-1, -2).contiguous() + for token in range(begin, end): + q_t = q[0, token] + k_t = k[0, token] + q_t = q_t / torch.sqrt((q_t * q_t).sum(-1, keepdim=True) + 1e-6) + k_t = k_t / torch.sqrt((k_t * k_t).sum(-1, keepdim=True) + 1e-6) + decay = lower_bound * torch.sigmoid( + torch.exp(a_log.float()).unsqueeze(-1) * (a[token] + dt_bias) + ) + state = state * torch.exp(decay).unsqueeze(-1) + delta = v[0, token] - torch.einsum("hkv,hk->hv", state, k_t) + delta = delta * torch.sigmoid(beta_logits[token].float()).unsqueeze(-1) + state = state + k_t.unsqueeze(-1) * delta.unsqueeze(-2) + outputs[0, token] = torch.einsum("hkv,hk->hv", state, q_t * scale) + expected_states[slot] = state.transpose(-1, -2) + return outputs, expected_states + + +@pytest.mark.parametrize("lengths", [(5,), (3, 2)]) +def test_glm5_bounded_kda_varlen_matches_reference(lengths): + from freetoken.models.glm5_next.kda import kda_decode + + torch.manual_seed(53) + device = torch.device("cuda") + heads, key_dim, value_dim = 2, 8, 8 + total = sum(lengths) + q = torch.randn(1, total, heads, key_dim, device=device, dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn(1, total, heads, value_dim, device=device, dtype=torch.bfloat16) + a = torch.randn(total, heads * key_dim, device=device, dtype=torch.bfloat16) + beta = torch.randn(total, heads, device=device, dtype=torch.bfloat16) + a_log = torch.randn(heads, device=device, dtype=torch.float32) / 4 + dt_bias = torch.randn(heads * key_dim, device=device, dtype=torch.float32) + cu = torch.tensor( + (0, *torch.tensor(lengths).cumsum(0).tolist()), device=device, dtype=torch.int32 + ) + indices = torch.arange(len(lengths), device=device, dtype=torch.int32) + initial = ( + torch.randn( + len(lengths), heads, value_dim, key_dim, device=device, dtype=torch.float32 + ) + / 10 + ) + expected_out, expected_state = _reference( + q, k, v, a, beta, a_log, dt_bias, cu, initial, indices, -5.0 + ) + + actual_state = initial.clone() + actual_out = kda_decode( + q, + k, + v, + a, + beta, + A_log=a_log, + dt_bias=dt_bias, + state_source=actual_state, + indices=indices, + cu_seqlens=cu, + gate_lower_bound=-5.0, + ) + + torch.testing.assert_close( + actual_out.float(), expected_out.float(), atol=3e-2, rtol=3e-2 + ) + torch.testing.assert_close(actual_state, expected_state, atol=3e-2, rtol=3e-2) diff --git a/tests/kvcache/test_kv_cache_rebuild.py b/tests/kvcache/test_kv_cache_rebuild.py index a5ec94a08..07a230770 100644 --- a/tests/kvcache/test_kv_cache_rebuild.py +++ b/tests/kvcache/test_kv_cache_rebuild.py @@ -1,7 +1,7 @@ from __future__ import annotations +import pytest import torch - from freetoken.distributed import set_tp_info, try_get_tp_info @@ -15,8 +15,13 @@ def _mha_pool(num_pages=4): _init_tp() return MHAKVCache( - num_kv_heads=8, num_layers=3, head_dim=64, - num_pages=num_pages, page_size=16, dtype=torch.float16, device=torch.device("cpu"), + num_kv_heads=8, + num_layers=3, + head_dim=64, + num_pages=num_pages, + page_size=16, + dtype=torch.float16, + device=torch.device("cpu"), ) @@ -50,26 +55,83 @@ def test_mla_and_dsa_rebuild_from_config_and_unit_bytes(): from freetoken.kvcache.dsa_pool import DSAKVCache, MLAKVCache latent, idx_dim, layers, n_idx = 80, 32, 2, 1 - mla = MLAKVCache(latent_dim=latent, num_layers=layers, num_pages=8, page_size=1, - dtype=torch.bfloat16, device=torch.device("cpu")) + mla = MLAKVCache( + latent_dim=latent, + num_layers=layers, + num_pages=8, + page_size=1, + dtype=torch.bfloat16, + device=torch.device("cpu"), + ) mla.rebuild_from_config(config=None, num_pages=20) assert mla.latent_rows(0).shape[0] == 21 # 20 usable + 1 dummy page assert mla.unit_bytes() == (layers * latent * 2, 0) - dsa = DSAKVCache(latent_dim=latent, num_layers=layers, num_pages=8, page_size=1, - dtype=torch.bfloat16, device=torch.device("cpu"), - index_head_dim=idx_dim, num_index_layers=n_idx) + dsa = DSAKVCache( + latent_dim=latent, + num_layers=layers, + num_pages=8, + page_size=1, + dtype=torch.bfloat16, + device=torch.device("cpu"), + index_head_dim=idx_dim, + num_index_layers=n_idx, + ) dsa.rebuild_from_config(config=None, num_pages=20) assert dsa.latent_rows(0).shape[0] == 21 and dsa.index_k_cache(0).shape[0] == 21 # the index slab's per-token bytes ride on top of the latent slab's, each floored on its own assert dsa.unit_bytes() == (layers * latent * 2 + n_idx * idx_dim * 2, 0) +def test_mla_hybrid_layer_map_allocates_only_paged_kv_layers(): + from freetoken.kvcache.dsa_pool import DSAKVCache, MLAKVCache + + latent, idx_dim = 80, 32 + layer_ids = (3, 7, 11) + mla = MLAKVCache( + latent_dim=latent, + num_layers=12, + num_pages=8, + page_size=1, + dtype=torch.bfloat16, + device=torch.device("cpu"), + layer_ids=layer_ids, + ) + assert mla._kv_buffer.shape == (1, len(layer_ids), 8, 1, 1, latent) + assert mla.num_layers == 12 + assert mla.latent_rows(7).shape == (8, latent) + with pytest.raises(KeyError, match="no paged KV storage"): + mla.latent_rows(4) + assert mla.unit_bytes() == (len(layer_ids) * latent * 2, 0) + + dsa = DSAKVCache( + latent_dim=latent, + num_layers=12, + num_pages=8, + page_size=1, + dtype=torch.bfloat16, + device=torch.device("cpu"), + index_head_dim=idx_dim, + num_index_layers=2, + layer_ids=layer_ids, + ) + assert dsa._kv_buffer.shape == (1, len(layer_ids), 8, 1, 1, latent) + assert dsa.index_k_cache(1).shape == (8, idx_dim) + assert dsa.unit_bytes() == ( + len(layer_ids) * latent * 2 + 2 * idx_dim * 2, + 0, + ) + + def _hybrid_groups(): from freetoken.models.config import KVCacheGroupSpec - full = KVCacheGroupSpec(name="full", layer_ids=(0, 2), num_kv_heads=8, head_dim=64, sliding_window=None) - swa = KVCacheGroupSpec(name="swa", layer_ids=(1,), num_kv_heads=8, head_dim=64, sliding_window=128) + full = KVCacheGroupSpec( + name="full", layer_ids=(0, 2), num_kv_heads=8, head_dim=64, sliding_window=None + ) + swa = KVCacheGroupSpec( + name="swa", layer_ids=(1,), num_kv_heads=8, head_dim=64, sliding_window=128 + ) return [full, swa] @@ -78,8 +140,13 @@ def test_hybrid_swa_rebuild_resizes_both_groups_preserves_identity(): _init_tp() pool = HybridSWAKVCache( - groups=_hybrid_groups(), num_layers=3, num_full_pages=4, page_size=16, - num_swa_tokens=32, dtype=torch.float16, device=torch.device("cpu"), + groups=_hybrid_groups(), + num_layers=3, + num_full_pages=4, + page_size=16, + num_swa_tokens=32, + dtype=torch.float16, + device=torch.device("cpu"), ) pool_id = id(pool) mapping_before = pool.layers_mapping @@ -123,8 +190,11 @@ def _swa_config(cache_type: str): def test_hybrid_swa_rebuild_from_config_derives_the_window_per_cache_type(): """The window size the engine used to compute for the pool: ratio x full for radix, concurrency x window for naive.""" - from freetoken.kvcache.hybrid_swa_pool import _naive_swa_num_tokens, _swa_paged_num_tokens - from freetoken.kvcache.hybrid_swa_pool import HybridSWAKVCache + from freetoken.kvcache.hybrid_swa_pool import ( + HybridSWAKVCache, + _naive_swa_num_tokens, + _swa_paged_num_tokens, + ) _init_tp() for cache_type, expected in ( @@ -132,8 +202,13 @@ def test_hybrid_swa_rebuild_from_config_derives_the_window_per_cache_type(): ("naive", _naive_swa_num_tokens), ): pool = HybridSWAKVCache( - groups=_hybrid_groups(), num_layers=3, num_full_pages=4, page_size=16, - num_swa_tokens=32, dtype=torch.float16, device=torch.device("cpu"), + groups=_hybrid_groups(), + num_layers=3, + num_full_pages=4, + page_size=16, + num_swa_tokens=32, + dtype=torch.float16, + device=torch.device("cpu"), ) config = _swa_config(cache_type) pool.rebuild_from_config(config, 10) @@ -149,15 +224,25 @@ def test_linear_state_pool_rebuild_resizes_preserves_identity_and_dtypes(): _init_tp() group = LinearGatedDeltaGroupConfig( - name="linear", layer_ids=(0, 1, 2), num_key_heads=4, num_value_heads=8, - key_head_dim=16, value_head_dim=16, conv_kernel_dim=4, output_gate="silu", + name="linear", + layer_ids=(0, 1, 2), + num_key_heads=4, + num_value_heads=8, + key_head_dim=16, + value_head_dim=16, + conv_kernel_dim=4, + output_gate="silu", + ) + pool = LinearStatePool( + group=group, num_slots=10, dtype=torch.bfloat16, device=torch.device("cpu") ) - pool = LinearStatePool(group=group, num_slots=10, dtype=torch.bfloat16, device=torch.device("cpu")) pid = id(pool) conv_dtype, rec_dtype = pool.conv_states.dtype, pool.recurrent_states.dtype _, _, conv_dim, km1 = pool.conv_states.shape _, _, v_heads, k_dim, v_dim = pool.recurrent_states.shape - assert pool.num_slots == 10 and pool.num_free_slots == 9 # slot 0 is the padding sink + assert ( + pool.num_slots == 10 and pool.num_free_slots == 9 + ) # slot 0 is the padding sink pool.rebuild(25) @@ -178,8 +263,8 @@ def test_dsv4_rebuild_from_config_builds_the_pool_sizes_and_attaches_the_table() the engine used to do -- and re-points full_loc_map at the shared page table.""" from types import SimpleNamespace - from freetoken.kvcache.dsv4_cost_model import _dsv4_pool_sizes from freetoken.kvcache.dsv4_cost_model import ( + _dsv4_pool_sizes, dsv4_kv_unit_bytes, dsv4_pool_sizes, dsv4_window_unit_bytes, @@ -189,17 +274,29 @@ def test_dsv4_rebuild_from_config_builds_the_pool_sizes_and_attaches_the_table() P, mrr = 128, 1 args = DeepseekV4Args( - n_layers=4, compress_ratios=(0, 4, 128, 4), max_seq_len=512, - head_dim=64, index_head_dim=32, window_size=P, + n_layers=4, + compress_ratios=(0, 4, 128, 4), + max_seq_len=512, + head_dim=64, + index_head_dim=32, + window_size=P, ) config = SimpleNamespace( - max_seq_len=512, page_size=P, max_running_req=mrr, cache_type="swa_radix", - swa_full_tokens_ratio=0.5, swa_num_pages_override=None, + max_seq_len=512, + page_size=P, + max_running_req=mrr, + cache_type="swa_radix", + swa_full_tokens_ratio=0.5, + swa_num_pages_override=None, model_config=SimpleNamespace(dsv4_args=args), ) pool = DSV4PagedKVCache( sizes=dsv4_pool_sizes(num_pages=4, args=args, swa_ratio=0.5, P=P), - args=args, device=torch.device("cpu"), dtype=torch.bfloat16, P=P, n_scratch=mrr + 1, + args=args, + device=torch.device("cpu"), + dtype=torch.bfloat16, + P=P, + n_scratch=mrr + 1, ) pool._init_paged_state(mrr, True) # the engine's create_kv_pool step pool.rebuild_from_config(config, 15) @@ -207,7 +304,10 @@ def test_dsv4_rebuild_from_config_builds_the_pool_sizes_and_attaches_the_table() assert pool.sizes == _dsv4_pool_sizes(config, 16) # 15 usable + 1 dummy page assert pool.sizes.full_token == 16 * P assert pool.window_pool[0].shape[0] == pool.sizes.n_win_slots - assert pool.unit_bytes() == (dsv4_kv_unit_bytes(args, P), dsv4_window_unit_bytes(args, P)) + assert pool.unit_bytes() == ( + dsv4_kv_unit_bytes(args, P), + dsv4_window_unit_bytes(args, P), + ) page_table = torch.zeros((mrr + 1, 64), dtype=torch.int32) pool.attach_page_table(page_table) @@ -221,13 +321,15 @@ def test_dsv4_refresh_seq_state_tracks_page_table_width(): from types import SimpleNamespace from freetoken.engine.engine import Engine - from freetoken.utils import align_ceil from freetoken.kvcache.dsv4_paged_pool import DSV4PagedKVCache from freetoken.scheduler.table import TableManager + from freetoken.utils import align_ceil P, mrr = 128, 4 config = SimpleNamespace( - max_seq_len=65536, page_size=P, max_running_req=mrr, + max_seq_len=65536, + page_size=P, + max_running_req=mrr, model_config=SimpleNamespace(dsv4_args=SimpleNamespace(window_size=P)), ) # attach_page_table only re-points the pool's full_loc_map; the decode snapshot belongs to @@ -236,13 +338,16 @@ def test_dsv4_refresh_seq_state_tracks_page_table_width(): pool = object.__new__(DSV4PagedKVCache) pool.full_loc_map = None eng = SimpleNamespace( - num_pages=489, device=torch.device("cpu"), + num_pages=489, + device=torch.device("cpu"), ctx=SimpleNamespace(page_table=None), dummy_req=SimpleNamespace(table_idx=mrr), kv_cache=pool, ) eng.max_seq_len = min(config.max_seq_len, eng.num_pages * P) - eng.page_table = torch.zeros((mrr + 1, align_ceil(eng.max_seq_len, 32)), dtype=torch.int32) + eng.page_table = torch.zeros( + (mrr + 1, align_ceil(eng.max_seq_len, 32)), dtype=torch.int32 + ) tm = TableManager(mrr, eng.page_table) for target in (519, 489, 400, 519): # grow, shrink back, shrink, grow again @@ -270,7 +375,9 @@ def test_every_kv_pool_answers_the_sizing_surface(): for cls in (MHAKVCache, MLAKVCache, DSAKVCache, HybridSWAKVCache, DSV4PagedKVCache): for hook in ("kv_cost", "solve_num_pages", "min_kv_tokens", "validate_rebuild"): - assert callable(getattr(cls, hook, None)), f"{cls.__name__} is missing {hook}" + assert callable(getattr(cls, hook, None)), ( + f"{cls.__name__} is missing {hook}" + ) for hook in ("kv_cost", "solve_num_pages", "min_kv_tokens", "validate_rebuild"): assert hook in DSV4PagedKVCache.__dict__, f"DSV4 lost its {hook} override" @@ -284,7 +391,9 @@ def test_every_kv_pool_answers_the_rebuild_surface(): from freetoken.kvcache.mha_pool import MHAKVCache for cls in (MHAKVCache, MLAKVCache, DSAKVCache, HybridSWAKVCache, DSV4PagedKVCache): - assert not cls.__abstractmethods__, f"{cls.__name__}: {sorted(cls.__abstractmethods__)}" + assert not cls.__abstractmethods__, ( + f"{cls.__name__}: {sorted(cls.__abstractmethods__)}" + ) # Only DSV4 rebinds the model and re-points a page table; the rest inherit the defaults. assert DSV4PagedKVCache.needs_rebind_on_rebuild assert not any( diff --git a/tests/models/test_glm5_next_config.py b/tests/models/test_glm5_next_config.py new file mode 100644 index 000000000..80fdce62d --- /dev/null +++ b/tests/models/test_glm5_next_config.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +from types import SimpleNamespace + + +def _release_config(): + """Small, self-contained projection of zai-org/GLM-5.3-Flash config.json.""" + num_layers = 45 + layer_types = [ + "deepseek_sparse_attention" if i % 4 == 3 else "linear_attention" + for i in range(num_layers) + ] + return SimpleNamespace( + model_type="glm5_next", + architectures=["Glm5NextForConditionalGeneration"], + quantization_config={ + "quant_method": "fp8", + "weight_block_size": [128, 128], + }, + text_config=SimpleNamespace( + num_hidden_layers=num_layers, + layer_types=layer_types, + mlp_layer_types=["dense"] * 3 + ["sparse"] * 42, + indexer_types=["full"] * num_layers, + hidden_size=4096, + vocab_size=154880, + intermediate_size=12288, + hidden_act="silu", + rms_norm_eps=1e-5, + tie_word_embeddings=False, + max_position_embeddings=1048576, + rope_theta=None, + num_attention_heads=64, + linear_attn_config={ + "num_heads": 64, + "head_dim": 128, + "short_conv_kernel_size": 4, + "gate_lower_bound": -5.0, + }, + qk_head_dim=256, + qk_nope_head_dim=256, + qk_rope_head_dim=0, + v_head_dim=256, + kv_lora_rank=512, + q_lora_rank=1536, + index_n_heads=32, + index_head_dim=128, + index_topk=2048, + index_kpool=4, + hc_mult=4, + hc_eps=1e-6, + hc_sinkhorn_iters=20, + n_routed_experts=288, + num_experts_per_tok=8, + moe_intermediate_size=2048, + norm_topk_prob=True, + first_k_dense_replace=3, + n_shared_experts=1, + routed_scaling_factor=2.5, + n_group=1, + topk_group=1, + swiglu_limit=10.0, + ), + ) + + +def test_official_fp8_config_maps_hybrid_geometry(): + from freetoken.models.glm5_next.config import parse_config + + config = parse_config(_release_config()) + + assert config.num_layers == 45 + assert config.num_moe_layers == 42 + assert config.expert_quant == "fp8_block" + assert config.weight_block_size == (128, 128) + assert config.attn_quant == "none" + assert config.glm5_args.hc_mult == 4 + assert config.glm5_args.index_kpool == 4 + assert config.linear_attention_group().layer_ids[:4] == (0, 1, 2, 4) + full = config.attention_group_for_layer(3) + assert full.layer_ids == tuple(range(3, 45, 4)) + assert full.mla and full.head_dim == 512 + assert full.index_ratio == 4 + + +def test_nvfp4_config_keeps_only_routed_experts_quantized(): + from freetoken.models.glm5_next.config import _quant_modes + + config = SimpleNamespace( + quantization_config={ + "quant_algo": "NVFP4", + "ignore": [ + "*.self_attn.q_proj", + "*.mlp.gate_proj", + "lm_head", + ], + } + ) + assert _quant_modes(config) == ("nvfp4", "none", "none", "none", None) + + +def test_sparse_block_propagates_release_swiglu_limit(monkeypatch): + from freetoken.models.glm5_next import moe + + captured = {} + + def fake_make_moe_layer(config, **kwargs): + captured.update(kwargs) + return SimpleNamespace() + + monkeypatch.setattr(moe, "make_moe_layer", fake_make_moe_layer) + monkeypatch.setattr( + moe, "LinearReplicated", lambda *args, **kwargs: SimpleNamespace() + ) + monkeypatch.setattr(moe, "Glm5NextMLP", lambda *args, **kwargs: SimpleNamespace()) + config = SimpleNamespace( + num_experts_per_tok=8, + num_experts=288, + norm_topk_prob=True, + routed_scaling_factor=2.5, + n_group=1, + topk_group=1, + hidden_size=128, + first_k_dense_replace=3, + n_shared_experts=1, + moe_intermediate_size=64, + swiglu_limit=10.0, + ) + + moe.Glm5NextSparseBlock(config, layer_id=3) + + assert captured["extra_attrs"] == {"swiglu_limit": 10.0} diff --git a/tests/models/test_glm5_next_hc.py b/tests/models/test_glm5_next_hc.py new file mode 100644 index 000000000..f76f702b2 --- /dev/null +++ b/tests/models/test_glm5_next_hc.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + + +def _reference(hidden, fn, base, scale, norm_eps=1e-5, hc_eps=1e-6, iters=20): + hc = hidden.shape[-2] + flat = hidden.flatten(-2).float() + normed = flat * torch.rsqrt(flat.square().mean(-1, keepdim=True) + norm_eps) + pre_w, post_w, comb_w = F.linear(normed, fn).split([hc, hc, hc * hc], dim=-1) + pre_b, post_b, comb_b = base.split([hc, hc, hc * hc]) + pre = torch.sigmoid(pre_w * scale[0] + pre_b) + hc_eps + post = 2 * torch.sigmoid(post_w * scale[1] + post_b) + comb = torch.softmax( + comb_w.reshape(*comb_w.shape[:-1], hc, hc) * scale[2] + comb_b.view(hc, hc), + -1, + ) + comb = comb + hc_eps + comb = comb / (comb.sum(-2, keepdim=True) + hc_eps) + for _ in range(iters - 1): + comb = comb / (comb.sum(-1, keepdim=True) + hc_eps) + comb = comb / (comb.sum(-2, keepdim=True) + hc_eps) + collapsed = (pre.unsqueeze(-1) * hidden.float()).sum(-2).to(hidden.dtype) + return collapsed, post, comb + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_mhc_fused_path_matches_release_equations_on_ampere(): + from freetoken.models.glm5_next.hc import HyperConnection, collapse_head + + torch.manual_seed(53) + device = torch.device("cuda") + module = HyperConnection(hidden_size=128).to(device) + with torch.no_grad(): + module.fn.normal_(std=0.02) + module.base.normal_(std=0.02) + module.scale.copy_(torch.tensor([0.7, 0.9, 1.1], device=device)) + hidden = torch.randn(2, 3, 4, 128, device=device, dtype=torch.bfloat16) + + collapsed, post, comb = module.mix(hidden) + ref_collapsed, ref_post, ref_comb = _reference( + hidden, module.fn, module.base, module.scale + ) + torch.testing.assert_close(collapsed, ref_collapsed, rtol=0, atol=0.02) + torch.testing.assert_close(post, ref_post, rtol=2e-5, atol=2e-5) + torch.testing.assert_close(comb, ref_comb, rtol=2e-5, atol=2e-5) + + block_output = torch.randn_like(collapsed) + got = module.combine(hidden, block_output, post, comb) + expected = post.to(hidden.dtype).unsqueeze(-1) * block_output.unsqueeze(-2) + expected += torch.matmul(comb.to(hidden.dtype).transpose(-1, -2), hidden) + # The fused kernel accumulates in fp32 before one bf16 store; the eager release + # expression rounds the multiply and matmul separately (at most two bf16 ulps here). + torch.testing.assert_close(got, expected, rtol=0, atol=0.04) + torch.testing.assert_close(collapse_head(hidden), hidden.mean(dim=-2), rtol=0, atol=0) diff --git a/tests/models/test_glm5_next_kda.py b/tests/models/test_glm5_next_kda.py new file mode 100644 index 000000000..3aaaa205d --- /dev/null +++ b/tests/models/test_glm5_next_kda.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import pytest +import torch + + +def _release_recurrence(q, k, v, gate_a, beta_logits, A_log, dt_bias, state): + q = q.float() + k = k.float() + q = q / torch.sqrt(q.square().sum(-1, keepdim=True) + 1e-6) + k = k / torch.sqrt(k.square().sum(-1, keepdim=True) + 1e-6) + q = q * (q.shape[-1] ** -0.5) + beta = beta_logits.float().sigmoid() + gate_a = gate_a.float().reshape(*q.shape[:-1], q.shape[-1]) + g = -5.0 * torch.sigmoid(A_log.float().view(1, 1, -1, 1).exp() * (gate_a + dt_bias)) + outputs = [] + state = state.float().clone() + for t in range(q.shape[1]): + state *= g[:, t].exp().unsqueeze(-1) + remembered = (state * k[:, t].unsqueeze(-1)).sum(-2) + delta = (v[:, t].float() - remembered) * beta[:, t].unsqueeze(-1) + state += k[:, t].unsqueeze(-1) * delta.unsqueeze(-2) + outputs.append((state * q[:, t].unsqueeze(-1)).sum(-2)) + return torch.stack(outputs, dim=1).to(q.dtype), state + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_bounded_kda_decode_matches_release_recurrence_on_ampere(): + from freetoken.models.glm5_next.kda import kda_decode + + torch.manual_seed(530) + device = torch.device("cuda") + batch, tokens, heads, dim = 1, 3, 2, 128 + q = torch.randn(batch, tokens, heads, dim, device=device, dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn_like(q) + gate_a = torch.randn(batch, tokens, heads * dim, device=device, dtype=torch.bfloat16) + beta_logits = torch.randn(batch, tokens, heads, device=device, dtype=torch.bfloat16) + A_log = torch.randn(heads, device=device, dtype=torch.float32) * 0.1 + dt_bias = torch.randn(heads * dim, device=device, dtype=torch.float32) * 0.1 + initial = torch.randn(batch, heads, dim, dim, device=device, dtype=torch.float32) * 0.01 + + # The vendored recurrent kernel stores each KxV state transposed for vectorized V tiles. + state_source = initial.transpose(-1, -2).contiguous() + got = kda_decode( + q, + k, + v, + gate_a, + beta_logits, + A_log=A_log, + dt_bias=dt_bias, + state_source=state_source, + indices=torch.tensor([0], device=device, dtype=torch.int32), + ) + expected, expected_state = _release_recurrence( + q, k, v, gate_a, beta_logits, A_log, dt_bias.view(1, 1, heads, dim), initial + ) + torch.testing.assert_close(got, expected.to(got.dtype), rtol=0, atol=0.02) + torch.testing.assert_close( + state_source.transpose(-1, -2), expected_state, rtol=3e-4, atol=3e-4 + ) diff --git a/tests/models/test_glm5_next_linear_attention.py b/tests/models/test_glm5_next_linear_attention.py new file mode 100644 index 000000000..ee983cb13 --- /dev/null +++ b/tests/models/test_glm5_next_linear_attention.py @@ -0,0 +1,82 @@ +from types import SimpleNamespace + +import torch + + +class _Op: + def __init__(self, fn): + self._fn = fn + + def forward(self, *args): + return self._fn(*args) + + +def test_glm5_kda_prefill_materializes_token_major_qkv(monkeypatch): + """KDA indexes key dimensions as contiguous after the channel-first conv.""" + import freetoken.models.glm5_next.linear_attention as module + + total, heads, head_dim = 4, 2, 3 + qkv_dim = heads * head_dim + projected = [ + torch.arange(total * qkv_dim, dtype=torch.float32).view(total, qkv_dim) + + 100 * i + for i in range(3) + ] + + op = module.Glm5NextLinearAttention.__new__(module.Glm5NextLinearAttention) + op.layer_id = 0 + op.num_heads = heads + op.head_dim = head_dim + op.qkv_dim = qkv_dim + op.conv_dim = 3 * qkv_dim + op.conv_kernel_size = 4 + op.gate_lower_bound = -5.0 + op.q_proj, op.k_proj, op.v_proj = (_Op(lambda _x, y=y: y) for y in projected) + op.conv1d = SimpleNamespace(weight=torch.ones(op.conv_dim, 1, 4)) + op.f_a_proj = _Op(lambda x: torch.zeros(total, head_dim)) + op.f_b_proj = _Op(lambda x: torch.zeros(total, qkv_dim)) + op.b_proj = _Op(lambda x: torch.zeros(total, heads)) + op.g_a_proj = _Op(lambda x: torch.zeros(total, head_dim)) + op.g_b_proj = _Op(lambda x: torch.zeros(total, qkv_dim)) + op.A_log = torch.zeros(heads) + op.dt_bias = torch.zeros(qkv_dim) + op.o_norm = _Op(lambda x, gate: x) + op.o_proj = _Op(lambda x: x) + + fla = SimpleNamespace( + cu_seqlens=torch.tensor([0, total], dtype=torch.int32), + cache_indices=torch.tensor([0], dtype=torch.int32), + has_initial_state=torch.tensor([False]), + fresh_state_indices=None, + ) + batch = SimpleNamespace(is_decode=False, fla_metadata=fla) + pool = SimpleNamespace( + local_index=lambda _layer: 0, + conv_states=[torch.zeros(1, op.conv_dim, 3)], + recurrent_states=[torch.zeros(1, heads, head_dim, head_dim)], + ) + monkeypatch.setattr( + module, + "get_global_ctx", + lambda: SimpleNamespace(batch=batch, linear_state_pool=pool), + ) + # The production conv returns channel-major storage. Its transpose has a + # non-unit feature stride until the model materializes token-major storage. + monkeypatch.setattr( + module, + "causal_conv1d_varlen", + lambda x, *_args: x, + ) + + def capture_kda(q, k, v, *_args, **_kwargs): + for tensor in (q, k, v): + assert tensor.stride(-1) == 1 + assert tensor.stride(-2) == head_dim + torch.testing.assert_close(q.reshape(total, qkv_dim), projected[0]) + torch.testing.assert_close(k.reshape(total, qkv_dim), projected[1]) + torch.testing.assert_close(v.reshape(total, qkv_dim), projected[2]) + return torch.zeros(1, total, heads, head_dim) + + monkeypatch.setattr(module, "kda_decode", capture_kda) + result = op.forward(torch.zeros(total, 8)) + assert result.shape == (total, qkv_dim) diff --git a/tests/models/test_glm5_next_weight.py b/tests/models/test_glm5_next_weight.py new file mode 100644 index 000000000..903b57ee5 --- /dev/null +++ b/tests/models/test_glm5_next_weight.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import torch + + +def test_glm5_next_weight_rename_covers_checkpoint_layout(): + from freetoken.models.glm5_next.weight import _rename + + prefix = "model.language_model.layers.3" + assert _rename(f"{prefix}.hc_attn_fn") == "model.layers.3.attn_hc.fn" + assert _rename(f"{prefix}.hc_ffn_base") == "model.layers.3.ffn_hc.base" + assert _rename(f"{prefix}.mlp.gate.e_score_correction_bias") == ( + "model.layers.3.mlp.e_score_correction_bias" + ) + assert _rename(f"{prefix}.self_attn.q_a_proj.weight") == ( + "model.layers.3.self_attn.q_a_proj.weight" + ) + assert _rename("model.visual.patch_embed.proj.weight") is None + assert _rename(f"{prefix}.mlp.experts.7.gate_proj.weight") is None + assert _rename(f"{prefix}.mlp.shared_experts.gate_proj.weight_scale") is None + assert _rename("model.language_model.layers.45.shared_head.norm.weight") is None + + +def test_glm5_next_fuses_kda_convs_per_layer_in_qkv_order(): + from freetoken.models.glm5_next.weight import _try_fuse_kda_conv + + buf: dict[str, dict[int, torch.Tensor]] = {} + base0 = "model.layers.0.self_attn" + base1 = "model.layers.1.self_attn" + q = torch.full((2, 1, 4), 1.0) + k = torch.full((2, 1, 4), 2.0) + v = torch.full((2, 1, 4), 3.0) + + assert _try_fuse_kda_conv(f"{base0}.q_conv1d.weight", q, buf) == () + assert _try_fuse_kda_conv(f"{base1}.v_conv1d.weight", v, buf) == () + assert _try_fuse_kda_conv(f"{base0}.v_conv1d.weight", v, buf) == () + fused = _try_fuse_kda_conv(f"{base0}.k_conv1d.weight", k, buf) + + assert fused is not None and fused != () + name, tensor = fused + assert name == f"{base0}.conv1d.weight" + assert tensor.shape == (6, 1, 4) + assert tensor[:, 0, 0].tolist() == [1.0, 1.0, 2.0, 2.0, 3.0, 3.0] + assert f"{base1}.conv1d.weight" in buf + assert _try_fuse_kda_conv(f"{base0}.q_proj.weight", q, buf) is None + + +def test_glm5_next_expert_bank_mapping_excludes_appended_mtp_layer(): + from freetoken.models.glm5_next.weight import _SOURCE_SPEC + + config = SimpleNamespace(first_k_dense_replace=3, num_layers=45) + assert _SOURCE_SPEC.layer_to_bank(2, config) is None + assert _SOURCE_SPEC.layer_to_bank(3, config) == 0 + assert _SOURCE_SPEC.layer_to_bank(44, config) == 41 + assert _SOURCE_SPEC.layer_to_bank(45, config) is None + + +def test_glm5_next_weight_loader_uses_safe_open_keys(monkeypatch): + from freetoken.models.glm5_next import weight + + class FakeSafeOpen: + def __enter__(self): + return self + + def __exit__(self, *args): + return False + + def keys(self): + return ["model.language_model.norm.weight"] + + def get_tensor(self, name): + assert name == "model.language_model.norm.weight" + return torch.ones(2) + + monkeypatch.setattr(weight, "get_tp_info", lambda: SimpleNamespace(size=1)) + monkeypatch.setattr( + weight, "iter_weight_files", lambda path: ["weights.safetensors"] + ) + monkeypatch.setattr( + weight.safetensors, "safe_open", lambda *args, **kwargs: FakeSafeOpen() + ) + monkeypatch.setattr(weight, "drop_page_cache", lambda path: None) + + loaded = list( + weight.iter_weights( + "/model", + torch.device("cpu"), + include_moe_experts=False, + include_non_moe=True, + ) + ) + assert len(loaded) == 1 + assert loaded[0][0] == "model.norm.weight" + assert torch.equal(loaded[0][1], torch.ones(2)) diff --git a/tests/models/test_nvfp4_banks.py b/tests/models/test_nvfp4_banks.py new file mode 100644 index 000000000..aaec95350 --- /dev/null +++ b/tests/models/test_nvfp4_banks.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import pytest +from freetoken.models import nvfp4_banks + + +@pytest.mark.parametrize("fail", [False, True]) +def test_single_threaded_torch_copies_restores_threads(monkeypatch, fail): + current = 32 + calls: list[int] = [] + + monkeypatch.setattr(nvfp4_banks.torch, "get_num_threads", lambda: current) + monkeypatch.setattr(nvfp4_banks.torch, "set_num_threads", calls.append) + + if fail: + with pytest.raises(RuntimeError): + with nvfp4_banks._single_threaded_torch_copies(): + raise RuntimeError("loader failed") + else: + with nvfp4_banks._single_threaded_torch_copies(): + pass + + assert calls == [1, current] diff --git a/tests/test_vast_serverless.py b/tests/test_vast_serverless.py new file mode 100644 index 000000000..6a24046cd --- /dev/null +++ b/tests/test_vast_serverless.py @@ -0,0 +1,126 @@ +import os +import sys +import types +from contextlib import contextmanager +from pathlib import Path +from unittest import mock + +import worker + +ROOT = Path(__file__).resolve().parents[1] + + +@contextmanager +def fake_vast_sdk(): + class Config: + def __init__(self, **kwargs): + self.__dict__.update(kwargs) + + module = types.SimpleNamespace( + BenchmarkConfig=Config, + HandlerConfig=Config, + LogActionConfig=Config, + WorkerConfig=Config, + ) + previous = sys.modules.get("vastai") + sys.modules["vastai"] = module + try: + yield + finally: + if previous is None: + del sys.modules["vastai"] + else: + sys.modules["vastai"] = previous + + +def test_vast_worker_defaults_to_serial(): + with mock.patch.dict(os.environ, {}, clear=True), fake_vast_sdk(): + config = worker.build_worker_config() + assert all(not handler.allow_parallel_requests for handler in config.handlers) + benchmark = next( + handler.benchmark_config + for handler in config.handlers + if hasattr(handler, "benchmark_config") + ) + assert benchmark.concurrency == 1 + + +def test_vast_worker_allows_bounded_parallelism(): + env = { + "TEKIZAI_ALLOW_PARALLEL_REQUESTS": "true", + "TEKIZAI_BENCHMARK_CONCURRENCY": "2", + } + with mock.patch.dict(os.environ, env, clear=True), fake_vast_sdk(): + config = worker.build_worker_config() + assert all(handler.allow_parallel_requests for handler in config.handlers) + benchmark = next( + handler.benchmark_config + for handler in config.handlers + if hasattr(handler, "benchmark_config") + ) + assert benchmark.concurrency == 2 + + +def test_vast_launcher_uses_validated_request_limit(): + source = (ROOT / "scripts/vast_glm53_start.sh").read_text(encoding="utf-8") + assert 'max_running_requests="${TEKIZAI_MAX_RUNNING_REQUESTS:-1}"' in source + assert '--max-running-requests "$max_running_requests"' in source + + +def test_vast_provisioner_installs_jit_headers_and_verifies_commit(): + source = (ROOT / "scripts/vast_glm53_provision.sh").read_text(encoding="utf-8") + assert "libcurand-dev-13-0" in source + assert "TEKIZAI_FREETOKEN_EXPECTED_COMMIT" in source + assert 'git -C "$checkout" rev-parse HEAD' in source + + +def test_vast_provisioner_has_validated_fast_resume_path(): + source = (ROOT / "scripts/vast_glm53_provision.sh").read_text(encoding="utf-8") + assert 'FREETOKEN_PROVISION_STAGE=fast_resume' in source + assert '[[ -s "$model_dir/config.json" ]]' in source + assert 'checkout_matches_expected_commit' in source + assert 'printf \'%s\\n\' "$provision_marker_value" >"$provision_marker"' in source + + +def test_vast_provisioner_supports_persistent_ftw_conversion(): + source = (ROOT / "scripts/vast_glm53_provision.sh").read_text(encoding="utf-8") + assert 'model_source_dir="${TEKIZAI_MODEL_SOURCE_PATH:-$model_dir}"' in source + assert 'convert_ftw="${TEKIZAI_CONVERT_FTW:-0}"' in source + assert 'FREETOKEN_PROVISION_STAGE=ftw_conversion' in source + assert '[[ ! -s "$model_dir/freetoken_weight.json" ]]' in source + assert '"$freetoken_dir/.venv/bin/ft" checkpoint' in source + assert '--model "$model_source_dir"' in source + assert '--out "$model_dir"' in source + assert 'TEKIZAI_PROVISION_MARKER' in source + assert '--dtype "$model_bench_dtype"' in source + + +def test_deepseek_v4_vast_provisioner_sets_model_contract(): + source = (ROOT / "scripts/vast_deepseek_v4_provision.sh").read_text( + encoding="utf-8" + ) + assert "deepseek-ai/DeepSeek-V4-Flash-0731" in source + assert "TEKIZAI_MODEL_BENCH_DTYPE" in source + assert "ds_fp4" in source + assert "TEKIZAI_SERVED_MODEL" in source + assert "raw.githubusercontent.com/earlvanze/FreeToken" in source + assert "TEKIZAI_FREETOKEN_REF" in source + assert "TEKIZAI_MODEL_SOURCE_PATH" in source + assert "DeepSeek-V4-Flash-0731-ftw" in source + assert 'TEKIZAI_CONVERT_FTW="${TEKIZAI_CONVERT_FTW:-1}"' in source + assert 'exec "$shared_provisioner"' in source + + +def test_qwen38_wrapper_selects_nvfp4_cache_contract(): + source = (ROOT / "scripts/vast_qwen38_27b_provision.sh").read_text( + encoding="utf-8" + ) + assert "set -Eeuo pipefail" in source + assert "RadixArk/Qwen3.8-27B-NVFP4" in source + assert "Qwen3.8-27B-NVFP4-hf" in source + assert "Qwen3.8-27B-NVFP4-ftw" in source + assert 'TEKIZAI_CONVERT_FTW="${TEKIZAI_CONVERT_FTW:-1}"' in source + assert 'TEKIZAI_MODEL_BENCH_DTYPE="${TEKIZAI_MODEL_BENCH_DTYPE:-nvfp4}"' in source + assert 'TEKIZAI_SERVED_MODEL="${TEKIZAI_SERVED_MODEL:-qwen3.8:27b}"' in source + assert 'TEKIZAI_MEMORY_RATIO="${TEKIZAI_MEMORY_RATIO:-0.90}"' in source + assert 'TEKIZAI_MAX_SEQ_LEN="${TEKIZAI_MAX_SEQ_LEN:-8192}"' in source diff --git a/worker.py b/worker.py new file mode 100644 index 000000000..5db3414ab --- /dev/null +++ b/worker.py @@ -0,0 +1,107 @@ +"""Single-request Vast PyWorker proxy for a loopback FreeToken server.""" + +from __future__ import annotations + +import os +from collections.abc import Mapping +from typing import Any + + +def _positive_int(value: Any, default: int = 0) -> int: + try: + return max(0, int(value)) + except (TypeError, ValueError): + return default + + +def _enabled(value: Any, default: bool = False) -> bool: + if value is None: + return default + return str(value).strip().lower() in {"1", "true", "yes", "on"} + + +def _text_size(value: Any) -> int: + if isinstance(value, str): + return len(value) + if isinstance(value, list): + return sum(_text_size(item) for item in value) + if isinstance(value, Mapping): + return sum(_text_size(item) for item in value.values()) + return 0 + + +def completion_workload(payload: Mapping[str, Any]) -> float: + prompt_tokens = max(1, (_text_size(payload.get("prompt", "")) + 3) // 4) + return float(prompt_tokens + _positive_int(payload.get("max_tokens"), 16)) + + +def chat_workload(payload: Mapping[str, Any]) -> float: + prompt_tokens = max(1, (_text_size(payload.get("messages", [])) + 3) // 4) + return float(prompt_tokens + _positive_int(payload.get("max_tokens"), 16)) + + +def benchmark_payload() -> dict[str, Any]: + return { + "model": os.environ.get("TEKIZAI_SERVED_MODEL", "glm-5.3-flash-nvfp4"), + "prompt": "Reply with exactly: Vast FreeToken worker ready", + "max_tokens": 16, + "temperature": 0, + "stream": False, + } + + +def build_worker_config() -> Any: + from vastai import ( # type: ignore[import-not-found] + BenchmarkConfig, + HandlerConfig, + LogActionConfig, + WorkerConfig, + ) + + port = _positive_int(os.environ.get("TEKIZAI_FREETOKEN_PORT"), 1919) or 1919 + allow_parallel = _enabled(os.environ.get("TEKIZAI_ALLOW_PARALLEL_REQUESTS")) + benchmark_concurrency = ( + _positive_int(os.environ.get("TEKIZAI_BENCHMARK_CONCURRENCY"), 1) or 1 + ) + common = { + "allow_parallel_requests": allow_parallel, + "max_queue_time": 180, + } + return WorkerConfig( + model_server_url="http://127.0.0.1", + model_server_port=port, + model_log_file=os.environ.get( + "TEKIZAI_FREETOKEN_LOG", "/workspace/logs/freetoken-glm53.log" + ), + handlers=[ + HandlerConfig( + route="/v1/completions", + **common, + workload_calculator=completion_workload, + benchmark_config=BenchmarkConfig( + generator=benchmark_payload, + runs=2, + concurrency=benchmark_concurrency, + ), + ), + HandlerConfig( + route="/v1/chat/completions", + **common, + workload_calculator=chat_workload, + ), + ], + log_action_config=LogActionConfig( + on_load=["FREETOKEN_SERVER_READY"], + on_error=["FREETOKEN_SERVER_EXITED", "Traceback"], + ), + ) + + +def main() -> None: + from vastai import Worker # type: ignore[import-not-found] + + Worker(build_worker_config()).run() + + +if __name__ == "__main__": + main()