Hosstia's picture
Rebuild experts directly from original BF16 (single quantization, no double quantization); remove ja language tag; update README
9c1e445 verified
Raw History Blame Contribute Delete
9.36 kB
#!/usr/bin/env python3
"""Direct BF16 -> mixed 2/4-bit MLX conversion of AliceAI-Foundation-80B-A3B-Base.
Source: yandex/AliceAI-Foundation-80B-A3B-Base (original BF16 safetensors,
~151 GiB, 49 shards). Unlike convert_f80b_mixed.py (which requantizes the
Yamada114514 MLX-4bit release and thus double-quantizes the experts), this
converter quantizes the original BF16 weights exactly once:
- MoE experts (model.layers.*.mlp.experts.gate_up_proj / down_proj, 3D
per-expert tensors): BF16 -> affine 2-bit g32 (single quantization).
- All other quantized tensors (self_attn, linear_attn projections,
shared_expert, embed_tokens, lm_head, attnres res_proj): BF16 ->
affine 4-bit g64 (single quantization).
- Norms, routers (mlp.gate), convs (q/k/v_conv1d), dt_bias, a_log_bias:
copied as-is (BF16).
- The MTP module (mtp.*) is dropped, matching the MLX release.
The exact target tensor set and per-tensor bit widths are derived from the
reference MLX-2bit model directory (its model.safetensors.index.json lists
every {name.weight, name.scales, name.biases} triple, and its config.json
holds the per-layer quantization map), so the output is drop-in compatible
with the existing model.py / inference.py / presets.
Usage:
python convert_f80b_direct.py \
--src /tmp/f80b-orig \
--ref literary-zh-ru/models/AliceAI-Foundation-80B-A3B-Base-MLX-2bit \
--dst /tmp/f80b-2bit-direct
"""
from __future__ import annotations
import argparse
import json
import os
import shutil
import time
import mlx.core as mx
EXPERT_MARKER = ".mlp.experts."
DEFAULT_GROUP = 64
def convert_model(src_dir: str, ref_dir: str, dst_dir: str) -> None:
os.makedirs(dst_dir, exist_ok=True)
t0 = time.time()
# --- reference: exact target tensor set + per-layer bits ---
ref_index = json.load(
open(os.path.join(ref_dir, "model.safetensors.index.json"))
)["weight_map"]
ref_config = json.load(open(os.path.join(ref_dir, "config.json")))
per_layer = {
k: v
for k, v in ref_config.get("quantization", {}).items()
if isinstance(v, dict)
}
def target_bits(base: str) -> int:
if base in per_layer:
return per_layer[base]["bits"]
return ref_config["quantization"]["bits"]
def target_group(base: str) -> int:
if base in per_layer:
return per_layer[base]["group_size"]
return ref_config["quantization"]["group_size"]
# Bases that must be quantized (have .scales in the reference)
quant_bases = sorted(
n[: -len(".scales")] for n in ref_index if n.endswith(".scales")
)
# Plain tensors copied as-is (everything in ref without .scales/.biases
# and not a packed-weight member of a quantized triple)
plain_names = sorted(
n
for n in ref_index
if not n.endswith(".scales") and not n.endswith(".biases")
and not (n.endswith(".weight") and n[: -len(".weight")] in quant_bases)
)
print(f"reference: {len(quant_bases)} quantized bases, {len(plain_names)} plain tensors")
# --- source index ---
src_index = json.load(
open(os.path.join(src_dir, "model.safetensors.index.json"))
)["weight_map"]
shard_names: dict[str, list[str]] = {}
for name, shard in src_index.items():
shard_names.setdefault(shard, []).append(name)
_cache: dict[str, dict] = {}
def load_tensor(name: str) -> mx.array:
shard = src_index[name]
if shard not in _cache:
if len(_cache) >= 2:
_cache.clear()
_cache[shard] = mx.load(os.path.join(src_dir, shard))
return _cache[shard][name]
def quantize_2d_or_3d(w: mx.array, group: int, bits: int):
"""Quantize a 2D matrix or a 3D (E, out, in) per-expert tensor.
For 3D tensors, flatten (E, out) -> rows; group_size applies along
the last (input) dim, identical to per-expert quantization. The
packed weights, scales and biases are all reshaped back to 3D
(E, out, packed_cols / cols//group) so SwitchLinear accepts them.
"""
orig_shape = w.shape
w2 = w.reshape(-1, orig_shape[-1]).astype(mx.float32)
qw, qs, qb = mx.quantize(w2, group_size=group, bits=bits)
qw = qw.reshape(*orig_shape[:-1], qw.shape[-1])
qs = qs.reshape(*orig_shape[:-1], qs.shape[-1])
qb = qb.reshape(*orig_shape[:-1], qb.shape[-1])
return qw, qs.astype(mx.bfloat16), qb.astype(mx.bfloat16)
# --- emit shards mirroring the reference layout, one ref shard at a time ---
# Build the ref shard -> tensor names map so each output shard is
# assembled, saved and freed before the next one (keeps RAM bounded).
ref_shard_tensors: dict[str, list[str]] = {}
for name, shard in ref_index.items():
ref_shard_tensors.setdefault(shard, []).append(name)
new_weight_map: dict[str, str] = {}
total_size = 0
n_quant, n_plain = 0, 0
missing = []
triple_cache: dict[str, tuple] = {} # base -> (qw, qs, qb), cleared per shard
def build_tensor(name: str) -> mx.array | None:
"""Produce one output tensor (quantized triple member or plain)."""
nonlocal n_quant, n_plain
base = member = None
if name.endswith(".scales") or name.endswith(".biases"):
base, member = name[: name.rindex(".")], name.rsplit(".", 1)[1]
elif name.endswith(".weight") and name[: -len(".weight")] in quant_bases:
# packed-weight member of a quantized triple (incl. experts,
# whose source tensors carry no ".weight" suffix)
base, member = name[: -len(".weight")], "weight"
if base is not None:
if base in triple_cache:
qw, qs, qb = triple_cache[base]
else:
src_name = base if EXPERT_MARKER in base else base + ".weight"
if src_name not in src_index:
missing.append(src_name)
return None
w = load_tensor(src_name)
qw, qs, qb = quantize_2d_or_3d(
w, target_group(base), target_bits(base)
)
del w
triple_cache[base] = (qw, qs, qb)
n_quant += 1
return {"weight": qw, "scales": qs, "biases": qb}[member]
# plain tensor (norms, routers, convs, dt_bias, ...)
if name not in src_index:
missing.append(name)
return None
n_plain += 1
return load_tensor(name)
for shard, names in sorted(ref_shard_tensors.items()):
ts = time.time()
triple_cache.clear()
tensors: dict[str, mx.array] = {}
for name in names:
arr = build_tensor(name)
if arr is not None:
tensors[name] = arr
new_weight_map[name] = shard
shard_bytes = sum(a.nbytes for a in tensors.values())
mx.save_safetensors(os.path.join(dst_dir, shard), tensors)
total_size += shard_bytes
del tensors
print(f"{shard}: {len(names)} tensors, "
f"{shard_bytes / 1024**3:.2f} GiB, {time.time() - ts:.0f}s")
if missing:
raise SystemExit(f"ERROR: {len(missing)} source tensors not found, e.g. {missing[:5]}")
with open(os.path.join(dst_dir, "model.safetensors.index.json"), "w") as f:
json.dump(
{"metadata": {"total_size": total_size}, "weight_map": new_weight_map},
f, indent=2,
)
# --- config / code / presets: copy from the reference model ---
for fname in os.listdir(ref_dir):
if fname.endswith(".safetensors") or fname in (
"model.safetensors.index.json", "conversion.json", "SHA256SUMS",
):
continue
src_f = os.path.join(ref_dir, fname)
if os.path.isfile(src_f):
shutil.copy2(src_f, os.path.join(dst_dir, fname))
# --- conversion.json ---
conv = {
"date": time.strftime("%Y-%m-%d"),
"source": "yandex/AliceAI-Foundation-80B-A3B-Base (original BF16)",
"quantization": {
"scheme": "mixed, direct single quantization from BF16",
"experts": {"bits": 2, "group_size": 32},
"other": {"bits": 4, "group_size": 64},
},
"precision": (
"Direct conversion from the original BF16 release: MoE experts "
"quantized once to affine 2-bit g32; attention, linear_attn, "
"shared_expert, embeddings and lm_head quantized once to affine "
"4-bit g64; norms/routers/convs kept BF16. No double "
"quantization. MTP module dropped (as in the MLX release)."
),
}
with open(os.path.join(dst_dir, "conversion.json"), "w") as f:
json.dump(conv, f, indent=2)
print(f"Done: {n_quant} quantized bases, {n_plain} plain tensors, "
f"{total_size / 1024**3:.2f} GiB, {time.time() - t0:.0f}s")
def main():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--src", required=True, help="original BF16 model dir")
p.add_argument("--ref", required=True, help="reference MLX-2bit model dir")
p.add_argument("--dst", required=True)
args = p.parse_args()
convert_model(args.src, args.ref, args.dst)
if __name__ == "__main__":
main()