#!/usr/bin/env python3 """Fine-tune Gemma 4 E2B on the Apochat chat dataset using PEFT/LoRA. Designed to run on a Hugging Face GPU Space or any CUDA machine with ≥16 GB VRAM. QLoRA mode (default) uses 4-bit quantization and should fit on a T4/V100. Usage: python scripts/finetune_apochat_peft.py \ --dataset apoapps/apochat-gemma4-e2b-chat-v1 \ --output-dir ./apochat-gemma4-e2b-peft \ --push-to-hub apoapps/apochat-gemma4-e2b-apochat-tuned-v2 \ --use-qlora """ from __future__ import annotations import argparse import os import sys from pathlib import Path import torch from datasets import load_dataset from peft import LoraConfig, TaskType, get_peft_model, prepare_model_for_kbit_training from transformers import ( AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, DataCollatorForLanguageModeling, TrainingArguments, ) from trl import SFTTrainer def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--base-model", default="google/gemma-4-E2B-it") parser.add_argument("--dataset", default="apoapps/apochat-gemma4-e2b-chat-v1") parser.add_argument("--output-dir", default="./apochat-gemma4-e2b-peft") parser.add_argument("--push-to-hub", default=None, help="HF repo to push the adapter") parser.add_argument("--use-qlora", action="store_true", help="Use 4-bit QLoRA (saves VRAM)") parser.add_argument("--epochs", type=float, default=1.0) parser.add_argument("--batch-size", type=int, default=1) parser.add_argument("--gradient-accumulation-steps", type=int, default=4) parser.add_argument("--learning-rate", type=float, default=2e-4) parser.add_argument("--lora-r", type=int, default=16) parser.add_argument("--lora-alpha", type=int, default=32) parser.add_argument("--max-seq-length", type=int, default=1024) return parser.parse_args() def formatting_prompts_func(examples: dict, tokenizer) -> dict: texts = [] for messages in examples["messages"]: # messages is a list of {"role": ..., "content": ...} text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=False, ) texts.append(text) return {"text": texts} def main() -> int: args = parse_args() tokenizer = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token bnb_config = None torch_dtype = torch.bfloat16 if args.use_qlora: bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_use_double_quant=True, ) torch_dtype = torch.bfloat16 model = AutoModelForCausalLM.from_pretrained( args.base_model, quantization_config=bnb_config, torch_dtype=torch_dtype if bnb_config is None else None, device_map="auto", trust_remote_code=True, attn_implementation="eager", # safer for Gemma 4 + gradient checkpointing ) lora_config = LoraConfig( r=args.lora_r, lora_alpha=args.lora_alpha, target_modules=[ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj", "per_layer_input_gate", "per_layer_projection", ], lora_dropout=0.05, bias="none", task_type=TaskType.CAUSAL_LM, ) if args.use_qlora: model = prepare_model_for_kbit_training(model) model = get_peft_model(model, lora_config) model.print_trainable_parameters() ds = load_dataset(args.dataset, data_files={"train": "train.jsonl", "valid": "valid.jsonl"}) train_ds = ds["train"] valid_ds = ds["valid"] if "valid" in ds else None train_ds = train_ds.map( lambda x: formatting_prompts_func(x, tokenizer), batched=True, remove_columns=train_ds.column_names, ) if valid_ds is not None: valid_ds = valid_ds.map( lambda x: formatting_prompts_func(x, tokenizer), batched=True, remove_columns=valid_ds.column_names, ) output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) training_args = TrainingArguments( output_dir=str(output_dir), num_train_epochs=args.epochs, per_device_train_batch_size=args.batch_size, gradient_accumulation_steps=args.gradient_accumulation_steps, learning_rate=args.learning_rate, bf16=True, logging_steps=10, evaluation_strategy="steps" if valid_ds is not None else "no", eval_steps=200, save_strategy="epoch", save_total_limit=2, push_to_hub=bool(args.push_to_hub), hub_model_id=args.push_to_hub, hub_private=True, gradient_checkpointing=True, optim="paged_adamw_8bit" if args.use_qlora else "adamw_torch", report_to="none", ) trainer = SFTTrainer( model=model, tokenizer=tokenizer, train_dataset=train_ds, eval_dataset=valid_ds, max_seq_length=args.max_seq_length, args=training_args, dataset_text_field="text", ) trainer.train() trainer.save_model(str(output_dir / "final_adapter")) if args.push_to_hub: model.push_to_hub(args.push_to_hub, private=True) tokenizer.push_to_hub(args.push_to_hub, private=True) print(f"Training complete. Adapter saved to {output_dir / 'final_adapter'}") return 0 if __name__ == "__main__": sys.exit(main())