File size: 6,203 Bytes
9118991 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | """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)
|