#!/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()