gemma4-training-scripts / merge_adapter.py
khalidFlex's picture
Upload merge_adapter.py with huggingface_hub
32ab603 verified
Raw
History Blame Contribute Delete
1.92 kB
# /// 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()