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