"""Thin Huginn architecture mapping for the isolated LoopQ baseline.""" 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 HUGINN_LOOP_COUNT = 32 HUGINN_TRANSITION_COUNT = 31 HUGINN_PHYSICAL_LAYERS = 4 PAPER_GROUPS = { "attention_qkv": { "hf_weights": ("attn.Wqkv",), "vllm_consumers": ("self_attn.qkv_proj",), "input_width": 5280, }, "attention_output": { "hf_weights": ("attn.proj",), "vllm_consumers": ("self_attn.o_proj",), "input_width": 5280, }, "mlp_up_gate": { "hf_weights": ("mlp.fc",), "vllm_consumers": ("mlp.gate_up_proj",), "input_width": 5280, }, "mlp_down": { "hf_weights": ("mlp.proj",), "vllm_consumers": ("mlp.down_proj",), "input_width": 17920, }, } @dataclass(frozen=True) class HuginnMappingValidation: revision: str config_sha256: str modeling_sha256: str checkpoint_shards: int 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 stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def validate_local_huginn_checkpoint(snapshot: str | Path) -> HuginnMappingValidation: """Validate the pinned Huginn config, code, shards, and mapped weights.""" root = Path(snapshot).resolve() config_path = root / "config.json" modeling_path = root / "raven_modeling_minimal.py" shards = sorted(root.glob("*.safetensors")) if not config_path.is_file() or not modeling_path.is_file() or not shards: raise FileNotFoundError( "Huginn snapshot must contain config, modeling code, and safetensors" ) config = json.loads(config_path.read_text()) required = { "model_type": "huginn_raven", "n_layers_in_recurrent_block": HUGINN_PHYSICAL_LAYERS, "mean_recurrence": HUGINN_LOOP_COUNT, "n_embd": 5280, "intermediate_size": 17920, } mismatches = { key: (config.get(key), expected) for key, expected in required.items() if config.get(key) != expected } if mismatches: raise ValueError(f"pinned Huginn config mismatch: {mismatches}") from safetensors import safe_open keys = set() for shard in shards: with safe_open(shard, framework="pt", device="cpu") as checkpoint: keys.update(checkpoint.keys()) required_weights = { f"transformer.core_block.{layer}.{projection}.weight" for layer in range(HUGINN_PHYSICAL_LAYERS) for group in PAPER_GROUPS.values() for projection in group["hf_weights"] } missing = sorted(required_weights - keys) if missing: raise ValueError( f"Huginn checkpoint is missing mapped projection weights: {missing[:8]}" ) return HuginnMappingValidation( revision=root.name, config_sha256=_sha256(config_path), modeling_sha256=_sha256(modeling_path), checkpoint_shards=len(shards), physical_layers=HUGINN_PHYSICAL_LAYERS, loop_count=HUGINN_LOOP_COUNT, hidden_size=config["n_embd"], intermediate_size=config["intermediate_size"], checked_weight_keys=len(required_weights), ) class HuginnLoopQAdapter: """Backend-neutral LoopQ boundaries for Huginn's true 32 recurrences.""" loop_count = HUGINN_LOOP_COUNT group_names = tuple(PAPER_GROUPS) @staticmethod def module_key(layer_index: int, group_name: str) -> str: if not 0 <= layer_index < HUGINN_PHYSICAL_LAYERS: raise IndexError("Huginn physical layer index must be in [0, 4)") if group_name not in PAPER_GROUPS: raise KeyError(f"unknown Huginn LoopQ group {group_name!r}") return f"transformer.core_block.{layer_index}.{group_name}" @staticmethod def validate_loop(loop_index: int) -> None: if not 0 <= loop_index < HUGINN_LOOP_COUNT: raise IndexError("Huginn LoopQ recurrence index must be in [0, 32)") 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) return ( quantize_weight(folded).dequantized if quantization_enabled else folded ) @staticmethod def apply_transition( hidden_state: torch.Tensor, *, completed_loop_index: int, cta: CrossLoopTransitionAdapter, ) -> torch.Tensor: if not 0 <= completed_loop_index < HUGINN_TRANSITION_COUNT: raise IndexError( "CTA is valid only after Huginn recurrences 0 through 30" ) if cta.transition_count != HUGINN_TRANSITION_COUNT: raise ValueError("Huginn CTA artifact must contain exactly 31 transitions") return cta(hidden_state, completed_loop_index)