""" 精度对齐:芯片适配的最后一个坎,在你自己的电脑上跑一遍。 同一个模型、同一份输入,在两块芯片上跑。两边的算子数学上等价, 但实现细节不同,结果不会逐位相同。适配工程师要判断:这点差别是正常的, 还是某个算子写错了。 这个脚本把三种情况放在同一套指标下看,看看哪个指标能分辨它们。 先说清楚这个脚本的四条边界。它们逐字印在输出的开头,也逐字贴在页面上: - 一台机器上用代码模拟的对照,不是两块真实芯片的对比。 「参考」和「被测」跑的是同一份 Python 代码,差别是代码里人为加进去的: 一条舍入输入表示,一条把方差分母写错,一条直接把输出整体放大。 - bf16 只对矩阵乘的输入做表示舍入,乘法与累加仍是双精度,不是硬件 bf16 内核。 累加器宽度、累加顺序与专用乘加单元一样都没模拟。 - 余弦相似度按六位小数显示,1.000000 是显示值,不表示完整精度下相等。 所以表里永远跟着一列「1-余弦」,把六位小数遮掉的量级摆出来。 - 门槛随张量尺度、层深和任务变,没有单一指标或统一阈值能判断所有算子。 这里的 0.999 只是这个玩具配置里的一个演示门槛。 情况三是人为构造的缩放反例,用来暴露余弦相似度看不见什么,不是哪块芯片真会犯的错。 只需要 python3,不用装任何东西: python3 align_check.py # 默认就跑这个,什么参数都不用给 python3 align_check.py --help # 想换配置(种子、层数、维度、行数)时再看 python3 align_check.py --json # 机器可读记录,附运行环境与脚本 sha256 脚本没有计时,所以可复现——但可复现到哪一位有边界,见输出里的「可复现」那一栏。 bf16 舍入、三个误差口径、缩放反例与方差分母错误的回归由 tests/demo/test_align_check.py 独立复核(那一组的逐行比对要求 CPython ≥ 3.12)。 """ import argparse import hashlib import json import math import platform import random import struct import sys import unicodedata LAYERS = 8 DIM = 64 ROWS = 4 SEED = 0 COS_THRESHOLD = 0.999 # 只是这个玩具配置里的演示门槛,不是通用验收线 SCALE_PROBE = 1.5 # 缩放反例:把参考输出整体乘以这个正数 EPS = 1e-5 # layernorm 里加在方差上的常数 COS_DIGITS = 6 # 余弦相似度显示几位小数 REPORT_SCHEMA = "ai-museum/align-check/1" # 这四句是这个实验的边界。它们同时出现在脚本输出和 demo 页面上, # 因为只看一眼表格的读者最容易把「余弦 1.000000」读成「两边算得一样」。 SCOPE_NOTES = ( "一台机器上用代码模拟的对照,不是两块真实芯片的对比", "bf16 只对矩阵乘的输入做表示舍入,乘法与累加仍是双精度,不是硬件 bf16 内核", "余弦相似度按六位小数显示,1.000000 是显示值,不表示完整精度下相等", "门槛随张量尺度、层深和任务变,没有单一指标或统一阈值能判断所有算子", ) # 「可复现」到哪一位是有边界的,这是作者逐个解释器跑出来的结果,不是推断。 # 情况三的真值是 0(被测 = 参考 × 一个正数,方向完全一致),所以那一列上的非零 # 全是求和方式带来的舍入噪声;CPython 3.12 起浮点 sum() 改用补偿求和,噪声就变了。 REPRODUCIBILITY = ( "固定种子、不测耗时。作者在一台 arm64 macOS 上逐个试过 6 个 CPython:", "3.12~3.14 与页面上贴的那份逐字节相同;3.9~3.11 只有情况三的「1-余弦」那列不同", "(1e-16 量级),因为 CPython 3.12 起浮点 sum() 改用了补偿求和。", "情况三的真值是 0,那一列上的非零全是求和方式带来的舍入噪声。换一种 CPU 架构没有试过。", ) def bf16(x): """把浮点数舍入到 bfloat16 的精度(就近舍入,取偶)。""" i = struct.unpack(">I", struct.pack(">f", x))[0] i = (i + 0x7FFF + ((i >> 16) & 1)) & 0xFFFF0000 return struct.unpack(">f", struct.pack(">I", i))[0] def matmul(x, w, low_precision=False): """一次矩阵乘。 low_precision=True 表示这条路径**只把两侧的输入舍入到 bf16 表示**; 乘法与累加照旧走 Python 的 float(双精度)。真实硬件的 bf16 内核还要 决定累加器宽度、累加顺序与分块方式,这里一样都没模拟。 """ n, k, m = len(x), len(w), len(w[0]) out = [[0.0] * m for _ in range(n)] for i in range(n): xi = x[i] for p in range(k): a = bf16(xi[p]) if low_precision else xi[p] if a == 0.0: continue wp, oi = w[p], out[i] for j in range(m): b = bf16(wp[j]) if low_precision else wp[j] oi[j] += a * b return out def layernorm(x, buggy=False): """归一化。buggy=True 模拟一个真实常见的算子 bug:方差除以 n-1 而不是 n。""" out = [] for row in x: n = len(row) mu = sum(row) / n denom = (n - 1) if buggy else n var = sum((v - mu) ** 2 for v in row) / denom s = math.sqrt(var + EPS) out.append([(v - mu) / s for v in row]) return out def forward(weights, x, low_precision=False, buggy=False): outs = [] for w in weights: x = matmul(x, w, low_precision) x = [[math.tanh(v) for v in row] for row in x] x = layernorm(x, buggy) outs.append(x) return outs def scaled(x, factor): """把一层输出整体乘以一个正数。缩放反例用的「被测」路径。""" return [[v * factor for v in row] for row in x] # -------------------------------------------------------------------------- # 四个指标。口径写在这里,页面和测试都引用这里的定义。 # -------------------------------------------------------------------------- def _flat(a, b): return [v for row in a for v in row], [v for row in b for v in row] def cosine(a, b): """余弦相似度:只看方向。对正比例缩放完全不敏感——情况三就是拿它做反例。""" fa, fb = _flat(a, b) dot = sum(p * q for p, q in zip(fa, fb)) na = math.sqrt(sum(p * p for p in fa)) nb = math.sqrt(sum(q * q for q in fb)) return dot / (na * nb) def max_abs_diff(a, b): """最大绝对误差 max|a-b|。带量纲,随张量整体尺度一起放大。""" fa, fb = _flat(a, b) return max(abs(p - q) for p, q in zip(fa, fb)) def rmse(a, b): """均方根误差 sqrt(mean((a-b)^2))。同样带量纲,但不被单个离群点主导。""" fa, fb = _flat(a, b) return math.sqrt(sum((p - q) ** 2 for p, q in zip(fa, fb)) / len(fa)) def rel_rmse(a, b): """相对均方根误差 ‖a-b‖₂ ÷ ‖a‖₂。无量纲。 被测输出等于参考的 c 倍时它精确等于 |1-c|,所以缩放看得见。 参考侧全零时没有分母,返回 None,不编一个数出来。 """ fa, fb = _flat(a, b) na = math.sqrt(sum(p * p for p in fa)) if na == 0.0: return None return math.sqrt(sum((p - q) ** 2 for p, q in zip(fa, fb))) / na def compare(ref, dev, threshold=COS_THRESHOLD): """逐层算四个指标。cosGap = 1 - 余弦,给出六位小数显示遮掉的那部分。""" layers = [] for i, (a, b) in enumerate(zip(ref, dev), start=1): cos = cosine(a, b) layers.append( { "layer": i, "cosine": cos, "cosGap": 1.0 - cos, "maxAbs": max_abs_diff(a, b), "rmse": rmse(a, b), "relRmse": rel_rmse(a, b), "cosinePass": cos >= threshold, } ) return layers # -------------------------------------------------------------------------- def build_report( seed=SEED, layers=LAYERS, dim=DIM, rows=ROWS, scale=SCALE_PROBE, threshold=COS_THRESHOLD ): """跑三种情况,返回一份纯数据的结果。渲染在 render_text / render_json 里。""" if layers < 1 or rows < 1 or dim < 2: raise ValueError(f"层数与行数至少是 1、维度至少是 2,收到 层={layers}、维度={dim}、行={rows}") if scale <= 0: raise ValueError(f"缩放反例的倍数要是正数(余弦对正比例缩放才不敏感),收到 {scale}") rnd = random.Random(seed) w_scale = dim**-0.5 weights = [ [[rnd.gauss(0, w_scale) for _ in range(dim)] for _ in range(dim)] for _ in range(layers) ] x0 = [[rnd.gauss(0, 1) for _ in range(dim)] for _ in range(rows)] ref = forward(weights, x0) cases = [ { "key": "lowPrecision", "title": "只有低精度输入表示(bf16 舍入),算子没写错", "what": "两侧输入舍入到 bf16 表示后再做同一次乘加,累加仍是双精度", "layers": compare(ref, forward(weights, x0, low_precision=True), threshold), }, { "key": "varianceDenominator", "title": "算子写错了:方差除以 n-1 而不是 n", "what": f"layernorm 的方差分母写成 {dim - 1} 而不是 {dim},其余一模一样", "layers": compare(ref, forward(weights, x0, buggy=True), threshold), }, { "key": "scaleProbe", "title": f"缩放反例:把参考输出整体乘以 {scale:g}(人为构造,不是芯片会犯的错)", "what": f"被测输出 = 参考输出 × {scale:g},方向一个字没变", "layers": compare(ref, [scaled(a, scale) for a in ref], threshold), }, ] for case in cases: case["cosinePassCount"] = sum(1 for r in case["layers"] if r["cosinePass"]) return { "schema": REPORT_SCHEMA, "config": { "seed": seed, "layers": layers, "dim": dim, "rows": rows, "scale": scale, "cosThreshold": threshold, "eps": EPS, "cosDigits": COS_DIGITS, }, "inputs": ( f"random.Random({seed}):按顺序先生成 {layers} 层 {dim}×{dim} 权重" f"(高斯,均值 0、标准差 {dim}^-0.5),再生成 {rows}×{dim} 输入(高斯,均值 0、标准差 1)" ), "reference": ( "参考路径是同一份代码在 Python float(IEEE 754 双精度)下跑出来的输出," "不是另一块芯片的输出,也不是更高精度的真值" ), "displayRounding": ( f"余弦相似度按 {COS_DIGITS} 位小数显示,其余三列按三位有效数字的科学计数法显示;" "「1-余弦」那一列专门给出六位小数遮掉的部分" ), "metrics": { "cosine": "余弦相似度 = a·b ÷ (‖a‖₂‖b‖₂),只看方向", "cosGap": "1-余弦 = 六位小数显示遮掉的那部分", "maxAbs": "最大绝对误差 = max|a-b|,带量纲,随张量整体尺度一起变", "rmse": "均方根误差 = sqrt(mean((a-b)²)),带量纲,不被单个离群点主导", "relRmse": "相对均方根误差 = ‖a-b‖₂ ÷ ‖a‖₂,无量纲;被测 = 参考 × c 时精确等于 |1-c|", }, "scopeNotes": list(SCOPE_NOTES), "reproducibility": list(REPRODUCIBILITY), "cases": cases, } def _wide(text): """东亚宽字符按两列算,好让表头和数字列在等宽字体里真的对齐。""" return sum(2 if unicodedata.east_asian_width(ch) in "WF" else 1 for ch in text) def _rpad(text, width): """按显示宽度右对齐。`f"{s:>12}"` 数的是字符数,中文表头会对不齐。""" return " " * max(0, width - _wide(text)) + text # 表头、列宽与取数方式写在一起,渲染和测试引用同一份定义 COLUMNS = ( ("层", 4, lambda r, cfg: str(r["layer"])), ("余弦相似度", 13, lambda r, cfg: f"{r['cosine']:.{cfg['cosDigits']}f}"), ("1-余弦", 12, lambda r, cfg: f"{r['cosGap']:.3e}"), ("最大绝对误差", 15, lambda r, cfg: f"{r['maxAbs']:.3e}"), ("均方根误差", 14, lambda r, cfg: f"{r['rmse']:.3e}"), ("相对均方根误差", 17, lambda r, cfg: "—" if r["relRmse"] is None else f"{r['relRmse']:.3e}"), ("余弦判定", 10, lambda r, cfg: "通过" if r["cosinePass"] else "不通过"), ) def render_text(report): """可读摘要。所有数字都来自同一份 report,不重跑。""" cfg = report["config"] notes = report["scopeNotes"] lines = [ f"精度对齐:同一份输入,参考路径和 {len(report['cases'])} 条被测路径差多少", f"这是什么 {notes[0]}", *[f" {note}" for note in notes[1:]], f"参数 {cfg['layers']} 层 · 维度 {cfg['dim']} · {cfg['rows']} 行输入 · 随机种子" f" {cfg['seed']} · 余弦门槛 {cfg['cosThreshold']} · 缩放反例 ×{cfg['scale']:g}" f" · layernorm 的 eps {cfg['eps']:g}", f"输入 {report['inputs']}", f"参考计算 {report['reference']}", f"显示舍入 {report['displayRounding']}", f"指标口径 {report['metrics']['maxAbs']}", f" {report['metrics']['rmse']}", f" {report['metrics']['relRmse']}", f"可复现 {report['reproducibility'][0]}", *[f" {line}" for line in report["reproducibility"][1:]], ] header = "".join(_rpad(name, width) for name, width, _ in COLUMNS) for index, case in enumerate(report["cases"]): lines += [ "", f"【情况{'一二三四五'[index]} · {case['title']}】", f" {case['what']}", header, "-" * _wide(header), ] for r in case["layers"]: lines.append("".join(_rpad(fmt(r, cfg), width) for _, width, fmt in COLUMNS)) lines.append( f" 余弦判定(门槛 {cfg['cosThreshold']}):{case['cosinePassCount']}/" f"{len(case['layers'])} 层通过" ) allpass = sum(1 for c in report["cases"] if c["cosinePassCount"] == len(c["layers"])) lines += [ "", f"三张表放在一起看:余弦判定在 {allpass}/{len(report['cases'])} 种情况里都是「全部通过」,", "包括第二种真写错的算子和第三种纯缩放——余弦只看方向,看不见整体缩放。", "带量纲的最大绝对误差和均方根误差看得见缩放,但数值随张量尺度一起变,换个模型、换一层就要换门槛;", f"无量纲的相对均方根误差在情况三里精确等于 |1-{cfg['scale']:g}| = {abs(1 - cfg['scale']):g}。", "所以真实的对齐脚本会同时看几个指标、逐层看趋势,而不是宣布某一个指标配某一个阈值就能定所有算子的对错。", ] return "\n".join(lines) def _environment(): """机器可读记录里的环境信息。脚本没有计时,这几项只用来标明这份记录是谁跑的。""" try: with open(__file__, "rb") as fh: digest = hashlib.sha256(fh.read()).hexdigest() except OSError: # 被 exec 进内存跑的时候没有文件可读,如实留空 digest = None return { "python": platform.python_version(), "implementation": platform.python_implementation(), "platform": platform.platform(), "scriptSha256": digest, } def render_json(report): """机器可读记录。和可读摘要是同一份 report 的两种渲染,额外附运行环境。""" return json.dumps({**report, "environment": _environment()}, ensure_ascii=False, indent=2) def _at_least(name, minimum): def parse(text): value = int(text) if value < minimum: raise argparse.ArgumentTypeError(f"{name}至少是 {minimum},收到 {value}") return value return parse def main(argv=None): parser = argparse.ArgumentParser( description="精度对齐的玩具实验:同一份输入,三种情况下几个误差指标分别看到了什么", epilog="不带参数直接跑就是页面上那一份默认配置。", ) parser.add_argument("--seed", type=int, default=SEED, help="随机种子") parser.add_argument("--layers", type=_at_least("层数", 1), default=LAYERS, help="跑几层") parser.add_argument("--dim", type=_at_least("维度", 2), default=DIM, help="每层的维度") parser.add_argument("--rows", type=_at_least("输入行数", 1), default=ROWS, help="一次过几行输入") parser.add_argument("--scale", type=float, default=SCALE_PROBE, help="缩放反例的倍数(正数)") parser.add_argument("--threshold", type=float, default=COS_THRESHOLD, help="余弦判定的演示门槛") parser.add_argument("--json", action="store_true", help="输出机器可读记录,附环境与脚本 sha256") args = parser.parse_args(argv) try: report = build_report( seed=args.seed, layers=args.layers, dim=args.dim, rows=args.rows, scale=args.scale, threshold=args.threshold, ) except ValueError as exc: parser.error(str(exc)) print(render_json(report) if args.json else render_text(report)) return 0 if __name__ == "__main__": sys.exit(main())