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