45 lines
2.0 KiB
Python
45 lines
2.0 KiB
Python
"""Convert a PEFT LoRA adapter (Unsloth / HF) to the mlx-lm adapter format.
|
|
|
|
train/.venv/bin/python train/peft_to_mlx.py <peft_adapter_dir> <mlx_adapter_dir>
|
|
|
|
PEFT key : base_model.model.model.language_model.layers.N.<mod>.lora_A.weight shape (r, in)
|
|
base_model.model.model.language_model.layers.N.<mod>.lora_B.weight shape (out, r)
|
|
mlx key : language_model.model.layers.N.<mod>.lora_a shape (in, r)
|
|
language_model.model.layers.N.<mod>.lora_b shape (r, out)
|
|
Scale : PEFT scaling = alpha / r; mlx `scale` is used directly (y + scale * x @ A @ B), so scale = alpha / r.
|
|
mlx loads only the modules listed in `lora_parameters.keys` (relative to the layer); they are taken from the file.
|
|
"""
|
|
import json, re, sys
|
|
from pathlib import Path
|
|
|
|
import mlx.core as mx
|
|
|
|
src, dst = Path(sys.argv[1]), Path(sys.argv[2])
|
|
cfg = json.load(open(src / "adapter_config.json"))
|
|
r, alpha = cfg["r"], cfg["lora_alpha"]
|
|
if cfg.get("use_rslora") or cfg.get("use_dora"):
|
|
sys.exit("rsLoRA / DoRA: scale is not alpha / r, not supported")
|
|
|
|
w = mx.load(str(src / "adapter_model.safetensors"))
|
|
pat = re.compile(r"^base_model\.model\.model\.language_model\.layers\.(\d+)\.(.+)\.lora_([AB])\.weight$")
|
|
out, keys, layers = {}, set(), set()
|
|
for k, v in w.items():
|
|
m = pat.match(k)
|
|
if not m:
|
|
sys.exit(f"unexpected key: {k}")
|
|
n, mod, ab = m.groups()
|
|
layers.add(int(n))
|
|
keys.add(mod)
|
|
out[f"language_model.model.layers.{n}.{mod}.lora_{ab.lower()}"] = mx.transpose(v).astype(mx.float32)
|
|
if ab == "A":
|
|
assert v.shape[0] == r, (k, v.shape)
|
|
|
|
dst.mkdir(parents=True, exist_ok=True)
|
|
mx.save_safetensors(str(dst / "adapters.safetensors"), out)
|
|
json.dump({
|
|
"fine_tune_type": "lora",
|
|
"num_layers": max(layers) + 1,
|
|
"lora_parameters": {"rank": r, "scale": alpha / r, "dropout": 0.0, "keys": sorted(keys)},
|
|
}, open(dst / "adapter_config.json", "w"), indent=1)
|
|
print(f"{len(out)} tensors, {len(layers)} layers, rank {r}, alpha {alpha} -> scale {alpha / r}, modules {sorted(keys)}")
|