#!/usr/bin/env python3 """Store the connector component (Qwen2 1.5B, shipped as float32) as bfloat16. infer.py loads the connector with torch_dtype=bfloat16, so the tensors it runs with are the fp32 values rounded to bf16 at load time. This does the same rounding once, offline, and proves every converted tensor equals `fp32_tensor.to(torch.bfloat16)` exactly — the runtime model is unchanged; only the download halves. usage: convert_connector.py SRC_CONNECTOR_DIR DST_CONNECTOR_DIR """ import json import shutil import sys from pathlib import Path import torch from safetensors import safe_open from safetensors.torch import load_file, save_file def main(): src, dst = Path(sys.argv[1]), Path(sys.argv[2]) dst.mkdir(parents=True, exist_ok=True) if any(dst.glob("*.safetensors")): sys.exit(f"refusing: {dst} already contains safetensors") index = json.loads((src / "model.safetensors.index.json").read_text()) shards = sorted(set(index["weight_map"].values())) def converted(tensor): return tensor.to(torch.bfloat16) if tensor.is_floating_point() else tensor out = {} for shard in shards: with safe_open(str(src / shard), "pt") as handle: for key in handle.keys(): if key in out: sys.exit(f"duplicate tensor {key}") out[key] = converted(handle.get_tensor(key)) if set(out) != set(index["weight_map"]): sys.exit("tensor set does not match the index weight_map") target = dst / "model.safetensors" save_file(out, str(target), metadata={"format": "pt"}) del out back = load_file(str(target)) checked = 0 for shard in shards: with safe_open(str(src / shard), "pt") as handle: for key in handle.keys(): reference = converted(handle.get_tensor(key)) if back[key].dtype != reference.dtype or not torch.equal(back[key], reference): sys.exit(f"MISMATCH {key}") checked += 1 if checked != len(back): sys.exit(f"checked {checked} tensors but the output holds {len(back)}") for path in src.iterdir(): if path.suffix == ".safetensors" or path.name == "model.safetensors.index.json": continue shutil.copy2(path, dst / path.name) config = json.loads((dst / "config.json").read_text()) key = "dtype" if "dtype" in config else "torch_dtype" previous = config.get(key) config[key] = "bfloat16" (dst / "config.json").write_text(json.dumps(config, indent=2) + "\n") dtypes = sorted({str(t.dtype) for t in back.values()}) print(f"CONNECTOR_OK tensors={checked} exact=all dtypes={dtypes} bytes={target.stat().st_size} " f"config.{key}: {previous} -> bfloat16") if __name__ == "__main__": main()