diff --git a/python/freetoken/attention/qsa_sparse.py b/python/freetoken/attention/qsa_sparse.py index 4a28dc852..85c901211 100644 --- a/python/freetoken/attention/qsa_sparse.py +++ b/python/freetoken/attention/qsa_sparse.py @@ -88,6 +88,10 @@ class QSASparseMetadata(BaseAttnMetadata): cmp_rows: torch.Tensor | None = None # [T] int32, compressed slab destination ring_rows: torch.Tensor | None = None # [T] int32, flat ring row or -1 positions: torch.Tensor | None = None # [T] int32, logical query positions + # mRoPE: rope position per token (positions + Req.mrope_delta), or -- when mrope_cos_sin is + # set (a prefill batch with image tokens) -- the token's row in that per-token cos|sin table. + rope_positions: torch.Tensor | None = None # [T] int32 + mrope_cos_sin: torch.Tensor | None = None # [T, rotary_dim] fp32 or None # fmt: on def get_last_indices(self, bs: int) -> torch.Tensor: @@ -296,6 +300,10 @@ def _plan_index_writes(self, md: QSASparseMetadata, batch: Batch) -> None: """Per-token slab row and ring row for this forward; the other QSA layers reuse it (it is layer-invariant). Pure device arithmetic: no host sync, graph-capturable.""" md.positions = batch.positions + md.rope_positions = ( + batch.positions if batch.rope_positions is None else batch.rope_positions + ) + md.mrope_cos_sin = batch.mrope_cos_sin out_loc = batch.out_loc.to(torch.int64) positions = batch.positions.to(torch.int64) rows = torch.arange(out_loc.numel(), device=self.device) @@ -338,10 +346,16 @@ def _update_index_cache(self, index, md: QSASparseMetadata, slot: int) -> None: pooled, first, ) + if md.rope_positions is not md.positions: + # Rope the pooled key at its group's FIRST token (HF: cos/sin indexed by group start). + # That token's rope position / table row sits at the same offset from this token's as + # the token indices do. Decode groups that close after an image are text (the chat + # template puts >= 4 tokens after <|vision_end|>), so ``first + delta`` is exact. + first = first + (md.rope_positions - md.positions) qsa_index_norm_rope( pooled, first, - self._index_rope_cache(), + self._index_rope_cache() if md.mrope_cos_sin is None else md.mrope_cos_sin, index.k_norm_weight, index.eps, self.kvcache.cmp_k_cache(slot), @@ -366,8 +380,8 @@ def _select(self, index, md: QSASparseMetadata, slot: int) -> torch.Tensor: ) qsa_index_norm_rope( index.q.view(-1, self.index_head_dim), - positions, - self._index_rope_cache(), + md.rope_positions, + self._index_rope_cache() if md.mrope_cos_sin is None else md.mrope_cos_sin, index.q_norm_weight, index.eps, q_index.view(-1, self.index_head_dim), diff --git a/python/freetoken/core.py b/python/freetoken/core.py index 281ff41a5..ab5e52b6e 100644 --- a/python/freetoken/core.py +++ b/python/freetoken/core.py @@ -42,6 +42,14 @@ class Req: # Optional precomputed multimodal soft-token embeddings (GPU, [num_image_tokens, # hidden]) scattered at image-token positions during this request's prefill. mm_embeds: torch.Tensor | None = None + # mRoPE (qwen4_exp image prompts): ``[3, prompt_len]`` T/H/W rope positions of the prompt + # tokens (CPU) and the offset decode adds to a token index (``max_pos + 1 - prompt_len``, + # <= 0). None / 0 for text-only prompts, where rope position == token index. + mrope_positions: torch.Tensor | None = None + mrope_delta: int = 0 + # Prefix-cache key ids (image prompts: placeholder runs replaced per image content hash); + # None when input_ids are the key. Same length as input_ids, sliced with it. + cache_ids: torch.Tensor | None = None # --- hybrid-radix (GDN linear-state) per-request slots; None for non-hybrid models or # until allocated from LinearStatePool. Set by the scheduler (P2). --- @@ -134,6 +142,12 @@ class Batch: attn_metadata: BaseAttnMetadata = field(init=False) # concatenated multimodal soft-token embeddings for a prefill batch (or None) mm_embeds: torch.Tensor | None = field(default=None, init=False) + # Rope positions per token (``positions + Req.mrope_delta``); the same tensor as ``positions`` + # when no request in the batch carries an image. Set by the scheduler / graph buffer. + rope_positions: torch.Tensor | None = field(default=None, init=False) + # Prefill batches with image tokens: per-token mRoPE cos|sin rows ``[T, rotary_dim]`` (fp32) + # that the attention layers use as the rope cache with ``positions = arange(T)``. + mrope_cos_sin: torch.Tensor | None = field(default=None, init=False) # Prefill log stats snapshotted at schedule time (before forward's complete_one() # advances cached_len), so the prefill log reports the tokens actually forwarded and # the prefix-cache hit -- matching SGLang's #new-token / #cached-token. Set by the diff --git a/python/freetoken/engine/config.py b/python/freetoken/engine/config.py index 7ab792f9e..e14fd3118 100644 --- a/python/freetoken/engine/config.py +++ b/python/freetoken/engine/config.py @@ -79,7 +79,7 @@ class EngineConfig: # ratio default above. A runtime cache rebuild sets this (num_swa_pages) to pin the window # regardless of the full anchor; the ratio is the startup default and the fallback. swa_num_pages_override: int | None = None - distributed_timeout: float = 60.0 + distributed_timeout: float = 1800.0 # ranks reach the first collective minutes apart on a 100+ GiB offload load use_dummy_weight: bool = False use_pynccl: bool = True max_seq_len_override: int | None = None diff --git a/python/freetoken/engine/graph.py b/python/freetoken/engine/graph.py index 4f2025025..38944a0b4 100644 --- a/python/freetoken/engine/graph.py +++ b/python/freetoken/engine/graph.py @@ -24,6 +24,7 @@ class GraphCaptureBuffer: input_ids: torch.Tensor out_loc: torch.Tensor positions: torch.Tensor + rope_positions: torch.Tensor # positions + per-request mRoPE delta (== positions, text-only) logits: torch.Tensor table_idx: torch.Tensor # per-request slot id for GatedDeltaNet state gather/scatter # Decode GDN query indptr = arange(bs+1); a constant per captured bs, filled once. @@ -35,6 +36,7 @@ def init(cls, bs: int, vocab_size: int, device: torch.device) -> GraphCaptureBuf input_ids=torch.zeros(bs, dtype=torch.int32, device=device), out_loc=torch.zeros(bs, dtype=torch.int32, device=device), positions=torch.zeros(bs, dtype=torch.int32, device=device), + rope_positions=torch.zeros(bs, dtype=torch.int32, device=device), logits=torch.empty(bs, vocab_size, dtype=torch.float32, device=device), table_idx=torch.zeros(bs, dtype=torch.int32, device=device), fla_cu_seqlens=torch.arange(bs + 1, dtype=torch.int32, device=device), @@ -48,6 +50,7 @@ def set_batch(self, batch: Batch) -> None: batch.input_ids = self.input_ids[_slice] batch.out_loc = self.out_loc[_slice] batch.positions = self.positions[_slice] + batch.rope_positions = self.rope_positions[_slice] batch.linear_table_idx = self.table_idx[_slice] # Decode GDN metadata reads the persistent cu_seqlens (constant arange) and the # persistent table_idx slot map, so the captured kernels see stable addresses. @@ -61,6 +64,8 @@ def copy_from(self, batch: Batch) -> None: if batch.out_loc is not None: self.out_loc[_slice] = batch.out_loc self.positions[_slice] = batch.positions + rope = batch.positions if batch.rope_positions is None else batch.rope_positions + self.rope_positions[_slice] = rope if batch.linear_table_idx is not None: self.table_idx[_slice] = batch.linear_table_idx diff --git a/python/freetoken/layers/linear.py b/python/freetoken/layers/linear.py index d8a3e8619..40ebe934c 100644 --- a/python/freetoken/layers/linear.py +++ b/python/freetoken/layers/linear.py @@ -83,13 +83,17 @@ def __init__( input_size: int, output_sizes: List[int], has_bias: bool, + local_output_sizes: List[int] | None = None, *, quant_config: QuantConfig | None = None, prefix: str = "", ): - # check that all output sizes are divisible by tp_size + # check that all output sizes are divisible by tp_size (a caller that replicates + # GQA kv heads across ranks passes the per-rank sizes explicitly) tp_info = get_tp_info() - tp_output_sizes = [div_even(size, tp_info.size) for size in output_sizes] + if local_output_sizes is None: + local_output_sizes = [div_even(size, tp_info.size) for size in output_sizes] + tp_output_sizes = local_output_sizes output_size = sum(output_sizes) tp_output_size = sum(tp_output_sizes) super().__init__( diff --git a/python/freetoken/layers/quantization/moe/nvfp4.py b/python/freetoken/layers/quantization/moe/nvfp4.py index bb341cdfb..e2dc1478c 100644 --- a/python/freetoken/layers/quantization/moe/nvfp4.py +++ b/python/freetoken/layers/quantization/moe/nvfp4.py @@ -41,14 +41,20 @@ class TritonNvfp4MoEKernel(MoEKernel): cpu_format = "nvfp4" def unusable_reason(self, cfg: MoEConfig) -> str | None: - reason = self._common_reject(cfg, resident_ok=False, tp_ok=False, cpu_ok=True, plain_silu_only=False) + # tp_ok: the source stream is sliced along the intermediate axis per rank + # (nvfp4_banks._tp_shard) and this kernel sizes its banks from + # cfg.local_intermediate, so a rank holds and reads only its own half. The routed + # output is then a partial sum, which the MoE layer reduces (_maybe_all_reduce, or + # one combined all-reduce in qwen4_exp's block). marlin and b12x stay off: their + # pack() repacks the rows and neither has been checked against a sharded bank. + reason = self._common_reject(cfg, resident_ok=False, tp_ok=True, cpu_ok=True, plain_silu_only=False) if reason: return reason reason = gated_epilogue_reason(cfg) return f"triton nvfp4 MoE kernel: {reason}" if reason else None def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: - i, h = cfg.intermediate, cfg.hidden + i, h = cfg.local_intermediate, cfg.hidden return { "gate_up": BankSpec((2 * i, h // 2), torch.uint8), "gate_up_scale": BankSpec((2 * i, h // GROUP), FP8), @@ -61,7 +67,7 @@ def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: def pack(self, pieces, cfg: MoEConfig, out): out["gate_up"].copy_(fused_piece(pieces, "gate_up")) out["gate_up_scale"].copy_(fused_piece(pieces, "gate_up_scale")) - out["gate_up_global"].copy_(fused_global(pieces, cfg.intermediate)) + out["gate_up_global"].copy_(fused_global(pieces, cfg.local_intermediate)) out["down"].copy_(pieces["down"]) out["down_scale"].copy_(pieces["down_scale"]) out["down_global"].copy_(global_rows(pieces["down_global"], cfg.hidden)) @@ -241,7 +247,7 @@ def worth_it(self, cfg: MoEConfig) -> bool: return (8, 0) <= backend.device_capability() < (10, 0) def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: - i, h = cfg.intermediate, cfg.hidden + i, h = cfg.local_intermediate, cfg.hidden return { "gate_up": BankSpec((h // GROUP, 4 * i), torch.int32), "gate_up_scale": BankSpec((h // GROUP, 2 * i), FP8), @@ -252,7 +258,7 @@ def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: } def pack(self, pieces, cfg: MoEConfig, out): - i, h = cfg.intermediate, cfg.hidden + i, h = cfg.local_intermediate, cfg.hidden device = torch.device("cuda") gu, gus, gug = fused_piece(pieces, "gate_up"), fused_piece(pieces, "gate_up_scale"), fused_global(pieces, i) dn, dns, dng = pieces["down"], pieces["down_scale"], global_rows(pieces["down_global"], h) @@ -503,7 +509,7 @@ def unusable_reason(self, cfg: MoEConfig) -> str | None: def worth_it(self, cfg: MoEConfig) -> bool: # NOTE: never auto-selected. flashinfer's cute launcher indexes each bank with int32 element offsets, so the GPU slot cache is capped at (2^31 - 1) / elements-per-slot (about 1000 slots for GLM-5.3-Flash's 12960 experts); until flashinfer lifts that, b12x stays behind triton in the table; the selection does not weigh the cap, cache-auto clamps to slot_limit() and OffloadMoeCache refuses a larger cache. - return cfg.intermediate >= B12X_MIN_INTERMEDIATE + return cfg.local_intermediate >= B12X_MIN_INTERMEDIATE def slot_limit(self, cfg: MoEConfig) -> int | None: # the cute launcher indexes each bank with int32 element offsets @@ -512,7 +518,7 @@ def slot_limit(self, cfg: MoEConfig) -> int | None: def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: # flashinfer's prepared tiles: the Marlin shapes, byte-identical to the native rows - i, h = cfg.intermediate, cfg.hidden + i, h = cfg.local_intermediate, cfg.hidden return { "gate_up": BankSpec((h // GROUP, 4 * i), torch.int32), "gate_up_scale": BankSpec((h // GROUP, 2 * i), FP8), @@ -525,7 +531,7 @@ def layout(self, cfg: MoEConfig) -> dict[str, BankSpec]: def pack(self, pieces, cfg: MoEConfig, out): from flashinfer.fused_moe.cute_dsl.blackwell_sm12x.moe_w4a16_prepare import prepare_w4a16_packed_weights - i, h = cfg.intermediate, cfg.hidden + i, h = cfg.local_intermediate, cfg.hidden device = torch.device("cuda") gu = fused_piece(pieces, "gate_up").to(device) gug = fused_global(pieces, i).to(device).float() diff --git a/python/freetoken/layers/rotary.py b/python/freetoken/layers/rotary.py index 3756a299b..e55aabe4f 100644 --- a/python/freetoken/layers/rotary.py +++ b/python/freetoken/layers/rotary.py @@ -73,13 +73,16 @@ def forward( positions: torch.Tensor, query: torch.Tensor, key: torch.Tensor, + cos_sin_cache: torch.Tensor | None = None, ) -> Tuple[torch.Tensor, torch.Tensor]: + # ``cos_sin_cache`` overrides the position-indexed cache with a per-row table (mRoPE: + # row i holds token i's cos|sin, positions = arange), same [rows, rotary_dim] layout. self.apply_rope_with_cos_sin_cache_inplace( positions=positions, query=query, key=key, head_size=self.head_size, - cos_sin_cache=self._cos_sin_cache, + cos_sin_cache=self._cos_sin_cache if cos_sin_cache is None else cos_sin_cache, is_neox=self.is_neox, ) return query, key diff --git a/python/freetoken/message/backend.py b/python/freetoken/message/backend.py index c42ecc5a6..3ef3be115 100644 --- a/python/freetoken/message/backend.py +++ b/python/freetoken/message/backend.py @@ -37,6 +37,15 @@ class UserMsg(BaseBackendMsg): # Optional precomputed multimodal soft-token embeddings (GPU tensor). Only used by # the in-process offline path; remains None for the (serialized) online path. mm_embeds: torch.Tensor | None = None + # Online image path: processor outputs from the tokenizer worker (CPU fp32 ``pixel_values`` + # [patches, C*T*P*P] + ``image_grid_thw`` [N, 3]); the scheduler encodes them on its rank + # into mm_embeds / mrope_positions / mrope_delta (see Req) before admission. + mm_inputs: dict | None = None + mrope_positions: torch.Tensor | None = None + mrope_delta: int = 0 + # Prefix-cache key ids for image prompts: placeholder runs replaced per image content + # hash (tokenizer.tokenize._image_cache_ids); None for text prompts (input_ids are the key). + cache_ids: torch.Tensor | None = None @dataclass diff --git a/python/freetoken/message/tokenizer.py b/python/freetoken/message/tokenizer.py index 33b75c785..2fb9000ae 100644 --- a/python/freetoken/message/tokenizer.py +++ b/python/freetoken/message/tokenizer.py @@ -72,6 +72,8 @@ class TokenizeMsg(BaseTokenizerMsg): sampling_params: SamplingParams chat_template_kwargs: Dict[str, Any] | None = None tools: List[Dict[str, Any]] | None = None + # Encoded image files, in the order their ``{"type": "image"}`` parts appear in ``text``. + images: List[bytes] | None = None @dataclass diff --git a/python/freetoken/message/utils.py b/python/freetoken/message/utils.py index ee92adf5d..e17be3596 100644 --- a/python/freetoken/message/utils.py +++ b/python/freetoken/message/utils.py @@ -32,10 +32,11 @@ def serialize_type(self) -> Dict: serialized = {} if isinstance(self, torch.Tensor): - assert self.dim() == 1, "we can only serialize 1D tensor for now" serialized["__type__"] = "Tensor" - serialized["buffer"] = self.numpy().tobytes() + serialized["buffer"] = self.contiguous().numpy().tobytes() serialized["dtype"] = str(self.dtype) + if self.dim() != 1: # 1-D stays shape-less (wire compatible); N-D carries its shape + serialized["shape"] = list(self.shape) return serialized # normal type @@ -64,14 +65,13 @@ def _deserialize_any(cls_map: Dict[str, Type], data: Any) -> Any: def deserialize_type(cls_map: Dict[str, Type], data: Dict) -> Any: type_name = data["__type__"] - # we can only serialize 1D tensor for now if type_name == "Tensor": buffer = data["buffer"] dtype_str = data["dtype"].replace("torch.", "") np_dtype = getattr(np, dtype_str) assert isinstance(buffer, bytes) np_tensor = np.frombuffer(buffer, dtype=np_dtype) - return torch.from_numpy(np_tensor.copy()) + return torch.from_numpy(np_tensor.copy()).reshape(data.get("shape", (-1,))) cls = cls_map.get(type_name) if cls is None: diff --git a/python/freetoken/models/nvfp4_banks.py b/python/freetoken/models/nvfp4_banks.py index 1d1b42859..0feddaf2a 100644 --- a/python/freetoken/models/nvfp4_banks.py +++ b/python/freetoken/models/nvfp4_banks.py @@ -60,6 +60,36 @@ def _bank_layer(spec: Nvfp4ExpertSourceSpec, layer: int, config) -> int | None: return bank_layer +def _tp_slice() -> tuple[int, int] | None: + """``(i_local, i_lo)``: this rank's slice of the expert intermediate axis, or None at TP=1. + + TP shards every routed expert along I (gate/up output rows, down input columns), so each + rank caches only its own half and the MoE layer all-reduces the partial sum. The expert + kernel sizes its banks from ``MoEConfig.local_intermediate``, which is the same split. + """ + from freetoken.distributed import try_get_tp_info + + tp = try_get_tp_info() + if tp is None or tp.size == 1: + return None + return tp.size, tp.rank + + +def _tp_shard(role: str, tensor: torch.Tensor, inter: int, tp_size: int, tp_rank: int): + """This rank's I-slice of one checkpoint expert tensor. ``_global`` is a per-tensor scalar + and has no I axis; gate/up carry I on the row axis, down on the column axis (halved for the + packed FP4 codes, sixteenthed for the fp8 block scales).""" + if role.endswith("_global"): + return tensor + n = inter // tp_size + assert n % 16 == 0, f"NVFP4 TP shard {n} must cover whole 16-wide scale blocks" + lo = tp_rank * n + if role.startswith("down"): + d = 2 if role == "down" else 16 + return tensor[:, lo // d : (lo + n) // d] + return tensor[lo : lo + n] + + def _kind_suffix(kind: str) -> str: return {"weight": "", "weight_scale": "_scale", "weight_scale_2": "_global"}[kind] @@ -133,7 +163,19 @@ def _parallel(): tensor = _ingest_global(spec, tensor) yield name, tensor - return per_expert_pieces(_parallel() if parallel else _serial(), wanted.get, tensors_per_expert=9) + stream = _parallel() if parallel else _serial() + shard = _tp_slice() + if shard is not None: + tp_size, tp_rank = shard + inter = config.moe_intermediate_size + base = stream + + def _sharded(): + for name, tensor in base: + yield name, _tp_shard(wanted[name][2], tensor, inter, tp_size, tp_rank) + + stream = _sharded() + return per_expert_pieces(stream, wanted.get, tensors_per_expert=9) __all__ = ["Nvfp4ExpertSourceSpec", "iter_nvfp4_expert_pieces"] diff --git a/python/freetoken/models/qwen4_exp/attention.py b/python/freetoken/models/qwen4_exp/attention.py index 66211c7c2..f68df445c 100644 --- a/python/freetoken/models/qwen4_exp/attention.py +++ b/python/freetoken/models/qwen4_exp/attention.py @@ -19,9 +19,16 @@ import torch from freetoken.core import get_global_ctx -from freetoken.layers import BaseOP, GemmaPlusOneRMSNorm, LinearColParallelMerged, LinearReplicated +from freetoken.distributed import get_tp_info +from freetoken.layers import ( + BaseOP, + GemmaPlusOneRMSNorm, + LinearColParallelMerged, + LinearOProj, + LinearReplicated, +) from freetoken.layers.rotary import get_rope -from freetoken.utils import nvtx_annotate +from freetoken.utils import div_even, nvtx_annotate if TYPE_CHECKING: from freetoken.core import Batch @@ -123,12 +130,25 @@ def __init__(self, config: ModelConfig, layer_id: int, *, prefix: str = "") -> N self.head_dim = config.head_dim self.qo_attn_dim = self.num_q * self.head_dim self.kv_attn_dim = self.num_kv * self.head_dim - self._qkv_split = [self.qo_attn_dim * 2, self.kv_attn_dim, self.kv_attn_dim] + # TP: q heads split across ranks, kv heads split or (num_kv < tp) replicated; the + # indexer stays replicated so every rank selects the same blocks. + tp = get_tp_info() + self._local_num_q = div_even(self.num_q, tp.size) + self._local_num_kv = div_even(self.num_kv, tp.size, allow_replicate=True) + self._local_qo_dim = self._local_num_q * self.head_dim + self._local_kv_dim = self._local_num_kv * self.head_dim + self._qkv_split = [self._local_qo_dim * 2, self._local_kv_dim, self._local_kv_dim] self.qkv_proj = LinearColParallelMerged( - config.hidden_size, self._qkv_split, has_bias=False, + config.hidden_size, + [self.qo_attn_dim * 2, self.kv_attn_dim, self.kv_attn_dim], + has_bias=False, + local_output_sizes=self._qkv_split, quant_config=config.quant, prefix=f"{prefix}.qkv_proj", ) - self.o_proj = LinearReplicated( + # row-parallel, NOT LinearReplicated: under TP>1 qkv_proj is column-parallel, so this + # rank's attention output is its own head slice and o_proj must consume the sharded + # input dim and all-reduce the partial sums. At TP=1 the two are equivalent. + self.o_proj = LinearOProj( self.qo_attn_dim, config.hidden_size, has_bias=False, quant_config=config.quant, prefix=f"{prefix}.o_proj", ) @@ -147,21 +167,27 @@ def __init__(self, config: ModelConfig, layer_id: int, *, prefix: str = "") -> N @nvtx_annotate("QSA") def forward(self, x: torch.Tensor, batch: Batch) -> torch.Tensor: qg, k, v = self.qkv_proj.forward(x).split(self._qkv_split, dim=-1) - qg = qg.view(-1, self.num_q, self.head_dim * 2) + qg = qg.view(-1, self._local_num_q, self.head_dim * 2) q = qg[..., : self.head_dim].contiguous() - gate = qg[..., self.head_dim :].reshape(-1, self.qo_attn_dim) - k = k.contiguous().view(-1, self.num_kv, self.head_dim) + gate = qg[..., self.head_dim :].reshape(-1, self._local_qo_dim) + k = k.contiguous().view(-1, self._local_num_kv, self.head_dim) v = v.contiguous() self.q_norm.forward_inplace(q) self.k_norm.forward_inplace(k) + # mRoPE: rope_positions is positions + the request's delta (text after an image), or + # row indices into batch.mrope_cos_sin for a prefill batch that carries image tokens. + rope_pos = batch.positions if batch.rope_positions is None else batch.rope_positions q, k = self.rotary.forward( - batch.positions, q.view(-1, self.qo_attn_dim), k.view(-1, self.kv_attn_dim) + rope_pos, + q.view(-1, self._local_qo_dim), + k.view(-1, self._local_kv_dim), + cos_sin_cache=batch.mrope_cos_sin, ) index = self.indexer.forward(x) o = get_global_ctx().attn_backend.qsa_forward( - q.view(-1, self.num_q, self.head_dim), k, v, index, self.layer_id, batch + q.view(-1, self._local_num_q, self.head_dim), k, v, index, self.layer_id, batch ) - gated = o.reshape(-1, self.qo_attn_dim) * torch.sigmoid(gate) + gated = o.reshape(-1, self._local_qo_dim) * torch.sigmoid(gate) return self.o_proj.forward(gated) diff --git a/python/freetoken/models/qwen4_exp/config.py b/python/freetoken/models/qwen4_exp/config.py index 4bcd28de3..1013d9bb1 100644 --- a/python/freetoken/models/qwen4_exp/config.py +++ b/python/freetoken/models/qwen4_exp/config.py @@ -12,6 +12,7 @@ ModelConfig, RotaryConfig, SlotStateSpec, + vision_load_enabled, ) @@ -40,6 +41,10 @@ class Qwen4ExpArgs: index_head_dim: int index_budget: int index_ratio: int + # mRoPE: which rotary frequencies follow the image H / W axes (HF mrope_section, interleaved). + mrope_section: Tuple[int, ...] = (11, 11, 10) + # Vision patch merge (2 -> one soft token per 2x2 patches); only read when vision is loaded. + spatial_merge_size: int = 2 @property def index_topk_blocks(self) -> int: @@ -196,6 +201,8 @@ def parse_config(hf_config: Any) -> ModelConfig: if isinstance(eos_token_id, (list, tuple)): eos_token_id = eos_token_id[0] + # Vision is opt-in (FREETOKEN_LOAD_VISION=1): the tower is ~0.9 GiB bf16 per rank. + vision_config = getattr(hf_config, "vision_config", None) if vision_load_enabled() else None qwen4_args = Qwen4ExpArgs( hidden_size=text.hidden_size, hc_count=int(text.hc_count), @@ -214,6 +221,8 @@ def parse_config(hf_config: Any) -> ModelConfig: index_head_dim=int(text.indexer_head_dim), index_budget=int(text.indexer_budget), index_ratio=int(text.indexer_compress_ratio), + mrope_section=tuple(int(v) for v in rope_params.get("mrope_section", (11, 11, 10))), + spatial_merge_size=int(getattr(vision_config, "spatial_merge_size", 2)), ) return ModelConfig( @@ -241,7 +250,7 @@ def parse_config(hf_config: Any) -> ModelConfig: use_qk_norm=True, model_type=getattr(hf_config, "model_type", "qwen4_exp"), architectures=getattr(hf_config, "architectures", ["Qwen4ExpForConditionalGeneration"]), - vision_config=None, # served text-only + vision_config=vision_config, image_token_id=getattr(hf_config, "image_token_id", None), attention_groups=groups, expert_quant=expert_quant, diff --git a/python/freetoken/models/qwen4_exp/gdn.py b/python/freetoken/models/qwen4_exp/gdn.py index ed49dd42f..3b2e02934 100644 --- a/python/freetoken/models/qwen4_exp/gdn.py +++ b/python/freetoken/models/qwen4_exp/gdn.py @@ -3,9 +3,12 @@ import torch import torch.nn.functional as F from freetoken.core import get_global_ctx +from freetoken.distributed import get_tp_info from freetoken.kernel.causal_conv1d import causal_conv1d_decode, causal_conv1d_varlen -from freetoken.layers import BaseOP, GatedRMSNorm, LinearColParallelMerged, LinearReplicated +from freetoken.distributed import get_tp_info +from freetoken.layers import BaseOP, GatedRMSNorm, LinearColParallelMerged, LinearOProj from freetoken.layers.quantization import QuantConfig +from freetoken.utils import div_even from freetoken.models.qwen3_5_moe.gdn_kernels import gdn_decode_fla, gdn_prefill_chunk_fla @@ -50,12 +53,29 @@ def __init__( self.value_dim = num_v_heads * head_v_dim self.conv_dim = 2 * self.key_dim + self.value_dim self.conv_kernel_size = conv_kernel_size + # TP-local head counts: k heads and their v-head groups split evenly across ranks + # (the fla kernels take the GQA ratio from the shapes); the state pool is sharded the + # same way (kvcache.linear_state_pool._linear_local_dims). + tp = get_tp_info() + self._local_num_k_heads = div_even(num_k_heads, tp.size, allow_replicate=True) + self._local_num_v_heads = div_even(num_v_heads, tp.size, allow_replicate=True) + self._local_key_dim = self._local_num_k_heads * head_k_dim + self._local_value_dim = self._local_num_v_heads * head_v_dim + self._local_conv_dim = 2 * self._local_key_dim + self._local_value_dim # quantized checkpoints quantize qkv|z but not b|a, so the fusion splits into a qkvz GEMM and a ba GEMM with their own schemes (matches sglang / vLLM) self._split_in_proj = ( quant_config is not None and quant_config.scheme_for(f"{prefix}.in_proj_qkvz") is not None ) - self._in_proj_split = [self.conv_dim, self.value_dim, num_v_heads, num_v_heads] + self._in_proj_split = [ + self._local_conv_dim, self._local_value_dim, + self._local_num_v_heads, self._local_num_v_heads, + ] + # The split (quantized) branch has no TP-aware variant: its two linears are built from + # the checkpoint's own scheme and are not sharded here. + assert not (tp.size > 1 and self._split_in_proj), ( + "qwen4_exp TP shards bf16 GDN projections only" + ) if self._split_in_proj: self.in_proj_qkvz = LinearColParallelMerged( hidden_size, [self.conv_dim, self.value_dim], has_bias=False, @@ -68,18 +88,23 @@ def __init__( else: # Fused input projection (one GEMM instead of four): qkv | z | b | a. self.in_proj = LinearColParallelMerged( - hidden_size, self._in_proj_split, has_bias=False, + hidden_size, [self.conv_dim, self.value_dim, num_v_heads, num_v_heads], + has_bias=False, local_output_sizes=self._in_proj_split, quant_config=quant_config, prefix=f"{prefix}.in_proj", ) - self.conv1d = _DepthwiseConv1d(self.conv_dim, conv_kernel_size) + self.conv1d = _DepthwiseConv1d(self._local_conv_dim, conv_kernel_size) # Recurrence-gating params kept in fp32 (exp/softplus is precision-sensitive, # and the fla kernel reads them as fp32) -- matches HF/sglang, and avoids a # per-call .float() upcast in the decode wrapper. The weight loader exempts # *.A_log / *.dt_bias from the model-dtype downcast. - self.dt_bias = torch.empty(num_v_heads, dtype=torch.float32) - self.A_log = torch.empty(num_v_heads, dtype=torch.float32) + self.dt_bias = torch.empty(self._local_num_v_heads, dtype=torch.float32) + self.A_log = torch.empty(self._local_num_v_heads, dtype=torch.float32) self.norm = GatedRMSNorm(head_v_dim, eps=rms_norm_eps, activation=output_gate) - self.out_proj = LinearReplicated( + # Row-parallel over the local v heads, all-reduce inside. At TP=1 this is exactly + # LinearReplicated (div_even(x, 1) == x and the reduction is skipped), so one class + # covers both -- and under TP>1 a replicated o_proj would silently take the wrong + # input width and skip the reduction. + self.out_proj = LinearOProj( self.value_dim, hidden_size, has_bias=False, quant_config=quant_config, prefix=f"{prefix}.out_proj", ) @@ -148,7 +173,9 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: else: proj = self.in_proj.forward(hidden_states) conv_in, z, b, a = torch.split(proj, self._in_proj_split, dim=-1) - z = z.reshape(total, self.num_v_heads, self.head_v_dim) + nk, nv = self._local_num_k_heads, self._local_num_v_heads + kd, vd = self._local_key_dim, self._local_value_dim + z = z.reshape(total, nv, self.head_v_dim) li = pool.local_index(self.layer_id) if batch.is_decode: @@ -157,10 +184,10 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: # no clone, no external l2norm). q/k stay at num_k_heads (kernel handles GQA). mixed = self._conv_decode(conv_in, fla.cache_indices, pool) # [B, conv_dim] B = mixed.shape[0] - qf, kf, vf = torch.split(mixed, [self.key_dim, self.key_dim, self.value_dim], dim=-1) - q = qf.reshape(1, B, self.num_k_heads, self.head_k_dim).to(dtype) - k = kf.reshape(1, B, self.num_k_heads, self.head_k_dim).to(dtype) - v = vf.reshape(1, B, self.num_v_heads, self.head_v_dim).to(dtype) + qf, kf, vf = torch.split(mixed, [kd, kd, vd], dim=-1) + q = qf.reshape(1, B, nk, self.head_k_dim).to(dtype) + k = kf.reshape(1, B, nk, self.head_k_dim).to(dtype) + v = vf.reshape(1, B, nv, self.head_v_dim).to(dtype) core_out = gdn_decode_fla( q, k, v, a, b, A_log=self.A_log, dt_bias=self.dt_bias, state_source=pool.recurrent_states[li], indices=fla.cache_indices, @@ -170,13 +197,13 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: mixed = self._conv_prefill( conv_in, pool, fla.cu_seqlens, fla.cache_indices, fla.has_initial_state) # fla chunk handles GQA in-kernel: q/k stay at num_k_heads, v at num_v_heads. - qf, kf, vf = torch.split(mixed, [self.key_dim, self.key_dim, self.value_dim], dim=-1) - q = qf.reshape(1, total, self.num_k_heads, self.head_k_dim).to(dtype) - k = kf.reshape(1, total, self.num_k_heads, self.head_k_dim).to(dtype) - v = vf.reshape(1, total, self.num_v_heads, self.head_v_dim).to(dtype) + qf, kf, vf = torch.split(mixed, [kd, kd, vd], dim=-1) + q = qf.reshape(1, total, nk, self.head_k_dim).to(dtype) + k = kf.reshape(1, total, nk, self.head_k_dim).to(dtype) + v = vf.reshape(1, total, nv, self.head_v_dim).to(dtype) g, beta = self._gate_params(a, b) - g = g.reshape(1, total, self.num_v_heads) - beta = beta.float().reshape(1, total, self.num_v_heads) + g = g.reshape(1, total, nv) + beta = beta.float().reshape(1, total, nv) # The chunk kernel reads + writes back initial_state[cache_indices] in place; # fresh sequences (cached_len==0) must start from a zeroed slot. if fla.fresh_state_indices is not None: diff --git a/python/freetoken/models/qwen4_exp/model.py b/python/freetoken/models/qwen4_exp/model.py index a4b63362d..14c47e8d5 100644 --- a/python/freetoken/models/qwen4_exp/model.py +++ b/python/freetoken/models/qwen4_exp/model.py @@ -1,4 +1,4 @@ -"""Qwen3.8-Flash-Next decoder stack (text-only). +"""Qwen3.8-Flash-Next decoder stack (text, plus images when the vision tower is loaded). The residual state is ``R [T, hc_count*hidden]`` end to end: the embedding is repeated over the ``hc_count`` streams, every layer mixes them down to one ``[T, hidden]`` block input and injects @@ -100,6 +100,7 @@ def __init__(self, config: ModelConfig, *, prefix: str = "model") -> None: self.hyper_connection_mixer = GatedResidual(config, use_combine=False, prefix=f"{prefix}.hyper_connection_mixer") # plain tuple (not an OP child), so it never shows up in the state dict self._ple = tuple(layer.ple for layer in self.layers.op_list if layer.ple is not None) + self._image_token_id = config.image_token_id @property def ple_layers(self) -> List[PLELayer]: @@ -107,7 +108,14 @@ def ple_layers(self) -> List[PLELayer]: return list(self._ple) def forward(self, input_ids: torch.Tensor, batch: Batch) -> torch.Tensor: - hidden = self.embed_tokens.forward(input_ids).repeat(1, self.hc_count) + hidden = self.embed_tokens.forward(input_ids) + if batch.mm_embeds is not None: + # image soft tokens replace the placeholder embeddings (HF order: before the + # hc_count repeat); a chunked prompt gets the rows of the placeholders inside + # this chunk (scheduler._gather_multimodal) + mask = input_ids == self._image_token_id + hidden = hidden.masked_scatter(mask.unsqueeze(-1), batch.mm_embeds.to(hidden.dtype)) + hidden = hidden.repeat(1, self.hc_count) meta = None if self._ple: from .ple import build_ple_metadata, commit_ngram_context @@ -127,7 +135,10 @@ def forward(self, input_ids: torch.Tensor, batch: Batch) -> torch.Tensor: class Qwen4ExpForCausalLM(BaseLLMModel): def __init__(self, config: ModelConfig) -> None: self._config = config + self._inv_freq: torch.Tensor | None = None self.model = Qwen4ExpModel(config) + # lm_head quantization now comes from the checkpoint's QuantConfig scheme, so the + # old lm_head_quant == "nvfp4" special case is gone. self.lm_head = ParallelLMHead( num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, @@ -136,8 +147,76 @@ def __init__(self, config: ModelConfig) -> None: quant_config=config.quant, prefix="lm_head", ) + if config.vision_config is not None: # FREETOKEN_LOAD_VISION=1 + from .vision import Qwen4ExpVisionTower + + self.visual = Qwen4ExpVisionTower(config.vision_config) super().__init__() + @property + def has_vision(self) -> bool: + return hasattr(self, "visual") + + @property + def image_token_id(self) -> int: + return self._config.image_token_id + + @torch.inference_mode() + def encode_images(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: + """Vision tower + merger on processor outputs: ``[num_image_tokens, hidden]`` (device).""" + if not self.has_vision: + raise RuntimeError("image inputs need the vision tower: start with FREETOKEN_LOAD_VISION=1") + return self.visual.forward(pixel_values, grid_thw) + + def prepare_mm_inputs( + self, input_ids: torch.Tensor, mm_inputs: dict + ) -> tuple[torch.Tensor, torch.Tensor, int]: + """One prompt's ``(mm_embeds [n, hidden] on device, mrope_positions [3, L] CPU, mrope_delta)``.""" + from .mrope import rope_index + + grid = mm_inputs["image_grid_thw"] + embeds = self.encode_images(mm_inputs["pixel_values"], grid) + image_token_id = self._config.image_token_id + n = int((input_ids == image_token_id).sum()) + if n != embeds.shape[0]: + raise ValueError(f"{n} image placeholder tokens for {embeds.shape[0]} image features") + pos, delta = rope_index( + input_ids, grid, image_token_id, self._config.qwen4_args.spatial_merge_size + ) + return embeds, pos, delta + + def mrope_table(self, reqs, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: + """``(rope_positions, mrope_cos_sin)`` for a prefill batch with image tokens; see + :func:`mrope.mrope_table`.""" + from .mrope import mrope_table + + return mrope_table( + reqs, + self._config.qwen4_args.index_ratio, + self._inv_freq_on(device), + self._config.qwen4_args.mrope_section, + device, + ) + + def _inv_freq_on(self, device: torch.device) -> torch.Tensor: + """Attention rope frequencies (fp32, the RotaryEmbedding formula) on ``device``.""" + if self._inv_freq is None or self._inv_freq.device != device: + rc = self._config.rotary_config + self._inv_freq = 1.0 / ( + rc.base + ** (torch.arange(0, rc.rotary_dim, 2, dtype=torch.float, device=device) / rc.rotary_dim) + ) + return self._inv_freq + + def mrope_cos_sin(self, positions: torch.Tensor) -> torch.Tensor: + """``[T, rotary_dim]`` fp32 cos|sin rows for 3-D ``positions [3, T]`` (same frequencies as + the attention rope cache, so text rows equal the cache rows bit for bit).""" + from .mrope import mrope_cos_sin + + return mrope_cos_sin( + positions, self._inv_freq_on(positions.device), self._config.qwen4_args.mrope_section + ) + def load_host_tables(self, engine_config) -> int: """Attach the PLE n-gram table (pinned checkpoint bank, or zeros for dummy weights); returns the pinned host bytes the engine reserves from its pin budget.""" ple_layers = self.model.ple_layers diff --git a/python/freetoken/models/qwen4_exp/moe.py b/python/freetoken/models/qwen4_exp/moe.py index 3687f2803..c10407861 100644 --- a/python/freetoken/models/qwen4_exp/moe.py +++ b/python/freetoken/models/qwen4_exp/moe.py @@ -3,7 +3,12 @@ from typing import TYPE_CHECKING import torch +import torch.nn.functional as F +from freetoken.core import get_global_ctx +from freetoken.distributed import DistributedCommunicator, get_tp_info from freetoken.kernel.triton.moe_shared_gate import shared_gate_mul_add, shared_gate_sigmoid +from freetoken.layers import LinearRowParallel, silu_and_mul +from freetoken.layers.moe import OffloadMoELayer from freetoken.models.qwen3_5_moe.moe import Qwen3_5MoE if TYPE_CHECKING: @@ -14,15 +19,39 @@ class Qwen4ExpMoE(Qwen3_5MoE): """Qwen3_5MoE with the shared-expert gate on triton instead of gemv + sigmoid + mul + add. Same weights, same state dict. The gate reduction stays ahead of the routed experts, which may write into ``hidden_states`` in place. + + TP: the offload experts are sharded along the intermediate axis and the bf16 shared expert is + row-parallel, so both produce partial sums; ``routed + gate * shared`` is linear in them and + is reduced once (one all-reduce per MoE layer instead of two). """ + def __init__( + self, config: ModelConfig, layer_id: int | None = None, *, prefix: str = "" + ) -> None: + super().__init__(config, layer_id=layer_id, prefix=prefix) + self._comm = DistributedCommunicator() + self._tp_size = get_tp_info().size + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: num_tokens, hidden_dim = hidden_states.shape hidden_states = hidden_states.view(-1, hidden_dim) router_logits = self.gate.forward(hidden_states) - shared = self.shared_expert.forward(hidden_states) gate = shared_gate_sigmoid(hidden_states, self.shared_expert_gate.weight.view(-1)) - routed = self.experts.forward(hidden_states=hidden_states, router_logits=router_logits) + se, ex = self.shared_expert, self.experts + if ( + self._tp_size > 1 + and isinstance(se.down_proj, LinearRowParallel) + and isinstance(ex, OffloadMoELayer) + ): + shared = F.linear(silu_and_mul(se.gate_up_proj.forward(hidden_states)), se.down_proj.weight) + if get_global_ctx().batch.is_prefill: + routed = ex.prefill_forward(hidden_states, router_logits) + else: + routed = ex.decode_forward(hidden_states, router_logits) + out = self._comm.all_reduce(shared_gate_mul_add(routed, shared, gate)) + return out.view(num_tokens, hidden_dim) + shared = se.forward(hidden_states) + routed = ex.forward(hidden_states=hidden_states, router_logits=router_logits) return shared_gate_mul_add(routed, shared, gate).view(num_tokens, hidden_dim) diff --git a/python/freetoken/models/qwen4_exp/mrope.py b/python/freetoken/models/qwen4_exp/mrope.py new file mode 100644 index 000000000..a9b9167f7 --- /dev/null +++ b/python/freetoken/models/qwen4_exp/mrope.py @@ -0,0 +1,112 @@ +"""mRoPE for Qwen3.8-Flash-Next image prompts. + +Ports two pieces of HF ``modeling_qwen4_exp``: ``get_rope_index`` (the 3-D T/H/W positions of a +prompt and the delta the text after the last image continues from) and +``Qwen4ExpTextRotaryEmbedding.apply_interleaved_mrope`` (the per-token cos|sin row). Text-only +prompts reduce to ``positions`` on all three axes and a zero delta, so the serving path only +builds a table for prefill batches that carry images; decode reads the normal cache at +``position + delta``. + +The cos|sin table has the same ``[rows, rotary_dim]`` layout as ``RotaryEmbedding._cos_sin_cache`` +(cos on the first half, sin on the second), so the flashinfer/triton rope kernels and the QSA +indexer's ``qsa_index_norm_rope`` take it as the cache with ``positions = arange(rows)``. +""" + +from __future__ import annotations + +import torch + + +def rope_index( + input_ids: torch.Tensor, + grid_thw: torch.Tensor | None, + image_token_id: int, + merge: int, +) -> tuple[torch.Tensor, int]: + """``([3, L] int64 positions, delta)`` for one prompt (HF ``get_rope_index``, images only). + + Text runs advance all three axes by their length. An image of ``t x h x w`` patches takes + ``t * (h // merge) * (w // merge)`` placeholder tokens at T/H/W offsets from the current + position and advances it by ``max(h, w) // merge``. ``delta`` is what decode adds to a + token index to get its rope position (``max_pos + 1 - L``, <= 0). + """ + ids = input_ids.tolist() + length = len(ids) + out = torch.empty(3, length, dtype=torch.int64) + cur = i = img = 0 + while i < length: + if ids[i] == image_token_id: + t, h, w = (int(v) for v in grid_thw[img]) + img += 1 + hh, ww = h // merge, w // merge + n = t * hh * ww + assert ids[i : i + n] == [image_token_id] * n, ( + "image placeholder run is too short" + ) + tt, hp, wp = torch.meshgrid( + torch.arange(t), torch.arange(hh), torch.arange(ww), indexing="ij" + ) + out[0, i : i + n] = tt.reshape(-1) + cur + out[1, i : i + n] = hp.reshape(-1) + cur + out[2, i : i + n] = wp.reshape(-1) + cur + cur += max(hh, ww) + i += n + continue + j = i + while j < length and ids[j] != image_token_id: + j += 1 + out[:, i:j] = torch.arange(cur, cur + (j - i)) + cur += j - i + i = j + assert grid_thw is None or img == len(grid_thw), "more images than placeholder runs" + return out, int(out.max()) + 1 - length + + +def mrope_cos_sin( + positions: torch.Tensor, inv_freq: torch.Tensor, section: tuple[int, ...] +) -> torch.Tensor: + """``[T, 2 * len(inv_freq)]`` fp32 cos|sin rows for 3-D ``positions [3, T]``. + + Interleaved layout (HF ``apply_interleaved_mrope``): frequency ``k`` rotates by the T axis, + except ``k = 1 + 3m`` (m < section[1]) which use H and ``k = 2 + 3m`` (m < section[2]) W. + """ + freqs = positions.to(inv_freq.dtype).unsqueeze(-1) * inv_freq # [3, T, n] + f = freqs[0].clone() + for axis in (1, 2): + idx = slice(axis, section[axis] * 3, 3) + f[:, idx] = freqs[axis][:, idx] + return torch.cat((f.cos(), f.sin()), dim=-1) + + +def mrope_table( + reqs, + ratio: int, + inv_freq: torch.Tensor, + section: tuple[int, ...], + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Rope rows for a prefill batch that carries image tokens: ``(rope_positions [T] int32 = each + token's row, table [rows, rotary_dim] fp32 cos|sin)``. + + Every request also gets rows for up to ``ratio - 1`` positions before its chunk: the QSA + indexer ropes a pooled key at its group's FIRST token, which for a chunked prompt may come + from the previous chunk. Requests without + ``mrope_positions`` use their token index on all three axes, so their rows equal the rope + cache rows exactly. + """ + segments, rows, base = [], [], 0 + for r in reqs: + pre = min(r.cached_len, ratio - 1) + start = r.cached_len - pre + if getattr(r, "mrope_positions", None) is not None: + seg = r.mrope_positions[:, start : r.device_len] + else: + seg = torch.arange(start, r.device_len).expand(3, -1) + segments.append(seg) + rows.append(torch.arange(base + pre, base + seg.shape[1], dtype=torch.int32)) + base += seg.shape[1] + table = mrope_cos_sin(torch.cat(segments, dim=1).to(device), inv_freq, section) + return torch.cat(rows).to(device), table + + +__all__ = ["mrope_cos_sin", "mrope_table", "rope_index"] diff --git a/python/freetoken/models/qwen4_exp/vision.py b/python/freetoken/models/qwen4_exp/vision.py new file mode 100644 index 000000000..23c1c47a7 --- /dev/null +++ b/python/freetoken/models/qwen4_exp/vision.py @@ -0,0 +1,80 @@ +"""Qwen3.8-Flash-Next vision tower: HF ``Qwen4ExpVisionModel`` behind the BaseOP state-dict contract. + +The tower is bf16, ~0.9 GiB, never quantized and identical on every TP rank (each rank encodes +its own copy of the request's images; the soft tokens are tiny next to the tower). Its tensors +travel under the ``visual.*`` keys of the model state dict so the engine loads them with the +dense weights: counted in the weight bytes before the expert-cache planner sizes the cache, +dummy-weight builds fill them like any other tensor, and TP ranks all see the same keys. + +Meta-device trap: ``create_model`` runs under ``torch.device("meta")``. The HF module built +there is shapes only; ``load_state_dict`` rebuilds a meta module, assigns the loaded tensors +(no second copy on the GPU) and re-creates the non-persistent rotary ``inv_freq`` buffer on the +device, which ``assign=True`` would otherwise leave on meta. +""" + +from __future__ import annotations + +from typing import Any + +import torch +from freetoken.layers import BaseOP +from freetoken.layers.base import _concat_prefix + + +def _vision_model_cls(): + from transformers.models.qwen4_exp.modeling_qwen4_exp import Qwen4ExpVisionModel + + return Qwen4ExpVisionModel + + +class Qwen4ExpVisionTower(BaseOP): + def __init__(self, vision_config: Any) -> None: + # The standalone vision config carries no attention backend; sdpa splits packed images + # by cu_seqlens, so several images per call are fine. + vision_config._attn_implementation = "sdpa" + self._config = vision_config + self._hf = _vision_model_cls()(vision_config) + self.spatial_merge_size = int(vision_config.spatial_merge_size) + + def state_dict(self, *, prefix: str = "", result=None): + result = result if result is not None else {} + for name, tensor in self._hf.state_dict().items(): + result[_concat_prefix(prefix, name)] = tensor + return result + + def load_state_dict( + self, state_dict, *, prefix: str = "", _internal: bool = False + ) -> None: + own = { + k: state_dict.pop(_concat_prefix(prefix, k)) for k in self._hf.state_dict() + } + device = next(iter(own.values())).device + with torch.device("meta"): + hf = _vision_model_cls()(self._config) + hf.load_state_dict(own, strict=True, assign=True) + head_dim = self._config.hidden_size // self._config.num_heads + with torch.device(device): + hf.rotary_pos_emb = type(hf.rotary_pos_emb)(head_dim // 2) + stuck = [n for n, b in hf.named_buffers() if b.is_meta] + assert not stuck, f"vision buffers left on meta: {stuck}" + self._hf = hf.eval().requires_grad_(False) + if not _internal and state_dict: + raise RuntimeError( + f"Unexpected keys in state_dict: {list(state_dict.keys())}" + ) + + @torch.inference_mode() + def forward( + self, pixel_values: torch.Tensor, grid_thw: torch.Tensor + ) -> torch.Tensor: + """``[sum(t*h*w) / merge**2, text_hidden]`` soft tokens for the packed images.""" + param = next(self._hf.parameters()) + out = self._hf( + pixel_values.to(param.device, param.dtype), + grid_thw=grid_thw.to(param.device), + return_dict=True, + ) + return out.pooler_output + + +__all__ = ["Qwen4ExpVisionTower"] diff --git a/python/freetoken/models/qwen4_exp/weight.py b/python/freetoken/models/qwen4_exp/weight.py index 2ce7462b8..e75255f13 100644 --- a/python/freetoken/models/qwen4_exp/weight.py +++ b/python/freetoken/models/qwen4_exp/weight.py @@ -6,7 +6,7 @@ * :func:`load_ple_table` -- the 47.7 GiB FP8 n-gram table, 128 checkpoint shards concatenated into one pinned :class:`HostBank`. * :func:`nvfp4_expert_spec` -- how the routed NVFP4 experts are named, for the offload cache's expert reader. -Dropped: ``mtp.*`` (speculative head, including its stacked ``mtp.layers.0.mlp.experts.*``) and ``model.visual.*`` (served text-only). +Dropped: ``mtp.*`` (speculative head, including its stacked ``mtp.layers.0.mlp.experts.*``); ``model.visual.*`` is dropped unless vision is opted in (``FREETOKEN_LOAD_VISION=1``), then it loads as ``visual.*``. """ from __future__ import annotations @@ -20,13 +20,14 @@ import safetensors import torch +from freetoken.models.config import vision_load_enabled from freetoken.distributed import get_tp_info -from freetoken.models.loader import drop_page_cache, iter_weight_files +from freetoken.models.loader import drop_page_cache, iter_weight_files, shard_tensor from freetoken.models.nvfp4_banks import ( Nvfp4ExpertSourceSpec, ) from freetoken.moe.host_banks import HostBank, read_range_into -from freetoken.utils import download_hf_weight +from freetoken.utils import cached_load_hf_config, div_even, download_hf_weight from freetoken.utils.progress import byte_bar from tqdm import tqdm @@ -100,8 +101,13 @@ def _rename(raw_name: str) -> str | None: """Checkpoint key -> FreeToken state-dict key, or None to skip.""" - if raw_name.startswith(("mtp.", "model.visual.", "visual.")): + if raw_name.startswith("mtp."): return None + if raw_name.startswith(("model.visual.", "visual.")): + # the tower's HF module keys, under the model's ``visual`` op (opt-in, else dropped) + if not vision_load_enabled(): + return None + return "visual." + raw_name.split("visual.", 1)[1] if _PLE_TABLE_INFIX in raw_name: return None # n-gram table + its scale: load_ple_table if _EXPERT_RE.search(raw_name): @@ -137,6 +143,58 @@ def _try_fuse( return None +def _shard_rows( + t: torch.Tensor, parts: list[tuple[int, int]], rank: int, world: int +) -> torch.Tensor: + """Column-parallel slice (dim 0) of a ``[part0 | part1 | ...]`` fusion; ``parts`` gives each + part as ``(heads, rows_per_head)``. Heads split evenly across ranks; a part with fewer heads + than ranks (GQA kv) replicates head ``rank * heads // world``, the ``div_even(..., + allow_replicate=True)`` convention of the TP-aware layers.""" + out, off = [], 0 + for heads, rows in parts: + local = div_even(heads, world, allow_replicate=True) + first = rank * heads // world + out.append(t[off + first * rows : off + (first + local) * rows]) + off += heads * rows + assert off == t.shape[0], f"fusion parts {parts} cover {off} rows, tensor has {t.shape[0]}" + return torch.cat(out, dim=0) + + +def _shard(name: str, t: torch.Tensor, config, rank: int, world: int) -> torch.Tensor: + """TP shard of one state-dict tensor (fused projections included); identity at TP=1. + + Column-parallel (dim 0, by head): attention ``qkv_proj`` [q|gate per head | k | v], GDN + ``in_proj`` [q | k | v | z | b | a] and the matching ``conv1d`` channels, ``A_log`` / + ``dt_bias``, shared-expert ``gate_up_proj``. Row-parallel (dim 1): ``o_proj``, + ``out_proj``, shared-expert ``down_proj``. Vocab rows: ``embed_tokens`` / ``lm_head``. + Everything else (router, indexer, norms, HC, PLE, shared-expert gate) is replicated. + """ + if world == 1 or name.startswith("visual."): + return t + if name.endswith(".self_attn.qkv_proj.weight"): + q = (config.num_qo_heads, 2 * config.head_dim) + kv = (config.num_kv_heads, config.head_dim) + return _shard_rows(t, [q, kv, kv], rank, world) + if ".linear_attn." in name: + g = config.linear_attention_group() + k = (g.num_key_heads, g.key_head_dim) + v = (g.num_value_heads, g.value_head_dim) + if name.endswith(".in_proj.weight"): + return _shard_rows(t, [k, k, v, v, (v[0], 1), (v[0], 1)], rank, world) + if name.endswith(".conv1d.weight"): + return _shard_rows(t, [k, k, v], rank, world) + if name.endswith((".A_log", ".dt_bias")): + return _shard_rows(t, [(v[0], 1)], rank, world) + if name.endswith(".out_proj.weight"): + return t.chunk(world, dim=1)[rank].clone() + return t + if name.endswith(".shared_expert.gate_up_proj.weight"): + half = t.shape[0] // 2 + return _shard_rows(t, [(half, 1), (half, 1)], rank, world) + # o_proj / down_proj: dim 1; embed_tokens / lm_head: vocab rows; others unchanged. + return shard_tensor(name, t, rank=rank, world_size=world, num_kv_heads=None) + + def iter_weights( model_path: str, device: torch.device, @@ -158,16 +216,20 @@ def iter_weights( ``include_moe_experts`` is accepted for the loader contract but never yields anything: the routed experts are NVFP4 and always come from the offload cache's expert reader. """ - if get_tp_info().size > 1: - raise NotImplementedError("qwen4_exp weight loading supports TP=1 only") if not include_non_moe: return + from .config import parse_config + + tp = get_tp_info() + # The sharding geometry needs the HF config; TP=1 never shards, so keep the plain path free + # of a config load (synthetic test checkpoints carry no model_type). + config = parse_config(cached_load_hf_config(model_path)) if tp.size > 1 else None fuse_buf: dict[str, dict[int, torch.Tensor]] = {} for file in tqdm( iter_weight_files(model_path), desc="Loading weights", - disable=not get_tp_info().is_primary(), + disable=not tp.is_primary(), ): with safetensors.safe_open(file, framework="pt", device=str(device)) as f: for raw_name in f.keys(): @@ -178,9 +240,10 @@ def iter_weights( fused = _try_fuse(name, tensor, fuse_buf) if fused is not None: if fused != (): # () means buffered, not yet complete - yield fused + name, tensor = fused + yield name, _shard(name, tensor, config, tp.rank, tp.size) continue - yield name, tensor + yield name, _shard(name, tensor, config, tp.rank, tp.size) assert not fuse_buf, f"Incomplete projection fusions: {sorted(fuse_buf)}" diff --git a/python/freetoken/scheduler/cache.py b/python/freetoken/scheduler/cache.py index 44adde42f..91db763ce 100644 --- a/python/freetoken/scheduler/cache.py +++ b/python/freetoken/scheduler/cache.py @@ -29,6 +29,13 @@ def _swa_eviction_interval() -> int: _SWA_RETAIN_GAP = 16 +def _key_ids(req) -> torch.Tensor: + """The ids the prefix cache keys a request on: ``cache_ids`` when set (image prompts, the + placeholder runs replaced per image content hash), else ``input_ids``.""" + ids = getattr(req, "cache_ids", None) + return req.input_ids if ids is None else ids + + class CacheManager: def __init__(self, num_pages: int, page_size: int, page_table: torch.Tensor, type: str, linear_state_pool=None, swa_pool=None, sliding_window_size=None): @@ -93,10 +100,10 @@ def _make_prefix_cache(self, device, page_size, type): def match_req(self, req: PendingReq) -> MatchResult: input_len = req.input_len assert input_len > 0, "Input length must be greater than 0." - # Multimodal requests must not reuse a shared prefix: image-placeholder tokens - # have identical ids across images but carry different content (and KV), so a - # match would serve the wrong image's KV. Match against the empty prefix. - ids = req.input_ids[:0] if req.mm_embeds is not None else req.input_ids[: input_len - 1] + # Image prompts key on cache_ids (placeholder runs per image content hash): the raw + # placeholder ids are shared across images, so keyed on them a hit could serve another + # image's KV. + ids = _key_ids(req)[: input_len - 1] if self.is_swa: from freetoken.kvcache.swa_radix_cache import SWACacheHandle m = self.prefix_cache.match_prefix(ids) @@ -299,18 +306,7 @@ def cache_req(self, req: Req, *, finished: bool) -> None: # We should free it if the request has finished. page_indices = self.page_table[req.table_idx, : req.cached_len] old_handle = req.cache_handle - # Multimodal requests are never inserted into the shared prefix cache (see - # ``match_req``). Their KV pages stay owned by the active request and are freed - # on completion; nothing is exposed for cross-request reuse. - if req.mm_embeds is not None: - self.unlock(old_handle) - if finished: - tail = self._padded_tail(req, old_handle.cached_len) - if self.swa_paged: - self._free_swa(tail) - self._free(tail) - return - insert_ids = req.input_ids[: req.cached_len] + insert_ids = _key_ids(req)[: req.cached_len] cached_len, new_handle = self.prefix_cache.insert_prefix(insert_ids, page_indices) # unlock until all operations on handle is done self.unlock(old_handle) @@ -350,13 +346,6 @@ def _cache_req_hybrid(self, req: Req, *, finished: bool) -> None: old_handle = req.cache_handle page_indices = self.page_table[req.table_idx, : req.cached_len] - if req.mm_embeds is not None: - self.unlock(old_handle) - if finished: - self._free(page_indices[old_handle.cached_len :]) - self._free_req_slots(req) - return - if finished: # A pending freeze (the tool-call anchor, or a prefill ×64 track the request # finished too early to chunk-commit) is a strictly shorter prefix than the live @@ -376,7 +365,7 @@ def _cache_req_hybrid(self, req: Req, *, finished: bool) -> None: frozen_idx = 1 - req.mamba_next_track_idx frozen = req.mamba_ping_pong[frozen_idx] prefix_len, mamba_exist = self.prefix_cache.insert( - req.input_ids[:L], page_indices[:L], frozen) + _key_ids(req)[:L], page_indices[:L], frozen) pool.free([s for s in req.mamba_ping_pong if mamba_exist or s != frozen]) req.mamba_ping_pong = None self._free(page_indices[free_upto : max(free_upto, prefix_len)]) @@ -390,7 +379,7 @@ def _cache_req_hybrid(self, req: Req, *, finished: bool) -> None: keep_live = False if insert_len == req.cached_len and insert_len > 0: prefix_len, mamba_exist = self.prefix_cache.insert( - req.input_ids[:insert_len], page_indices[:insert_len], req.linear_slot_idx) + _key_ids(req)[:insert_len], page_indices[:insert_len], req.linear_slot_idx) self.unlock(old_handle) self._free(page_indices[free_upto : max(free_upto, prefix_len)]) keep_live = not mamba_exist # tree now owns linear_slot_idx @@ -413,13 +402,13 @@ def _cache_req_hybrid(self, req: Req, *, finished: bool) -> None: frozen_idx = 1 - req.mamba_next_track_idx # the slot the forward just wrote frozen = req.mamba_ping_pong[frozen_idx] prefix_len, mamba_exist = self.prefix_cache.insert( - req.input_ids[:L], page_indices[:L], frozen) + _key_ids(req)[:L], page_indices[:L], frozen) self.unlock(old_handle) self._free(page_indices[old_handle.cached_len : prefix_len]) # Lock the committed snapshot node FIRST: the replacement-slot alloc below can trigger # evict_mamba (via ensure_mamba_slots), which would otherwise reclaim this still-unlocked # just-donated node -- freeing its KV pages under the still-decoding request. - m = self.prefix_cache.match_prefix(req.input_ids[:L]) + m = self.prefix_cache.match_prefix(_key_ids(req)[:L]) # Same re-point as the generic path: the dedup free above returned this request's own # pages for [old_handle.cached_len, prefix_len) while its row still named them. if prefix_len > old_handle.cached_len: @@ -444,14 +433,6 @@ def _cache_req_swa(self, req: Req, *, finished: bool) -> None: old_handle = req.cache_handle page_indices = self.page_table[req.table_idx, : req.cached_len] - if req.mm_embeds is not None: - self.unlock(old_handle) - if finished: - tail = self._padded_tail(req, old_handle.cached_len) - self._free_swa(tail) - self._free(tail) - return - insert_len = align_down(req.cached_len, self.page_size) freed = page_indices[:0] if insert_len > 0: @@ -464,7 +445,7 @@ def _cache_req_swa(self, req: Req, *, finished: bool) -> None: # unfinished chunk's frontier is already > 0 and must be honored (else insert adopts # sentinel slots -> the request's later SWA gathers read slot 0 -> corruption). _, freed = self.prefix_cache.insert( - req.input_ids[:insert_len], page_indices[:insert_len], + _key_ids(req)[:insert_len], page_indices[:insert_len], swa_evicted_seqlen=req.swa_evicted_seqlen, update_kv_after_len=old_handle.cached_len) self.unlock(old_handle) @@ -492,8 +473,8 @@ def _cache_req_swa(self, req: Req, *, finished: bool) -> None: ) if keep_from > 0: self._free_swa( - self.prefix_cache.trim_head_swa(req.input_ids[:prompt_len], keep_from)) - self.prefix_cache.match_prefix(req.input_ids[:prompt_len]) + self.prefix_cache.trim_head_swa(_key_ids(req)[:prompt_len], keep_from)) + self.prefix_cache.match_prefix(_key_ids(req)[:prompt_len]) else: # inc_lock is node-granular, and the suffix insert just made this chunk's whole # extend one node: locking it would pin the entire chunk's swa for all of decode, @@ -504,8 +485,8 @@ def _cache_req_swa(self, req: Req, *, finished: bool) -> None: keep_from = align_down( max(insert_len - self.sliding_window_size - _SWA_RETAIN_GAP, 0), self.page_size) if keep_from > 0: - self.prefix_cache.match_prefix(req.input_ids[:keep_from]) - m = self.prefix_cache.match_prefix(req.input_ids[:insert_len]) + self.prefix_cache.match_prefix(_key_ids(req)[:keep_from]) + m = self.prefix_cache.match_prefix(_key_ids(req)[:insert_len]) # Re-point the page table to the tree's live slots for the committed region. Any dup # slots insert reclaimed had their full->swa mapping reset to the 0 sentinel; unlike the # full pool (KV survives in place until realloc), a stale swa mapping would make the diff --git a/python/freetoken/scheduler/prefill.py b/python/freetoken/scheduler/prefill.py index f5bc8f7a3..0bf5bf6db 100644 --- a/python/freetoken/scheduler/prefill.py +++ b/python/freetoken/scheduler/prefill.py @@ -171,11 +171,6 @@ def _add_one_req( _slice = slice(cached_len, cached_len + chunk_size) device_ids = self.table_manager.token_pool[table_idx, _slice] device_ids.copy_(_maybe_pinned(pending_req.input_ids[_slice]), non_blocking=True) - if is_chunked and pending_req.mm_embeds is not None: - raise NotImplementedError( - "Multimodal prompts must fit in a single prefill chunk; increase " - "--max-extend-tokens or shrink the prompt." - ) req = CLS( input_ids=pending_req.input_ids[: cached_len + chunk_size], table_idx=table_idx, @@ -185,6 +180,13 @@ def _add_one_req( cache_handle=cache_handle, sampling_params=pending_req.sampling_params, mm_embeds=pending_req.mm_embeds, + mrope_positions=pending_req.mrope_positions, + mrope_delta=pending_req.mrope_delta, + cache_ids=( + None + if pending_req.cache_ids is None + else pending_req.cache_ids[: cached_len + chunk_size] + ), ) # Hybrid GDN per-request state slots (None for non-hybrid). On a fresh admit these are # freshly allocated; on a chunked continuation they are inherited from the prior chunk. @@ -245,7 +247,15 @@ class PrefillManager: def add_one_req(self, req: UserMsg) -> None: self.pending_list.append( - PendingReq(req.uid, req.input_ids, req.sampling_params, mm_embeds=req.mm_embeds) + PendingReq( + req.uid, + req.input_ids, + req.sampling_params, + mm_embeds=req.mm_embeds, + mrope_positions=req.mrope_positions, + mrope_delta=req.mrope_delta, + cache_ids=req.cache_ids, + ) ) def schedule_next_batch(self, prefill_budget: int) -> Batch | None: diff --git a/python/freetoken/scheduler/scheduler.py b/python/freetoken/scheduler/scheduler.py index 48923e3b0..e1a695624 100644 --- a/python/freetoken/scheduler/scheduler.py +++ b/python/freetoken/scheduler/scheduler.py @@ -520,6 +520,12 @@ def _process_one_msg(self, msg: BaseBackendMsg) -> None: logger.warning_rank0( f"Adjust max_tokens to {max_output_len} for request {msg.uid}." ) + if msg.mm_inputs is not None: + error = self._prepare_multimodal(msg) + if error is not None: + logger.warning_rank0(f"Rejecting request {msg.uid}: {error}") + self.send_result([ErrorReplyMsg(uid=msg.uid, error=error)]) + return self.prefill_manager.add_one_req(msg) elif isinstance(msg, AbortBackendMsg): logger.debug_rank0("Aborting request %d", msg.uid) @@ -759,6 +765,23 @@ def _log_cache_geometry(self, event: str) -> None: except Exception as e: # noqa: BLE001 logger.warning(f"could not log cache geometry: {e!r}") + def _prepare_multimodal(self, msg: UserMsg) -> str | None: + """Encode the request's images on this rank (every TP rank sees the same UserMsg) into + ``mm_embeds`` + mRoPE positions; returns the client-facing error instead of raising.""" + model = self.engine.model + if not getattr(model, "has_vision", False): + return "image inputs need the vision tower: start the server with FREETOKEN_LOAD_VISION=1" + try: + with self.engine_stream_ctx: # ordered before the prefill that reads mm_embeds + msg.mm_embeds, msg.mrope_positions, msg.mrope_delta = model.prepare_mm_inputs( + msg.input_ids, msg.mm_inputs + ) + except Exception as exc: # noqa: BLE001 -- a bad image must not take the scheduler down + logger.warning_rank0(f"image encoding failed for request {msg.uid}: {exc!r}") + return f"could not encode images: {exc}" + msg.mm_inputs = None + return None + def _prepare_batch(self, batch: Batch) -> ForwardInput: self.engine.graph_runner.pad_batch(batch) self._forward_iter += 1 @@ -781,6 +804,9 @@ def _prepare_batch(self, batch: Batch) -> ForwardInput: if batch.is_prefill: self._gather_multimodal(batch) batch.positions = _make_positions(batch, self.device) + batch.rope_positions, batch.mrope_cos_sin = _make_rope_positions( + batch, self.device, self.engine.model + ) input_mapping = _make_input_tuple(batch, self.device) write_mapping = _make_write_tuple(batch, self.device) batch.out_loc = self.engine.page_table[input_mapping] @@ -820,11 +846,18 @@ def _prepare_batch(self, batch: Batch) -> ForwardInput: def _gather_multimodal(self, batch: Batch) -> None: """Concatenate per-request vision soft tokens (in request order) for a prefill - batch so the model can scatter them at image-token positions. ``req.mm_embeds`` - is kept (not cleared) so the cache manager can recognize multimodal requests and - keep them out of the shared prefix cache (image placeholders share a token id but - carry per-image content).""" - parts = [req.mm_embeds for req in batch.reqs if req.mm_embeds is not None] + batch so the model can scatter them at image-token positions. A chunked prompt + contributes only the placeholders inside its current chunk; ``req.mm_embeds`` is + kept (not cleared) so later chunks can slice theirs. Image prompts take part in the + prefix cache through ``req.cache_ids`` (per-image content hash, see + ``tokenizer.tokenize._image_cache_ids``): a hit skips the cached placeholders too.""" + image_token_id = getattr(self.engine.model, "image_token_id", None) + parts = [ + _mm_embeds_window(req, image_token_id) + for req in batch.reqs + if req.mm_embeds is not None + ] + parts = [p for p in parts if p.shape[0]] if parts: batch.mm_embeds = torch.cat(parts, dim=0) @@ -890,6 +923,35 @@ def _make_positions(batch: Batch, device: torch.device) -> torch.Tensor: return indices_host.to(device, non_blocking=True) +def _mm_embeds_window(req, image_token_id: int) -> torch.Tensor: + """The rows of ``req.mm_embeds`` for this prefill window: the image placeholders in + ``input_ids[cached_len:]``, skipping those an earlier chunk already scattered.""" + is_image = req.input_ids == image_token_id + before = int(is_image[: req.cached_len].sum()) + within = int(is_image[req.cached_len :].sum()) + return req.mm_embeds[before : before + within] + + +def _make_rope_positions( + batch: Batch, device: torch.device, model +) -> tuple[torch.Tensor, torch.Tensor | None]: + """``(rope_positions, mrope_cos_sin)`` for the batch (see ``Batch``). Text-only batches + alias ``positions`` (no extra work); requests decoding after an image add their + ``mrope_delta``; a prefill batch holding image tokens gets a per-token cos|sin table from + the model (``mrope_table``) and each token's table row as its rope position.""" + reqs = batch.padded_reqs + if batch.is_prefill and any(r.mrope_positions is not None for r in reqs): + return model.mrope_table(reqs, device) + if all(r.mrope_delta == 0 for r in reqs): + return batch.positions, None + delta_host = torch.empty(len(batch.positions), dtype=torch.int32, pin_memory=True) + offset = 0 + for req in reqs: + delta_host[offset : offset + req.extend_len].fill_(req.mrope_delta) + offset += req.extend_len + return batch.positions + delta_host.to(device, non_blocking=True), None + + def _make_input_tuple(batch: Batch, device: torch.device) -> Indice2D: mapping_host = torch.empty(len(batch.positions), dtype=torch.int64, pin_memory=True) offset = 0 diff --git a/python/freetoken/scheduler/utils.py b/python/freetoken/scheduler/utils.py index 0d6dd5129..4e0128d8e 100644 --- a/python/freetoken/scheduler/utils.py +++ b/python/freetoken/scheduler/utils.py @@ -18,6 +18,11 @@ class PendingReq: sampling_params: SamplingParams chunked_req: ChunkedReq | None = None mm_embeds: torch.Tensor | None = None + mrope_positions: torch.Tensor | None = None + mrope_delta: int = 0 + cache_ids: torch.Tensor | None = ( + None # prefix-cache key ids (image prompts), else input_ids + ) @property def input_len(self) -> int: diff --git a/python/freetoken/server/generation.py b/python/freetoken/server/generation.py index be05d908a..9e14e980a 100644 --- a/python/freetoken/server/generation.py +++ b/python/freetoken/server/generation.py @@ -14,6 +14,8 @@ from __future__ import annotations import asyncio +import base64 +import binascii import json import time from collections.abc import AsyncIterator @@ -139,6 +141,7 @@ class GenSpec: chat_template_kwargs: dict[str, Any] = field(default_factory=dict) template_tools: list[dict[str, Any]] | None = None # tools the model sees (TokenizeMsg.tools) parser_tools: list[dict[str, Any]] | None = None # tools for FunctionCallParser; None disables parsing + images: list[bytes] = field(default_factory=list) # encoded image files, in message order @property def parse_tools(self) -> bool: @@ -187,18 +190,23 @@ def pick(value, key, framework): ) -def render_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: +def render_messages( + messages: list[dict[str, Any]], images: list[bytes] | None = None +) -> list[dict[str, Any]]: """Normalize OpenAI-shaped message dicts for the chat template: flatten text - content parts to a string and decode tool-call arguments from JSON. Raises - ValueError on a non-text content part (text-only server). Shared by all adapters.""" - return [_render_message(m) for m in messages] + content parts to a string and decode tool-call arguments from JSON. Shared by all adapters. + ``images`` (OpenAI chat only): accept ``image_url`` parts, decode them (``data:`` URLs) into + the list in message order and keep ``{"type": "image"}`` parts in the content for the + template. Without it a non-text part raises ValueError (text-only adapters).""" + return [_render_message(m, images) for m in messages] -def _render_message(message: dict[str, Any]) -> dict[str, Any]: + +def _render_message(message: dict[str, Any], images: list[bytes] | None = None) -> dict[str, Any]: m = dict(message) content = m.get("content") if isinstance(content, list): - m["content"] = _flatten_text_parts(content) + m["content"] = _render_parts(content, images) # Templates read different reasoning keys (reasoning_content: most; reasoning: # gemma4; thinking: gpt-oss) — accept any, emit both. reasoning = m.get("reasoning_content") or m.get("reasoning") or m.get("thinking") @@ -230,6 +238,48 @@ def _render_message(message: dict[str, Any]) -> dict[str, Any]: return m +# Decoded image bytes above this are refused (a 4000x3000 JPEG is ~3 MiB; PNGs a few times that). +IMAGE_MAX_BYTES = 16 << 20 + + +def _render_parts(parts: list[Any], images: list[bytes] | None) -> str | list[dict[str, Any]]: + """Text-only part lists flatten to a string (unchanged behaviour); a list carrying images keeps + ``{"type": "text"}`` / ``{"type": "image"}`` parts, which the chat template renders as text and + one ``<|image_pad|>`` placeholder per image (expanded by the tokenizer worker).""" + if images is None or not any( + isinstance(p, dict) and p.get("type") == "image_url" for p in parts + ): + return _flatten_text_parts(parts) + out: list[dict[str, Any]] = [] + for part in parts: + ptype = part.get("type") if isinstance(part, dict) else None + if ptype == "text": + out.append({"type": "text", "text": part.get("text") or ""}) + elif ptype == "image_url": + images.append(_decode_image_url(part.get("image_url"))) + out.append({"type": "image"}) + else: + raise ValueError(f"Unsupported content part type: {ptype}") + return out + + +def _decode_image_url(image_url: Any) -> bytes: + """OpenAI ``image_url`` part -> image file bytes. Only inline ``data:`` URLs: fetching a remote + URL from the server would be an outbound request on the client's behalf.""" + url = image_url.get("url") if isinstance(image_url, dict) else image_url + if not isinstance(url, str) or not url.startswith("data:"): + raise ValueError("image_url must be a data: URL (base64-encoded image); remote URLs are not fetched") + header, sep, payload = url.partition(",") + if not sep or ";base64" not in header: + raise ValueError("image_url data: URL must be base64-encoded") + if len(payload) > IMAGE_MAX_BYTES * 4 // 3 + 4: + raise ValueError(f"image exceeds {IMAGE_MAX_BYTES >> 20} MiB") + try: + return base64.b64decode(payload, validate=True) + except (binascii.Error, ValueError) as exc: + raise ValueError(f"image_url data: URL is not valid base64: {exc}") from exc + + def _flatten_text_parts(parts: list[Any]) -> str: texts: list[str] = [] for part in parts: @@ -269,6 +319,7 @@ async def submit_generation(spec: GenSpec, state: Any) -> int: sampling_params=spec.sampling_params, chat_template_kwargs=spec.chat_template_kwargs, tools=spec.template_tools, + images=spec.images or None, ) ) return uid diff --git a/python/freetoken/server/openai_api.py b/python/freetoken/server/openai_api.py index b4becd263..89c6582ab 100644 --- a/python/freetoken/server/openai_api.py +++ b/python/freetoken/server/openai_api.py @@ -66,8 +66,9 @@ def chat_request_to_genspec( thinking_type = _thinking_type(req) if req.reasoning_effort or thinking_type: ctk = effort_toggle_kwargs(req.reasoning_effort, ctk, thinking_type=thinking_type) + images: list[bytes] = [] return GenSpec( - messages=render_messages([m.model_dump(exclude_none=True) for m in req.messages]), + messages=render_messages([m.model_dump(exclude_none=True) for m in req.messages], images), sampling_params=resolve_sampling( temperature=req.temperature, top_k=req.top_k, @@ -80,6 +81,7 @@ def chat_request_to_genspec( chat_template_kwargs=ctk, template_tools=_tools_for_template(req), parser_tools=(_all_tool_dicts(req.tools) if _should_parse_tools(req) else None), + images=images, ) diff --git a/python/freetoken/tokenizer/server.py b/python/freetoken/tokenizer/server.py index 530e862d0..2b4446676 100644 --- a/python/freetoken/tokenizer/server.py +++ b/python/freetoken/tokenizer/server.py @@ -80,12 +80,26 @@ def _send_generation_replies( _put_user_replies(send_frontend, terminal_errors) +def _backend_msg(msg, t) -> UserMsg: + mm_inputs = getattr(t, "mm_inputs", None) + cache_ids = mm_inputs.pop("cache_ids", None) if mm_inputs else None + return UserMsg( + uid=msg.uid, + input_ids=t, + sampling_params=msg.sampling_params, + mm_inputs=mm_inputs, + cache_ids=cache_ids, + ) + + def _tokenize_requests( tokenize_manager: Any, messages: List[TokenizeMsg], logger: Any, ) -> tuple[List[TokenizeMsg], List[torch.Tensor], List[UserReply]]: - """Tokenize independently, returning backend work plus terminal frontend errors. + """Tokenize independently, returning backend work plus terminal frontend errors. The + returned tensors carry the request's vision-processor outputs as ``.mm_inputs`` (None for + text-only requests). Successful tokenization deliberately emits no prompt-token reply: accounting starts only when the scheduler later confirms first-prefill admission. @@ -95,7 +109,8 @@ def _tokenize_requests( errors: List[UserReply] = [] for msg in messages: try: - tokens = tokenize_manager.tokenize([msg])[0] + tokens, mm_inputs = tokenize_manager.tokenize_one(msg) + tokens.mm_inputs = mm_inputs except Exception as exc: # noqa: BLE001 — isolate, never crash the worker logger.warning(f"tokenization failed for request {msg.uid}: {exc!r}") errors.append( @@ -147,7 +162,7 @@ def tokenize_worker( from .detokenize import DetokenizeManager from .tokenize import TokenizeManager - tokenize_manager = TokenizeManager(tokenizer) + tokenize_manager = TokenizeManager(tokenizer, model_path=tokenizer_path) detokenize_manager = DetokenizeManager( tokenizer, load_eos_token_ids(tokenizer_path, tokenizer) ) @@ -254,8 +269,7 @@ def tokenize_worker( ) if ok_msgs: backend = [ - UserMsg(uid=msg.uid, input_ids=t, sampling_params=msg.sampling_params) - for msg, t in zip(ok_msgs, ok_tensors, strict=True) + _backend_msg(msg, t) for msg, t in zip(ok_msgs, ok_tensors, strict=True) ] send_backend.put(backend[0] if len(backend) == 1 else BatchBackendMsg(data=backend)) if len(abort_msg) > 0: diff --git a/python/freetoken/tokenizer/tokenize.py b/python/freetoken/tokenizer/tokenize.py index 0636b3428..af724a341 100644 --- a/python/freetoken/tokenizer/tokenize.py +++ b/python/freetoken/tokenizer/tokenize.py @@ -2,6 +2,7 @@ import importlib.util import json +import hashlib import os import threading from types import ModuleType @@ -45,10 +46,61 @@ def resolve_thinking_mode(chat_template_kwargs: dict[str, Any] | None, tools: An _EFFORT_PROBE_MESSAGES = [{"role": "user", "content": "ping"}] +# Image resize bounds in pixels (Qwen2-VL processor ``size``). The cap bounds the prompt: at +# patch 16 / merge 2 one soft token covers 32x32 pixels, so the default is ~1000 tokens per +# image and a 4000x3000 photo is downscaled instead of costing 11k tokens. +IMAGE_MIN_PIXELS = 4 * 28 * 28 +IMAGE_MAX_PIXELS = int(os.getenv("FREETOKEN_IMAGE_MAX_PIXELS", str(1280 * 28 * 28))) + + +_IMAGE_KEY_BASE = 1 << 30 # image cache-key ids start above every vocabulary + + +def _image_cache_ids( + input_ids: torch.Tensor, image_token_id: int, counts: list[int], hashes: list[int] +) -> torch.Tensor: + """Prefix-cache key ids for an image prompt: ``input_ids`` with the i-th placeholder run + replaced by ids derived from the i-th image's content hash. Placeholders share one token id + across images, so keyed on raw ids a prefix hit could serve another image's KV; the derived + ids (``>= 2**30``, above any vocabulary, ``hash + offset`` within the run) make each run + unique to its image and the prompt safe to share. The model still reads ``input_ids``.""" + ids = input_ids.clone() + hits = (ids == image_token_id).nonzero().flatten() + mask = _IMAGE_KEY_BASE - 1 + pos = 0 + for n, h in zip(counts, hashes, strict=True): + run = (torch.arange(n, dtype=torch.int64) + (h & mask)) & mask + ids[hits[pos : pos + n]] = (run + _IMAGE_KEY_BASE).to(ids.dtype) + pos += n + return ids + + +def _expand_image_tokens( + input_ids: torch.Tensor, image_token_id: int, counts: list[int] +) -> torch.Tensor: + """Replace the i-th ``image_token_id`` placeholder with ``counts[i]`` copies.""" + hits = (input_ids == image_token_id).nonzero().flatten().tolist() + if len(hits) != len(counts): + raise ValueError( + f"{len(hits)} image placeholders in the rendered prompt for {len(counts)} images" + ) + pieces, prev = [], 0 + for hit, n in zip(hits, counts, strict=True): + pieces.append(input_ids[prev:hit]) + pieces.append(torch.full((n,), image_token_id, dtype=input_ids.dtype)) + prev = hit + 1 + pieces.append(input_ids[prev:]) + return torch.cat(pieces) + class TokenizeManager: - def __init__(self, tokenizer: PreTrainedTokenizerBase) -> None: + def __init__( + self, tokenizer: PreTrainedTokenizerBase, model_path: str | None = None + ) -> None: self.tokenizer = tokenizer + # image requests need the checkpoint's image processor + placeholder id (worker only) + self._model_path = model_path + self._vision: tuple[Any, int] | None = None self._dsv4_encoder = _load_dsv4_encoder_if_needed(tokenizer) self._effort_profile: EffortProfile | None = None self._thinking_profile: ThinkingProfile | None = None @@ -74,6 +126,51 @@ def tokenize(self, msgs: List[TokenizeMsg]) -> List[torch.Tensor]: results.append(input_ids.view(-1).to(torch.int32)) return results + def tokenize_one(self, msg: TokenizeMsg) -> tuple[torch.Tensor, dict[str, Any] | None]: + """``tokenize`` for one message plus, when it carries images, the vision processor's + outputs (``pixel_values`` fp32 [patches, C*T*P*P], ``image_grid_thw`` [N, 3]) with every + ``<|image_pad|>`` placeholder expanded to its image's soft-token count.""" + input_ids = self.tokenize([msg])[0] + if not msg.images: + return input_ids, None + import io + + from PIL import Image + + processor, image_token_id = self._vision_processor() + images = [Image.open(io.BytesIO(data)).convert("RGB") for data in msg.images] + out = processor(images=images, return_tensors="pt") + grid = out["image_grid_thw"] + counts = (grid.prod(-1) // processor.merge_size**2).tolist() + hashes = [ + int.from_bytes(hashlib.blake2b(data, digest_size=8).digest(), "big") + for data in msg.images + ] + input_ids = _expand_image_tokens(input_ids, image_token_id, counts) + # pixel_values stay fp32 on the wire (numpy cannot carry bf16); the tower casts them + return input_ids, { + "pixel_values": out["pixel_values"], + "image_grid_thw": grid, + "cache_ids": _image_cache_ids(input_ids, image_token_id, counts, hashes), + } + + def _vision_processor(self) -> tuple[Any, int]: + if self._vision is None: + if self._model_path is None: + raise ValueError("image inputs are not supported here") + from freetoken.utils.hf import cached_load_hf_config + from transformers import AutoImageProcessor + + image_token_id = getattr(cached_load_hf_config(self._model_path), "image_token_id", None) + if image_token_id is None: + raise ValueError("this model has no image placeholder token") + processor = AutoImageProcessor.from_pretrained( + self._model_path, + size={"shortest_edge": IMAGE_MIN_PIXELS, "longest_edge": IMAGE_MAX_PIXELS}, + ) + self._vision = (processor, int(image_token_id)) + return self._vision + def render_prompt(self, msg: TokenizeMsg) -> str: """The template/encoder half of ``tokenize``, exposed so the frontend can validate a request before committing an SSE stream. Sanitizes diff --git a/tests/models/qwen4_exp/common.py b/tests/models/qwen4_exp/common.py index 1f9c117bd..03d8f3534 100644 --- a/tests/models/qwen4_exp/common.py +++ b/tests/models/qwen4_exp/common.py @@ -206,6 +206,8 @@ def req(self, table_idx: int, cached_len: int, device_len: int) -> SimpleNamespa cached_len=cached_len, device_len=device_len, extend_len=device_len - cached_len, + mrope_positions=None, + mrope_delta=0, ) def step(self, req: SimpleNamespace) -> None: @@ -231,6 +233,8 @@ def batch(self, reqs, phase: str) -> SimpleNamespace: is_prefill=phase == "prefill", is_decode=phase == "decode", positions=positions, + rope_positions=None, # text-only: the layer falls back to positions + mrope_cos_sin=None, out_loc=out_loc, attn_metadata=None, active_table_idx=torch.tensor( diff --git a/tests/models/qwen4_exp/test_build_cpu.py b/tests/models/qwen4_exp/test_build_cpu.py new file mode 100644 index 000000000..0cf2027d1 --- /dev/null +++ b/tests/models/qwen4_exp/test_build_cpu.py @@ -0,0 +1,43 @@ +"""The whole model must CONSTRUCT without a GPU. + +Every other model test that builds a decoder layer is behind requires_cuda, so a constructor +signature mismatch between Qwen4ExpDecoderLayer and the ops it builds is invisible to a +CPU-only run -- it only shows up when a server boots. That happened: Qwen4ExpMoE did not +accept the `prefix` kwarg model.py passes, and a full CPU suite still went green. + +Building on the meta device costs no memory and no GPU, so there is no reason not to. +""" + +import torch +from freetoken.distributed import set_tp_info, try_get_tp_info +from freetoken.layers import set_rope_device +from freetoken.models.qwen4_exp.config import parse_config +from freetoken.models.qwen4_exp.model import Qwen4ExpForCausalLM + +from .common import toy_hf_config + + +def _build(): + if try_get_tp_info() is None: + set_tp_info(rank=0, size=1) + set_rope_device(torch.device("cpu")) + with torch.device("meta"): + return Qwen4ExpForCausalLM(parse_config(toy_hf_config())) + + +def test_model_constructs_on_meta_without_cuda(): + model = _build() + sd = model.state_dict() + assert sd, "state dict is empty" + # both layer families and the head must be present + assert any(".self_attn." in k for k in sd), "no full-attention layer built" + assert any(".linear_attn." in k for k in sd), "no GDN layer built" + assert any(k.startswith("lm_head") for k in sd), "no lm_head built" + + +def test_every_layer_gets_its_own_prefixed_weights(): + # a prefix that is dropped or shared silently collapses layers onto one another + sd = _build().state_dict() + mlp_keys = {k for k in sd if ".mlp." in k} + layers = {k.split(".layers.")[1].split(".")[0] for k in mlp_keys if ".layers." in k} + assert len(layers) > 1, f"MoE weights landed on a single layer: {sorted(layers)}" diff --git a/tests/models/qwen4_exp/test_mrope.py b/tests/models/qwen4_exp/test_mrope.py new file mode 100644 index 000000000..333ba3cbf --- /dev/null +++ b/tests/models/qwen4_exp/test_mrope.py @@ -0,0 +1,183 @@ +"""mRoPE for qwen4_exp image prompts: FreeToken's port against the HF reference (CPU). + +The HF checks need a transformers with ``qwen4_exp``; the pure-FreeToken checks always run. +""" + +from __future__ import annotations + +import base64 +from types import MethodType, SimpleNamespace + +import pytest +import torch +from freetoken.models.qwen4_exp.mrope import mrope_cos_sin, rope_index +from freetoken.server.generation import render_messages +from freetoken.tokenizer.tokenize import _expand_image_tokens + +IMAGE = 248056 +MERGE = 2 +SECTION = (11, 11, 10) + + +def _prompt( + grids: list[tuple[int, int, int]], text_len: int = 5, tail: int = 7 +) -> torch.Tensor: + ids = list(range(1, text_len + 1)) + for t, h, w in grids: + ids += [IMAGE] * (t * (h // MERGE) * (w // MERGE)) + [50, 51] + ids += list(range(100, 100 + tail)) + return torch.tensor(ids, dtype=torch.int32) + + +def test_rope_index_text_only(): + ids = _prompt([]) + pos, delta = rope_index(ids, None, IMAGE, MERGE) + assert delta == 0 + assert torch.equal(pos, torch.arange(len(ids)).expand(3, -1)) + + +def test_rope_index_image_layout(): + grid = torch.tensor([[1, 4, 6]]) + ids = _prompt([(1, 4, 6)], text_len=3, tail=2) + pos, delta = rope_index(ids, grid, IMAGE, MERGE) + # text 0..2, image at start 3: T=3, H in 3..4, W in 3..5, then text resumes at 3 + max(2, 3) + img = pos[:, 3:9] + assert img[0].tolist() == [3] * 6 + assert img[1].tolist() == [3, 3, 3, 4, 4, 4] + assert img[2].tolist() == [3, 4, 5, 3, 4, 5] + assert pos[:, 9].tolist() == [6, 6, 6] + assert delta == int(pos.max()) + 1 - len(ids) + assert delta <= 0 + + +def test_mrope_cos_sin_text_rows_equal_cache(): + inv_freq = 1.0 / (1e7 ** (torch.arange(0, 64, 2, dtype=torch.float) / 64)) + pos = torch.arange(0, 37).expand(3, -1) + table = mrope_cos_sin(pos, inv_freq, SECTION) + freqs = torch.einsum("i,j -> ij", torch.arange(37, dtype=torch.float), inv_freq) + assert torch.equal(table, torch.cat((freqs.cos(), freqs.sin()), dim=-1)) + + +def test_expand_image_tokens(): + ids = torch.tensor([1, IMAGE, 2, IMAGE, 3], dtype=torch.int32) + out = _expand_image_tokens(ids, IMAGE, [3, 2]) + assert out.tolist() == [1, IMAGE, IMAGE, IMAGE, 2, IMAGE, IMAGE, 3] + with pytest.raises(ValueError): + _expand_image_tokens(ids, IMAGE, [3]) + + +def test_render_messages_images(): + png = b"\x89PNG fake" + url = "data:image/png;base64," + base64.b64encode(png).decode() + messages = [ + {"role": "system", "content": [{"type": "text", "text": "sys"}]}, + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": url}}, + {"type": "text", "text": "hi"}, + ], + }, + ] + images: list[bytes] = [] + out = render_messages(messages, images) + assert out[0]["content"] == "sys" # text-only lists still flatten + assert out[1]["content"] == [{"type": "image"}, {"type": "text", "text": "hi"}] + assert images == [png] + with pytest.raises(ValueError): # text-only adapters keep refusing image parts + render_messages(messages) + with pytest.raises(ValueError): # remote URLs are never fetched + render_messages( + [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": "https://x/y.png"}} + ], + } + ], + [], + ) + + +# ----------------------------------------------------------------------------- HF reference +hf = pytest.importorskip("transformers.models.qwen4_exp.modeling_qwen4_exp") + + +def _hf_rope_index(ids: torch.Tensor, grid: torch.Tensor | None): + cfg = SimpleNamespace( + image_token_id=IMAGE, + video_token_id=248057, + vision_start_token_id=248053, + vision_end_token_id=248054, + vision_config=SimpleNamespace(spatial_merge_size=MERGE), + ) + stub = SimpleNamespace(config=cfg) + stub.get_vision_position_ids = MethodType( + hf.Qwen4ExpModel.get_vision_position_ids, stub + ) + ids = ids.view(1, -1).long() + pos, delta = hf.Qwen4ExpModel.get_rope_index( + stub, ids, mm_token_type_ids=(ids == IMAGE).int(), image_grid_thw=grid + ) + return pos[:, 0], int(delta.reshape(-1)[0]) + + +@pytest.mark.parametrize( + "grids", + [[(1, 4, 6)], [(1, 8, 8), (1, 2, 12)], [(1, 6, 2)]], +) +def test_rope_index_matches_hf(grids): + ids = _prompt(grids) + grid = torch.tensor(grids) + pos, delta = rope_index(ids, grid, IMAGE, MERGE) + ref_pos, ref_delta = _hf_rope_index(ids, grid) + assert torch.equal(pos, ref_pos), (pos, ref_pos) + assert delta == ref_delta + + +def test_mrope_cos_sin_matches_hf(): + text_cfg = SimpleNamespace( + head_dim=256, + hidden_size=2560, + num_attention_heads=24, + max_position_embeddings=262144, + rope_parameters={ + "rope_type": "default", + "rope_theta": 1e7, + "partial_rotary_factor": 0.25, + "mrope_section": list(SECTION), + "mrope_interleaved": True, + }, + ) + rotary = hf.Qwen4ExpTextRotaryEmbedding(text_cfg) + pos = torch.stack( + [torch.randint(0, 5000, (3, 29)) for _ in range(1)], dim=1 + ) # [3, 1, 29] + x = torch.zeros(1, 29, 256) + cos, sin = rotary(x, pos) + table = mrope_cos_sin(pos[:, 0], rotary.inv_freq.float(), SECTION) + torch.testing.assert_close(table[:, :32], cos[0, :, :32], atol=1e-5, rtol=1e-5) + torch.testing.assert_close(table[:, 32:], sin[0, :, :32], atol=1e-5, rtol=1e-5) + + +def test_image_cache_ids(): + from freetoken.tokenizer.tokenize import _image_cache_ids + + ids = _expand_image_tokens( + torch.tensor([1, IMAGE, 2, IMAGE, 3], dtype=torch.int32), IMAGE, [3, 2] + ) + a = _image_cache_ids(ids, IMAGE, [3, 2], [0xAAAA, 0xBBBB]) + b = _image_cache_ids(ids, IMAGE, [3, 2], [0xAAAA, 0xCCCC]) + text = ids != IMAGE + assert a.dtype == ids.dtype and torch.equal( + a[text], ids[text] + ) # text ids untouched + assert bool((a[~text] >= 1 << 30).all()) # never a vocabulary id + assert torch.equal(a[:5], b[:5]) and not torch.equal( + a[5:7], b[5:7] + ) # keyed per image + assert torch.equal( + a, _image_cache_ids(ids, IMAGE, [3, 2], [0xAAAA, 0xBBBB]) + ) # deterministic + assert len(set(a[1:4].tolist())) == 3 # one run, distinct ids per position diff --git a/tests/models/qwen4_exp/test_mrope_gpu.py b/tests/models/qwen4_exp/test_mrope_gpu.py new file mode 100644 index 000000000..33f847023 --- /dev/null +++ b/tests/models/qwen4_exp/test_mrope_gpu.py @@ -0,0 +1,79 @@ +"""mRoPE plumbing through the QSA layer (GPU): the per-token table path (what an image-carrying +prefill batch uses) must reproduce the position-indexed cache path exactly for text, including +a chunked continuation sharing the batch (the indexer ropes a straddling group at its first +token, which lives in the table's lead rows) and the decode `positions + delta` path. +""" + +from __future__ import annotations + +import torch +from freetoken.models.qwen4_exp.mrope import mrope_table + +from .common import Fixture, parsed_config, requires_cuda + +QSA_LAYER = 3 +SECTION = (11, 11, 10) + + +def _inv_freq(config, device): + rc = config.rotary_config + return 1.0 / ( + rc.base + ** ( + torch.arange(0, rc.rotary_dim, 2, dtype=torch.float, device=device) + / rc.rotary_dim + ) + ) + + +def _run(config, table_path: bool, seed: int = 5): + """Prefill B (chunk 1) -> prefill [A, B chunk 2] -> one decode step; return the two outputs.""" + fixture = Fixture(config, num_pages=64) + attn = fixture.layer(QSA_LAYER) + ratio = config.qwen4_args.index_ratio + gen = torch.Generator(device=fixture.device).manual_seed(seed) + hidden = config.hidden_size + x_b1 = torch.randn( + 41, hidden, device=fixture.device, dtype=fixture.dtype, generator=gen + ) + x_ab = torch.randn( + 37 + 30, hidden, device=fixture.device, dtype=fixture.dtype, generator=gen + ) + x_dec = torch.randn( + 2, hidden, device=fixture.device, dtype=fixture.dtype, generator=gen + ) + assert 41 % ratio, ( + "chunk 1 must end mid-group so a group straddles the chunk boundary" + ) + + b = fixture.req(1, 0, 41) + # chunk 1 always takes the cache path + attn.forward(x_b1, fixture.batch([b], "prefill")) + b.cached_len, b.device_len, b.extend_len = 41, 71, 30 + fixture.allocate(1, 41, 71) + a = fixture.req(0, 0, 37) + batch = fixture.batch([a, b], "prefill") + if table_path: + batch.rope_positions, batch.mrope_cos_sin = mrope_table( + [a, b], ratio, _inv_freq(config, fixture.device), SECTION, fixture.device + ) + assert batch.mrope_cos_sin.shape[0] == 37 + 30 + (ratio - 1) + out_prefill = attn.forward(x_ab, batch).clone() + + fixture.step(a) + fixture.step(b) + batch = fixture.batch([a, b], "decode") + if table_path: + # decode never has a table; exercise the explicit rope_positions path (delta 0) + batch.rope_positions = batch.positions + 0 + out_decode = attn.forward(x_dec, batch).clone() + return out_prefill, out_decode + + +@requires_cuda +def test_table_path_matches_cache_path(): + config = parsed_config() + ref_prefill, ref_decode = _run(config, table_path=False) + tab_prefill, tab_decode = _run(config, table_path=True) + assert torch.equal(ref_prefill, tab_prefill) + assert torch.equal(ref_decode, tab_decode) diff --git a/tests/models/qwen4_exp/test_qsa_backend.py b/tests/models/qwen4_exp/test_qsa_backend.py index 1d3b944ce..da2c8027c 100644 --- a/tests/models/qwen4_exp/test_qsa_backend.py +++ b/tests/models/qwen4_exp/test_qsa_backend.py @@ -174,6 +174,7 @@ def test_decode_graph_replay_matches_eager(): capture_batch = SimpleNamespace( padded_reqs=[dummy] * bs, reqs=[dummy] * bs, phase="decode", size=bs, padded_size=bs, is_prefill=False, is_decode=True, positions=static["positions"], + rope_positions=None, mrope_cos_sin=None, out_loc=static["out_loc"], attn_metadata=None, active_table_idx=None, ) fixture.backend.prepare_for_capture(capture_batch) diff --git a/tests/models/qwen4_exp/test_skeleton.py b/tests/models/qwen4_exp/test_skeleton.py index fa75d6ee6..b645fa042 100644 --- a/tests/models/qwen4_exp/test_skeleton.py +++ b/tests/models/qwen4_exp/test_skeleton.py @@ -346,7 +346,9 @@ def test_qsa_layer_matches_hf_dense(): x = (torch.randn(seq_len, config.hidden_size, device=device, dtype=dtype) * 0.5) positions = torch.arange(seq_len, device=device, dtype=torch.int64) req = SimpleNamespace(extend_len=seq_len, cached_len=0, table_idx=1) - batch = SimpleNamespace(padded_reqs=[req], reqs=[req], positions=positions) + batch = SimpleNamespace( + padded_reqs=[req], reqs=[req], positions=positions, rope_positions=None, mrope_cos_sin=None + ) backend = TorchDenseQSAReference(config, num_slots=4, max_len=64, device=device, dtype=dtype) _fresh_ctx(attn_backend=backend) @@ -477,6 +479,7 @@ def test_decoder_stack_prefill_and_decode(monkeypatch): padded_reqs=reqs, reqs=reqs, size=len(reqs), is_prefill=True, is_decode=False, input_ids=torch.tensor(flat, dtype=torch.int64, device=device), positions=torch.cat([torch.arange(len(p)) for p in prompts]).to(device), + rope_positions=None, mrope_cos_sin=None, mm_embeds=None, attn_metadata=SimpleNamespace(get_last_indices=lambda bs: last[:bs]), ) with ctx.forward_batch(batch): @@ -492,6 +495,7 @@ def test_decoder_stack_prefill_and_decode(monkeypatch): padded_reqs=reqs, reqs=reqs, size=len(reqs), is_prefill=False, is_decode=True, input_ids=torch.tensor([14] * len(reqs), dtype=torch.int64, device=device), positions=torch.tensor([len(p) for p in prompts], dtype=torch.int64, device=device), + rope_positions=None, mrope_cos_sin=None, mm_embeds=None, attn_metadata=None, ) with ctx.forward_batch(decode): diff --git a/tests/models/qwen4_exp/test_tp_shard.py b/tests/models/qwen4_exp/test_tp_shard.py new file mode 100644 index 000000000..87a305ef0 --- /dev/null +++ b/tests/models/qwen4_exp/test_tp_shard.py @@ -0,0 +1,92 @@ +"""TP sharding of the qwen4_exp dense weights (pure tensor math, no TP runtime needed).""" + +from types import SimpleNamespace + +import torch +from freetoken.models.qwen4_exp.weight import _shard, _shard_rows + + +def _cfg(): + g = SimpleNamespace( + num_key_heads=4, key_head_dim=8, num_value_heads=12, value_head_dim=8 + ) + return SimpleNamespace( + num_qo_heads=6, num_kv_heads=2, head_dim=8, linear_attention_group=lambda: g + ) + + +def _gather(name, t, cfg, world, dim=0): + return torch.cat([_shard(name, t, cfg, r, world) for r in range(world)], dim=dim) + + +def test_shard_rows_splits_each_part_by_head(): + t = torch.arange(4 * 4 + 2 * 4).reshape( + -1, 1 + ) # part A: 4 heads x 4 rows, part B: 2 heads x 4 + s = [_shard_rows(t, [(4, 4), (2, 4)], r, 4) for r in range(4)] + assert all(x.shape[0] == 8 for x in s) + assert s[0][:4].flatten().tolist() == [0, 1, 2, 3] + assert s[3][:4].flatten().tolist() == [12, 13, 14, 15] + # 2 kv heads over 4 ranks replicate: rank * heads // world -> heads 0, 0, 1, 1 + assert s[0][4:].equal(s[1][4:]) and s[2][4:].flatten().tolist() == [20, 21, 22, 23] + + +def test_shard_round_trips_the_fused_projections(): + cfg, world = _cfg(), 2 + qkv = torch.randn(6 * 16 + 2 * 8 + 2 * 8, 5) + assert ( + _gather("model.layers.0.self_attn.qkv_proj.weight", qkv, cfg, world).shape + == qkv.shape + ) + q0 = _shard("model.layers.0.self_attn.qkv_proj.weight", qkv, cfg, 0, world) + assert q0.shape[0] == 3 * 16 + 8 + 8 + torch.testing.assert_close(q0[:48], qkv[:48]) # q heads 0-2 (16 rows each: q|gate) + torch.testing.assert_close(q0[48:56], qkv[96:104]) # kv head 0 of k + torch.testing.assert_close(q0[56:64], qkv[112:120]) # kv head 0 of v + kd, vd, nv = 4 * 8, 12 * 8, 12 + in_proj = torch.randn(kd + kd + vd + vd + nv + nv, 5) + s = _shard("model.layers.1.linear_attn.in_proj.weight", in_proj, cfg, 1, world) + assert s.shape[0] == (kd + kd + vd + vd + nv + nv) // 2 + torch.testing.assert_close(s[:16], in_proj[16:32]) # q: k heads 2,3 + torch.testing.assert_close(s[-6:], in_proj[-6:]) # a: v heads 6-11 + conv = torch.randn(kd + kd + vd, 1, 4) + assert _shard( + "model.layers.1.linear_attn.conv1d.weight", conv, cfg, 0, world + ).shape == ((kd + kd + vd) // 2, 1, 4) + a_log = torch.randn(nv) + torch.testing.assert_close( + _shard("model.layers.1.linear_attn.A_log", a_log, cfg, 1, world), a_log[6:] + ) + + +def test_shard_row_parallel_and_vocab_and_replicated(): + cfg, world = _cfg(), 2 + for name in ( + "model.layers.0.self_attn.o_proj.weight", + "model.layers.1.linear_attn.out_proj.weight", + "model.layers.0.mlp.shared_expert.down_proj.weight", + ): + t = torch.randn(3, 8) + torch.testing.assert_close(_gather(name, t, cfg, world, dim=1), t) + gate_up = torch.randn(2 * 6, 3) + s1 = _shard( + "model.layers.0.mlp.shared_expert.gate_up_proj.weight", gate_up, cfg, 1, world + ) + torch.testing.assert_close(s1, torch.cat([gate_up[3:6], gate_up[9:12]])) + emb = torch.randn(10, 3) + torch.testing.assert_close( + _gather("model.embed_tokens.weight", emb, cfg, world), emb + ) + torch.testing.assert_close(_gather("lm_head.weight", emb, cfg, world), emb) + for name in ( + "model.layers.0.mlp.gate.weight", + "model.layers.0.self_attn.indexer.index_qk_proj.weight", + "model.layers.0.attn_hyper_connection.input_mix_weight_down_block_inject.weight", + "model.layers.1.ple.value_proj.weight", + "model.layers.0.mlp.shared_expert_gate.weight", + ): + t = torch.randn(4, 6) + assert _shard(name, t, cfg, 1, world).equal(t) + assert _shard("model.layers.0.self_attn.qkv_proj.weight", gate_up, cfg, 0, 1).equal( + gate_up + ) diff --git a/tests/models/test_nvfp4_banks_tp.py b/tests/models/test_nvfp4_banks_tp.py new file mode 100644 index 000000000..b943cff1c --- /dev/null +++ b/tests/models/test_nvfp4_banks_tp.py @@ -0,0 +1,64 @@ +"""TP sharding of the NVFP4 routed-expert stream. + +Every rank reads the same checkpoint tensors and keeps only its own slice of the expert +intermediate axis, so the ranks' slices must partition the original exactly -- the MoE layer +all-reduces the partial sums, which is only correct if nothing is dropped or double-counted. +The expert kernel sizes its banks from ``MoEConfig.local_intermediate``, the same split. +""" + +import pytest +import torch +from freetoken.layers.quantization.moe.base import MoEConfig +from freetoken.models.nvfp4_banks import _tp_shard + +INTER, HIDDEN = 640, 2560 + + +def _roles(): + """One expert's checkpoint tensors, in the modelopt layout.""" + return { + "gate": torch.arange(INTER * (HIDDEN // 2), dtype=torch.uint8).reshape(INTER, HIDDEN // 2), + "gate_scale": torch.arange(INTER * (HIDDEN // 16), dtype=torch.uint8).reshape(INTER, HIDDEN // 16), + "gate_global": torch.tensor([2.5]), + "up": torch.arange(INTER * (HIDDEN // 2), dtype=torch.uint8).reshape(INTER, HIDDEN // 2), + "up_scale": torch.arange(INTER * (HIDDEN // 16), dtype=torch.uint8).reshape(INTER, HIDDEN // 16), + "up_global": torch.tensor([2.5]), + "down": torch.arange(HIDDEN * (INTER // 2), dtype=torch.uint8).reshape(HIDDEN, INTER // 2), + "down_scale": torch.arange(HIDDEN * (INTER // 16), dtype=torch.uint8).reshape(HIDDEN, INTER // 16), + "down_global": torch.tensor([3.5]), + } + + +@pytest.mark.parametrize("tp_size", [1, 2, 4]) +def test_rank_slices_partition_the_original_tensor(tp_size): + full = _roles() + for role, tensor in full.items(): + parts = [_tp_shard(role, tensor, INTER, tp_size, r) for r in range(tp_size)] + if role.endswith("_global"): + # per-tensor scalar: no I axis, every rank keeps it whole + assert all(torch.equal(p, tensor) for p in parts) + continue + axis = 1 if role.startswith("down") else 0 + assert torch.equal(torch.cat(parts, dim=axis), tensor), role + assert all(p.shape[axis] == tensor.shape[axis] // tp_size for p in parts), role + + +def test_tp1_is_the_identity(): + for role, tensor in _roles().items(): + assert _tp_shard(role, tensor, INTER, 1, 0) is tensor or torch.equal( + _tp_shard(role, tensor, INTER, 1, 0), tensor + ) + + +def test_a_shard_that_splits_a_scale_block_is_rejected(): + # scales cover 16 values; a rank slice that is not a whole number of blocks cannot be + # sliced consistently across the weight and its scale + with pytest.raises(AssertionError, match="16-wide scale blocks"): + _tp_shard("gate", torch.zeros(24, 8, dtype=torch.uint8), 24, 3, 0) + + +def test_moe_config_local_intermediate_matches_the_shard(): + for tp_size in (1, 2, 4): + cfg = MoEConfig(num_experts=8, hidden=HIDDEN, intermediate=INTER, top_k=2, tp_size=tp_size) + assert cfg.local_intermediate == INTER // tp_size + assert _tp_shard("gate", _roles()["gate"], INTER, tp_size, 0).shape[0] == cfg.local_intermediate diff --git a/tests/scheduler/test_cost_accounting_core.py b/tests/scheduler/test_cost_accounting_core.py index 8942c4cac..9ea850a11 100644 --- a/tests/scheduler/test_cost_accounting_core.py +++ b/tests/scheduler/test_cost_accounting_core.py @@ -50,8 +50,8 @@ def _tokenize_msg(uid: int) -> TokenizeMsg: def test_successful_tokenization_does_not_account_prompt_before_admission(): class Tokenizer: - def tokenize(self, messages): - return [torch.tensor([10, 11, 12], dtype=torch.int32)] + def tokenize_one(self, msg): + return torch.tensor([10, 11, 12], dtype=torch.int32), None ok, tensors, errors = _tokenize_requests(Tokenizer(), [_tokenize_msg(1)], _Logger()) assert [msg.uid for msg in ok] == [1] @@ -61,11 +61,10 @@ def tokenize(self, messages): def test_tokenization_failure_and_empty_prompt_are_terminal_without_usage(): class Tokenizer: - def tokenize(self, messages): - uid = messages[0].uid - if uid == 2: + def tokenize_one(self, msg): + if msg.uid == 2: raise ValueError("bad template") - return [torch.empty(0, dtype=torch.int32)] + return torch.empty(0, dtype=torch.int32), None logger = _Logger() ok, tensors, errors = _tokenize_requests( diff --git a/tests/scheduler/test_mm_cache_key.py b/tests/scheduler/test_mm_cache_key.py new file mode 100644 index 000000000..dc7415a60 --- /dev/null +++ b/tests/scheduler/test_mm_cache_key.py @@ -0,0 +1,60 @@ +"""Image prompts take part in the prefix cache through ``cache_ids``: the same image hits, a +different image (same placeholder ids) matches only the text prefix, and a text prompt never +matches a placeholder run.""" + +import torch + +from freetoken.core import Req, SamplingParams +from freetoken.scheduler.cache import CacheManager +from freetoken.scheduler.utils import PendingReq + +IMAGE = 9 +KEY = 1 << 30 + + +def _ids(*toks): + return torch.tensor(toks, dtype=torch.int32) + + +def _finish(cm, uid, input_ids, cache_ids, table_idx): + """Admit, 'prefill', and commit one request into the tree.""" + sp = SamplingParams(max_tokens=1) + mr = cm.match_req(PendingReq(uid, input_ids, sp, cache_ids=cache_ids)) + req = Req( + input_ids=input_ids, + table_idx=table_idx, + cached_len=mr.cuda_handle.cached_len, + output_len=1, + uid=uid, + sampling_params=sp, + cache_handle=mr.cuda_handle, + mm_embeds=None if cache_ids is None else torch.zeros(3, 2), + cache_ids=cache_ids, + ) + cm.lock(mr.cuda_handle) + cm.allocate_paged([req]) + req.complete_one() + cm.cache_req(req, finished=True) + return req + + +def _cached(cm, input_ids, cache_ids=None): + return cm.match_req( + PendingReq(99, input_ids, SamplingParams(max_tokens=1), cache_ids=cache_ids) + ).cuda_handle.cached_len + + +def test_image_prompt_hits_only_for_the_same_image(): + cm = CacheManager(64, 1, torch.zeros(4, 64, dtype=torch.int32), "radix") + prompt = _ids(1, 2, IMAGE, IMAGE, IMAGE, 3) + same = _ids(1, 2, KEY + 7, KEY + 8, KEY + 9, 3) + other = _ids(1, 2, KEY + 70, KEY + 71, KEY + 72, 3) + _finish(cm, 1, prompt, same, table_idx=0) + # match_req keys on all but the last token, so a follow-up turn hits the whole prompt + follow = torch.cat([prompt, _ids(4, 5)]) + assert _cached(cm, follow, torch.cat([same, _ids(4, 5)])) == len(prompt) + # a different image behind the same placeholder ids only matches the text before it + assert _cached(cm, follow, torch.cat([other, _ids(4, 5)])) == 2 + # a text prompt with the raw placeholder ids never matches a keyed run + assert _cached(cm, follow) == 2 + cm.check_integrity() diff --git a/tests/scheduler/test_mm_window.py b/tests/scheduler/test_mm_window.py new file mode 100644 index 000000000..c81f6d3e9 --- /dev/null +++ b/tests/scheduler/test_mm_window.py @@ -0,0 +1,29 @@ +"""Chunked prefill of an image prompt: each chunk scatters only the placeholder rows it holds.""" + +from types import SimpleNamespace + +import torch + +from freetoken.scheduler.scheduler import _mm_embeds_window + +IMG = 7 + + +def test_window_rows_follow_the_placeholders_across_chunks(): + ids = torch.tensor([1, 2, IMG, IMG, IMG, 3, IMG, IMG, 4, 5]) + embeds = ( + torch.arange(5).unsqueeze(1).float() + ) # one row per placeholder, prompt order + chunks = [ + (0, 4), + (4, 8), + (8, 10), + ] # a cut inside the first image run, one after the second + got = [] + for start, end in chunks: + req = SimpleNamespace(input_ids=ids[:end], cached_len=start, mm_embeds=embeds) + got.append(_mm_embeds_window(req, IMG)) + assert [g.shape[0] for g in got] == [2, 3, 0] + assert torch.equal(torch.cat(got), embeds) + whole = SimpleNamespace(input_ids=ids, cached_len=0, mm_embeds=embeds) + assert torch.equal(_mm_embeds_window(whole, IMG), embeds)