| |
| """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) |
| |
| 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() |
|
|