JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
1.86 kB
"""Executable numerical semantics, not a claim of unpublished author settings."""
from collections.abc import Mapping
CONTRACT_VERSION = 1
def numerical_contract(activation_bits: int, *, dynamic_las: bool = True) -> dict:
if activation_bits not in (4, 8):
raise ValueError("activation bits must be 4 or 8")
return {
"version": CONTRACT_VERSION,
"weight_bits": 4,
"activation_bits": activation_bits,
"weight_range": [-8, 7],
"activation_range": [-2 ** (activation_bits - 1), 2 ** (activation_bits - 1) - 1],
"group_size": 32,
"axis": -1,
"rounding": "nearest_ties_away_from_zero",
"scale_dtype": "float32",
"scale_estimator": "group_absmax_over_positive_qmax; zero_group_scale_1",
"tail_group": "independent_unpadded_output",
"activation_scale": "token_group_absmax_times_module_loop_clip" if dynamic_las else "static_group_scale",
"las_multiplier": "exp_log_scale_clamp_max_1_with_identity_clamp_ste" if dynamic_las else "exp_log_scale",
"transform": "X_P__W_inverse_transpose_P",
"cta_boundary": "post_boundary_norm_before_next_loop; no_final_cta",
"status": "explicit_local_interpretation_of_paper",
}
def validate_numerical_contract(contract, activation_bits: int, *, dynamic_las: bool) -> None:
expected = numerical_contract(activation_bits, dynamic_las=dynamic_las)
if not isinstance(contract, Mapping):
raise ValueError("numerical contract is missing or invalid; do not relabel legacy artifacts")
if contract != expected:
differing = sorted(key for key in set(expected) | set(contract or {})
if (contract or {}).get(key) != expected.get(key))
raise ValueError(f"unsupported numerical contract: {differing}; do not relabel legacy artifacts")