#!/usr/bin/env python3 """Пересчёт параметров по заголовкам shard-файлов. Веса не читаются. Зачем. Заголовочное число карточки — 26 895 998 464 — не имело первоисточника: его нет ни в конфиге, ни в родословных, ни в одной квитанции. Правило выпуска запрещает в карточке цифры без замера, поэтому число надо либо подтвердить, либо исправить. Почему «в лоб» не считается. Веса упакованы, и упакованы по-разному: * `weight_packed` встречается и в `U8`, и в `I32` — то есть число элементов тензора само по себе о числе параметров не говорит; * ширина тоже разная: NVFP4 кладёт по два значения в байт, эмбеддинги квантованы в 8 бит, зрение — `SELECTIVE_W8_W4_A16`, то есть вперемешку. Первая попытка удвоила все `U8` подряд и дала 27 553 793 896 — мимо на 658 миллионов. Вторая читала число элементов вместо байт и промахнулась в другую сторону. Как считается здесь. Исходная ширина восстанавливается из таблицы масштабов: при групповом квантовании `weight_scale` имеет форму [out, in/group_size], а `weight_packed` занимает `in * bits / 8` байт на строку. Перебирая размер группы из конфига (16, 32, 64, 128), берём тот, при котором ширина выходит ровно 4 или 8 бит. Неразрешённых тензоров при этом не остаётся ни одного — это и есть проверка, что подбор не выдумывает. Служебные тензоры (масштабы, нули, формы) в счёт не идут; нормировки, смещения и параметры Gated DeltaNet (`A_log`, `dt_bias`, `altay_alpha`) идут. """ from __future__ import annotations import argparse import collections import glob import json import math import os import struct BYTES = {"U8": 1, "I8": 1, "I32": 4, "I64": 8, "F8_E4M3": 1, "BF16": 2, "F16": 2, "F32": 4} # Не параметры: таблицы масштабов, нулей и сохранённые формы. SKIP = ("weight_scale", "weight_shape", "weight_zero_point", "weight_global_scale", "input_scale", "input_global_scale", "weight_g_idx") def bucket(name: str) -> str: if name.startswith("model.visual"): return "vision_tower" if "altay" in name.lower(): return "altay_overlay" if "embed_tokens" in name: return "embed_tokens" if name.startswith("lm_head") or ".lm_head" in name: return "lm_head" return "language_model" def header(path: str) -> dict: with open(path, "rb") as fh: n = struct.unpack(" dict: per = collections.Counter() by_file = collections.Counter() by_width = collections.Counter() unresolved: list = [] for path in sorted(glob.glob(os.path.join(directory, "*.safetensors"))): hdr = header(path) shapes = {k: (v.get("shape") or [], v.get("dtype")) for k, v in hdr.items() if k != "__metadata__"} for name, (shape, dtype) in shapes.items(): if name.endswith(SKIP): continue n = 0 if name.endswith(".weight_packed"): scale = shapes.get(name[: -len(".weight_packed")] + ".weight_scale") if not scale or len(shape) < 2 or len(scale[0]) < 2: unresolved.append(name) continue out_features = shape[0] row_bytes = shape[1] * BYTES[dtype] groups = scale[0][1] for group_size in (16, 32, 64, 128): in_features = groups * group_size bits = row_bytes * 8 / in_features if abs(bits - 4) < 1e-9 or abs(bits - 8) < 1e-9: n = out_features * in_features by_width[int(round(bits))] += n break else: unresolved.append(name) continue elif dtype in ("F8_E4M3", "BF16", "F32", "F16"): n = math.prod(shape) if shape else 1 per[bucket(name)] += n by_file[os.path.basename(path)] += n packed_total = sum(by_width.values()) total = sum(per.values()) language = per["language_model"] + per["embed_tokens"] + per["lm_head"] return { "schema": "lomonosov_zenit_parameter_count_v1", "method": "заголовки safetensors; ширина восстановлена из формы weight_scale", "unresolved_tensors": unresolved, "by_purpose": dict(per), "by_file": dict(by_file), "quantised_by_width_bits": {str(k): v for k, v in sorted(by_width.items())}, "packed_total": packed_total, "unpacked_total": total - packed_total, "language_model_total": language, "vision_tower": per["vision_tower"], "altay_overlay": per["altay_overlay"], "total_all_parts": total, } def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--model", required=True) ap.add_argument("--out") args = ap.parse_args() result = count(args.model) cfg = json.load(open(os.path.join(args.model, "config.json"), encoding="utf-8")) result["tie_word_embeddings"] = cfg.get("tie_word_embeddings") result["note_on_tying"] = ( "эмбеддинги НЕ связаны, поэтому embed_tokens и lm_head считаются отдельно " "и двойного счёта нет" ) result["card_headline_figure"] = 26_895_998_464 result["card_figure_equals"] = ( "language_model_total: заголовочное число карточки — счёт ТОЛЬКО языковой " "части, без зрения и без оверлея ALTAY" ) result["card_figure_matches_language_model"] = ( result["language_model_total"] == 26_895_998_464 ) if args.out: with open(args.out, "w", encoding="utf-8") as fh: json.dump(result, fh, ensure_ascii=False, indent=1) print(json.dumps(result, ensure_ascii=False, indent=1)) return 0 if __name__ == "__main__": raise SystemExit(main())