| |
| """LoRA fine-tune t5-smaller from JSONL rows containing input and target.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
|
|
| from datasets import load_dataset |
| from peft import LoraConfig, TaskType, get_peft_model, prepare_model_for_kbit_training |
| from transformers import ( |
| AutoModelForSeq2SeqLM, |
| AutoTokenizer, |
| DataCollatorForSeq2Seq, |
| Seq2SeqTrainer, |
| Seq2SeqTrainingArguments, |
| ) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--train-file", required=True) |
| parser.add_argument("--model", default="ShinpacheShimura/t5-smaller") |
| parser.add_argument("--subfolder", default="optimized-flan-t5-small") |
| parser.add_argument("--output-dir", default="t5-smaller-lora") |
| parser.add_argument("--epochs", type=float, default=3.0) |
| parser.add_argument("--batch-size", type=int, default=4) |
| parser.add_argument("--learning-rate", type=float, default=2e-4) |
| args = parser.parse_args() |
|
|
| common = {"subfolder": args.subfolder} if args.subfolder else {} |
| tokenizer = AutoTokenizer.from_pretrained(args.model, **common) |
| model = AutoModelForSeq2SeqLM.from_pretrained(args.model, device_map="auto", **common) |
| model = prepare_model_for_kbit_training(model) |
| model = get_peft_model( |
| model, |
| LoraConfig( |
| task_type=TaskType.SEQ_2_SEQ_LM, |
| target_modules=["q", "v"], |
| r=8, |
| lora_alpha=16, |
| lora_dropout=0.05, |
| ), |
| ) |
|
|
| dataset = load_dataset("json", data_files=args.train_file, split="train") |
| missing = {"input", "target"} - set(dataset.column_names) |
| if missing: |
| raise SystemExit(f"Missing fields: {', '.join(sorted(missing))}") |
|
|
| def tokenize(batch): |
| encoded = tokenizer(batch["input"], truncation=True, max_length=256) |
| encoded["labels"] = tokenizer( |
| text_target=batch["target"], truncation=True, max_length=128 |
| )["input_ids"] |
| return encoded |
|
|
| tokenized = dataset.map(tokenize, batched=True, remove_columns=dataset.column_names) |
| train_args = Seq2SeqTrainingArguments( |
| output_dir=args.output_dir, |
| num_train_epochs=args.epochs, |
| per_device_train_batch_size=args.batch_size, |
| gradient_accumulation_steps=4, |
| learning_rate=args.learning_rate, |
| logging_steps=10, |
| save_strategy="epoch", |
| report_to="none", |
| fp16=True, |
| ) |
| trainer = Seq2SeqTrainer( |
| model=model, |
| args=train_args, |
| train_dataset=tokenized, |
| processing_class=tokenizer, |
| data_collator=DataCollatorForSeq2Seq(tokenizer, model=model), |
| ) |
| trainer.train() |
| trainer.save_model(args.output_dir) |
| tokenizer.save_pretrained(args.output_dir) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|