| |
| """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"]: |
| |
| 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", |
| ) |
|
|
| 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()) |
|
|