File size: 1,864 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
"""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")