File size: 6,156 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 | """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)
|