Ming-Image-0.1-Design-ROCm-INT8 / code /tools /convert_connector.py
kingjones777's picture
Add files using upload-large-folder tool
da1a4ff verified
Raw History Blame Contribute Delete
2.82 kB
#!/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()