Download loopq_quantization/scripts/adapters/ouro.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 6.16 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/adapters/ouro.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/adapters/ouro.py
-
curl -L -o ouro.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/adapters/ouro.py
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", | |
| }, | |
| } | |
| 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) | |
| 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}" | |
| 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 | |
| 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 | |
| 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) | |