summaryrefslogtreecommitdiff
path: root/finetune.py
diff options
context:
space:
mode:
authorvin <git@vineetk.net>2025-08-23 17:23:59 -0400
committervin <git@vineetk.net>2025-08-23 17:23:59 -0400
commit48cdc65c0eafa08fc39b5bfa49ef7c42943ac91d (patch)
tree2cc12f42fbcb4de515513b30c51fc9139c034324 /finetune.py
parent4ec10c33159f435963b73da5a9bd12f25f65cda9 (diff)
add finetuning scripts
Diffstat (limited to 'finetune.py')
-rw-r--r--finetune.py99
1 files changed, 99 insertions, 0 deletions
diff --git a/finetune.py b/finetune.py
new file mode 100644
index 0000000..bca7ac1
--- /dev/null
+++ b/finetune.py
@@ -0,0 +1,99 @@
1import torch
2from datasets import load_dataset
3from peft import LoraConfig
4from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer
5from trl import SFTTrainer, SFTConfig
6import os
7
8#MODEL_ID = "google/gemma-3-270m-it"
9MODEL_ID = "google/gemma-3-1b-it"
10DATASET_TRAIN_PATH = "train.jsonl"
11DATASET_VAL_PATH = "validation.jsonl"
12OUTPUT_DIR = "./gemma3-1b-it-sponsors"
13LORA_R = 16
14LORA_ALPHA = 32
15LORA_DROPOUT = 0.05
16LORA_TARGET_MODULES = [
17 "q_proj", "k_proj", "v_proj", "o_proj",
18 "gate_proj", "up_proj", "down_proj",
19]
20
21def main():
22 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
23 # A pad token is required for training
24 if tokenizer.pad_token is None:
25 tokenizer.pad_token = tokenizer.eos_token
26
27 model = AutoModelForCausalLM.from_pretrained(
28 MODEL_ID,
29 device_map="auto",
30 torch_dtype=torch.bfloat16,
31 attn_implementation='eager',
32 offload_folder="./offload",
33 offload_state_dict=True,
34 )
35 model.config.use_cache = False # disable kv-cache during training
36 model.gradient_checkpointing_enable() # enable gradient checkpointing
37
38 train_dataset = load_dataset("json", data_files=DATASET_TRAIN_PATH, split="train")
39 validation_dataset = load_dataset("json", data_files=DATASET_VAL_PATH, split="train")
40
41 peft_config = LoraConfig(
42 r=LORA_R,
43 lora_alpha=LORA_ALPHA,
44 target_modules=LORA_TARGET_MODULES,
45 lora_dropout=LORA_DROPOUT,
46 bias="none",
47 task_type="CAUSAL_LM"
48 )
49
50 print("Model and LoRA config loaded. Setting up trainer...")
51
52 training_args = SFTConfig(
53 output_dir=OUTPUT_DIR,
54
55 # Training parameters
56 num_train_epochs=3,
57 per_device_eval_batch_size=1,
58 per_device_train_batch_size=1,
59 gradient_accumulation_steps=8,
60 gradient_checkpointing=True,
61 learning_rate=1e-4,
62 lr_scheduler_type="linear",
63 optim="adamw_torch",
64
65 # Model and data parameters
66 max_length=5120,
67 dataset_text_field="text",
68
69 # Technical parameters
70 bf16=True,
71 max_grad_norm=1.0,
72 warmup_ratio=0.03,
73
74 # Logging, saving, and evaluation
75 logging_steps=25,
76 eval_strategy="epoch",
77 save_strategy="epoch",
78 load_best_model_at_end=True,
79 )
80 training_args.eval_accumulation_steps = 4
81
82 trainer = SFTTrainer(
83 model=model,
84 processing_class=tokenizer,
85 args=training_args,
86 peft_config=peft_config,
87 train_dataset=train_dataset,
88 eval_dataset=validation_dataset,
89 )
90
91 print("Trainer initialized. Starting fine-tuning...")
92 trainer.train()
93
94 print("Training complete. Saving the best model...")
95 trainer.save_model(os.path.join(OUTPUT_DIR, "final_checkpoint"))
96
97
98if __name__ == "__main__":
99 main()