LOMONOSOV-ZENIT-27B-1M-INDEV / count_parameters.py
Ddavidich's picture
Сплошная сверка чисел карточки: заголовочное число было подписано неверно
f2dce28 verified
Raw
History Blame Contribute Delete
6.99 kB
#!/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("<Q", fh.read(8))[0]
return json.loads(fh.read(n))
def count(directory: str) -> 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())