File size: 1,916 Bytes
32ab603 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 | # /// script
# dependencies = [
# "transformers>=5.14.0",
# "peft>=0.19.0",
# "torch>=2.5",
# "torchvision>=0.20",
# "accelerate>=1.0",
# "num2words",
# ]
# ///
"""Merge a trained LoRA adapter into base Gemma 4 and push the merged model.
The MLX runtime (gemma4/server.py) can't load PEFT adapters directly, so after
training we bake the adapter into the base weights and push a standalone model.
The Mac then converts that to a quantized MLX build:
# on HF Jobs (this script):
merged = base ⊕ adapter → push to --merged-repo
# locally on the Mac afterwards:
venus/.venv/bin/python -m mlx_vlm convert \
--hf-path khalidFlex/gemma4-gui-agent-merged \
--mlx-path ~/.cache/gemma4-gui-agent-mlx-8bit -q --q-bits 8
"""
import argparse
import torch
from peft import PeftModel
from transformers import AutoModelForImageTextToText, AutoProcessor
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--base", default="google/gemma-4-E4B-it")
ap.add_argument("--adapter", required=True, help="PEFT adapter repo")
ap.add_argument("--merged-repo", required=True, help="where to push the merged model")
ap.add_argument("--public", action="store_true")
args = ap.parse_args()
print(f"[merge] loading base {args.base} (bf16, CPU is fine)")
model = AutoModelForImageTextToText.from_pretrained(args.base, dtype=torch.bfloat16)
processor = AutoProcessor.from_pretrained(args.base)
print(f"[merge] applying adapter {args.adapter}")
model = PeftModel.from_pretrained(model, args.adapter)
model = model.merge_and_unload()
print(f"[merge] pushing merged model to {args.merged_repo}")
model.push_to_hub(args.merged_repo, private=not args.public, max_shard_size="4GB")
processor.push_to_hub(args.merged_repo, private=not args.public)
print("[merge] done.")
if __name__ == "__main__":
main()
|