diff options
| -rw-r--r-- | finetune.py | 99 | ||||
| -rw-r--r-- | length.py | 23 |
2 files changed, 122 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 @@ | |||
| 1 | import torch | ||
| 2 | from datasets import load_dataset | ||
| 3 | from peft import LoraConfig | ||
| 4 | from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer | ||
| 5 | from trl import SFTTrainer, SFTConfig | ||
| 6 | import os | ||
| 7 | |||
| 8 | #MODEL_ID = "google/gemma-3-270m-it" | ||
| 9 | MODEL_ID = "google/gemma-3-1b-it" | ||
| 10 | DATASET_TRAIN_PATH = "train.jsonl" | ||
| 11 | DATASET_VAL_PATH = "validation.jsonl" | ||
| 12 | OUTPUT_DIR = "./gemma3-1b-it-sponsors" | ||
| 13 | LORA_R = 16 | ||
| 14 | LORA_ALPHA = 32 | ||
| 15 | LORA_DROPOUT = 0.05 | ||
| 16 | LORA_TARGET_MODULES = [ | ||
| 17 | "q_proj", "k_proj", "v_proj", "o_proj", | ||
| 18 | "gate_proj", "up_proj", "down_proj", | ||
| 19 | ] | ||
| 20 | |||
| 21 | def 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 | |||
| 98 | if __name__ == "__main__": | ||
| 99 | main() | ||
diff --git a/length.py b/length.py new file mode 100644 index 0000000..06f8055 --- /dev/null +++ b/length.py | |||
| @@ -0,0 +1,23 @@ | |||
| 1 | import json | ||
| 2 | from transformers import AutoTokenizer | ||
| 3 | |||
| 4 | MODEL_ID = "google/gemma-3-270m-it" | ||
| 5 | DATASET_PATH = "train.jsonl" | ||
| 6 | |||
| 7 | tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) | ||
| 8 | token_lengths = [] | ||
| 9 | |||
| 10 | print(f"Loading dataset from {DATASET_PATH}...") | ||
| 11 | with open(DATASET_PATH, 'r', encoding='utf-8') as f: | ||
| 12 | for line in f: | ||
| 13 | entry = json.loads(line) | ||
| 14 | text = entry.get("text", "") | ||
| 15 | # Tokenize the text and get the number of tokens | ||
| 16 | length = len(tokenizer(text).input_ids) | ||
| 17 | token_lengths.append(length) | ||
| 18 | |||
| 19 | print(f"Number of examples: {len(token_lengths)}") | ||
| 20 | if token_lengths: | ||
| 21 | print(f"Min length: {min(token_lengths)}") | ||
| 22 | print(f"Max length: {max(token_lengths)}") | ||
| 23 | print(f"Average length: {sum(token_lengths) / len(token_lengths):.2f}") | ||
