CoolFace
Apppublic

animesh08/schema-sage

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
train_lora.py121 linesDownload Raw Back to scripts
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