diff --git a/docs/api/model_package.md b/docs/api/model_package.md index db4718e6c..44c33abe5 100644 --- a/docs/api/model_package.md +++ b/docs/api/model_package.md @@ -114,6 +114,35 @@ pkg.save("output/llama/", external_data="safetensors") pkg.save("output/llama-serial/", max_workers=1) ``` +### Generic backbone-and-head packages + +Use `MultiComponentModelTask` when one non-generative package contains exactly +one backbone or encoder and one or more named head graphs: + +```python +from mobius import ( + ComponentConfig, + ComponentRole, + ComponentSpec, + MultiComponentModelTask, +) + + +class EncoderWithHeadsTask(MultiComponentModelTask): + components = ComponentSpec( + encoder=ComponentConfig("encoder", ComponentRole.ENCODER), + classifier=ComponentConfig("heads.classifier", ComponentRole.HEAD), + ) + + def build_component(self, name, component, module, config): + ... +``` + +The component names become `ModelPackage` keys. Roles are also exposed through +`model_roles`, so existing inspection and optimization code treats head graphs +as heads rather than decoder models. Dotted module paths allow components such +as `heads.classifier` to be resolved from the root module. + ## Output Layout - **Single model**: `directory/model.onnx` + `directory/model.onnx.data` diff --git a/docs/model-catalog.md b/docs/model-catalog.md index c42960214..5e9c882a7 100644 --- a/docs/model-catalog.md +++ b/docs/model-catalog.md @@ -1,5 +1,92 @@ # Model Catalog +## Non-generative decision models + +Mobius exports `Contrastive-LM/CLM-v0.1-8B` with one `ModelPackage` +containing its headless Qwen3 encoder, state/action projection heads, and +scaled-cosine scorer (`CLMModel` + `CLMTask`). The CLM release did **not** pin a +Qwen3-8B base revision: generic construction records the base as `unpinned`; +`build_clm_package(..., base_revision=)` requires the caller's selected +revision and records it as reproducible provenance. + +The v0.1 heads first L2-normalize each selected Qwen embedding, then use depth +3: input `Linear + GELU`, one +`Linear + LayerNorm + GELU` hidden block (optionally residual), output +`Linear`, then L2 normalization. The headless encoder emits the complete +`token_hidden_states` tensor, contains no vocabulary projection, and has no +generation-cache inputs or outputs. CLM uses +last-token pooling, but neither the upstream padding side nor its EOS behavior +is pinned. Callers must select the last attended position using their +attention mask (the deterministic `clm_last_token_indices` helper is provided) +before passing `[items, 4096]` embeddings to either projection head. Mobius +does not assume that sequence position `-1` is a real token. + +`jaredpalmer/kev-4b` similarly exports its headless Qwen3.5 backbone and grouped +pointer head (`KevModel` + `KevTask`). Its published Qwen3.5 base revision is +pinned. The backbone package exposes only input IDs, attention/position inputs, +and token hidden states; KV, convolution, and recurrent generation caches are +not part of this scoring export. `build_kev_package` accepts a dense state dict +after the published PEFT +adapter has been merged. Mobius does not currently merge arbitrary PEFT LoRA +artifacts into Qwen3.5: passing separate base/adapter mappings fails with an +actionable instruction to use PEFT `merge_and_unload()` first, rather than +silently exporting the unadapted base. + +The deterministic request adapters in `mobius.models.decision` implement each +model's verified structured rendering, candidate mapping, stable grouped +softmax, confidence, and ranking contracts. They require a caller-provided +tokenizer and local tensors; export and preprocessing do not use an HTTP or +session API. Checkpoint mapping helpers accept already-deserialized mappings, +so applications retain control over pickle deserialization. Production build +helpers validate the complete published checkpoint metadata, exact head key +sets, tensor shapes, and scalar calibration before constructing any graph. +Flexible `synthetic_*_checkpoint` helpers are test fixtures and intentionally +do not pass these production contracts. + +The Kev adapter enforces the published 8,192-token causal-row limit, 1–255 +option bounds, and a non-empty question set. Non-strict mode truncates state +user tokens to 8,191 before adding the state control token; strict mode rejects +that overflow. Every question remains an independent state-plus-question row; +`batch_kev_rows` produces padded backbone inputs and matching local +decide/option indices for the pointer graph. Score confidence is +`max(0, 1 - Σ pᵢ|i-mode|/(L-1))`, matching the public adapter. + +### Decision-model parity tests + +`tests/kev_parity_test.py` and `tests/clm_parity_test.py` contain fast synthetic +ONNX-versus-PyTorch tests that run in the normal CPU test lane: + +```bash +pytest tests/kev_parity_test.py tests/clm_parity_test.py -m "not integration" +``` + +The full tests are opt-in because they load the exported FP32 package and its +reference model on CUDA. Keep user-site packages disabled when the selected +Conda environment contains the intended ONNX Runtime GPU build. + +```bash +# Pinned KEV base, adapter, preprocessing, hidden-state readouts, pointer head, +# probabilities, and typed answers. +PYTHONNOUSERSITE=1 \ +MOBIUS_DECISION_EXPORT_ROOT=/path/to/kev-4b-fp32 \ +MOBIUS_KEV_ALLOW_DOWNLOAD=1 \ +pytest tests/kev_parity_test.py -m integration + +# CLM with an explicitly selected Qwen3 revision. +PYTHONNOUSERSITE=1 \ +CLM_RUN_FULL_PARITY=1 \ +CLM_BASE_REVISION= \ +MOBIUS_DECISION_EXPORT_ROOT=/path/to/clm-v0.1-8b-fp32 \ +pytest tests/clm_parity_test.py -m integration -k exported_clm +``` + +CLM's official serving path remains a separate qualification because upstream +did not pin Qwen3, vLLM, or tokenizer/pooling versions. Start and record an +external `vllm serve ... --runner pooling` instance, then set +`CLM_RUN_VLLM_PARITY=1`, `CLM_EMBED_URL`, `CLM_EMBED_MODEL`, and +`CLM_VLLM_METADATA` to compare its normalized embeddings and probabilities +against the same exported package. + This user-facing catalog groups registered HuggingFace model types by task and lists their module classes and example model IDs. Diffusers pipeline components are described separately because they are not entries in the model registry. Neither list implies GGUF diff --git a/src/mobius/__init__.py b/src/mobius/__init__.py index e4c802f1d..1e7489340 100644 --- a/src/mobius/__init__.py +++ b/src/mobius/__init__.py @@ -21,6 +21,12 @@ "BaseModelConfig", "CausalLMConfig", "CausalLMTask", + "CLMModel", + "CLMProjectionHead", + "CLMTask", + "ComponentConfig", + "ComponentRole", + "ComponentSpec", "ComponentInfo", "SharedWeightEndpoint", "SharedWeightInfo", @@ -42,6 +48,10 @@ "ModelRegistration", "ModelRegistry", "ModelTask", + "MultiComponentModelTask", + "KevModel", + "KevPointerHead", + "KevTask", "MLPWorldModel", "MMSConfig", "OPSET_VERSION", @@ -58,11 +68,13 @@ "adapter_source_from_onnx_adapter", "attach_peft_adapter", "build", + "build_clm_package", "build_context", "build_diffusers_pipeline", "build_from_gguf", "build_from_module", "build_from_nemo", + "build_kev_package", "compose_adapter_deltas", "components", "ep_capabilities", @@ -154,5 +166,23 @@ from mobius.integrations.gguf import build_from_gguf from mobius.integrations.nemo import build_from_nemo from mobius.integrations.transformers import build -from mobius.models import MLPWorldModel -from mobius.tasks import CausalLMTask, ModelTask, WorldModelTask +from mobius.models import ( + CLMModel, + CLMProjectionHead, + KevModel, + KevPointerHead, + MLPWorldModel, + build_clm_package, + build_kev_package, +) +from mobius.tasks import ( + CausalLMTask, + CLMTask, + ComponentConfig, + ComponentRole, + ComponentSpec, + KevTask, + ModelTask, + MultiComponentModelTask, + WorldModelTask, +) diff --git a/src/mobius/_model_package.py b/src/mobius/_model_package.py index e80dcdb00..6aee34c28 100644 --- a/src/mobius/_model_package.py +++ b/src/mobius/_model_package.py @@ -1415,7 +1415,7 @@ def apply_weights( prefix_map: dict[str, str] | None = None, *, fold_constants: bool = True, - ) -> None: + ) -> set[str]: """Apply weights from a state dict across component models. For single-component packages, all weights are applied to the sole @@ -1432,6 +1432,9 @@ def apply_weights( Weights whose name starts with a prefix are applied to the named component (with the prefix stripped). Unmatched weights are applied to all components. + + Returns: + The original state-dict names that matched package initializers. """ applied: set[str] = set() @@ -1448,7 +1451,7 @@ def apply_weights( routed: dict[str, dict[str, torch.Tensor]] = {name: {} for name in self.data} unmatched: dict[str, torch.Tensor] = {} # Track original HF names for weights that get stripped - stripped_to_original: dict[str, str] = {} + stripped_to_original: dict[str, dict[str, str]] = {name: {} for name in self.data} for weight_name, tensor in state_dict.items(): matched = False @@ -1456,7 +1459,7 @@ def apply_weights( if weight_name.startswith(prefix): stripped = weight_name[len(prefix) :].lstrip(".") routed[component][stripped] = tensor - stripped_to_original[stripped] = weight_name + stripped_to_original[component][stripped] = weight_name matched = True break if not matched: @@ -1467,7 +1470,7 @@ def apply_weights( self.data[component_name], component_weights ) for s in applied_stripped: - applied.add(stripped_to_original.get(s, s)) + applied.add(stripped_to_original[component_name].get(s, s)) # Try unmatched weights against all models if unmatched: @@ -1477,13 +1480,14 @@ def apply_weights( _log_weight_mapping(state_dict, applied) if not fold_constants: - return + return applied # Fold constants now that weights have been loaded. # PackQKV emits Concat(w_q, w_k, w_v) in the graph; those nodes can only # be constant-folded once the weight tensors carry their const_value. for model in self.data.values(): fold_initializers_after_weights(model) + return applied @contextmanager diff --git a/src/mobius/_model_package_test.py b/src/mobius/_model_package_test.py index fff6cdcad..3427eadb4 100644 --- a/src/mobius/_model_package_test.py +++ b/src/mobius/_model_package_test.py @@ -1477,6 +1477,26 @@ def test_multi_component_with_prefix_map(self): ) assert model1.graph.initializers[init_name].const_value is not None + def test_prefix_routes_report_colliding_stripped_names(self): + config = make_config() + pkg1 = build_from_module(CausalLMModel(config), config) + pkg2 = build_from_module(CausalLMModel(config), config) + model1 = pkg1["model"] + model2 = pkg2["model"] + pkg = ModelPackage({"text": model1, "vision": model2}) + + init_name = next(iter(model1.graph.initializers.keys())) + shape = list(model1.graph.initializers[init_name].shape) + applied = pkg.apply_weights( + { + f"text.{init_name}": torch.ones(shape), + f"vision.{init_name}": torch.zeros(shape), + }, + prefix_map={"text.": "text", "vision.": "vision"}, + ) + + assert applied == {f"text.{init_name}", f"vision.{init_name}"} + class TestBuildPackageFromModule: def test_returns_model_package(self): diff --git a/src/mobius/components/_gated_deltanet.py b/src/mobius/components/_gated_deltanet.py index 44be2185d..b2eac4c8d 100644 --- a/src/mobius/components/_gated_deltanet.py +++ b/src/mobius/components/_gated_deltanet.py @@ -164,8 +164,8 @@ def forward( self, op: OpBuilder, hidden_states: ir.Value, - conv_state: ir.Value, - recurrent_state: ir.Value, + conv_state: ir.Value | None, + recurrent_state: ir.Value | None, ): """Forward pass for the Gated DeltaNet layer. @@ -184,6 +184,38 @@ def forward( new_recurrent_state: (batch, num_v_heads, k_dim, v_dim) """ batch_dim = op.Shape(hidden_states, start=0, end=1) + if conv_state is None: + conv_state = op.Expand( + op.CastLike(op.Constant(value_float=0.0), hidden_states), + op.Concat( + batch_dim, + op.Constant( + value_ints=[ + self.conv_dim, + self.conv_kernel_size - 1, + ] + ), + axis=0, + ), + ) + if recurrent_state is None: + recurrent_state = op.Expand( + op.Cast( + op.Constant(value_float=0.0), + to=self._stash_type, + ), + op.Concat( + batch_dim, + op.Constant( + value_ints=[ + self.num_v_heads, + self.head_k_dim, + self.head_v_dim, + ] + ), + axis=0, + ), + ) # === Projections === mixed_qkv = self.in_proj_qkv(op, hidden_states) diff --git a/src/mobius/components/_gated_deltanet_test.py b/src/mobius/components/_gated_deltanet_test.py index fb72fee63..c3753b5ae 100644 --- a/src/mobius/components/_gated_deltanet_test.py +++ b/src/mobius/components/_gated_deltanet_test.py @@ -87,6 +87,19 @@ def test_forward_builds_graph(self): assert count_op_type(graph, "Scan") == 0 assert count_op_type(graph, "Conv") == 0 + def test_forward_initializes_missing_states(self): + config = self._make_deltanet_config() + dn = GatedDeltaNet(config) + builder, op, graph = create_test_builder() + hidden = create_test_input(builder, "hidden", [1, 4, 64]) + + output, new_conv, new_rec = dn(op, hidden, None, None) + builder._adapt_outputs([output, new_conv, new_rec], "") + + assert count_op_type(graph, "Expand") >= 2 + assert count_op_type(graph, "CausalConvWithState") == 1 + assert count_op_type(graph, "LinearAttention") == 1 + def test_forward_has_sigmoid_for_beta(self): config = self._make_deltanet_config() dn = GatedDeltaNet(config) diff --git a/src/mobius/models/__init__.py b/src/mobius/models/__init__.py index 11b9ef922..d5e5ca552 100644 --- a/src/mobius/models/__init__.py +++ b/src/mobius/models/__init__.py @@ -23,6 +23,12 @@ "ChatGLMCausalLMModel", "CodeGenCausalLMModel", "CodeShellCausalLMModel", + "CLMModel", + "CLMProjectionHead", + "CLM_PROVENANCE", + "CheckpointContractError", + "ModelProvenance", + "Qwen3EncoderModel", "AutoencoderKLCogVideoXModel", "CogVideoXTransformer3DModel", "CogVideoXVAEConfig", @@ -102,6 +108,31 @@ "JetMoeCausalLMModel", "KimiK3CausalLMModel", "KimiLinearCausalLMModel", + "KevModel", + "KevPointerHead", + "KEV_PROVENANCE", + "Qwen35EncoderModel", + "batch_kev_rows", + "build_clm_package", + "build_kev_package", + "clm_answer", + "clm_candidates", + "clm_last_token_indices", + "clm_pairs", + "clm_state_text", + "encode_kev_rows", + "grouped_softmax", + "kev_answer", + "map_clm_checkpoint", + "map_kev_checkpoint", + "rank_probabilities", + "render_clm", + "render_kev", + "synthetic_clm_checkpoint", + "synthetic_kev_checkpoint", + "validate_temperature", + "validate_clm_checkpoint", + "validate_kev_checkpoint", "Llama4CausalLMModel", "LlamaEmbedGGUFModel", "DreamModel", @@ -252,6 +283,39 @@ from mobius.models.cosmos import Cosmos3EdgeTextModel, Cosmos3EdgeVLModel from mobius.models.cosmos3_omni import Cosmos3OmniReasonerModel from mobius.models.ctrl import CTRLCausalLMModel +from mobius.models.decision import ( + CLM_PROVENANCE, + KEV_PROVENANCE, + CheckpointContractError, + CLMModel, + CLMProjectionHead, + KevModel, + KevPointerHead, + ModelProvenance, + Qwen3EncoderModel, + Qwen35EncoderModel, + batch_kev_rows, + build_clm_package, + build_kev_package, + clm_answer, + clm_candidates, + clm_last_token_indices, + clm_pairs, + clm_state_text, + encode_kev_rows, + grouped_softmax, + kev_answer, + map_clm_checkpoint, + map_kev_checkpoint, + rank_probabilities, + render_clm, + render_kev, + synthetic_clm_checkpoint, + synthetic_kev_checkpoint, + validate_clm_checkpoint, + validate_kev_checkpoint, + validate_temperature, +) from mobius.models.deepseek import DeepSeekV3CausalLMModel from mobius.models.deepseek_ocr2 import DeepSeekOCR2CausalLMModel from mobius.models.deepseek_v4 import DeepSeekV4CausalLMModel diff --git a/src/mobius/models/decision.py b/src/mobius/models/decision.py new file mode 100644 index 000000000..384cff550 --- /dev/null +++ b/src/mobius/models/decision.py @@ -0,0 +1,1217 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Non-generative CLM and Kev model components and deterministic adapters.""" + +from __future__ import annotations + +import dataclasses +import json +import math +import re +from collections.abc import Callable, Mapping, Sequence +from typing import Any, Literal + +import onnx_ir as ir +import torch +from onnxscript import OpBuilder, nn + +from mobius._configs import ArchitectureConfig +from mobius.components._common import LayerNorm, Linear +from mobius.models.qwen import Qwen3CausalLMModel +from mobius.models.qwen35 import Qwen35CausalLMModel + +JSONContent = str | int | float | bool | None | list["JSONContent"] | dict[str, "JSONContent"] + +CLM_MODEL_ID = "Contrastive-LM/CLM-v0.1-8B" +CLM_REVISION = "e939398d4556fcd9400c76fa8c5a513202f42b0a" +CLM_BASE_MODEL_ID = "Qwen/Qwen3-8B" +KEV_MODEL_ID = "jaredpalmer/kev-4b" +KEV_REVISION = "139fdd94f1b6a6ad80cc15e08fcb99cac885a101" +KEV_BASE_MODEL_ID = "Qwen/Qwen3.5-4B-Base" +KEV_BASE_REVISION = "1001bb4d826a52d1f399e183466143f4da7b741b" +KEV_HIDDEN_SIZE = 2560 +KEV_POINTER_SIZE = 256 +KEV_TEMPERATURE = 2.406050072164233 +KEV_MAX_OPTIONS = 255 +KEV_MAX_STATE_TOKENS = 8192 +KEV_MAX_ROW_TOKENS = 8192 +CLM_HIDDEN_SIZE = 4096 +CLM_PROJECTION_DIM = 512 +CLM_WIDTH = 1536 + + +@dataclasses.dataclass(frozen=True) +class ModelProvenance: + """Immutable identities needed to reproduce an exported decision model.""" + + model_id: str + revision: str + base_model_id: str + base_revision: str | None + + def as_metadata(self) -> dict[str, str]: + """Return string metadata, marking an absent base revision unpinned.""" + return { + "model_id": self.model_id, + "revision": self.revision, + "base_model_id": self.base_model_id, + "base_revision": self.base_revision or "unpinned", + "reproducible": str(self.base_revision is not None).lower(), + } + + +CLM_PROVENANCE = ModelProvenance(CLM_MODEL_ID, CLM_REVISION, CLM_BASE_MODEL_ID, None) +KEV_PROVENANCE = ModelProvenance( + KEV_MODEL_ID, KEV_REVISION, KEV_BASE_MODEL_ID, KEV_BASE_REVISION +) + + +class CheckpointContractError(ValueError): + """A checkpoint does not match the pinned model-specific export contract.""" + + +def _metadata_matches(actual: Any, expected: Any) -> bool: + """Compare checkpoint metadata without accepting bool/int coercions.""" + if isinstance(expected, bool): + return actual is expected + if isinstance(expected, int): + return isinstance(actual, int) and not isinstance(actual, bool) and actual == expected + if isinstance(expected, float): + return isinstance(actual, float) and actual == expected + return isinstance(actual, type(expected)) and actual == expected + + +def _validate_weight_mapping(weights: Mapping[str, Any], *, owner: str) -> None: + """Require a non-empty string-to-tensor weight mapping.""" + if not isinstance(weights, Mapping) or not weights: + raise CheckpointContractError(f"{owner} weights must be a non-empty mapping") + invalid = [ + key + for key, value in weights.items() + if not isinstance(key, str) or not isinstance(value, torch.Tensor) + ] + if invalid: + raise CheckpointContractError( + f"{owner} weights must contain only string-to-tensor entries" + ) + + +def _apply_package_weights_strict( + package, weights: dict[str, torch.Tensor], *, owner: str +) -> None: + """Apply every supplied weight and require every initializer to be bound.""" + applied = package.apply_weights(weights) + unapplied = sorted(set(weights) - applied) + if unapplied: + examples = ", ".join(repr(name) for name in unapplied[:5]) + suffix = f" and {len(unapplied) - 5} more" if len(unapplied) > 5 else "" + raise CheckpointContractError( + f"{owner} contains unapplied weights: {examples}{suffix}" + ) + missing = [ + f"{component}.{name}" + for component, model in package.items() + for name, initializer in model.graph.initializers.items() + if initializer.const_value is None + ] + if missing: + examples = ", ".join(repr(name) for name in missing[:5]) + suffix = f" and {len(missing) - 5} more" if len(missing) > 5 else "" + raise CheckpointContractError( + f"{owner} leaves initializers without weights: {examples}{suffix}" + ) + + +def _require_tensor_shape( + state_dict: Mapping[str, Any], key: str, shape: tuple[int, ...], *, owner: str +) -> None: + """Require one checkpoint tensor with an exact name, type, and shape.""" + value = state_dict.get(key) + if not isinstance(value, torch.Tensor): + raise CheckpointContractError(f"{owner}.{key} must be a torch.Tensor") + if tuple(value.shape) != shape: + raise CheckpointContractError( + f"{owner}.{key} shape must be {shape}, got {tuple(value.shape)}" + ) + + +def validate_clm_checkpoint(checkpoint: Mapping[str, Any]) -> None: + """Fail closed unless *checkpoint* is the published CLM-v0.1 head layout.""" + if not isinstance(checkpoint, Mapping): + raise CheckpointContractError("CLM checkpoint must be a mapping") + cfg = checkpoint.get("cfg") + if not isinstance(cfg, Mapping): + raise CheckpointContractError("CLM checkpoint must contain a cfg mapping") + expected = { + "model": CLM_BASE_MODEL_ID, + "hidden_size": CLM_HIDDEN_SIZE, + "width": CLM_WIDTH, + "depth": 3, + "activation": "gelu", + "layernorm": True, + "residual": False, + } + for key, value in expected.items(): + if key not in cfg or not _metadata_matches(cfg[key], value): + raise CheckpointContractError( + f"CLM checkpoint cfg.{key} must be {value!r}, got {cfg.get(key)!r}" + ) + projection_dim = checkpoint.get("projection_dim", cfg.get("projection_dim")) + if not _metadata_matches(projection_dim, CLM_PROJECTION_DIM): + raise CheckpointContractError("CLM checkpoint projection_dim must be 512") + shapes = { + "inp.weight": (CLM_WIDTH, CLM_HIDDEN_SIZE), + "inp.bias": (CLM_WIDTH,), + "hidden.0.weight": (CLM_WIDTH, CLM_WIDTH), + "hidden.0.bias": (CLM_WIDTH,), + "norms.0.weight": (CLM_WIDTH,), + "norms.0.bias": (CLM_WIDTH,), + "out.weight": (CLM_PROJECTION_DIM, CLM_WIDTH), + "out.bias": (CLM_PROJECTION_DIM,), + } + for name in ("state_head", "action_head"): + head = checkpoint.get(name) + if not isinstance(head, Mapping): + raise CheckpointContractError(f"CLM checkpoint must contain a {name} state dict") + if set(head) != set(shapes): + missing = sorted(map(str, set(shapes) - set(head))) + extra = sorted(map(str, set(head) - set(shapes))) + raise CheckpointContractError( + f"CLM {name} keys mismatch; missing={missing}, extra={extra}" + ) + for key, shape in shapes.items(): + _require_tensor_shape(head, key, shape, owner=name) + logit_scale = checkpoint.get("logit_scale") + if isinstance(logit_scale, bool) or not isinstance( + logit_scale, (int, float, torch.Tensor) + ): + raise CheckpointContractError("CLM logit_scale must be a scalar") + if isinstance(logit_scale, torch.Tensor) and logit_scale.numel() != 1: + raise CheckpointContractError("CLM logit_scale must be a scalar") + + +def validate_kev_checkpoint(checkpoint: Mapping[str, Any]) -> None: + """Fail closed unless *checkpoint* matches the published Kev-4B head.""" + if not isinstance(checkpoint, Mapping): + raise CheckpointContractError("Kev checkpoint must be a mapping") + expected = { + "base": KEV_BASE_MODEL_ID, + "base_revision": KEV_BASE_REVISION, + "head_dim": KEV_POINTER_SIZE, + "option_isolation": False, + "temperature": KEV_TEMPERATURE, + } + for key, value in expected.items(): + actual = checkpoint.get(key) + if not _metadata_matches(actual, value): + raise CheckpointContractError( + f"Kev checkpoint {key} must be {value!r}, got {actual!r}" + ) + head = checkpoint.get("head") + if not isinstance(head, Mapping): + raise CheckpointContractError("Kev checkpoint must contain a head state dict") + shapes = { + "q.weight": (KEV_POINTER_SIZE, KEV_HIDDEN_SIZE), + "q.bias": (KEV_POINTER_SIZE,), + "k.weight": (KEV_POINTER_SIZE, KEV_HIDDEN_SIZE), + "k.bias": (KEV_POINTER_SIZE,), + } + if set(head) != set(shapes): + missing = sorted(map(str, set(shapes) - set(head))) + extra = sorted(map(str, set(head) - set(shapes))) + raise CheckpointContractError( + f"Kev head keys mismatch; missing={missing}, extra={extra}" + ) + for key, shape in shapes.items(): + _require_tensor_shape(head, key, shape, owner="head") + + +def render_clm(value: JSONContent, indent: int = 0) -> str: + """Render structured CLM input exactly as the reference preprocessing does.""" + if value is None: + return "" + if isinstance(value, str): + return value + if isinstance(value, bool): + return "true" if value else "false" + if isinstance(value, (int, float)): + return str(value) + pad = " " * indent + if isinstance(value, dict): + parts = [] + for key, item in value.items(): + if isinstance(item, (dict, list)) and item: + parts.append(f"{pad}{key}:\n{render_clm(item, indent + 2)}") + else: + parts.append(f"{pad}{key}: {render_clm(item)}") + return ("\n\n" if indent == 0 else "\n").join(parts) + parts = [] + for item in value: + if isinstance(item, (dict, list)) and item: + parts.append(f"{pad}-\n{render_clm(item, indent + 2)}") + else: + parts.append(f"{pad}- {render_clm(item)}") + return "\n".join(parts) + + +def clm_state_text(state: JSONContent, instructions: JSONContent) -> str: + """Render state and instructions with the reference blank-line separator.""" + state_value = render_clm(state).strip() + instruction_value = render_clm(instructions).strip() + return ( + f"{state_value}\n\n{instruction_value}" + if state_value and instruction_value + else state_value or instruction_value + ) + + +def clm_candidates(question: Mapping[str, Any]) -> tuple[list[str], list[str]]: + """Return answer keys and exact action-head candidate strings.""" + question_type = question.get("type") + criteria = question.get("criteria") + instructions = render_clm(question.get("instructions")).strip() + if question_type == "choice": + if not isinstance(criteria, Mapping) or not criteria: + raise ValueError("choice question needs a non-empty 'criteria' object") + keys = list(criteria) + return keys, [ + render_clm(criteria[key]) if criteria[key] not in (None, "") else key + for key in keys + ] + if question_type == "score": + if not isinstance(criteria, list) or len(criteria) < 2: + raise ValueError("score question needs an ordered list of at least two levels") + return [str(i) for i in range(len(criteria))], [render_clm(item) for item in criteria] + if question_type != "noul": + raise ValueError(f"unknown question type {question_type!r}") + criteria = criteria or {} + values = [] + for key in ("false", "true"): + description = criteria.get(key) if isinstance(criteria, Mapping) else None + if description in (None, ""): + if instructions: + description = ( + f"Yes. This is true: {instructions}" + if key == "true" + else f"No. This is false: {instructions}" + ) + else: + description = key + values.append(f"{key}: {render_clm(description)}") + return ["false", "true"], values + + +def clm_pairs( + state: JSONContent, questions: Mapping[str, Mapping[str, Any]] +) -> dict[str, tuple[str, list[str], list[str]]]: + """Build each question's state text, answer keys, and action texts.""" + return { + key: ( + clm_state_text(state, question.get("instructions")), + *clm_candidates(question), + ) + for key, question in questions.items() + } + + +def clm_last_token_indices(attention_mask: Sequence[Sequence[int]]) -> list[int]: + """Return each row's last attended token for padding-aware CLM pooling.""" + indices = [] + for row_index, row in enumerate(attention_mask): + attended = [index for index, value in enumerate(row) if value] + if not attended: + raise ValueError(f"attention_mask row {row_index} has no attended token") + indices.append(attended[-1]) + return indices + + +def validate_temperature(temperature: float) -> float: + """Return a valid request temperature or reject values outside ``(0, 100]``.""" + value = float(temperature) + if not 0.0 < value <= 100.0: + raise ValueError("temperature must be in (0, 100]") + return value + + +def stable_softmax(values: Sequence[float]) -> list[float]: + """Compute overflow-safe softmax for one non-empty logit group.""" + if not values: + raise ValueError("softmax group must not be empty") + maximum = max(values) + exponentials = [math.exp(float(value) - maximum) for value in values] + total = sum(exponentials) + return [value / total for value in exponentials] + + +def grouped_softmax( + logits: Sequence[float], group_sizes: Sequence[int], *, temperature: float = 1.0 +) -> list[list[float]]: + """Split flat logits into validated groups and softmax each independently.""" + temperature = validate_temperature(temperature) + if any(size <= 0 for size in group_sizes) or sum(group_sizes) != len(logits): + raise ValueError("group_sizes must be positive and consume every logit") + result, offset = [], 0 + for size in group_sizes: + result.append( + stable_softmax([value / temperature for value in logits[offset : offset + size]]) + ) + offset += size + return result + + +def rank_probabilities( + candidates: Sequence[str], probabilities: Sequence[float] +) -> list[dict]: + """Rank candidates by descending probability with stable input-order ties.""" + if len(candidates) != len(probabilities): + raise ValueError("candidates and probabilities must have equal lengths") + order = sorted(range(len(candidates)), key=lambda index: (-probabilities[index], index)) + return [ + { + "rank": rank + 1, + "candidate": candidates[index], + "prob": float(probabilities[index]), + } + for rank, index in enumerate(order) + ] + + +def clm_answer( + question: Mapping[str, Any], keys: Sequence[str], probabilities: Sequence[float] +) -> dict[str, Any]: + """Map one CLM distribution to its exact noul/choice/score answer.""" + probs = [float(value) for value in probabilities] + if len(keys) != len(probs) or not probs: + raise ValueError("keys and probabilities must have the same non-zero length") + question_type = question["type"] + distribution = dict(zip(keys, probs)) + if question_type == "noul": + return {"type": "noul", "noul": distribution["true"]} + confidence = _clm_confidence(probs) + if question_type == "choice": + selected = max(range(len(probs)), key=probs.__getitem__) + return { + "type": "choice", + "choice": keys[selected], + "confidence": confidence, + "probabilities": distribution, + } + levels = question["criteria"] + return { + "type": "score", + "score": sum(index * value for index, value in enumerate(probs)), + "confidence": confidence, + "legend": {str(i): render_clm(level) for i, level in enumerate(levels)}, + "probabilities": distribution, + } + + +def _clm_confidence(probabilities: Sequence[float]) -> float: + """Compute CLM's top probability minus the mean competing probability.""" + if len(probabilities) < 2: + return 1.0 + selected = max(range(len(probabilities)), key=probabilities.__getitem__) + rest = [value for index, value in enumerate(probabilities) if index != selected] + return max(0.0, min(1.0, probabilities[selected] - sum(rest) / len(rest))) + + +class CLMProjectionHead(nn.Module): + """The checkpoint-compatible CLM ``hidden -> projection`` MLP.""" + + def __init__( + self, + *, + hidden_size: int = 4096, + width: int, + depth: int = 3, + projection_dim: int = 512, + activation: Literal["gelu", "relu", "silu"] = "gelu", + layernorm: bool = True, + residual: bool = False, + ): + """Construct the checkpoint-shaped projection MLP. + + ``depth`` includes input and output projections, so depth three creates + exactly one hidden block. Unsupported depths and activations fail + before or during graph construction. + """ + super().__init__() + if depth < 2: + raise ValueError("CLM projection depth must be at least 2") + self.inp = Linear(hidden_size, width) + self.hidden = nn.ModuleList([Linear(width, width) for _ in range(depth - 2)]) + self.norms = nn.ModuleList( + [ + LayerNorm(width, eps=1e-5) if layernorm else _Identity() + for _ in range(depth - 2) + ] + ) + self.out = Linear(width, projection_dim) + self.activation = activation + self.residual = residual + + def _activate(self, op: OpBuilder, value: ir.Value) -> ir.Value: + """Apply the configured reference activation to one graph value.""" + if self.activation == "gelu": + return op.Gelu(value) + if self.activation == "relu": + return op.Relu(value) + if self.activation == "silu": + return op.Mul(value, op.Sigmoid(value)) + raise ValueError(f"unsupported CLM activation {self.activation!r}") + + def forward(self, op: OpBuilder, embeddings: ir.Value) -> ir.Value: + """Project ``[items, hidden]`` embeddings and L2-normalize each row.""" + embeddings_float = op.Cast(embeddings, to=ir.DataType.FLOAT) + input_norm = op.ReduceL2(embeddings_float, [-1], keepdims=1) + normalized = op.CastLike( + op.Div(embeddings_float, op.Max(input_norm, op.CastLike(1e-12, input_norm))), + embeddings, + ) + value = self._activate(op, self.inp(op, normalized)) + for linear, norm in zip(self.hidden, self.norms): + projected = self._activate(op, norm(op, linear(op, value))) + value = op.Add(value, projected) if self.residual else projected + value = self.out(op, value) + value_float = op.Cast(value, to=ir.DataType.FLOAT) + norm = op.ReduceL2(value_float, [-1], keepdims=1) + return op.CastLike( + op.Div(value_float, op.Max(norm, op.CastLike(1e-12, norm))), + value, + ) + + +class _Identity(nn.Module): + """Graph identity used where a checkpoint omits hidden LayerNorm.""" + + def forward(self, op: OpBuilder, value: ir.Value) -> ir.Value: + """Return an ONNX identity of *value*.""" + return op.Identity(value) + + +class CLMScorer(nn.Module): + """Scaled-cosine scorer; inputs are already L2-normalized projections.""" + + def __init__(self): + """Create the learned scalar log-temperature parameter.""" + super().__init__() + self.logit_scale = nn.Parameter([1]) + + def forward( + self, + op: OpBuilder, + states: ir.Value, + actions: ir.Value, + candidate_owners: ir.Value, + temperature: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + """Score owned state/action rows and return grouped logits/probabilities.""" + scale = op.Min(op.Exp(self.logit_scale), op.CastLike(100.0, self.logit_scale)) + state_rows = op.Gather(states, candidate_owners, axis=0) + scores = op.ReduceSum(op.Mul(state_rows, actions), [-1], keepdims=0) + logits = op.Div(op.Mul(scores, scale), temperature) + return logits, _grouped_softmax_graph(op, logits, candidate_owners, states) + + +class CLMModel(nn.Module): + """Qwen3-8B encoder and both CLM-v0.1 projection heads.""" + + default_task = "clm-scoring" + category = "Text Ranking" + provenance = CLM_PROVENANCE + + def __init__( + self, + config: ArchitectureConfig, + *, + width: int, + depth: int = 3, + projection_dim: int = 512, + activation: Literal["gelu", "relu", "silu"] = "gelu", + layernorm: bool = True, + residual: bool = False, + base_revision: str | None = None, + ): + """Construct a headless Qwen3 encoder and exact CLM-v0.1 heads. + + ``base_revision`` is optional for manual graph construction; package + provenance explicitly reports such a graph as unpinned. The public + production builder requires a concrete revision. + """ + super().__init__() + if config.hidden_size != 4096: + raise ValueError("CLM-v0.1-8B requires Qwen3-8B hidden_size=4096") + self.config = config + self.provenance = dataclasses.replace(CLM_PROVENANCE, base_revision=base_revision) + self.encoder = Qwen3EncoderModel(config) + options = { + "hidden_size": config.hidden_size, + "width": width, + "depth": depth, + "projection_dim": projection_dim, + "activation": activation, + "layernorm": layernorm, + "residual": residual, + } + self.state_head = CLMProjectionHead(**options) + self.action_head = CLMProjectionHead(**options) + self.scorer = CLMScorer() + + def preprocess_weights(self, state_dict: Mapping[str, Any]) -> dict[str, torch.Tensor]: + """Map combined base/head tensors into this module's component paths.""" + mapped = map_clm_checkpoint(state_dict) + base = { + key: value + for key, value in state_dict.items() + if isinstance(key, str) + and isinstance(value, torch.Tensor) + and not key.startswith(("state_head.", "action_head.", "scorer.")) + and key != "logit_scale" + } + for key, value in self.encoder.preprocess_weights(base).items(): + mapped.setdefault(f"encoder.{key}", value) + return mapped + + +def map_clm_checkpoint(checkpoint: Mapping[str, Any]) -> dict[str, torch.Tensor]: + """Flatten a validated in-memory CLM checkpoint without deserializing code.""" + mapped: dict[str, torch.Tensor] = {} + for source, target in (("state_head", "state_head"), ("action_head", "action_head")): + values = checkpoint.get(source) + if values is None: + continue + if not isinstance(values, Mapping): + raise TypeError(f"{source} must be a state-dict mapping") + for key, value in values.items(): + if not isinstance(key, str) or not isinstance(value, torch.Tensor): + raise TypeError(f"{source} must contain string-to-tensor entries") + mapped[f"{target}.{key}"] = value + scale = checkpoint.get("logit_scale") + if scale is not None: + mapped["scorer.logit_scale"] = torch.as_tensor(scale).reshape(1) + for key, value in checkpoint.items(): + if not isinstance(key, str) or not isinstance(value, torch.Tensor): + continue + if key.startswith(("state_head.", "action_head.", "scorer.", "encoder.")): + mapped[key] = value + return mapped + + +KEV_CONTROL_TOKEN_IDS = { + "state": 248060, + "question": 248061, + "option_start": 248049, + "option_end": 248050, + "decide": 248062, +} +_KEV_DELIMITER_RE = re.compile(r"<\|([A-Za-z0-9_]+)\|>") + + +def escape_kev_delimiters(text: str) -> str: + """Rewrite Qwen control-token spellings so user text cannot forge boundaries.""" + return _KEV_DELIMITER_RE.sub(r"<¦\1¦>", text) + + +def render_kev(value: JSONContent, indent: int = 0) -> str: + """Kev rendering, including Python's capitalized bool spelling.""" + pad = " " * indent + if value is None: + return "" + if isinstance(value, (str, int, float, bool)): + return str(value) + if isinstance(value, list): + return "\n".join(f"{pad}- {render_kev(item, indent + 1).lstrip()}" for item in value) + return "\n".join( + ( + f"{pad}{key}:\n{render_kev(item, indent + 1)}" + if isinstance(item, (dict, list)) + else f"{pad}{key}: {render_kev(item)}" + ) + for key, item in value.items() + ) + + +def kev_question_options(question: Mapping[str, Any]) -> tuple[list[str], list[str]]: + """Return exact Kev answer keys/texts, enforcing type and 1-255 bounds.""" + question_type = question.get("type") + criteria = question.get("criteria") + if question_type == "noul": + criteria = criteria or {} + return ["false", "true"], [ + _kev_option_text("no", criteria.get("false")), + _kev_option_text("yes", criteria.get("true")), + ] + if question_type == "choice": + if not isinstance(criteria, Mapping) or not 1 <= len(criteria) <= KEV_MAX_OPTIONS: + raise ValueError("choice question needs 1..255 criteria") + keys = list(criteria) + return keys, [_kev_option_text(key, criteria[key]) for key in keys] + if question_type == "score": + if not isinstance(criteria, list) or not 1 <= len(criteria) <= KEV_MAX_OPTIONS: + raise ValueError("score question needs 1..255 levels") + return [str(i) for i in range(len(criteria))], [render_kev(item) for item in criteria] + raise ValueError(f"unknown question type {question_type!r}") + + +def _kev_option_text(name: str, description: JSONContent) -> str: + """Render an option key alone or followed by its structured description.""" + return name if description in (None, "") else f"{name}: {render_kev(description)}" + + +@dataclasses.dataclass(frozen=True) +class KevQuestionRow: + """One independent causal state+question row and its readout indices.""" + + question_id: str + input_ids: tuple[int, ...] + position_ids: tuple[int, ...] + decide_index: int + option_indices: tuple[int, ...] + keys: tuple[str, ...] + + +@dataclasses.dataclass(frozen=True) +class KevBatch: + """Right-padded backbone inputs and pointer-head indices.""" + + input_ids: tuple[tuple[int, ...], ...] + attention_mask: tuple[tuple[int, ...], ...] + position_ids: tuple[tuple[int, ...], ...] + decide_indices: tuple[int, ...] + option_indices: tuple[int, ...] + option_owners: tuple[int, ...] + + +def encode_kev_rows( + tokenizer: Callable[..., Any], + state: JSONContent, + questions: Mapping[str, Mapping[str, Any]], + *, + strict: bool = False, + max_state_tokens: int = KEV_MAX_STATE_TOKENS, + max_row_tokens: int = KEV_MAX_ROW_TOKENS, +) -> tuple[KevQuestionRow, ...]: + """Encode one causal row/question with the verified serving limits. + + The state control token counts toward ``max_state_tokens``. In normal + serving mode user state tokens are truncated to ``max_state_tokens - 1``; + strict mode rejects the same overflow instead. + """ + if not questions: + raise ValueError("questions must not be empty") + if max_state_tokens < 1 or max_row_tokens < max_state_tokens: + raise ValueError("token limits must be positive and row >= state") + + def tokens(text: str) -> list[int]: + """Tokenize escaped caller text without tokenizer-added special tokens.""" + encoded = tokenizer(escape_kev_delimiters(text), add_special_tokens=False) + values = encoded.input_ids if hasattr(encoded, "input_ids") else encoded["input_ids"] + return list(values) + + user_state_ids = tokens(render_kev(state)) + if strict and len(user_state_ids) + 1 > max_state_tokens: + raise ValueError(f"state exceeds {max_state_tokens} tokens") + state_ids = [ + KEV_CONTROL_TOKEN_IDS["state"], + *user_state_ids[: max_state_tokens - 1], + ] + rows = [] + for question_id, question in questions.items(): + keys, options = kev_question_options(question) + # Each question starts a fresh causal row that repeats the bounded state. + ids = [ + *state_ids, + KEV_CONTROL_TOKEN_IDS["question"], + *tokens(render_kev(question.get("instructions"))), + ] + option_indices = [] + for option in options: + ids.extend([KEV_CONTROL_TOKEN_IDS["option_start"], *tokens(option)]) + ids.append(KEV_CONTROL_TOKEN_IDS["option_end"]) + # Pointer K reads the closing boundary, matching Kev training. + option_indices.append(len(ids) - 1) + ids.append(KEV_CONTROL_TOKEN_IDS["decide"]) + if len(ids) > max_row_tokens: + raise ValueError(f"state+question row exceeds {max_row_tokens} tokens: {len(ids)}") + # The final decide token is the pointer Q location for this row. + rows.append( + KevQuestionRow( + question_id, + tuple(ids), + tuple(range(len(ids))), + len(ids) - 1, + tuple(option_indices), + tuple(keys), + ) + ) + return tuple(rows) + + +def batch_kev_rows(rows: Sequence[KevQuestionRow], *, pad_token_id: int) -> KevBatch: + """Right-pad encoded rows and flatten their aligned pointer readouts.""" + if not rows: + raise ValueError("rows must not be empty") + width = max(len(row.input_ids) for row in rows) + input_ids, attention_mask, position_ids = [], [], [] + option_indices, option_owners = [], [] + for owner, row in enumerate(rows): + padding = width - len(row.input_ids) + input_ids.append((*row.input_ids, *((pad_token_id,) * padding))) + attention_mask.append((*((1,) * len(row.input_ids)), *((0,) * padding))) + position_ids.append((*row.position_ids, *((0,) * padding))) + option_indices.extend(row.option_indices) + option_owners.extend([owner] * len(row.option_indices)) + return KevBatch( + tuple(input_ids), + tuple(attention_mask), + tuple(position_ids), + tuple(row.decide_index for row in rows), + tuple(option_indices), + tuple(option_owners), + ) + + +def _normalize_probabilities(values: Sequence[float]) -> list[float]: + """Normalize values, using a uniform distribution when their sum is zero.""" + total = sum(values) + return ( + [1.0 / len(values)] * len(values) + if total == 0 + else [float(value) / total for value in values] + ) + + +def kev_choice_confidence(probabilities: Sequence[float]) -> float: + """Compute Kev's uniform-adjusted maximum-probability confidence.""" + probs = _normalize_probabilities(probabilities) + count = len(probs) + return 1.0 if count == 1 else (max(probs) - 1 / count) / (1 - 1 / count) + + +def kev_score_confidence(probabilities: Sequence[float]) -> float: + """Compute public score confidence from expected distance to the mode.""" + probs = _normalize_probabilities(probabilities) + count = len(probs) + if count == 1: + return 1.0 + mode = max(range(count), key=probs.__getitem__) + return max( + 0.0, + 1.0 + - sum(value * abs(index - mode) for index, value in enumerate(probs)) / (count - 1), + ) + + +def _round4(value: float) -> float: + """Round a response scalar to Kev's four-decimal wire precision.""" + return round(float(value), 4) + + +def kev_answer( + question: Mapping[str, Any], keys: Sequence[str], probabilities: Sequence[float] +) -> dict[str, Any]: + """Convert a Kev option distribution into exact noul/choice/score output.""" + probs = [float(value) for value in probabilities] + question_type = question["type"] + if question_type == "noul": + return {"type": "noul", "noul": _round4(probs[1])} + distribution = {key: _round4(value) for key, value in zip(keys, probs)} + if question_type == "choice": + index = max(range(len(probs)), key=probs.__getitem__) + return { + "type": "choice", + "choice": keys[index], + "confidence": _round4(kev_choice_confidence(probs)), + "probabilities": distribution, + } + return { + "type": "score", + "score": _round4(sum(index * value for index, value in enumerate(probs))), + "legend": {str(i): render_kev(item) for i, item in enumerate(question["criteria"])}, + "probabilities": distribution, + "confidence": _round4(kev_score_confidence(probs)), + } + + +class KevPointerHead(nn.Module): + """Kev-4B checkpoint-compatible 256-dimensional pointer head.""" + + def __init__(self, hidden_size: int = KEV_HIDDEN_SIZE): + """Create checkpoint-compatible Q/K projections and fixed calibration.""" + super().__init__() + self.q = Linear(hidden_size, KEV_POINTER_SIZE) + self.k = Linear(hidden_size, KEV_POINTER_SIZE) + self.temperature = KEV_TEMPERATURE + + def forward( + self, + op: OpBuilder, + hidden_states: ir.Value, + decide_indices: ir.Value, + option_indices: ir.Value, + option_owners: ir.Value, + ) -> tuple[ir.Value, ir.Value]: + """Read indexed row tokens and return calibrated grouped scores. + + ``decide_indices`` and ``option_indices`` are positions local to each + independently padded question row; ``option_owners`` maps each option + to the corresponding row. + """ + # Form [row, token] coordinates so GatherND cannot cross question rows. + question_ids = op.Range( + op.Constant(value_int=0), + op.Gather(op.Shape(decide_indices), 0), + op.Constant(value_int=1), + ) + decide_coordinates = op.Concat( + op.Unsqueeze(question_ids, [1]), + op.Unsqueeze(decide_indices, [1]), + axis=1, + ) + option_coordinates = op.Concat( + op.Unsqueeze(option_owners, [1]), + op.Unsqueeze(option_indices, [1]), + axis=1, + ) + decide = op.GatherND(hidden_states, decide_coordinates) + options = op.GatherND(hidden_states, option_coordinates) + queries = op.Gather(self.q(op, decide), option_owners, axis=0) + logits = op.ReduceSum(op.Mul(self.k(op, options), queries), [-1], keepdims=0) + logits = op.Div(logits, op.CastLike(math.sqrt(KEV_POINTER_SIZE), logits)) + logits = op.Div(logits, op.CastLike(self.temperature, logits)) + probabilities = _grouped_softmax_graph(op, logits, option_owners, decide_indices) + return logits, probabilities + + +def _grouped_softmax_graph( + op: OpBuilder, + logits: ir.Value, + owners: ir.Value, + group_rows: ir.Value, +) -> ir.Value: + """Build stable softmax over flat logits independently for every owner. + + Owner ids must be contiguous zero-based group indices. A one-hot ownership + matrix computes each group's maximum and denominator without mixing + questions. + """ + question_count = op.Gather(op.Shape(group_rows), 0) + mask = op.OneHot( + owners, + question_count, + [0.0, 1.0], + axis=-1, + ) + # Subtract each owner's maximum before Exp, then gather its own denominator. + expanded = op.Unsqueeze(logits, [1]) + masked = op.Where( + op.Cast(mask, to=ir.DataType.BOOL), + expanded, + op.CastLike(-3.4028234663852886e38, logits), + ) + maxima = op.ReduceMax(masked, [0], keepdims=0) + stable = op.Sub(logits, op.Gather(maxima, owners)) + exponentials = op.Exp(stable) + weighted_mask = op.CastLike(mask, exponentials) + totals = op.MatMul( + op.Transpose(weighted_mask, perm=[1, 0]), + op.Unsqueeze(exponentials, [1]), + ) + totals = op.Squeeze(totals, [1]) + return op.Div(exponentials, op.Gather(totals, owners)) + + +class KevModel(nn.Module): + """Qwen3.5-4B backbone and Kev pointer readout in one exportable model.""" + + default_task = "kev-scoring" + category = "Text Classification" + provenance = KEV_PROVENANCE + + def __init__(self, config: ArchitectureConfig): + """Construct the verified-width headless Qwen3.5 and pointer head.""" + super().__init__() + if config.hidden_size != KEV_HIDDEN_SIZE: + raise ValueError("kev-4b requires Qwen3.5 hidden_size=2560") + self.config = config + self.backbone = Qwen35EncoderModel(config) + self.pointer_head = KevPointerHead(config.hidden_size) + + def preprocess_weights(self, state_dict: Mapping[str, Any]) -> dict[str, torch.Tensor]: + """Map merged backbone and pointer tensors into component paths.""" + mapped = map_kev_checkpoint(state_dict) + base = { + key.removeprefix("backbone."): value + for key, value in state_dict.items() + if isinstance(key, str) + and isinstance(value, torch.Tensor) + and not key.startswith("pointer_head.") + } + for key, value in self.backbone.preprocess_weights(base).items(): + mapped.setdefault(f"backbone.{key}", value) + return mapped + + +def map_kev_checkpoint(checkpoint: Mapping[str, Any]) -> dict[str, torch.Tensor]: + """Map controlled head/base tensors; callers deserialize ``head.pt`` themselves.""" + mapped: dict[str, torch.Tensor] = {} + head = checkpoint.get("head") + if head is not None: + if not isinstance(head, Mapping): + raise TypeError("Kev checkpoint head must be a state-dict mapping") + for key, value in head.items(): + if key not in {"q.weight", "q.bias", "k.weight", "k.bias"}: + continue + if not isinstance(value, torch.Tensor): + raise TypeError("Kev head values must be tensors") + mapped[f"pointer_head.{key}"] = value + for key, value in checkpoint.items(): + if isinstance(key, str) and isinstance(value, torch.Tensor): + if key.startswith("backbone."): + mapped[key] = value + elif key.startswith("model."): + mapped[f"backbone.{key}"] = value + return mapped + + +def synthetic_clm_checkpoint( + *, + hidden_size: int, + width: int, + depth: int = 3, + projection_dim: int = 512, + seed: int = 0, +) -> dict[str, Any]: + """Create a flexible deterministic test fixture, not a production checkpoint.""" + generator = torch.Generator().manual_seed(seed) + + def tensor(*shape: int) -> torch.Tensor: + """Draw one deterministic tensor from the fixture generator.""" + return torch.randn(shape, generator=generator) + + def head() -> dict[str, torch.Tensor]: + """Create one fixture head using checkpoint-compatible key names.""" + values = { + "inp.weight": tensor(width, hidden_size), + "inp.bias": tensor(width), + "out.weight": tensor(projection_dim, width), + "out.bias": tensor(projection_dim), + } + for index in range(depth - 2): + values[f"hidden.{index}.weight"] = tensor(width, width) + values[f"hidden.{index}.bias"] = tensor(width) + values[f"norms.{index}.weight"] = tensor(width) + values[f"norms.{index}.bias"] = tensor(width) + return values + + return { + "state_head": head(), + "action_head": head(), + "logit_scale": torch.tensor(0.0), + "cfg": { + "hidden_size": hidden_size, + "width": width, + "depth": depth, + "projection_dim": projection_dim, + }, + } + + +def synthetic_kev_checkpoint( + *, hidden_size: int = KEV_HIDDEN_SIZE, seed: int = 0 +) -> dict[str, Any]: + """Create flexible deterministic Kev test weights, not a production artifact.""" + generator = torch.Generator().manual_seed(seed) + return { + "head": { + "q.weight": torch.randn(KEV_POINTER_SIZE, hidden_size, generator=generator), + "q.bias": torch.randn(KEV_POINTER_SIZE, generator=generator), + "k.weight": torch.randn(KEV_POINTER_SIZE, hidden_size, generator=generator), + "k.bias": torch.randn(KEV_POINTER_SIZE, generator=generator), + }, + "temperature": KEV_TEMPERATURE, + } + + +def provenance_json(provenance: ModelProvenance) -> str: + """Serialize deterministic compact provenance for ONNX metadata.""" + return json.dumps(provenance.as_metadata(), sort_keys=True, separators=(",", ":")) + + +def build_clm_package( + config: ArchitectureConfig, + *, + base_weights: Mapping[str, torch.Tensor], + head_checkpoint: Mapping[str, Any], + base_revision: str, +): + """Build and populate CLM-v0.1 from separate base and head checkpoints. + + ``head_checkpoint`` must already be loaded by trusted caller code. The + upstream CLM artifact does not pin Qwen3-8B, so a concrete base revision is + mandatory here and is embedded into every component's provenance. + + Validation is fail-closed and occurs before graph construction: the exact + published metadata, head key sets, tensor shapes, and scalar logit scale + are required. Base weights must be a non-empty tensor mapping. + """ + if not base_revision or not base_revision.strip(): + raise ValueError( + "CLM export requires the caller to supply the exact Qwen3-8B base_revision" + ) + _validate_weight_mapping(base_weights, owner="CLM base") + # Validate the entire artifact contract before allocating a large graph. + validate_clm_checkpoint(head_checkpoint) + cfg = head_checkpoint.get("cfg") + assert isinstance(cfg, Mapping) + module = CLMModel( + config, + width=CLM_WIDTH, + depth=3, + projection_dim=CLM_PROJECTION_DIM, + activation="gelu", + layernorm=True, + residual=False, + base_revision=base_revision, + ) + from mobius._builder import build_from_module + from mobius.tasks._decision import CLMTask + + package = build_from_module(module, config, task=CLMTask(module.provenance)) + weights = { + f"encoder.{key}": value + for key, value in module.encoder.preprocess_weights(dict(base_weights)).items() + } + weights.update(map_clm_checkpoint(head_checkpoint)) + # The component graphs retain their root module paths (for example, + # ``encoder.model.layers...`` and ``state_head.inp...``). Trying to route by + # those same prefixes would strip them and leave every initializer unbound. + # Let ModelPackage match the fully qualified names across all components. + _apply_package_weights_strict(package, weights, owner="CLM checkpoint") + return package + + +def build_kev_package( + config: ArchitectureConfig, + *, + head_checkpoint: Mapping[str, Any], + merged_base_weights: Mapping[str, torch.Tensor] | None = None, + base_weights: Mapping[str, torch.Tensor] | None = None, + peft_adapter_weights: Mapping[str, torch.Tensor] | None = None, +): + """Build Kev from a correctly PEFT-merged Qwen3.5 state dict plus ``head.pt``. + + Mobius does not currently merge arbitrary PEFT LoRA artifacts into dense + Qwen3.5 tensors. Supplying separate base/adapter mappings therefore fails + explicitly; merge with PEFT at the pinned base revision first and pass the + resulting state dict as ``merged_base_weights``. + + The published base identity, pointer metadata, calibration, exact Q/K key + set, and tensor shapes are validated before graph construction. Separate + PEFT/base inputs fail explicitly rather than producing an unadapted model. + """ + # A valid head cannot make an unmerged PEFT base safe; enforce both gates. + validate_kev_checkpoint(head_checkpoint) + if merged_base_weights is None: + if base_weights is not None or peft_adapter_weights is not None: + raise NotImplementedError( + "Kev export cannot merge separate PEFT LoRA weights yet; load " + f"{KEV_BASE_MODEL_ID}@{KEV_BASE_REVISION}, merge the adapter with " + "PEFT merge_and_unload(), and pass merged_base_weights" + ) + raise ValueError("merged_base_weights is required for Kev export") + if base_weights is not None or peft_adapter_weights is not None: + raise ValueError( + "pass either merged_base_weights or separate base/adapter inputs, not both" + ) + _validate_weight_mapping(merged_base_weights, owner="Kev merged base") + module = KevModel(config) + from mobius._builder import build_from_module + from mobius.tasks._decision import KevTask + + package = build_from_module(module, config, task=KevTask()) + weights = { + f"backbone.{key}": value + for key, value in module.backbone.preprocess_weights(dict(merged_base_weights)).items() + } + weights.update(map_kev_checkpoint(head_checkpoint)) + # As with CLM, these graphs retain ``backbone.`` and ``pointer_head.`` in + # their initializer names, so applying the fully qualified mapping directly + # is required. + _apply_package_weights_strict(package, weights, owner="Kev checkpoint") + return package + + +class Qwen3EncoderModel(Qwen3CausalLMModel): + """Headless Qwen3 backbone returning every token hidden state and KV state.""" + + def __init__(self, config: ArchitectureConfig): + """Initialize Qwen3 then remove the unused vocabulary projection.""" + super().__init__(config) + del self.lm_head + + def forward( + self, + op: OpBuilder, + input_ids: ir.Value, + attention_mask: ir.Value | None, + position_ids: ir.Value, + past_key_values: list | None = None, + ): + """Return full ``[batch, sequence, hidden]`` states and KV cache. + + No pooling is performed: callers must select the last attended token, + because right padding makes a fixed position ``-1`` incorrect. + """ + # Keep every token state; CLM pooling depends on the caller's padding mask. + return self.model( + op, + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + ) + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Apply Qwen3 preprocessing while discarding any tied LM-head tensor.""" + mapped = super().preprocess_weights(state_dict) + mapped.pop("lm_head.weight", None) + return mapped + + +class Qwen35EncoderModel(Qwen35CausalLMModel): + """Headless Qwen3.5 backbone returning token hidden states and hybrid state.""" + + def __init__(self, config: ArchitectureConfig): + """Initialize Qwen3.5 then remove the unused vocabulary projection.""" + super().__init__(config) + del self.lm_head + + def forward( + self, + op: OpBuilder, + input_ids: ir.Value, + attention_mask: ir.Value | None, + position_ids: ir.Value, + past_key_values: list | None = None, + ): + """Return full token hidden states and Qwen3.5 hybrid cache state.""" + # Kev reads arbitrary decide/option positions, so token states stay unpooled. + return self.model( + op, + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + ) + + def preprocess_weights( + self, state_dict: dict[str, torch.Tensor] + ) -> dict[str, torch.Tensor]: + """Apply Qwen3.5 preprocessing while discarding LM-head weights.""" + mapped = super().preprocess_weights(state_dict) + mapped.pop("lm_head.weight", None) + return mapped diff --git a/src/mobius/models/decision_test.py b/src/mobius/models/decision_test.py new file mode 100644 index 000000000..9246986aa --- /dev/null +++ b/src/mobius/models/decision_test.py @@ -0,0 +1,442 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +from types import SimpleNamespace + +import onnx_ir as ir +import pytest +import torch + +from mobius._configs import ArchitectureConfig +from mobius._model_package import ModelPackage +from mobius._testing import make_config +from mobius.models.decision import ( + CLM_PROVENANCE, + KEV_CONTROL_TOKEN_IDS, + KEV_TEMPERATURE, + CheckpointContractError, + CLMModel, + CLMProjectionHead, + CLMScorer, + KevModel, + KevPointerHead, + Qwen35EncoderModel, + _apply_package_weights_strict, + batch_kev_rows, + build_clm_package, + build_kev_package, + clm_answer, + clm_candidates, + clm_last_token_indices, + clm_state_text, + encode_kev_rows, + grouped_softmax, + kev_answer, + map_clm_checkpoint, + map_kev_checkpoint, + render_clm, + render_kev, + synthetic_clm_checkpoint, + synthetic_kev_checkpoint, + validate_clm_checkpoint, + validate_kev_checkpoint, + validate_temperature, +) +from mobius.tasks._decision import CLMTask, HeadlessBackboneTask, KevTask + + +class _Tokenizer: + def __call__(self, text, add_special_tokens=False): + assert add_special_tokens is False + return SimpleNamespace(input_ids=[ord(char) for char in text]) + + +def test_clm_rendering_and_candidates_match_verified_contract(): + assert render_clm(None) == "" + assert render_clm(True) == "true" + assert render_clm({"a": 1, "b": {"c": False}, "d": ["x", {"y": 2}]}) == ( + "a: 1\n\nb:\n c: false\n\nd:\n - x\n -\n y: 2" + ) + assert clm_state_text("state", "question") == "state\n\nquestion" + assert clm_candidates( + { + "type": "choice", + "criteria": {"plain": None, "rich": "description"}, + } + ) == (["plain", "rich"], ["plain", "description"]) + assert clm_candidates({"type": "noul", "instructions": "Ready?"}) == ( + ["false", "true"], + ["false: No. This is false: Ready?", "true: Yes. This is true: Ready?"], + ) + + +def test_clm_grouping_answers_and_temperature(): + grouped = grouped_softmax([1001.0, 1000.0, -1000.0], [2, 1]) + assert grouped[0] == pytest.approx([0.7310585786, 0.2689414214]) + assert grouped[1] == [1.0] + with pytest.raises(ValueError, match=r"\(0, 100\]"): + validate_temperature(0) + answer = clm_answer( + {"type": "score", "criteria": ["low", "mid", "high"]}, + ["0", "1", "2"], + [0.1, 0.2, 0.7], + ) + assert answer["score"] == pytest.approx(1.6) + assert answer["confidence"] == pytest.approx(0.55) + assert clm_last_token_indices([[1, 1, 0], [1, 1, 1]]) == [1, 2] + with pytest.raises(ValueError, match="no attended token"): + clm_last_token_indices([[0, 0]]) + + +def test_kev_render_rows_and_indices_are_exact(): + assert render_kev({"enabled": True, "items": [1, False]}) == ( + "enabled: True\nitems:\n - 1\n - False" + ) + rows = encode_kev_rows( + _Tokenizer(), + "<|box_end|>", + { + "q": { + "type": "choice", + "instructions": "pick", + "criteria": {"a": None, "b": "Bee"}, + } + }, + ) + row = rows[0] + assert len(rows) == 1 + assert row.input_ids[0] == KEV_CONTROL_TOKEN_IDS["state"] + assert row.input_ids[row.decide_index] == KEV_CONTROL_TOKEN_IDS["decide"] + assert all( + row.input_ids[index] == KEV_CONTROL_TOKEN_IDS["option_end"] + for index in row.option_indices + ) + escaped = "".join(map(chr, row.input_ids[1 : 1 + len("<¦box_end¦>")])) + assert escaped == "<¦box_end¦>" + assert row.keys == ("a", "b") + batch = batch_kev_rows(rows, pad_token_id=0) + assert batch.decide_indices == (row.decide_index,) + assert batch.option_indices == row.option_indices + assert batch.option_owners == (0, 0) + + +def test_kev_rows_validate_request_and_serving_limits(): + with pytest.raises(ValueError, match="questions must not be empty"): + encode_kev_rows(_Tokenizer(), "state", {}) + with pytest.raises(ValueError, match=r"1\.\.255"): + encode_kev_rows( + _Tokenizer(), + "state", + {"q": {"type": "choice", "criteria": {}}}, + ) + with pytest.raises(ValueError, match="state exceeds 8 tokens"): + encode_kev_rows( + _Tokenizer(), + "12345678", + {"q": {"type": "noul"}}, + strict=True, + max_state_tokens=8, + max_row_tokens=32, + ) + with pytest.raises(ValueError, match=r"state\+question row exceeds 12"): + encode_kev_rows( + _Tokenizer(), + "123456789", + {"q": {"type": "noul", "instructions": "long"}}, + max_state_tokens=8, + max_row_tokens=12, + ) + + +def test_kev_answer_mapping_rounding_and_confidence(): + choice = kev_answer({"type": "choice"}, ["a", "b"], [0.123456, 0.876544]) + assert choice == { + "type": "choice", + "choice": "b", + "confidence": 0.7531, + "probabilities": {"a": 0.1235, "b": 0.8765}, + } + score = kev_answer( + {"type": "score", "criteria": ["bad", "ok", "good"]}, + ["0", "1", "2"], + [0.1, 0.2, 0.7], + ) + assert score["score"] == pytest.approx(1.6) + assert score["confidence"] == pytest.approx(0.8) + assert kev_answer({"type": "noul"}, ["false", "true"], [0.2, 0.8]) == { + "type": "noul", + "noul": 0.8, + } + + +def test_checkpoint_helpers_are_deterministic_and_mappable(): + first = synthetic_clm_checkpoint(hidden_size=4, width=3, projection_dim=2) + second = synthetic_clm_checkpoint(hidden_size=4, width=3, projection_dim=2) + assert first["state_head"].keys() == second["state_head"].keys() + assert first["state_head"]["inp.weight"].equal(second["state_head"]["inp.weight"]) + assert "hidden.0.weight" in first["state_head"] + assert "norms.0.weight" in first["state_head"] + assert "hidden.1.weight" not in first["state_head"] + kev = synthetic_kev_checkpoint(hidden_size=4) + assert kev["temperature"] == KEV_TEMPERATURE + assert kev["head"]["q.weight"].shape == (256, 4) + assert CLM_PROVENANCE.base_revision is None + assert CLM_PROVENANCE.as_metadata()["reproducible"] == "false" + + +def test_head_graphs_export_without_model_downloads(): + clm_config = ArchitectureConfig(hidden_size=4096, dtype=ir.DataType.FLOAT) + state = CLMTask().build_component( + "state_head", + None, + CLMProjectionHead(hidden_size=4096, width=8), + clm_config, + ) + scorer = CLMTask().build_component("scorer", None, CLMScorer(), clm_config) + kev_config = ArchitectureConfig(hidden_size=2560, dtype=ir.DataType.FLOAT) + pointer = KevTask().build_component("pointer_head", None, KevPointerHead(), kev_config) + assert [output.name for output in state.graph.outputs] == ["projections"] + assert [output.name for output in scorer.graph.outputs] == [ + "logits", + "probabilities", + ] + assert [output.name for output in pointer.graph.outputs] == [ + "logits", + "probabilities", + ] + assert "mobius.provenance" in pointer.metadata_props + + with pytest.raises(CheckpointContractError, match="unapplied weights"): + _apply_package_weights_strict( + ModelPackage({"state_head": state}, config=clm_config), + {"garbage": torch.ones(1)}, + owner="test checkpoint", + ) + with pytest.raises(CheckpointContractError, match="without weights"): + _apply_package_weights_strict( + ModelPackage({"state_head": state}, config=clm_config), + {}, + owner="test checkpoint", + ) + + +def test_headless_backbone_contract_has_no_generation_cache(): + class Backbone: + def __call__( + self, + op, + input_ids, + attention_mask, + position_ids, + past_key_values, + ): + assert attention_mask is not None + assert position_ids is not None + assert past_key_values is None + return op.Cast(input_ids, to=ir.DataType.FLOAT), ["unused-cache"] + + package = HeadlessBackboneTask().build(Backbone(), ArchitectureConfig()) + graph = package["model"].graph + + assert [value.name for value in graph.inputs] == [ + "input_ids", + "attention_mask", + "position_ids", + ] + assert [value.name for value in graph.outputs] == ["token_hidden_states"] + + +def test_qwen35_headless_backbone_initializes_hybrid_state(): + config = make_config( + hidden_size=64, + num_hidden_layers=2, + intermediate_size=128, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + layer_types=["linear_attention", "full_attention"], + linear_num_value_heads=4, + linear_num_key_heads=2, + linear_key_head_dim=16, + linear_value_head_dim=16, + linear_conv_kernel_dim=4, + ) + package = HeadlessBackboneTask().build(Qwen35EncoderModel(config), config) + graph = package["model"].graph + + assert [value.name for value in graph.outputs] == ["token_hidden_states"] + assert sum(node.op_type == "Expand" for node in graph) >= 2 + + +def test_clm_head_depth_three_has_one_normalized_hidden_block(): + head = CLMProjectionHead(hidden_size=4, width=3, projection_dim=2) + assert len(head.hidden) == 1 + assert len(head.norms) == 1 + assert hasattr(head.norms[0], "weight") + + +def test_clm_fp16_head_normalizes_in_float32(): + config = ArchitectureConfig(hidden_size=4, dtype=ir.DataType.FLOAT16) + model = CLMTask().build_component( + "state_head", + None, + CLMProjectionHead(hidden_size=4, width=3, projection_dim=2), + config, + ) + casts_to_float = [ + node + for node in model.graph.all_nodes() + if node.op_type == "Cast" and node.attributes["to"].value == ir.DataType.FLOAT + ] + + assert len(casts_to_float) == 2 + + +def test_public_build_helpers_reject_unsafe_or_inexact_inputs(): + config = ArchitectureConfig(hidden_size=4096) + checkpoint = _valid_clm_checkpoint() + with pytest.raises(ValueError, match="base_revision"): + build_clm_package( + config, + base_weights={}, + head_checkpoint=checkpoint, + base_revision="", + ) + with pytest.raises(CheckpointContractError, match="non-empty mapping"): + build_clm_package( + config, + base_weights={}, + head_checkpoint=checkpoint, + base_revision="caller-pin", + ) + bad = {**checkpoint, "cfg": {**checkpoint["cfg"], "depth": 2}} + with pytest.raises(CheckpointContractError, match=r"cfg\.depth"): + build_clm_package( + config, + base_weights={"model.embed_tokens.weight": _meta(1)}, + head_checkpoint=bad, + base_revision="caller-pin", + ) + with pytest.raises(NotImplementedError, match="merge_and_unload"): + build_kev_package( + ArchitectureConfig(hidden_size=2560), + head_checkpoint=_valid_kev_checkpoint(), + base_weights={}, + peft_adapter_weights={}, + ) + + +def _meta(*shape): + return torch.empty(shape, device="meta") + + +def _valid_clm_checkpoint(): + shapes = { + "inp.weight": (1536, 4096), + "inp.bias": (1536,), + "hidden.0.weight": (1536, 1536), + "hidden.0.bias": (1536,), + "norms.0.weight": (1536,), + "norms.0.bias": (1536,), + "out.weight": (512, 1536), + "out.bias": (512,), + } + return { + "cfg": { + "model": "Qwen/Qwen3-8B", + "hidden_size": 4096, + "projection_dim": 512, + "width": 1536, + "depth": 3, + "activation": "gelu", + "layernorm": True, + "residual": False, + }, + "state_head": {key: _meta(*shape) for key, shape in shapes.items()}, + "action_head": {key: _meta(*shape) for key, shape in shapes.items()}, + "logit_scale": _meta(), + } + + +def _valid_kev_checkpoint(): + return { + "base": "Qwen/Qwen3.5-4B-Base", + "base_revision": "1001bb4d826a52d1f399e183466143f4da7b741b", + "head_dim": 256, + "option_isolation": False, + "temperature": KEV_TEMPERATURE, + "head": { + "q.weight": _meta(256, 2560), + "q.bias": _meta(256), + "k.weight": _meta(256, 2560), + "k.bias": _meta(256), + }, + } + + +def test_production_checkpoint_contracts_fail_closed_and_map_metadata(): + clm = _valid_clm_checkpoint() + validate_clm_checkpoint(clm) + mapped_clm = map_clm_checkpoint(clm) + assert mapped_clm["state_head.hidden.0.weight"].shape == (1536, 1536) + assert mapped_clm["scorer.logit_scale"].shape == (1,) + invalid_clm = { + **clm, + "state_head": {**clm["state_head"], "unexpected": _meta(1)}, + } + with pytest.raises(CheckpointContractError, match="keys mismatch"): + validate_clm_checkpoint(invalid_clm) + invalid_scale = {**clm, "logit_scale": _meta(2)} + with pytest.raises(CheckpointContractError, match="scalar"): + validate_clm_checkpoint(invalid_scale) + + kev = _valid_kev_checkpoint() + validate_kev_checkpoint(kev) + mapped_kev = map_kev_checkpoint(kev) + assert mapped_kev["pointer_head.q.weight"].shape == (256, 2560) + with pytest.raises(CheckpointContractError, match="base_revision"): + validate_kev_checkpoint({**kev, "base_revision": "wrong"}) + with pytest.raises(CheckpointContractError, match=r"q\.weight shape"): + validate_kev_checkpoint({**kev, "head": {**kev["head"], "q.weight": _meta(1, 1)}}) + + +def test_complete_packages_use_headless_hidden_state_backbones(): + clm_config = ArchitectureConfig( + vocab_size=32, + hidden_size=4096, + intermediate_size=16, + num_hidden_layers=0, + num_attention_heads=1, + num_key_value_heads=1, + head_dim=4096, + max_position_embeddings=8, + rms_norm_eps=1e-6, + rope_theta=10_000.0, + dtype=ir.DataType.FLOAT, + ) + clm = CLMTask().build(CLMModel(clm_config, width=8), clm_config) + assert tuple(clm) == ("encoder", "state_head", "action_head", "scorer") + assert clm["encoder"].graph.outputs[0].name == "token_hidden_states" + assert not any(name.startswith("lm_head.") for name in clm["encoder"].graph.initializers) + + kev_config = ArchitectureConfig( + vocab_size=32, + hidden_size=2560, + intermediate_size=16, + num_hidden_layers=0, + num_attention_heads=1, + num_key_value_heads=1, + head_dim=2560, + max_position_embeddings=8, + rms_norm_eps=1e-6, + rope_theta=10_000.0, + rope_type="default", + dtype=ir.DataType.FLOAT, + layer_types=[], + ) + kev = KevTask().build(KevModel(kev_config), kev_config) + assert tuple(kev) == ("backbone", "pointer_head") + assert kev["backbone"].graph.outputs[0].name == "token_hidden_states" + assert not any(name.startswith("lm_head.") for name in kev["backbone"].graph.initializers) diff --git a/src/mobius/models/qwen35.py b/src/mobius/models/qwen35.py index 72cd58d12..b80a9ce83 100644 --- a/src/mobius/models/qwen35.py +++ b/src/mobius/models/qwen35.py @@ -193,7 +193,9 @@ def forward( if self.layer_type == "linear_attention": # DeltaNet states are passed through past_key_value as # (conv_state, recurrent_state), same tuple pattern as KV cache - conv_state, recurrent_state = past_key_value + conv_state, recurrent_state = ( + past_key_value if past_key_value is not None else (None, None) + ) attn_output, new_conv_state, new_recurrent_state = self.linear_attn( op, hidden_states, conv_state, recurrent_state diff --git a/src/mobius/tasks/__init__.py b/src/mobius/tasks/__init__.py index 1677cd01e..2d25bb410 100644 --- a/src/mobius/tasks/__init__.py +++ b/src/mobius/tasks/__init__.py @@ -23,11 +23,14 @@ "AudioCTCTask", "AudioFeatureExtractionTask", "CausalLMTask", + "CLMTask", "SmallThinkerGGUFCausalLMTask", "CTCAsrTask", "FeatureCTCAsrTask", "RNNTTask", "CodecTask", + "ComponentConfig", + "ComponentRole", "ComponentSpec", "ControlNetTask", "DeepSeekV4Task", @@ -68,8 +71,10 @@ "ImageClassificationTask", "KimiK3CausalLMTask", "KimiLinearCausalLMTask", + "KevTask", "Lfm2VlTask", "ModelTask", + "MultiComponentModelTask", "MllamaVisionLanguageTask", "MageVLTask", "Mistral4GGUFCausalLMTask", @@ -125,8 +130,11 @@ from mobius.tasks._audio_ctc import AudioCTCTask from mobius.tasks._audio_feature_extraction import AudioFeatureExtractionTask from mobius.tasks._base import ( + ComponentConfig, + ComponentRole, ComponentSpec, ModelTask, + MultiComponentModelTask, build_decoder_from_embeds, build_embedding_from_features, ) @@ -138,6 +146,7 @@ from mobius.tasks._codec import CodecTask from mobius.tasks._controlnet import ControlNetTask from mobius.tasks._ctc_asr import CTCAsrTask, FeatureCTCAsrTask +from mobius.tasks._decision import CLMTask, KevTask from mobius.tasks._deepseek_v4 import DeepSeekV4Task from mobius.tasks._denoising import DenoisingTask from mobius.tasks._dflash import DFlashDraftTask @@ -243,6 +252,7 @@ "ctc-asr": CTCAsrTask, "feature-ctc-asr": FeatureCTCAsrTask, "codec": CodecTask, + "clm-scoring": CLMTask, "controlnet": ControlNetTask, "denoising": DenoisingTask, "diarization": DiarizationTask, @@ -269,6 +279,7 @@ "hy-v3-mtp": HyV3MtpTask, "kimi-k3-text-generation": KimiK3CausalLMTask, "kimi-linear-text-generation": KimiLinearCausalLMTask, + "kev-scoring": KevTask, "falcon-h1-text-generation": FalconH1CausalLMTask, "plamo-text-generation": PlamoCausalLMTask, "plamo2-text-generation": Plamo2CausalLMTask, diff --git a/src/mobius/tasks/_base.py b/src/mobius/tasks/_base.py index f5847f894..2f46c10e6 100644 --- a/src/mobius/tasks/_base.py +++ b/src/mobius/tasks/_base.py @@ -5,7 +5,9 @@ from __future__ import annotations +import dataclasses from abc import ABC, abstractmethod +from enum import Enum from typing import TYPE_CHECKING, ClassVar import onnx_ir as ir @@ -20,6 +22,49 @@ from mobius._component_manifest import ComponentManifest +class ComponentRole(str, Enum): + """Neutral roles for non-generative package components.""" + + BACKBONE = "backbone" + ENCODER = "encoder" + HEAD = "head" + + +@dataclasses.dataclass(frozen=True) +class ComponentConfig: + """Configuration for one named component of a multi-component task. + + ``role`` accepts a :class:`ComponentRole` or one of its string values: + ``"backbone"``, ``"encoder"``, and ``"head"``. + """ + + module_attribute_path: str + role: str | ComponentRole | None = None + + def __post_init__(self) -> None: + """Validate and normalize the dotted path and optional neutral role.""" + path = self.module_attribute_path + if not isinstance(path, str): + raise TypeError("component module_attribute_path must be a string") + if not path or any(not part or not part.strip() for part in path.split(".")): + raise ValueError( + "component module_attribute_path must be a non-empty dotted attribute path" + ) + role = self.role + if role is None: + return + if not isinstance(role, (str, ComponentRole)): + raise TypeError("component role must be a ComponentRole or string") + try: + normalized_role = ComponentRole(role) + except ValueError: + supported = ", ".join(role.value for role in ComponentRole) + raise ValueError( + f"unsupported component role {role!r}; expected one of: {supported}" + ) from None + object.__setattr__(self, "role", normalized_role.value) + + class ComponentSpec: """Declares which sub-module attributes a multi-component task requires. @@ -29,28 +74,29 @@ class ComponentSpec: cryptic ``AttributeError`` that would otherwise surface deep inside ``build()``. - Map output model names to the module attribute that builds each component:: + Map output model names to the module attribute that builds each component. + A :class:`ComponentConfig` additionally declares a neutral component role:: ComponentSpec( - decoder="decoder", - vision_encoder="vision_encoder", - embedding="embedding", + encoder=ComponentConfig("encoder", ComponentRole.ENCODER), + classifier=ComponentConfig("heads.classifier", ComponentRole.HEAD), ) The keys are the names used in the output :class:`ModelPackage`; the - values are the attribute names on the ``nn.Module`` passed to - ``task.build()``. Dot notation is supported for nested attributes - (e.g. ``"model.encoder"``). + values are attribute names or component configurations for the ``nn.Module`` + passed to ``task.build()``. Dot notation is supported for nested attributes. Args: - **components: Keyword arguments mapping output name → module attribute - name. For example, ``vision_encoder="vision_encoder"`` means - the task expects ``module.vision_encoder`` and will store the - result as ``package["vision_encoder"]``. + **components: Keyword arguments mapping output name to a module + attribute name or :class:`ComponentConfig`. """ - def __init__(self, **components: str) -> None: - self._components: dict[str, str] = dict(components) + def __init__(self, **components: str | ComponentConfig) -> None: + """Store component declarations, normalizing plain paths to configs.""" + self._components = { + name: value if isinstance(value, ComponentConfig) else ComponentConfig(value) + for name, value in components.items() + } def validate(self, module: nn.Module, task_name: str) -> None: """Check that all required sub-module attributes exist on *module*. @@ -64,6 +110,7 @@ def validate(self, module: nn.Module, task_name: str) -> None: """ def _has_nested(obj: object, dotted: str) -> bool: + """Return whether *obj* exposes every segment of a dotted path.""" for part in dotted.split("."): if not hasattr(obj, part): return False @@ -71,9 +118,9 @@ def _has_nested(obj: object, dotted: str) -> bool: return True missing = [ - (output_name, attr_name) - for output_name, attr_name in self._components.items() - if not _has_nested(module, attr_name) + (output_name, component.module_attribute_path) + for output_name, component in self._components.items() + if not _has_nested(module, component.module_attribute_path) ] if not missing: return @@ -89,8 +136,30 @@ def _has_nested(obj: object, dotted: str) -> bool: def items(self): """Iterate over ``(output_name, attribute_name)`` pairs.""" + return ( + (name, component.module_attribute_path) + for name, component in self._components.items() + ) + + def configs(self): + """Iterate over ``(output_name, component_config)`` pairs.""" return self._components.items() + def roles(self) -> dict[str, str]: + """Return roles explicitly declared by component configurations.""" + return { + name: str(component.role) + for name, component in self._components.items() + if component.role is not None + } + + def resolve(self, module: nn.Module, name: str) -> object: + """Resolve a declared component module from the root module.""" + value: object = module + for part in self._components[name].module_attribute_path.split("."): + value = getattr(value, part) + return value + def keys(self): """Return the output model names declared by this spec.""" return self._components.keys() @@ -100,7 +169,13 @@ def __contains__(self, item: str) -> bool: return item in self._components def __repr__(self) -> str: - parts = ", ".join(f"{k}={v!r}" for k, v in self._components.items()) + """Return a constructor-like representation of this component spec.""" + parts = ", ".join( + f"{name}={component.module_attribute_path!r}" + if component.role is None + else f"{name}={component!r}" + for name, component in self._components.items() + ) return f"ComponentSpec({parts})" @@ -212,6 +287,84 @@ def build( ... +class MultiComponentModelTask(ModelTask): + """Base task for a named backbone/encoder and one or more named heads. + + Subclasses declare :attr:`components` with :class:`ComponentConfig` values + and implement :meth:`build_component`. The common build implementation + validates the layout and returns every graph in one :class:`ModelPackage`. + """ + + model_roles: ClassVar[dict[str, str]] = {} + components: ClassVar[ComponentSpec | None] = None + + def __init_subclass__(cls, **kwargs) -> None: + """Derive optimization roles from the subclass component declaration.""" + super().__init_subclass__(**kwargs) + components = cls.components + if components is None: + return + declared_roles = components.roles() + explicit_roles = cls.__dict__.get("model_roles") + if explicit_roles is not None: + declared_roles.update(explicit_roles) + cls.model_roles = declared_roles + + def build( + self, + module: nn.Module, + config: BaseModelConfig, + ) -> ModelPackage: + """Build every declared component into one atomic model package. + + The declaration must contain exactly one backbone or encoder and at + least one head. Component paths are validated before graph creation; + each :meth:`build_component` result must be an ``ir.Model``. + """ + components = self.components + if components is None: + raise TypeError(f"{type(self).__name__} must declare components") + roles = components.roles() + backbone_names = [ + name + for name, role in roles.items() + if role in {ComponentRole.BACKBONE.value, ComponentRole.ENCODER.value} + ] + head_names = [name for name, role in roles.items() if role == ComponentRole.HEAD.value] + if len(backbone_names) != 1 or not head_names: + raise ValueError( + f"{type(self).__name__} components must declare exactly one " + "backbone/encoder and at least one head" + ) + + self._validate_components(module) + models: dict[str, ir.Model] = {} + for name, component in components.configs(): + model = self.build_component( + name, + component, + components.resolve(module, name), + config, + ) + if not isinstance(model, ir.Model): + raise TypeError( + f"{type(self).__name__}.build_component({name!r}) " + "must return an onnx_ir.Model" + ) + models[name] = model + return ModelPackage(models, config=config) + + @abstractmethod + def build_component( + self, + name: str, + component: ComponentConfig, + module: object, + config: BaseModelConfig, + ) -> ir.Model: + """Build one graph for a declared component.""" + + # --------------------------------------------------------------------------- # Shared graph-builder helpers for multi-component tasks # --------------------------------------------------------------------------- diff --git a/src/mobius/tasks/_decision.py b/src/mobius/tasks/_decision.py new file mode 100644 index 000000000..a11a998af --- /dev/null +++ b/src/mobius/tasks/_decision.py @@ -0,0 +1,184 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Export tasks for the CLM-v0.1-8B and Kev-4B decision models.""" + +from __future__ import annotations + +import onnx_ir as ir + +from mobius._configs import BaseModelConfig +from mobius._model_package import ModelPackage +from mobius.models.decision import ( + CLM_PROVENANCE, + KEV_PROVENANCE, + CLMProjectionHead, + CLMScorer, + KevPointerHead, + provenance_json, +) +from mobius.tasks._base import ( + ComponentConfig, + ComponentRole, + ComponentSpec, + ModelTask, + MultiComponentModelTask, + _make_graph, + _make_model, +) + + +def _stamp(model: ir.Model, provenance, contract: str) -> ir.Model: + """Attach reproducibility and component-contract metadata to a model.""" + model.metadata_props["mobius.provenance"] = provenance_json(provenance) + model.metadata_props["mobius.decision_contract"] = contract + return model + + +class HeadlessBackboneTask(ModelTask): + """Export one full-sequence backbone without generation cache state.""" + + def build(self, module, config): + """Expose token IDs, masks, positions, and token hidden states only.""" + batch = ir.SymbolicDim("batch") + sequence = ir.SymbolicDim("sequence_length") + graph, builder = _make_graph("backbone") + input_ids = builder.input("input_ids", ir.DataType.INT64, [batch, sequence]) + attention_mask = builder.input("attention_mask", ir.DataType.INT64, [batch, sequence]) + position_ids = builder.input("position_ids", ir.DataType.INT64, [batch, sequence]) + outputs = module( + builder.op, + input_ids, + attention_mask, + position_ids, + past_key_values=None, + ) + hidden_states = outputs[0] if isinstance(outputs, tuple) else outputs + builder.add_output(hidden_states, "token_hidden_states") + return ModelPackage({"model": _make_model(graph)}, config=config) + + +class CLMTask(MultiComponentModelTask): + """Export Qwen3 encoder, two projection heads, and scorer as one package.""" + + components = ComponentSpec( + encoder=ComponentConfig("encoder", ComponentRole.ENCODER), + state_head=ComponentConfig("state_head", ComponentRole.HEAD), + action_head=ComponentConfig("action_head", ComponentRole.HEAD), + scorer=ComponentConfig("scorer", ComponentRole.HEAD), + ) + + def __init__(self, provenance=CLM_PROVENANCE): + """Create the task with pinned or explicitly unpinned provenance.""" + self.provenance = provenance + + def build(self, module, config): + """Build one package while honoring provenance selected by the module.""" + self.provenance = getattr(module, "provenance", self.provenance) + return super().build(module, config) + + def build_component(self, name, component, module, config): + """Build a headless encoder, projection head, or grouped scorer graph.""" + if name == "encoder": + package = HeadlessBackboneTask().build(module, config) + model = package["model"] + model.graph.name = name + return _stamp( + model, + self.provenance, + "qwen3-token-hidden-states;caller-selects-last-token;" + "projection-head-normalizes-input", + ) + if isinstance(module, CLMProjectionHead): + graph, builder = _make_graph(name) + embeddings = builder.input( + "embeddings", + config.dtype, + [ir.SymbolicDim("items"), config.hidden_size], + ) + builder.add_output(module(builder.op, embeddings), "projections") + return _stamp(_make_model(graph), self.provenance, "l2-normalized-projection") + if isinstance(module, CLMScorer): + graph, builder = _make_graph(name) + states = builder.input( + "state_projections", + config.dtype, + [ir.SymbolicDim("questions"), ir.SymbolicDim("projection_dim")], + ) + actions = builder.input( + "action_projections", + config.dtype, + [ir.SymbolicDim("candidates"), ir.SymbolicDim("projection_dim")], + ) + temperature = builder.input("temperature", config.dtype, []) + candidate_owners = builder.input( + "candidate_owners", + ir.DataType.INT64, + [ir.SymbolicDim("candidates")], + ) + logits, probabilities = module( + builder.op, states, actions, candidate_owners, temperature + ) + builder.add_output(logits, "logits") + builder.add_output(probabilities, "probabilities") + return _stamp(_make_model(graph), self.provenance, "scaled-cosine") + raise TypeError(f"unsupported CLM component {name!r}") + + +class KevTask(MultiComponentModelTask): + """Export the Qwen3.5 hybrid backbone and grouped pointer head together.""" + + components = ComponentSpec( + backbone=ComponentConfig("backbone", ComponentRole.BACKBONE), + pointer_head=ComponentConfig("pointer_head", ComponentRole.HEAD), + ) + + def build_component( + self, + name: str, + component: ComponentConfig, + module: object, + config: BaseModelConfig, + ) -> ir.Model: + """Build the headless hybrid backbone or Kev pointer graph. + + The backbone emits token-level hidden states. The pointer component + consumes per-row decide/option indices and returns flat logits plus a + stable probability distribution within each question. + """ + if name == "backbone": + package = HeadlessBackboneTask().build(module, config) + model = package["model"] + model.graph.name = name + return _stamp(model, KEV_PROVENANCE, "one-causal-row-per-question") + if not isinstance(module, KevPointerHead): + raise TypeError("Kev pointer_head has an unexpected module type") + graph, builder = _make_graph(name) + hidden_states = builder.input( + "hidden_states", + config.dtype, + [ + ir.SymbolicDim("questions"), + ir.SymbolicDim("sequence_length"), + config.hidden_size, + ], + ) + decide_indices = builder.input( + "decide_indices", ir.DataType.INT64, [ir.SymbolicDim("questions")] + ) + option_indices = builder.input( + "option_indices", ir.DataType.INT64, [ir.SymbolicDim("options")] + ) + option_owners = builder.input( + "option_owners", ir.DataType.INT64, [ir.SymbolicDim("options")] + ) + logits, probabilities = module( + builder.op, + hidden_states, + decide_indices, + option_indices, + option_owners, + ) + builder.add_output(logits, "logits") + builder.add_output(probabilities, "probabilities") + return _stamp(_make_model(graph), KEV_PROVENANCE, "grouped-pointer-softmax") diff --git a/src/mobius/tasks/_multi_component_test.py b/src/mobius/tasks/_multi_component_test.py new file mode 100644 index 000000000..0df80518e --- /dev/null +++ b/src/mobius/tasks/_multi_component_test.py @@ -0,0 +1,137 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Tests for generic non-generative multi-component tasks.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import onnx_ir as ir +import pytest + +from mobius._builder import build_from_module +from mobius._configs import BaseModelConfig +from mobius._model_package import ModelPackage +from mobius.tasks import ( + ComponentConfig, + ComponentRole, + ComponentSpec, + MultiComponentModelTask, +) +from mobius.tasks._base import _make_graph, _make_model + + +class _BackboneAndHeadsTask(MultiComponentModelTask): + components = ComponentSpec( + feature_extractor=ComponentConfig( + "encoder", + ComponentRole.ENCODER, + ), + category_scores=ComponentConfig( + "heads.category", + ComponentRole.HEAD, + ), + quality_score=ComponentConfig( + "heads.quality", + ComponentRole.HEAD, + ), + ) + + def __init__(self) -> None: + self.built_modules: dict[str, object] = {} + + def build_component(self, name, component, module, config): + self.built_modules[name] = module + graph, builder = _make_graph(name=name) + value = builder.input("input", ir.DataType.FLOAT, [1]) + builder.add_output(builder.op.Identity(value), "output") + return _make_model(graph) + + +def _module(): + return SimpleNamespace( + encoder=object(), + heads=SimpleNamespace(category=object(), quality=object()), + ) + + +def test_manifest_resolves_names_paths_and_neutral_roles(): + manifest = _BackboneAndHeadsTask().component_manifest() + + assert _BackboneAndHeadsTask.model_roles == { + "feature_extractor": "encoder", + "category_scores": "head", + "quality_score": "head", + } + assert manifest.names == ( + "feature_extractor", + "category_scores", + "quality_score", + ) + assert manifest["feature_extractor"].module_attribute_path == "encoder" + assert manifest["feature_extractor"].role == "encoder" + assert manifest["category_scores"].module_attribute_path == "heads.category" + assert manifest["category_scores"].role == "head" + assert manifest["quality_score"].role == "head" + + +def test_build_packages_backbone_and_all_heads_together(): + module = _module() + task = _BackboneAndHeadsTask() + + package = task.build(module, BaseModelConfig()) + + assert isinstance(package, ModelPackage) + assert tuple(package) == ( + "feature_extractor", + "category_scores", + "quality_score", + ) + assert task.built_modules == { + "feature_extractor": module.encoder, + "category_scores": module.heads.category, + "quality_score": module.heads.quality, + } + assert all(model.graph.name == name for name, model in package.items()) + + +def test_build_optimization_uses_declared_roles(monkeypatch): + optimized_roles = {} + + def record_role(model, **kwargs): + optimized_roles[model.graph.name] = kwargs["model_role"] + + monkeypatch.setattr("mobius._builder.optimize_model", record_role) + + build_from_module( + _module(), + BaseModelConfig(), + task=_BackboneAndHeadsTask(), + ) + + assert optimized_roles == { + "feature_extractor": "encoder", + "category_scores": "head", + "quality_score": "head", + } + assert "decoder" not in optimized_roles.values() + + +@pytest.mark.parametrize("path", [None, 3]) +def test_component_config_rejects_non_string_paths(path): + with pytest.raises(TypeError, match="must be a string"): + ComponentConfig(path, ComponentRole.HEAD) + + +@pytest.mark.parametrize("path", ["", ".", ".head", "heads.", "heads..classifier", " "]) +def test_component_config_rejects_empty_attribute_path_segments(path): + with pytest.raises(ValueError, match="non-empty dotted attribute path"): + ComponentConfig(path, ComponentRole.HEAD) + + +@pytest.mark.parametrize("role", ["decoder", "", "classifier", 3]) +def test_component_config_rejects_unsupported_roles(role): + error = TypeError if role == 3 else ValueError + with pytest.raises(error, match="component role"): + ComponentConfig("head", role) diff --git a/tests/clm_parity_test.py b/tests/clm_parity_test.py new file mode 100644 index 000000000..3d8032164 --- /dev/null +++ b/tests/clm_parity_test.py @@ -0,0 +1,551 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Focused numerical parity tests for the CLM decision-model package. + +The two integration tests are intentionally opt-in. They never download a +checkpoint merely because the integration marker was selected. +""" + +from __future__ import annotations + +import json +import os +import urllib.request +from collections.abc import Mapping, Sequence +from pathlib import Path +from typing import Any, ClassVar + +import numpy as np +import onnx_ir as ir +import pytest +import torch +from torch.nn import functional + +from mobius._configs import ArchitectureConfig +from mobius.integrations._weight_loading import apply_weights +from mobius.models.decision import ( + CLM_BASE_MODEL_ID, + CLM_MODEL_ID, + CLM_REVISION, + CLMProjectionHead, + CLMScorer, + clm_answer, + clm_candidates, + clm_last_token_indices, + clm_state_text, + render_clm, +) +from mobius.tasks._decision import CLMTask + +_FULL_PARITY_FLAG = "CLM_RUN_FULL_PARITY" +_VLLM_PARITY_FLAG = "CLM_RUN_VLLM_PARITY" +_QUESTIONS = { + "approval": { + "type": "noul", + "instructions": "Approve this request?", + }, + "route": { + "type": "choice", + "instructions": "Choose a route.", + "criteria": {"safe": "Use the safe route", "fast": "Use the fast route"}, + }, +} +_STATE = {"ready": True, "attempt": 2, "items": ["a", {"nested": False}]} + + +def _torch_projection( + embeddings: torch.Tensor, + weights: Mapping[str, torch.Tensor], +) -> torch.Tensor: + """Independent checkpoint-order reference for the published depth-three head.""" + embeddings = functional.normalize(embeddings, dim=-1) + value = functional.gelu( + functional.linear(embeddings, weights["inp.weight"], weights["inp.bias"]) + ) + value = functional.linear(value, weights["hidden.0.weight"], weights["hidden.0.bias"]) + value = functional.layer_norm( + value, + (value.shape[-1],), + weights["norms.0.weight"], + weights["norms.0.bias"], + eps=1e-5, + ) + value = functional.gelu(value) + value = functional.linear(value, weights["out.weight"], weights["out.bias"]) + return value / torch.linalg.vector_norm(value, dim=-1, keepdim=True).clamp_min(1e-12) + + +def _torch_score( + states: torch.Tensor, + actions: torch.Tensor, + owners: torch.Tensor, + logit_scale: torch.Tensor, + temperature: float, +) -> tuple[torch.Tensor, torch.Tensor]: + """Independent scaled-cosine and grouped-softmax reference.""" + scale = torch.exp(logit_scale).clamp_max(100.0) + logits = (states.index_select(0, owners) * actions).sum(-1) + logits = logits * scale / temperature + probabilities = torch.empty_like(logits) + for owner in range(states.shape[0]): + selected = owners == owner + probabilities[selected] = torch.softmax(logits[selected], dim=0) + return logits, probabilities + + +def _weights(prefix: float, hidden: int, width: int, projection: int) -> dict: + """Create deterministic, non-symmetric weights that expose ordering errors.""" + sizes = { + "inp.weight": (width, hidden), + "inp.bias": (width,), + "hidden.0.weight": (width, width), + "hidden.0.bias": (width,), + "norms.0.weight": (width,), + "norms.0.bias": (width,), + "out.weight": (projection, width), + "out.bias": (projection,), + } + result = {} + offset = prefix + for name, shape in sizes.items(): + count = int(np.prod(shape)) + values = torch.linspace(offset, offset + 0.3, count, dtype=torch.float32) + result[name] = values.reshape(shape) + offset += 0.17 + return result + + +def _session(model: ir.Model): + ort = pytest.importorskip("onnxruntime") + return ort.InferenceSession( + ir.to_proto(model).SerializeToString(), + providers=["CPUExecutionProvider"], + ) + + +def _head_model(weights: Mapping[str, torch.Tensor], hidden: int): + config = ArchitectureConfig(hidden_size=hidden, dtype=ir.DataType.FLOAT) + head = CLMProjectionHead( + hidden_size=hidden, + width=next(iter(weights.values())).shape[0], + projection_dim=weights["out.weight"].shape[0], + ) + model = CLMTask().build_component("state_head", None, head, config) + apply_weights(model, dict(weights)) + return model + + +def test_synthetic_clm_onnx_heads_and_scorer_match_direct_torch(): + """Exercise both projection graphs and scorer in the fast CPU CI lane.""" + torch.manual_seed(7) + hidden, width, projection = 4, 5, 3 + state_weights = _weights(-0.42, hidden, width, projection) + action_weights = _weights(0.09, hidden, width, projection) + state_input = torch.tensor([[0.2, -0.4, 0.8, 1.1], [-0.3, 0.7, 0.5, -0.9]]) + action_input = torch.tensor( + [ + [0.1, 0.9, -0.2, 0.4], + [-0.5, 0.3, 1.2, -0.7], + [0.6, -0.1, 0.2, 0.8], + [0.4, 0.5, -0.6, 0.3], + ] + ) + owners = torch.tensor([0, 0, 1, 1], dtype=torch.int64) + + state_ort = _session(_head_model(state_weights, hidden)).run( + None, {"embeddings": state_input.numpy()} + )[0] + action_ort = _session(_head_model(action_weights, hidden)).run( + None, {"embeddings": action_input.numpy()} + )[0] + state_ref = _torch_projection(state_input, state_weights) + action_ref = _torch_projection(action_input, action_weights) + np.testing.assert_allclose(state_ort, state_ref.numpy(), rtol=2e-5, atol=2e-6) + np.testing.assert_allclose(action_ort, action_ref.numpy(), rtol=2e-5, atol=2e-6) + np.testing.assert_allclose(np.linalg.norm(state_ort, axis=-1), 1.0, atol=2e-6) + np.testing.assert_allclose(np.linalg.norm(action_ort, axis=-1), 1.0, atol=2e-6) + + config = ArchitectureConfig(hidden_size=hidden, dtype=ir.DataType.FLOAT) + scorer = CLMTask().build_component("scorer", None, CLMScorer(), config) + logit_scale = torch.tensor([8.0]) + apply_weights(scorer, {"logit_scale": logit_scale}) + temperature = np.asarray(2.75, dtype=np.float32) + logits, probabilities = _session(scorer).run( + None, + { + "state_projections": state_ort, + "action_projections": action_ort, + "candidate_owners": owners.numpy(), + "temperature": temperature, + }, + ) + reference_logits, reference_probabilities = _torch_score( + state_ref, action_ref, owners, logit_scale, float(temperature) + ) + np.testing.assert_allclose(logits, reference_logits.numpy(), rtol=2e-5, atol=2e-5) + np.testing.assert_allclose( + probabilities, reference_probabilities.numpy(), rtol=2e-5, atol=2e-6 + ) + assert probabilities[:2].sum() == pytest.approx(1.0) + assert probabilities[2:].sum() == pytest.approx(1.0) + assert clm_answer({"type": "choice"}, ["first", "second"], [0.5, 0.5])["choice"] == "first" + assert ( + clm_answer({"type": "choice"}, ["first", "second"], [0.25, 0.75])["choice"] == "second" + ) + + +class _GoldenTokenizer: + """Small deterministic tokenizer oracle; no network or model dependency.""" + + _pieces: ClassVar[dict[str, list[int]]] = { + "ready: true\n\nitems:\n - red\n - blue\n\nChoose.": [11, 12, 21, 22, 31], + "alpha": [41, 42], + "A letter": [43, 44, 45], + } + + def __call__(self, texts: Sequence[str], *, padding: bool) -> dict[str, list]: + assert padding + rows = [self._pieces[text] for text in texts] + length = max(map(len, rows)) + return { + "input_ids": [row + [0] * (length - len(row)) for row in rows], + "attention_mask": [[1] * len(row) + [0] * (length - len(row)) for row in rows], + } + + +def test_clm_render_tokenization_and_last_token_golden(): + """Lock exact rendering, candidate text, token ids, and padding-aware pooling.""" + state_text = clm_state_text({"ready": True, "items": ["red", "blue"]}, "Choose.") + keys, actions = clm_candidates( + {"type": "choice", "criteria": {"alpha": None, "letter": "A letter"}} + ) + assert state_text == "ready: true\n\nitems:\n - red\n - blue\n\nChoose." + assert keys == ["alpha", "letter"] + assert actions == ["alpha", "A letter"] + batch = _GoldenTokenizer()([state_text, *actions], padding=True) + assert batch == { + "input_ids": [ + [11, 12, 21, 22, 31], + [41, 42, 0, 0, 0], + [43, 44, 45, 0, 0], + ], + "attention_mask": [ + [1, 1, 1, 1, 1], + [1, 1, 0, 0, 0], + [1, 1, 1, 0, 0], + ], + } + assert clm_last_token_indices(batch["attention_mask"]) == [4, 1, 2] + assert render_clm({"empty": None, "enabled": False}) == "empty: \n\nenabled: false" + + +def _require_enabled(flag: str) -> None: + if os.getenv(flag) != "1": + pytest.skip(f"set {flag}=1 to enable this heavyweight parity test") + + +def _required_env(name: str) -> str: + value = os.getenv(name) + if not value: + pytest.fail(f"{name} must be set when full CLM parity is enabled") + return value + + +def _component_path(root: Path, name: str) -> Path: + nested = root / name / "model.onnx" + direct = root / f"{name}.onnx" + path = nested if nested.is_file() else direct + if not path.is_file(): + pytest.fail(f"missing exported CLM {name} graph; tried {nested} and {direct}") + return path + + +def _cuda_sessions(root: Path) -> dict[str, Any]: + ort = pytest.importorskip("onnxruntime") + if "CUDAExecutionProvider" not in ort.get_available_providers(): + pytest.skip("CLM full parity requires ONNX Runtime CUDAExecutionProvider") + options = ["CUDAExecutionProvider", "CPUExecutionProvider"] + return { + name: ort.InferenceSession(str(_component_path(root, name)), providers=options) + for name in ("encoder", "state_head", "action_head", "scorer") + } + + +def _load_reference(): + """Load pinned Qwen and official CLM heads only after explicit opt-in.""" + revision = _required_env("CLM_BASE_REVISION") + transformers = pytest.importorskip("transformers") + hub = pytest.importorskip("huggingface_hub") + tokenizer = transformers.AutoTokenizer.from_pretrained( + CLM_BASE_MODEL_ID, revision=revision + ) + encoder = ( + transformers.AutoModel.from_pretrained( + CLM_BASE_MODEL_ID, + revision=revision, + dtype=torch.float32, + ) + .eval() + .cuda() + ) + filename = os.getenv("CLM_HEAD_FILENAME", "CLM_v0.1-8B.pt") + local = os.getenv("CLM_HEAD_CHECKPOINT_PATH") + checkpoint_path = ( + local if local else hub.hf_hub_download(CLM_MODEL_ID, filename, revision=CLM_REVISION) + ) + checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True) + return tokenizer, encoder, checkpoint + + +def _texts() -> tuple[list[str], list[str], list[str], np.ndarray]: + states, actions, keys, owners = [], [], [], [] + for owner, question in enumerate(_QUESTIONS.values()): + state_text = clm_state_text(_STATE, question["instructions"]) + question_keys, question_actions = clm_candidates(question) + states.append(state_text) + actions.extend(question_actions) + keys.extend(question_keys) + owners.extend([owner] * len(question_actions)) + return states, actions, keys, np.asarray(owners, dtype=np.int64) + + +def _tokenize(tokenizer, texts: Sequence[str]) -> dict[str, torch.Tensor]: + return tokenizer( + list(texts), + add_special_tokens=False, + padding=True, + return_tensors="pt", + ) + + +def _last_hidden(reference, tokens: Mapping[str, torch.Tensor]) -> torch.Tensor: + cuda_tokens = {name: value.cuda() for name, value in tokens.items()} + with torch.inference_mode(): + hidden = reference(**cuda_tokens, return_dict=True).last_hidden_state + indices = torch.tensor( + clm_last_token_indices(tokens["attention_mask"].tolist()), + device=hidden.device, + ) + return hidden[torch.arange(hidden.shape[0], device=hidden.device), indices].float() + + +def _encoder_ort(session, tokens: Mapping[str, torch.Tensor]) -> np.ndarray: + ids = tokens["input_ids"].numpy().astype(np.int64) + mask = tokens["attention_mask"].numpy().astype(np.int64) + position_ids = np.maximum(np.cumsum(mask, axis=1) - 1, 0).astype(np.int64) + values = {"input_ids": ids, "attention_mask": mask, "position_ids": position_ids} + feeds = {} + for item in session.get_inputs(): + if item.name in values: + feeds[item.name] = values[item.name] + continue + if not item.name.startswith("past_key_values."): + pytest.fail(f"unsupported encoder input {item.name!r} in CLM parity fixture") + shape = [] + for axis, dimension in enumerate(item.shape): + if isinstance(dimension, int): + shape.append(dimension) + elif axis == 0: + shape.append(ids.shape[0]) + else: + shape.append(0) + dtype = np.float16 if item.type == "tensor(float16)" else np.float32 + feeds[item.name] = np.zeros(shape, dtype=dtype) + hidden = session.run(None, feeds)[0] + indices = clm_last_token_indices(mask.tolist()) + return hidden[np.arange(len(indices)), indices] + + +def _checkpoint_head(checkpoint: Mapping[str, Any], name: str) -> dict[str, torch.Tensor]: + value = checkpoint.get(name) + if not isinstance(value, Mapping): + pytest.fail(f"official CLM checkpoint has no {name!r} state dict") + return dict(value) + + +def _run_full_parity(root: Path) -> dict[str, Any]: + sessions = _cuda_sessions(root) + tokenizer, encoder, checkpoint = _load_reference() + state_texts, action_texts, keys, owners = _texts() + assert state_texts == [ + ( + "ready: true\n\nattempt: 2\n\nitems:\n - a\n -\n nested: false" + "\n\nApprove this request?" + ), + ( + "ready: true\n\nattempt: 2\n\nitems:\n - a\n -\n nested: false" + "\n\nChoose a route." + ), + ], f"CLM state rendering drifted: {state_texts!r}" + assert action_texts == [ + "false: No. This is false: Approve this request?", + "true: Yes. This is true: Approve this request?", + "Use the safe route", + "Use the fast route", + ], f"CLM action rendering drifted: {action_texts!r}" + state_tokens = _tokenize(tokenizer, state_texts) + action_tokens = _tokenize(tokenizer, action_texts) + assert state_tokens["input_ids"].ndim == action_tokens["input_ids"].ndim == 2 + state_ids_unbatched = [ + tokenizer(text, add_special_tokens=False)["input_ids"] for text in state_texts + ] + state_ids_batched = [ + ids[mask.bool()].tolist() + for ids, mask in zip(state_tokens["input_ids"], state_tokens["attention_mask"]) + ] + assert state_ids_batched == state_ids_unbatched, ( + "batched CLM token IDs differ from per-text tokenization: " + f"batched={state_ids_batched!r}, unbatched={state_ids_unbatched!r}" + ) + state_hidden_ref = functional.normalize(_last_hidden(encoder, state_tokens), dim=-1) + action_hidden_ref = functional.normalize(_last_hidden(encoder, action_tokens), dim=-1) + state_hidden_ort = _encoder_ort(sessions["encoder"], state_tokens) + action_hidden_ort = _encoder_ort(sessions["encoder"], action_tokens) + state_hidden_ort /= np.maximum( + np.linalg.norm(state_hidden_ort, axis=-1, keepdims=True), 1e-12 + ) + action_hidden_ort /= np.maximum( + np.linalg.norm(action_hidden_ort, axis=-1, keepdims=True), 1e-12 + ) + np.testing.assert_allclose(state_hidden_ort, state_hidden_ref.cpu(), rtol=3e-3, atol=3e-3) + np.testing.assert_allclose( + action_hidden_ort, action_hidden_ref.cpu(), rtol=3e-3, atol=3e-3 + ) + + state_weights = _checkpoint_head(checkpoint, "state_head") + action_weights = _checkpoint_head(checkpoint, "action_head") + state_ref = _torch_projection(state_hidden_ref.cpu(), state_weights) + action_ref = _torch_projection(action_hidden_ref.cpu(), action_weights) + state_ort = sessions["state_head"].run( + None, {"embeddings": state_hidden_ort.astype(np.float32)} + )[0] + action_ort = sessions["action_head"].run( + None, {"embeddings": action_hidden_ort.astype(np.float32)} + )[0] + np.testing.assert_allclose(state_ort, state_ref, rtol=2e-3, atol=2e-3) + np.testing.assert_allclose(action_ort, action_ref, rtol=2e-3, atol=2e-3) + + temperature = np.asarray(float(os.getenv("CLM_TEST_TEMPERATURE", "1.7")), np.float32) + logits, probabilities = sessions["scorer"].run( + None, + { + "state_projections": state_ort, + "action_projections": action_ort, + "candidate_owners": owners, + "temperature": temperature, + }, + ) + scale = torch.as_tensor(checkpoint["logit_scale"]).reshape(1) + logits_ref, probabilities_ref = _torch_score( + state_ref, action_ref, torch.from_numpy(owners), scale, float(temperature) + ) + np.testing.assert_allclose(logits, logits_ref, rtol=2e-3, atol=2e-3) + np.testing.assert_allclose(probabilities, probabilities_ref, rtol=2e-3, atol=2e-4) + answers, offset = [], 0 + for question in _QUESTIONS.values(): + question_keys, _ = clm_candidates(question) + size = len(question_keys) + answers.append( + clm_answer(question, question_keys, probabilities[offset : offset + size]) + ) + reference_answer = clm_answer( + question, question_keys, probabilities_ref[offset : offset + size] + ) + assert answers[-1]["type"] == reference_answer["type"] + if question["type"] == "choice": + assert answers[-1]["choice"] == reference_answer["choice"], ( + f"answer mismatch for {question!r}: " + f"onnx={answers[-1]!r}, torch={reference_answer!r}" + ) + elif question["type"] == "noul": + assert answers[-1]["noul"] == pytest.approx(reference_answer["noul"], abs=2e-4) + offset += size + return { + "sessions": sessions, + "tokenizer": tokenizer, + "checkpoint": checkpoint, + "state_texts": state_texts, + "action_texts": action_texts, + "state_hidden": state_hidden_ort, + "action_hidden": action_hidden_ort, + "probabilities": probabilities, + "owners": owners, + "answers": answers, + "keys": keys, + "temperature": temperature, + } + + +@pytest.mark.integration +def test_exported_clm_package_matches_pinned_pytorch_reference(): + """Compare every externally visible full-model stage with diagnostics.""" + _require_enabled(_FULL_PARITY_FLAG) + root = Path(_required_env("MOBIUS_DECISION_EXPORT_ROOT")).expanduser() + result = _run_full_parity(root) + assert result["answers"], ( + f"no answers produced; root={root}, base_revision=" + f"{os.getenv('CLM_BASE_REVISION')}, clm_revision={CLM_REVISION}" + ) + + +def _post_embeddings(url: str, model: str, texts: Sequence[str]) -> np.ndarray: + body = json.dumps({"model": model, "input": list(texts), "encoding_format": "float"}) + request = urllib.request.Request( + url.rstrip("/") + "/v1/embeddings", + data=body.encode(), + headers={"Content-Type": "application/json"}, + method="POST", + ) + with urllib.request.urlopen(request, timeout=120) as response: + payload = json.load(response) + rows = sorted(payload["data"], key=lambda row: row["index"]) + return np.asarray([row["embedding"] for row in rows], dtype=np.float32) + + +@pytest.mark.integration +def test_official_clm_vllm_embeddings_endpoint_matches_onnx(): + """Compare an externally managed official vLLM server with the export.""" + _require_enabled(_VLLM_PARITY_FLAG) + _require_enabled(_FULL_PARITY_FLAG) + url = _required_env("CLM_EMBED_URL") + model = _required_env("CLM_EMBED_MODEL") + metadata = _required_env("CLM_VLLM_METADATA") + root = Path(_required_env("MOBIUS_DECISION_EXPORT_ROOT")).expanduser() + result = _run_full_parity(root) + texts = [*result["state_texts"], *result["action_texts"]] + endpoint = _post_embeddings(url, model, texts) + state_count = len(result["state_texts"]) + onnx_hidden = np.concatenate([result["state_hidden"], result["action_hidden"]], axis=0) + np.testing.assert_allclose( + endpoint, + onnx_hidden, + rtol=4e-3, + atol=4e-3, + err_msg=f"external vLLM metadata: {metadata}", + ) + state_endpoint = _torch_projection( + torch.from_numpy(endpoint[:state_count]), + _checkpoint_head(result["checkpoint"], "state_head"), + ) + action_endpoint = _torch_projection( + torch.from_numpy(endpoint[state_count:]), + _checkpoint_head(result["checkpoint"], "action_head"), + ) + _, endpoint_probabilities = _torch_score( + state_endpoint, + action_endpoint, + torch.from_numpy(result["owners"]), + torch.as_tensor(result["checkpoint"]["logit_scale"]).reshape(1), + float(result["temperature"]), + ) + np.testing.assert_allclose( + endpoint_probabilities, + result["probabilities"], + rtol=4e-3, + atol=4e-4, + err_msg=( + "vLLM version and served revision are not discoverable reliably from " + f"/v1/embeddings; preserve them as external CLM_VLLM_METADATA: {metadata}" + ), + ) diff --git a/tests/kev_parity_test.py b/tests/kev_parity_test.py new file mode 100644 index 000000000..039fcd4b4 --- /dev/null +++ b/tests/kev_parity_test.py @@ -0,0 +1,540 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Focused synthetic, preprocessing, and full-model parity tests for Kev-4B.""" + +from __future__ import annotations + +import importlib.util +import os +from pathlib import Path +from types import SimpleNamespace + +import numpy as np +import onnx_ir as ir +import onnxruntime as ort +import pytest +import torch +import torch.nn.functional as torch_functional + +from mobius._testing import create_test_builder, create_test_input +from mobius.models.decision import ( + KEV_BASE_MODEL_ID, + KEV_BASE_REVISION, + KEV_CONTROL_TOKEN_IDS, + KEV_MODEL_ID, + KEV_POINTER_SIZE, + KEV_REVISION, + KEV_TEMPERATURE, + KevPointerHead, + batch_kev_rows, + encode_kev_rows, + kev_answer, + validate_kev_checkpoint, +) + +_EXPORT_ROOT_ENV = "MOBIUS_DECISION_EXPORT_ROOT" +_ALLOW_DOWNLOAD_ENV = "MOBIUS_KEV_ALLOW_DOWNLOAD" + + +def _assert_close( + name: str, + actual: np.ndarray, + expected: np.ndarray, + *, + rtol: float, + atol: float, +) -> None: + maximum = float(np.max(np.abs(actual - expected))) if actual.size else 0.0 + np.testing.assert_allclose( + actual, + expected, + rtol=rtol, + atol=atol, + err_msg=f"{name} max_abs_error={maximum:.8g}", + ) + + +def _assert_answers_close( + actual: dict[str, dict], expected: dict[str, dict], *, atol: float +) -> None: + actual_numbers: dict[str, float] = {} + expected_numbers: dict[str, float] = {} + + def structure(value, path: str, numbers: dict[str, float]): + if isinstance(value, dict): + return { + key: structure(item, f"{path}.{key}", numbers) for key, item in value.items() + } + if isinstance(value, (int, float)) and not isinstance(value, bool): + numbers[path] = float(value) + return "" + return value + + assert structure(actual, "answers", actual_numbers) == structure( + expected, "answers", expected_numbers + ) + assert actual_numbers.keys() == expected_numbers.keys() + paths = sorted(actual_numbers) + _assert_close( + f"final answer fields {paths}", + np.asarray([actual_numbers[path] for path in paths]), + np.asarray([expected_numbers[path] for path in paths]), + rtol=0, + atol=atol, + ) + + +def _pointer_reference( + hidden_states: torch.Tensor, + decide_indices: torch.Tensor, + option_indices: torch.Tensor, + option_owners: torch.Tensor, + state: dict[str, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + rows = torch.arange(decide_indices.numel(), device=hidden_states.device) + decide = hidden_states[rows, decide_indices] + options = hidden_states[option_owners, option_indices] + queries = torch_functional.linear(decide, state["q.weight"], state["q.bias"])[ + option_owners + ] + keys = torch_functional.linear(options, state["k.weight"], state["k.bias"]) + logits = (queries * keys).sum(dim=-1) + logits = logits / (KEV_POINTER_SIZE**0.5) / KEV_TEMPERATURE + probabilities = torch.empty_like(logits) + for owner in range(decide_indices.numel()): + selected = option_owners == owner + probabilities[selected] = torch.softmax(logits[selected], dim=0) + return logits, probabilities + + +def _synthetic_pointer_graph( + hidden_states: np.ndarray, + decide_indices: np.ndarray, + option_indices: np.ndarray, + option_owners: np.ndarray, +) -> tuple[list[np.ndarray], dict[str, torch.Tensor]]: + builder, op, graph = create_test_builder() + hidden = create_test_input(builder, "hidden_states", list(hidden_states.shape)) + decide = create_test_input( + builder, "decide_indices", list(decide_indices.shape), ir.DataType.INT64 + ) + options = create_test_input( + builder, "option_indices", list(option_indices.shape), ir.DataType.INT64 + ) + owners = create_test_input( + builder, "option_owners", list(option_owners.shape), ir.DataType.INT64 + ) + head = KevPointerHead(hidden_size=hidden_states.shape[-1]) + logits, probabilities = head(op, hidden, decide, options, owners) + logits.name = "logits" + probabilities.name = "probabilities" + graph.outputs.extend((logits, probabilities)) + + generator = torch.Generator().manual_seed(20260928) + state = {} + for name, parameter in head.named_parameters(): + value = torch.randn(tuple(parameter.shape), generator=generator) / 8 + parameter.const_value = ir.tensor(value.numpy()) + state[name] = value + + model = ir.serde.serialize_model(ir.Model(graph, ir_version=11)) + session = ort.InferenceSession( + model.SerializeToString(), providers=["CPUExecutionProvider"] + ) + outputs = session.run( + None, + { + "hidden_states": hidden_states, + "decide_indices": decide_indices, + "option_indices": option_indices, + "option_owners": option_owners, + }, + ) + return outputs, state + + +def test_kev_pointer_head_matches_direct_torch_reference() -> None: + """Exercise indexed Q/K projection and independent variable-size groups.""" + generator = np.random.default_rng(20260928) + hidden = generator.normal(0, 0.25, (3, 8, 5)).astype(np.float32) + decide = np.array([7, 6, 5], dtype=np.int64) + option_indices = np.array([2, 4, 1, 3, 5, 2, 4], dtype=np.int64) + option_owners = np.array([0, 0, 1, 1, 1, 2, 2], dtype=np.int64) + + (actual_logits, actual_probabilities), state = _synthetic_pointer_graph( + hidden, decide, option_indices, option_owners + ) + expected_logits, expected_probabilities = _pointer_reference( + torch.from_numpy(hidden), + torch.from_numpy(decide), + torch.from_numpy(option_indices), + torch.from_numpy(option_owners), + state, + ) + _assert_close( + "synthetic logits", + actual_logits, + expected_logits.numpy(), + rtol=1e-5, + atol=1e-6, + ) + _assert_close( + "synthetic probabilities", + actual_probabilities, + expected_probabilities.numpy(), + rtol=1e-5, + atol=1e-6, + ) + np.testing.assert_allclose( + [ + actual_probabilities[:2].sum(), + actual_probabilities[2:5].sum(), + actual_probabilities[5:].sum(), + ], + 1.0, + rtol=0, + atol=1e-6, + ) + + questions = ( + ({"type": "choice"}, ("red", "blue")), + ({"type": "score", "criteria": ["low", "medium", "high"]}, ("0", "1", "2")), + ({"type": "noul"}, ("false", "true")), + ) + offsets = (0, 2, 5, 7) + actual_answers = [ + kev_answer(question, keys, actual_probabilities[start:stop]) + for (question, keys), start, stop in zip(questions, offsets[:-1], offsets[1:]) + ] + expected_answers = [ + kev_answer(question, keys, expected_probabilities[start:stop].tolist()) + for (question, keys), start, stop in zip(questions, offsets[:-1], offsets[1:]) + ] + assert actual_answers == expected_answers + assert [answer["type"] for answer in actual_answers] == [ + "choice", + "score", + "noul", + ] + + +class _CharacterTokenizer: + def __call__(self, text: str, *, add_special_tokens: bool): + assert add_special_tokens is False + return SimpleNamespace(input_ids=[ord(character) for character in text]) + + +def _characters(text: str) -> list[int]: + return [ord(character) for character in text] + + +def test_kev_preprocessing_golden_rows_padding_and_indices() -> None: + """Lock down every boundary token and flattened readout coordinate.""" + questions = { + "pick": { + "type": "choice", + "instructions": "Pick <|state|>", + "criteria": { + "alpha": "A<|option_end|>", + "b": None, + }, + }, + "ready": { + "type": "noul", + "instructions": "Ready?", + "criteria": {"false": "N", "true": "Y"}, + }, + } + rows = encode_kev_rows(_CharacterTokenizer(), "S<|decide|>", questions) + control = KEV_CONTROL_TOKEN_IDS + state = [control["state"], *_characters("S<¦decide¦>")] + expected_pick = [ + *state, + control["question"], + *_characters("Pick <¦state¦>"), + control["option_start"], + *_characters("alpha: A<¦option_end¦>"), + control["option_end"], + control["option_start"], + *_characters("b"), + control["option_end"], + control["decide"], + ] + expected_ready = [ + *state, + control["question"], + *_characters("Ready?"), + control["option_start"], + *_characters("no: N"), + control["option_end"], + control["option_start"], + *_characters("yes: Y"), + control["option_end"], + control["decide"], + ] + assert [row.question_id for row in rows] == ["pick", "ready"] + assert rows[0].input_ids == tuple(expected_pick) + assert rows[1].input_ids == tuple(expected_ready) + assert rows[0].position_ids == tuple(range(len(expected_pick))) + assert rows[1].position_ids == tuple(range(len(expected_ready))) + assert rows[0].keys == ("alpha", "b") + assert rows[1].keys == ("false", "true") + + pick_options = tuple( + index for index, token in enumerate(expected_pick) if token == control["option_end"] + ) + ready_options = tuple( + index for index, token in enumerate(expected_ready) if token == control["option_end"] + ) + assert rows[0].decide_index == len(expected_pick) - 1 + assert rows[1].decide_index == len(expected_ready) - 1 + assert rows[0].option_indices == pick_options + assert rows[1].option_indices == ready_options + + batch = batch_kev_rows(rows, pad_token_id=99) + width = len(expected_pick) + ready_padding = width - len(expected_ready) + assert batch.input_ids == ( + tuple(expected_pick), + (*expected_ready, *((99,) * ready_padding)), + ) + assert batch.attention_mask == ( + (1,) * width, + (*((1,) * len(expected_ready)), *((0,) * ready_padding)), + ) + assert batch.position_ids == ( + tuple(range(width)), + (*range(len(expected_ready)), *((0,) * ready_padding)), + ) + assert batch.decide_indices == ( + len(expected_pick) - 1, + len(expected_ready) - 1, + ) + assert batch.option_indices == (*pick_options, *ready_options) + assert batch.option_owners == (0, 0, 1, 1) + + +def _require_integration_inputs() -> tuple[Path, Path, Path]: + root_value = os.environ.get(_EXPORT_ROOT_ENV) + if not root_value: + pytest.skip(f"set {_EXPORT_ROOT_ENV} to an exported Kev package") + root = Path(root_value) + required = (root / "backbone/model.onnx", root / "pointer_head/model.onnx") + if not root.is_dir() or not all(path.is_file() for path in required): + pytest.skip(f"{_EXPORT_ROOT_ENV} is not a complete Kev package: {root}") + if not torch.cuda.is_available(): + pytest.skip("PyTorch CUDA is unavailable") + if "CUDAExecutionProvider" not in ort.get_available_providers(): + pytest.skip("ONNX Runtime CUDAExecutionProvider is unavailable") + if importlib.util.find_spec("peft") is None: + pytest.skip("peft is required for the pinned Kev adapter reference") + + from huggingface_hub import snapshot_download + + allow_download = os.environ.get(_ALLOW_DOWNLOAD_ENV) == "1" + try: + base = Path( + snapshot_download( + KEV_BASE_MODEL_ID, + revision=KEV_BASE_REVISION, + local_files_only=not allow_download, + ) + ) + adapter = Path( + snapshot_download( + KEV_MODEL_ID, + revision=KEV_REVISION, + local_files_only=not allow_download, + ) + ) + except Exception as error: + pytest.skip( + "pinned Kev reference is not cached; set " + f"{_ALLOW_DOWNLOAD_ENV}=1 to permit download ({error})" + ) + if not (adapter / "head.pt").is_file(): + pytest.skip(f"pinned Kev snapshot has no head.pt: {adapter}") + return root, base, adapter + + +def _empty_cache_feeds( + session: ort.InferenceSession, batch_size: int +) -> dict[str, np.ndarray]: + feeds = {} + for value in session.get_inputs()[3:]: + shape = [] + for dimension in value.shape: + if isinstance(dimension, int): + shape.append(dimension) + elif isinstance(dimension, str) and "batch" in dimension: + shape.append(batch_size) + elif isinstance(dimension, str) and "past" in dimension: + shape.append(0) + else: + raise AssertionError( + f"unsupported symbolic cache dimension {dimension!r} in {value.name}" + ) + feeds[value.name] = np.zeros(shape, dtype=np.float32) + return feeds + + +@pytest.mark.integration +def test_kev_export_matches_pinned_cuda_reference() -> None: + """Compare the exported package with the pinned merged PEFT reference.""" + root, base_path, adapter_path = _require_integration_inputs() + from peft import PeftModel + from transformers import AutoModelForCausalLM, AutoTokenizer + + checkpoint = torch.load(adapter_path / "head.pt", map_location="cpu", weights_only=True) + validate_kev_checkpoint(checkpoint) + export_tokenizer = AutoTokenizer.from_pretrained(root, local_files_only=True) + reference_tokenizer = AutoTokenizer.from_pretrained(base_path, local_files_only=True) + questions = { + "choice": { + "type": "choice", + "instructions": "Choose the best color.", + "criteria": {"red": "warm", "blue": "cool", "green": "natural"}, + }, + "truth": { + "type": "noul", + "instructions": "The user requested a cool color.", + }, + } + export_rows = encode_kev_rows(export_tokenizer, {"request": "blue"}, questions) + reference_rows = encode_kev_rows(reference_tokenizer, {"request": "blue"}, questions) + export_batch = batch_kev_rows(export_rows, pad_token_id=export_tokenizer.pad_token_id) + reference_batch = batch_kev_rows( + reference_rows, pad_token_id=reference_tokenizer.pad_token_id + ) + assert export_batch == reference_batch, ( + "export/reference token IDs or decide/option indices differ" + ) + + device = torch.device("cuda:0") + base_container = AutoModelForCausalLM.from_pretrained( + base_path, + local_files_only=True, + dtype=torch.float32, + low_cpu_mem_usage=True, + ) + reference_model = ( + PeftModel.from_pretrained(base_container.model, adapter_path, local_files_only=True) + .eval() + .to(device) + ) + input_ids = np.asarray(export_batch.input_ids, dtype=np.int64) + attention_mask = np.asarray(export_batch.attention_mask, dtype=np.int64) + position_ids = np.asarray(export_batch.position_ids, dtype=np.int64) + with torch.inference_mode(): + reference_output = reference_model( + input_ids=torch.from_numpy(input_ids).to(device), + attention_mask=torch.from_numpy(attention_mask).to(device), + position_ids=torch.from_numpy(position_ids).to(device), + use_cache=False, + return_dict=True, + ) + reference_hidden = reference_output.last_hidden_state + + backbone = ort.InferenceSession( + str(root / "backbone/model.onnx"), + providers=["CUDAExecutionProvider"], + ) + assert backbone.get_providers()[0] == "CUDAExecutionProvider" + feeds = { + "input_ids": input_ids, + "attention_mask": attention_mask, + "position_ids": position_ids, + **_empty_cache_feeds(backbone, input_ids.shape[0]), + } + ort_hidden = backbone.run(["token_hidden_states"], feeds)[0] + + decide = np.asarray(export_batch.decide_indices, dtype=np.int64) + options = np.asarray(export_batch.option_indices, dtype=np.int64) + owners = np.asarray(export_batch.option_owners, dtype=np.int64) + rows = np.arange(len(decide)) + torch_rows = torch.arange(len(decide), device=device) + torch_decide = torch.from_numpy(decide).to(device) + torch_options = torch.from_numpy(options).to(device) + torch_owners = torch.from_numpy(owners).to(device) + reference_positions = ( + torch.cat( + ( + reference_hidden[torch_rows, torch_decide], + reference_hidden[torch_owners, torch_options], + ) + ) + .float() + .cpu() + .numpy() + ) + ort_positions = np.concatenate((ort_hidden[rows, decide], ort_hidden[owners, options])) + _assert_close( + "decide/option hidden states", + ort_positions, + reference_positions, + rtol=2e-3, + atol=2e-3, + ) + + head_state = { + name: tensor.to(device=device, dtype=torch.float32) + for name, tensor in checkpoint["head"].items() + } + reference_logits, reference_probabilities = _pointer_reference( + reference_hidden.float(), + torch_decide, + torch_options, + torch_owners, + head_state, + ) + pointer = ort.InferenceSession( + str(root / "pointer_head/model.onnx"), + providers=["CUDAExecutionProvider"], + ) + assert pointer.get_providers()[0] == "CUDAExecutionProvider" + actual_logits, actual_probabilities = pointer.run( + None, + { + "hidden_states": ort_hidden, + "decide_indices": decide, + "option_indices": options, + "option_owners": owners, + }, + ) + expected_logits = reference_logits.cpu().numpy() + expected_probabilities = reference_probabilities.cpu().numpy() + _assert_close( + "full-model raw logits", + actual_logits, + expected_logits, + rtol=3e-3, + atol=3e-3, + ) + _assert_close( + "full-model probabilities", + actual_probabilities, + expected_probabilities, + rtol=3e-3, + atol=3e-3, + ) + + counts = [len(row.keys) for row in export_rows] + offsets = np.cumsum([0, *counts]) + actual_answers = { + row.question_id: kev_answer( + questions[row.question_id], + row.keys, + actual_probabilities[offset : offset + count], + ) + for row, offset, count in zip(export_rows, offsets, counts) + } + expected_answers = { + row.question_id: kev_answer( + questions[row.question_id], + row.keys, + expected_probabilities[offset : offset + count], + ) + for row, offset, count in zip(export_rows, offsets, counts) + } + _assert_answers_close(actual_answers, expected_answers, atol=5e-3)