"""对称 per-tensor / per-row int8 量化与误差比较。"""
import numpy as np


def quantize_symmetric(x: np.ndarray, axis=None):
    max_abs = np.max(np.abs(x), axis=axis, keepdims=True)
    scale = np.maximum(max_abs / 127.0, 1e-12)
    q = np.clip(np.round(x / scale), -127, 127).astype(np.int8)
    return q, scale


def dequantize(q: np.ndarray, scale: np.ndarray) -> np.ndarray:
    return q.astype(np.float32) * scale


if __name__ == "__main__":
    rng = np.random.default_rng(0)
    weights = rng.normal(size=(8, 32)).astype(np.float32)
    weights[0, 0] = 20.0  # 人为离群值。
    q_tensor, s_tensor = quantize_symmetric(weights)
    q_row, s_row = quantize_symmetric(weights, axis=1)
    err_tensor = np.mean((weights - dequantize(q_tensor, s_tensor)) ** 2)
    err_row = np.mean((weights - dequantize(q_row, s_row)) ** 2)
    print("per-tensor MSE:", err_tensor, "per-row MSE:", err_row)
    assert err_row <= err_tensor + 1e-12
