JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
6.2 kB
"""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)