animesh08/schema-sage
0
1from __future__ import annotations2 3import argparse4from pathlib import Path5 6from datasets import load_dataset7from peft import LoraConfig8from peft import get_peft_model9from transformers import AutoModelForCausalLM, AutoTokenizer, DataCollatorForLanguageModeling, Trainer, TrainingArguments10 11 12def format_example(example: dict[str, str]) -> str:13 return format_prompt(example) + example["sql"]14 15 16def format_prompt(example: dict[str, str]) -> str:17 return f"""### Instruction18{example["instruction"]}19 20### Schema21{example["schema"]}22 23### Question24{example["question"]}25 26### SQL27"""28 29 30def main() -> None:31 parser = argparse.ArgumentParser(description="Train a SchemaSage LoRA Text2SQL adapter.")32 parser.add_argument("--base-model", default="Qwen/Qwen2.5-Coder-0.5B-Instruct")33 parser.add_argument("--train-file", default="data/training/demo_text2sql.jsonl")34 parser.add_argument("--output-dir", default="artifacts/lora/schemasage-qwen-0.5b-adapter")35 parser.add_argument("--epochs", type=float, default=3.0)36 parser.add_argument("--learning-rate", type=float, default=2e-4)37 parser.add_argument("--rank", type=int, default=8)38 parser.add_argument("--max-seq-length", type=int, default=768)39 parser.add_argument("--limit", type=int, default=0, help="Optional cap for quick training runs.")40 parser.add_argument("--seed", type=int, default=42)41 parser.add_argument("--save-strategy", choices=["no", "epoch"], default="epoch")42 parser.add_argument("--save-total-limit", type=int, default=1)43 args = parser.parse_args()44 45 train_path = Path(args.train_file)46 if not train_path.exists():47 raise FileNotFoundError(f"Training file not found: {train_path}")48 49 dataset = load_dataset("json", data_files=str(train_path), split="train")50 if args.limit:51 dataset = dataset.shuffle(seed=args.seed).select(range(min(args.limit, len(dataset))))52 tokenizer = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=True)53 if tokenizer.pad_token is None:54 tokenizer.pad_token = tokenizer.eos_token55 model = AutoModelForCausalLM.from_pretrained(args.base_model, trust_remote_code=True)56 model.config.use_cache = False57 58 peft_config = LoraConfig(59 r=args.rank,60 lora_alpha=args.rank * 2,61 lora_dropout=0.05,62 bias="none",63 task_type="CAUSAL_LM",64 target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],65 )66 model = get_peft_model(model, peft_config)67 68 def tokenize_example(example: dict[str, str]) -> dict[str, list[int]]:69 prompt = format_prompt(example)70 text = prompt + example["sql"] + tokenizer.eos_token71 tokenized = tokenizer(72 text,73 truncation=True,74 max_length=args.max_seq_length,75 padding=False,76 )77 prompt_token_count = len(78 tokenizer(79 prompt,80 truncation=True,81 max_length=args.max_seq_length,82 padding=False,83 )["input_ids"]84 )85 labels = tokenized["input_ids"].copy()86 labels[:prompt_token_count] = [-100] * min(prompt_token_count, len(labels))87 tokenized["labels"] = labels88 return tokenized89 90 tokenized_dataset = dataset.map(91 tokenize_example,92 remove_columns=dataset.column_names,93 )94 training_args = TrainingArguments(95 output_dir=args.output_dir,96 num_train_epochs=args.epochs,97 per_device_train_batch_size=1,98 gradient_accumulation_steps=4,99 learning_rate=args.learning_rate,100 logging_steps=1,101 save_strategy=args.save_strategy,102 save_total_limit=args.save_total_limit,103 report_to=[],104 seed=args.seed,105 dataloader_pin_memory=False,106 )107 108 trainer = Trainer(109 model=model,110 train_dataset=tokenized_dataset,111 args=training_args,112 data_collator=DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False),113 )114 trainer.train()115 trainer.model.save_pretrained(args.output_dir)116 tokenizer.save_pretrained(args.output_dir)117 118 119if __name__ == "__main__":120 main()121 