jenny / finetune_lora.py
ars4eh's picture
Upload 3 files
96812a6 verified
Raw
History Blame Contribute Delete
9.82 kB
#!/usr/bin/env python3
"""Fine-tune TinyLlama with LoRA for Generation Alpha slang translation."""
from __future__ import annotations
import argparse
import json
import random
from dataclasses import dataclass
from pathlib import Path
from typing import List, Tuple
import pandas as pd
import torch
from datasets import Dataset
from evaluate import load as load_metric
from peft import LoraConfig, get_peft_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
DataCollatorForLanguageModeling,
Trainer,
TrainingArguments,
set_seed,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="LoRA fine-tune TinyLlama for Gen Alpha slang translation."
)
parser.add_argument(
"--corpus_path",
type=str,
default="data/jenny_genA_corpus.csv",
help="CSV with columns gen_a_slang/gena_slang and normal_sentence/plain_english.",
)
parser.add_argument(
"--output_dir",
type=str,
default="outputs",
help="Directory to save the LoRA adapter and evaluation files.",
)
parser.add_argument(
"--model_id",
type=str,
default="TinyLlama/TinyLlama-1.1B-Chat-v1.0",
help="Base chat model checkpoint.",
)
parser.add_argument(
"--max_length",
type=int,
default=256,
help="Sequence length for tokenization.",
)
parser.add_argument("--lr", type=float, default=2e-4, help="Learning rate.")
parser.add_argument(
"--num_train_epochs", type=int, default=3, help="Training epochs."
)
parser.add_argument(
"--train_batch_size",
type=int,
default=4,
help="Per-device train batch size.",
)
parser.add_argument(
"--grad_accum",
type=int,
default=4,
help="Gradient accumulation steps.",
)
parser.add_argument(
"--seed", type=int, default=42, help="Random seed for reproducibility."
)
parser.add_argument(
"--eval_split",
type=float,
default=0.1,
help="Test split fraction from the corpus.",
)
parser.add_argument(
"--save_eval_jsonl",
action="store_true",
help="Export test split to data/gena_test.jsonl for lm-eval.",
)
parser.add_argument(
"--bf16",
action="store_true",
help="Use bfloat16 when available.",
)
parser.add_argument(
"--fp16",
action="store_true",
help="Use float16 when bfloat16 is unavailable.",
)
return parser.parse_args()
@dataclass
class EncodedExample:
input_ids: List[int]
attention_mask: List[int]
labels: List[int]
def load_corpus(path: str) -> Dataset:
df = pd.read_csv(path)
# Normalize expected column names
rename_map = {
"gen_a_slang": "gena_slang",
"normal_sentence": "plain_english",
}
df = df.rename(columns=rename_map)
missing = {"gena_slang", "plain_english"} - set(df.columns)
if missing:
raise ValueError(f"Missing required columns: {missing}")
df = df[["gena_slang", "plain_english"]].dropna()
df["gena_slang"] = df["gena_slang"].str.strip()
df["plain_english"] = df["plain_english"].str.strip()
return Dataset.from_pandas(df)
def split_dataset(ds: Dataset, test_size: float, seed: int) -> Tuple[Dataset, Dataset]:
ds = ds.shuffle(seed=seed)
split = ds.train_test_split(test_size=test_size, seed=seed)
return split["train"], split["test"]
def build_tokenizer(model_id: str):
tokenizer = AutoTokenizer.from_pretrained(model_id)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
return tokenizer
def format_and_tokenize(example: dict, tokenizer, max_length: int) -> EncodedExample:
messages = [
{
"role": "system",
"content": "Translate the following Generation Alpha slang sentence to plain English.",
},
{"role": "user", "content": example["gena_slang"]},
{"role": "assistant", "content": example["plain_english"]},
]
full_text = tokenizer.apply_chat_template(messages, tokenize=False)
tokenized = tokenizer(
full_text,
truncation=True,
max_length=max_length,
padding="max_length",
)
prompt = tokenizer.apply_chat_template(
messages[:-1], tokenize=False, add_generation_prompt=True
)
prompt_len = len(tokenizer(prompt)["input_ids"])
labels = [-100] * prompt_len + tokenized["input_ids"][prompt_len:]
tokenized["labels"] = labels
return tokenized
def tokenize_dataset(dataset: Dataset, tokenizer, max_length: int) -> Dataset:
return dataset.map(
lambda ex: format_and_tokenize(ex, tokenizer, max_length),
remove_columns=dataset.column_names,
)
def create_lora_model(model_id: str):
if torch.cuda.is_available() and torch.cuda.is_bf16_supported():
dtype = torch.bfloat16
elif torch.cuda.is_available():
dtype = torch.float16
else:
dtype = torch.float32
model = AutoModelForCausalLM.from_pretrained(
model_id, torch_dtype=dtype, device_map="auto"
)
lora_config = LoraConfig(
r=16,
lora_alpha=32,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
return model
def save_test_jsonl(test_ds: Dataset, tokenizer, path: Path) -> None:
records = []
for ex in test_ds:
prompt = tokenizer.apply_chat_template(
[
{
"role": "system",
"content": "Translate the following Generation Alpha slang sentence to plain English.",
},
{"role": "user", "content": ex["gena_slang"]},
],
tokenize=False,
add_generation_prompt=True,
)
records.append({"input": prompt, "outputs": [ex["plain_english"]]})
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as handle:
for row in records:
handle.write(json.dumps(row) + "\n")
def evaluate_model(model, tokenizer, test_ds: Dataset, max_length: int) -> dict:
sacrebleu = load_metric("sacrebleu")
rouge = load_metric("rouge")
preds, refs = [], []
model.eval()
for ex in test_ds:
messages = [
{
"role": "system",
"content": "Translate the following Generation Alpha slang sentence to plain English.",
},
{"role": "user", "content": ex["gena_slang"]},
]
prompt = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
encoded = tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=max_length,
).to(model.device)
with torch.no_grad():
output_ids = model.generate(
**encoded,
max_new_tokens=64,
do_sample=False,
)
text = tokenizer.decode(output_ids[0], skip_special_tokens=True).strip()
preds.append(text)
refs.append(ex["plain_english"])
bleu = sacrebleu.compute(predictions=preds, references=[[r] for r in refs])
rouge_res = rouge.compute(predictions=preds, references=refs, use_stemmer=True)
rouge_l = rouge_res["rougeL"]
if hasattr(rouge_l, "mid"):
rouge_l = rouge_l.mid.fmeasure
elif isinstance(rouge_l, dict) and "mid" in rouge_l:
rouge_l = rouge_l["mid"].get("fmeasure", rouge_l["mid"].get("f"))
rouge_l = float(rouge_l)
return {
"sacrebleu": bleu["score"],
"rougeL": rouge_l,
"n_samples": len(refs),
}
def main() -> None:
args = parse_args()
set_seed(args.seed)
random.seed(args.seed)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
tokenizer = build_tokenizer(args.model_id)
raw_ds = load_corpus(args.corpus_path)
train_ds, test_ds = split_dataset(raw_ds, test_size=args.eval_split, seed=args.seed)
tokenized_train = tokenize_dataset(train_ds, tokenizer, args.max_length)
model = create_lora_model(args.model_id)
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
training_args = TrainingArguments(
output_dir=str(output_dir / "checkpoints"),
num_train_epochs=args.num_train_epochs,
per_device_train_batch_size=args.train_batch_size,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
bf16=args.bf16,
fp16=args.fp16 and not args.bf16,
save_steps=500,
logging_steps=25,
weight_decay=0.01,
optim="paged_adamw_8bit",
report_to="none",
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_train,
data_collator=data_collator,
)
trainer.train()
adapter_dir = output_dir / "jenny_lora_adapter"
trainer.save_model(str(adapter_dir))
tokenizer.save_pretrained(adapter_dir)
metrics = evaluate_model(model, tokenizer, test_ds, args.max_length)
with (output_dir / "eval_metrics.json").open("w", encoding="utf-8") as handle:
json.dump(metrics, handle, indent=2)
print("Eval metrics:", metrics)
if args.save_eval_jsonl:
save_test_jsonl(test_ds, tokenizer, Path("data/gena_test.jsonl"))
print("Saved data/gena_test.jsonl for lm-eval.")
if __name__ == "__main__":
main()