"""Thin Ouro mapping for LoopQ artifacts and future runtime integration.""" from __future__ import annotations import hashlib import json from dataclasses import dataclass from pathlib import Path import torch from loopq.cta import CrossLoopTransitionAdapter from loopq.las import LoopAwareActivationScales from loopq.quantization import quantize_weight from loopq.transforms import SharedKroneckerTransform OURO_LOOP_COUNT = 4 OURO_TRANSITION_COUNT = 3 PAPER_GROUPS = { "attention_qkv": { "hf_weights": ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj"), "vllm_consumers": ("self_attn.qkv_proj",), "transform_site": "self_attn.ln_trans", }, "attention_output": { "hf_weights": ("self_attn.o_proj",), "vllm_consumers": ("self_attn.o_proj",), "transform_site": "self_attn.o_trans", }, "mlp_up_gate": { "hf_weights": ("mlp.gate_proj", "mlp.up_proj"), "vllm_consumers": ("mlp.gate_up_proj",), "transform_site": "mlp.up_gate_trans", }, "mlp_down": { "hf_weights": ("mlp.down_proj",), "vllm_consumers": ("mlp.down_proj",), "transform_site": "mlp.down_trans", }, } @dataclass(frozen=True) class OuroMappingValidation: revision: str config_sha256: str modeling_sha256: str physical_layers: int loop_count: int hidden_size: int intermediate_size: int checked_weight_keys: int def to_dict(self) -> dict[str, str | int]: return self.__dict__.copy() def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def validate_local_ouro_checkpoint(snapshot: str | Path) -> OuroMappingValidation: """Validate config and safetensors keys without loading parameter tensors.""" root = Path(snapshot).resolve() config_path = root / "config.json" modeling_path = root / "modeling_ouro.py" checkpoint_path = root / "model.safetensors" if not all(path.is_file() for path in (config_path, modeling_path, checkpoint_path)): raise FileNotFoundError("Ouro snapshot must contain config, modeling code, and model.safetensors") config = json.loads(config_path.read_text()) required_config = { "model_type": "ouro", "num_hidden_layers": 24, "total_ut_steps": OURO_LOOP_COUNT, "hidden_size": 2048, "intermediate_size": 5632, "torch_dtype": "bfloat16", } mismatches = { key: (config.get(key), expected) for key, expected in required_config.items() if config.get(key) != expected } if mismatches: raise ValueError(f"pinned Ouro config mismatch: {mismatches}") from safetensors import safe_open with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint: keys = set(checkpoint.keys()) required_weights = { f"model.layers.{layer}.{projection}.weight" for layer in range(config["num_hidden_layers"]) for group in PAPER_GROUPS.values() for projection in group["hf_weights"] } missing = sorted(required_weights.difference(keys)) if missing: raise ValueError(f"Ouro checkpoint is missing mapped projection weights: {missing[:8]}") return OuroMappingValidation( revision=root.name, config_sha256=_sha256(config_path), modeling_sha256=_sha256(modeling_path), physical_layers=config["num_hidden_layers"], loop_count=config["total_ut_steps"], hidden_size=config["hidden_size"], intermediate_size=config["intermediate_size"], checked_weight_keys=len(required_weights), ) class OuroLoopQAdapter: """Backend-neutral hook contract; it does not patch vLLM by itself.""" loop_count = OURO_LOOP_COUNT group_names = tuple(PAPER_GROUPS) @staticmethod def module_key(layer_index: int, group_name: str) -> str: if not 0 <= layer_index < 24: raise IndexError("Ouro physical layer index must be in [0, 24)") if group_name not in PAPER_GROUPS: raise KeyError(f"unknown Ouro LoopQ group {group_name!r}") return f"model.layers.{layer_index}.{group_name}" @staticmethod def validate_loop(loop_index: int) -> None: if not 0 <= loop_index < OURO_LOOP_COUNT: raise IndexError("Ouro LoopQ loop index must be in [0, 4)") def prepare_activation( self, activation: torch.Tensor, *, layer_index: int, loop_index: int, group_name: str, transform: SharedKroneckerTransform, las: LoopAwareActivationScales | None, activation_bits: int, quantization_enabled: bool, ) -> torch.Tensor: self.validate_loop(loop_index) module_key = self.module_key(layer_index, group_name) transformed = transform(activation) if not quantization_enabled: return transformed if las is None: raise ValueError("quantized LoopQ activation requires LAS") return las.quantize(module_key, loop_index, transformed, bits=activation_bits).dequantized @staticmethod def prepare_weight( weight: torch.Tensor, *, transform: SharedKroneckerTransform, quantization_enabled: bool, ) -> torch.Tensor: folded = transform.fold_weight(weight) if not quantization_enabled: return folded return quantize_weight(folded).dequantized @staticmethod def apply_transition( hidden_state: torch.Tensor, *, completed_loop_index: int, cta: CrossLoopTransitionAdapter, ) -> torch.Tensor: if not 0 <= completed_loop_index < OURO_TRANSITION_COUNT: raise IndexError("CTA is valid only after Ouro loops 0, 1, and 2") if cta.transition_count != OURO_TRANSITION_COUNT: raise ValueError("Ouro CTA artifact must contain exactly 3 transitions") return cta(hidden_state, completed_loop_index)