#!/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()