Spaces:
Sleeping
Sleeping
File size: 3,792 Bytes
0279c66 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | import os
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, BitsAndBytesConfig
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
from datasets import load_dataset
from trl import SFTTrainer
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
def train_lora(
base_model_name: str = "TinyLlama/TinyLlama-1.1B-Chat-v1.0",
dataset_path: str = "./data/reflective_feedback.jsonl",
output_dir: str = "./param_mem_lora",
epochs: int = 3,
batch_size: int = 4
):
"""
Fine-tunes the base language model using Low-Rank Adaptation (LoRA)
on a curated dataset of reflective feedback.
"""
logger.info("Initializing LoRA Training Pipeline for Parametric Memory...")
# Load Tokenizer
tokenizer = AutoTokenizer.from_pretrained(base_model_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# Setup Quantization for local laptop environment
logger.info("Loading base model in 4-bit...")
# In a real deployment, paths should be injected via environment variables
# We default to the local mock base model setup
base_model = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4"
)
model = AutoModelForCausalLM.from_pretrained(
base_model_name,
quantization_config=quantization_config,
device_map="auto"
)
# Prepare model for PEFT
model = prepare_model_for_kbit_training(model)
# Define LoRA Config
peft_config = LoraConfig(
r=16,
lora_alpha=32,
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
target_modules=["q_proj", "v_proj"]
)
# model = get_peft_model(model, peft_config) # Removed to let SFTTrainer handle it
# Load Dataset
# Assuming the dataset has a "text" field with the formatted reflective prompt
if not os.path.exists(dataset_path):
logger.error(f"Dataset not found at {dataset_path}. Please generate dummy data first.")
return
dataset = load_dataset("json", data_files=dataset_path, split="train")
# Define Training Arguments
training_args = TrainingArguments(
output_dir=output_dir,
num_train_epochs=epochs,
per_device_train_batch_size=batch_size,
gradient_accumulation_steps=4,
optim="paged_adamw_32bit",
save_steps=100,
logging_steps=10,
learning_rate=2e-4,
weight_decay=0.001,
fp16=True,
bf16=False,
max_grad_norm=0.3,
warmup_ratio=0.03,
lr_scheduler_type="constant"
)
# Initialize SFTTrainer
trainer = SFTTrainer(
model=model,
train_dataset=dataset,
peft_config=peft_config,
formatting_func=lambda x: x["text"],
args=training_args,
)
# Train
logger.info("Starting Training...")
trainer.train()
# Save
logger.info(f"Saving LoRA adapter to {output_dir}")
trainer.model.save_pretrained(output_dir)
tokenizer.save_pretrained(output_dir)
logger.info("Training complete.")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, default="meta-llama/Meta-Llama-3-8B-Instruct")
parser.add_argument("--dataset", type=str, default="./data/reflective_feedback.jsonl")
parser.add_argument("--output", type=str, default="./param_mem_lora")
args = parser.parse_args()
train_lora(args.model, args.dataset, args.output)
|