apochat-gemma4-e2b-peft-training / finetune_apochat_peft.py
apoapps's picture
Upload finetune_apochat_peft.py with huggingface_hub
10c3f31 verified
Raw
History Blame Contribute Delete
5.78 kB
#!/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())