commit 48cdc65c0eafa08fc39b5bfa49ef7c42943ac91d
parent 4ec10c33159f435963b73da5a9bd12f25f65cda9
Author: vin <git@vineetk.net>
Date: Sat, 23 Aug 2025 17:23:59 -0400
add finetuning scripts
Diffstat:
| A | finetune.py | | | 99 | +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ |
| A | length.py | | | 23 | +++++++++++++++++++++++ |
2 files changed, 122 insertions(+), 0 deletions(-)
diff --git a/finetune.py b/finetune.py
@@ -0,0 +1,99 @@
+import torch
+from datasets import load_dataset
+from peft import LoraConfig
+from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer
+from trl import SFTTrainer, SFTConfig
+import os
+
+#MODEL_ID = "google/gemma-3-270m-it"
+MODEL_ID = "google/gemma-3-1b-it"
+DATASET_TRAIN_PATH = "train.jsonl"
+DATASET_VAL_PATH = "validation.jsonl"
+OUTPUT_DIR = "./gemma3-1b-it-sponsors"
+LORA_R = 16
+LORA_ALPHA = 32
+LORA_DROPOUT = 0.05
+LORA_TARGET_MODULES = [
+ "q_proj", "k_proj", "v_proj", "o_proj",
+ "gate_proj", "up_proj", "down_proj",
+]
+
+def main():
+ tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
+ # A pad token is required for training
+ if tokenizer.pad_token is None:
+ tokenizer.pad_token = tokenizer.eos_token
+
+ model = AutoModelForCausalLM.from_pretrained(
+ MODEL_ID,
+ device_map="auto",
+ torch_dtype=torch.bfloat16,
+ attn_implementation='eager',
+ offload_folder="./offload",
+ offload_state_dict=True,
+ )
+ model.config.use_cache = False # disable kv-cache during training
+ model.gradient_checkpointing_enable() # enable gradient checkpointing
+
+ train_dataset = load_dataset("json", data_files=DATASET_TRAIN_PATH, split="train")
+ validation_dataset = load_dataset("json", data_files=DATASET_VAL_PATH, split="train")
+
+ peft_config = LoraConfig(
+ r=LORA_R,
+ lora_alpha=LORA_ALPHA,
+ target_modules=LORA_TARGET_MODULES,
+ lora_dropout=LORA_DROPOUT,
+ bias="none",
+ task_type="CAUSAL_LM"
+ )
+
+ print("Model and LoRA config loaded. Setting up trainer...")
+
+ training_args = SFTConfig(
+ output_dir=OUTPUT_DIR,
+
+ # Training parameters
+ num_train_epochs=3,
+ per_device_eval_batch_size=1,
+ per_device_train_batch_size=1,
+ gradient_accumulation_steps=8,
+ gradient_checkpointing=True,
+ learning_rate=1e-4,
+ lr_scheduler_type="linear",
+ optim="adamw_torch",
+
+ # Model and data parameters
+ max_length=5120,
+ dataset_text_field="text",
+
+ # Technical parameters
+ bf16=True,
+ max_grad_norm=1.0,
+ warmup_ratio=0.03,
+
+ # Logging, saving, and evaluation
+ logging_steps=25,
+ eval_strategy="epoch",
+ save_strategy="epoch",
+ load_best_model_at_end=True,
+ )
+ training_args.eval_accumulation_steps = 4
+
+ trainer = SFTTrainer(
+ model=model,
+ processing_class=tokenizer,
+ args=training_args,
+ peft_config=peft_config,
+ train_dataset=train_dataset,
+ eval_dataset=validation_dataset,
+ )
+
+ print("Trainer initialized. Starting fine-tuning...")
+ trainer.train()
+
+ print("Training complete. Saving the best model...")
+ trainer.save_model(os.path.join(OUTPUT_DIR, "final_checkpoint"))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/length.py b/length.py
@@ -0,0 +1,23 @@
+import json
+from transformers import AutoTokenizer
+
+MODEL_ID = "google/gemma-3-270m-it"
+DATASET_PATH = "train.jsonl"
+
+tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
+token_lengths = []
+
+print(f"Loading dataset from {DATASET_PATH}...")
+with open(DATASET_PATH, 'r', encoding='utf-8') as f:
+ for line in f:
+ entry = json.loads(line)
+ text = entry.get("text", "")
+ # Tokenize the text and get the number of tokens
+ length = len(tokenizer(text).input_ids)
+ token_lengths.append(length)
+
+print(f"Number of examples: {len(token_lengths)}")
+if token_lengths:
+ print(f"Min length: {min(token_lengths)}")
+ print(f"Max length: {max(token_lengths)}")
+ print(f"Average length: {sum(token_lengths) / len(token_lengths):.2f}")