Stage 1 on HF Jobs: Unsloth job script, PEFT to MLX converter, base valid loss 0.849, Qwen base model docs
Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
44
train/peft_to_mlx.py
Normal file
44
train/peft_to_mlx.py
Normal file
@@ -0,0 +1,44 @@
|
||||
"""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)}")
|
||||
Reference in New Issue
Block a user