DevOps / train_lora.py
PrithviRana's picture
Upload folder using huggingface_hub
9a40336 verified
Raw History Blame Contribute Delete
11.6 kB
import os
import torch
from datasets import load_dataset
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
TrainingArguments,
Trainer,
DataCollatorForLanguageModeling,
)
from peft import (
LoraConfig,
get_peft_model,
)
# ============================================================
# CONFIGURATION
# ============================================================
MODEL_NAME = "Qwen/Qwen2.5-3B-Instruct"
TRAIN_FILE = "/root/ai-tuning/train.jsonl"
VAL_FILE = "/root/ai-tuning/validation.jsonl"
OUTPUT_DIR = "/root/ai-tuning/lora-output"
# 256 is faster for CPU training.
# Use 512 later if your dataset contains long logs/configs.
MAX_LENGTH = 256
# ============================================================
# CPU CONFIGURATION
# ============================================================
CPU_COUNT = os.cpu_count() or 4
# Use available CPU threads
torch.set_num_threads(CPU_COUNT)
# Keep inter-op threads smaller to avoid CPU overhead
torch.set_num_interop_threads(2)
print("\n==========================================")
print("SYSTEM INFORMATION")
print("==========================================")
print("CPU threads :", CPU_COUNT)
print("PyTorch threads :", torch.get_num_threads())
print("CUDA available :", torch.cuda.is_available())
print("PyTorch version :", torch.__version__)
# ============================================================
# LOAD TOKENIZER
# ============================================================
print("\n==========================================")
print("LOADING TOKENIZER")
print("==========================================")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_NAME,
use_fast=True,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
print("Tokenizer loaded")
# ============================================================
# LOAD DATASET
# ============================================================
print("\n==========================================")
print("LOADING DATASET")
print("==========================================")
dataset = load_dataset(
"json",
data_files={
"train": TRAIN_FILE,
"validation": VAL_FILE,
},
)
print(dataset)
print(
"Training examples :",
len(dataset["train"])
)
print(
"Validation examples :",
len(dataset["validation"])
)
# ============================================================
# FORMAT DATASET
# ============================================================
def format_prompt(example):
instruction = example.get(
"instruction",
""
)
user_input = example.get(
"input",
""
)
output = example.get(
"output",
""
)
if user_input:
text = (
"### Instruction:\n"
+ instruction
+ "\n\n"
"### Input:\n"
+ user_input
+ "\n\n"
"### Response:\n"
+ output
)
else:
text = (
"### Instruction:\n"
+ instruction
+ "\n\n"
"### Response:\n"
+ output
)
return text
# ============================================================
# TOKENIZATION
# ============================================================
def tokenize_function(example):
text = format_prompt(example)
tokens = tokenizer(
text,
truncation=True,
max_length=MAX_LENGTH,
padding=False,
)
return tokens
print("\n==========================================")
print("TOKENIZING DATASET")
print("==========================================")
tokenized_dataset = dataset.map(
tokenize_function,
remove_columns=dataset["train"].column_names,
desc="Tokenizing",
)
print(tokenized_dataset)
# ============================================================
# LOAD BASE MODEL
# ============================================================
print("\n==========================================")
print("LOADING QWEN2.5-3B-INSTRUCT")
print("==========================================")
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
torch_dtype=torch.float32,
)
model.config.pad_token_id = tokenizer.pad_token_id
print("Base model loaded")
# ============================================================
# LOADING MODEL INFORMATION
# ============================================================
print("\n==========================================")
print("MODEL INFORMATION")
print("==========================================")
total_parameters = sum(
p.numel()
for p in model.parameters()
)
print(
"Total parameters:",
f"{total_parameters:,}"
)
# ============================================================
# LORA CONFIGURATION
# ============================================================
print("\n==========================================")
print("APPLYING LoRA")
print("==========================================")
lora_config = LoraConfig(
# LoRA rank
r=8,
# LoRA scaling
lora_alpha=16,
# Dropout
lora_dropout=0.05,
# Transformer modules to train
target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
],
# Do not train bias
bias="none",
# Causal language model
task_type="CAUSAL_LM",
)
model = get_peft_model(
model,
lora_config,
)
print("\nTrainable parameters:")
model.print_trainable_parameters()
# ============================================================
# TRAINING ARGUMENTS
# ============================================================
print("\n==========================================")
print("TRAINING CONFIGURATION")
print("==========================================")
training_args = TrainingArguments(
# Output
output_dir=OUTPUT_DIR,
# --------------------------------------------------------
# CPU
# --------------------------------------------------------
use_cpu=True,
# --------------------------------------------------------
# Training
# --------------------------------------------------------
num_train_epochs=1,
per_device_train_batch_size=1,
per_device_eval_batch_size=1,
# CPU speed optimization
gradient_accumulation_steps=1,
# LoRA learning rate
learning_rate=2e-4,
# --------------------------------------------------------
# Logging
# --------------------------------------------------------
logging_strategy="steps",
logging_steps=1,
# --------------------------------------------------------
# Evaluation
# --------------------------------------------------------
# Disabled for first fast training run
eval_strategy="no",
# --------------------------------------------------------
# Checkpoints
# --------------------------------------------------------
# Disabled to reduce disk I/O
save_strategy="no",
# --------------------------------------------------------
# Optimizer
# --------------------------------------------------------
optim="adamw_torch",
# --------------------------------------------------------
# Precision
# --------------------------------------------------------
fp16=False,
bf16=False,
# --------------------------------------------------------
# DataLoader
# --------------------------------------------------------
dataloader_num_workers=0,
# --------------------------------------------------------
# Reporting
# --------------------------------------------------------
report_to="none",
# --------------------------------------------------------
# Other
# --------------------------------------------------------
remove_unused_columns=False,
)
print("Epochs :", 1)
print("Train batch size :", 1)
print("Gradient accumulation :", 1)
print("Learning rate :", 2e-4)
print("Max sequence length :", MAX_LENGTH)
print("Evaluation : Disabled")
print("Checkpoint saving : Disabled")
# ============================================================
# DATA COLLATOR
# ============================================================
print("\n==========================================")
print("CREATING DATA COLLATOR")
print("==========================================")
data_collator = DataCollatorForLanguageModeling(
tokenizer=tokenizer,
mlm=False,
)
# ============================================================
# TRAINER
# ============================================================
print("\n==========================================")
print("CREATING TRAINER")
print("==========================================")
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized_dataset["train"],
# No evaluation during this fast first run
eval_dataset=None,
tokenizer=tokenizer,
data_collator=data_collator,
)
# ============================================================
# START TRAINING
# ============================================================
print("\n")
print("==========================================")
print("STARTING LoRA TRAINING")
print("==========================================")
print("")
print("Model :", MODEL_NAME)
print("Dataset :", TRAIN_FILE)
print("Output :", OUTPUT_DIR)
print("CPU threads :", CPU_COUNT)
print("Max length :", MAX_LENGTH)
print("")
print("Training started...")
print("")
train_result = trainer.train()
# ============================================================
# TRAINING RESULTS
# ============================================================
print("\n==========================================")
print("TRAINING RESULTS")
print("==========================================")
print(
"Training loss:",
train_result.training_loss
)
print(
"Training time:",
train_result.metrics.get(
"train_runtime",
"N/A"
),
"seconds"
)
# ============================================================
# SAVE LoRA ADAPTER
# ============================================================
print("\n==========================================")
print("SAVING LoRA ADAPTER")
print("==========================================")
os.makedirs(
OUTPUT_DIR,
exist_ok=True
)
trainer.save_model(
OUTPUT_DIR
)
tokenizer.save_pretrained(
OUTPUT_DIR
)
# ============================================================
# VERIFY OUTPUT
# ============================================================
print("\n==========================================")
print("VERIFYING OUTPUT")
print("==========================================")
for root, dirs, files in os.walk(OUTPUT_DIR):
for file in files:
path = os.path.join(
root,
file
)
size = os.path.getsize(path)
print(
f"{path} "
f"({size / 1024 / 1024:.2f} MB)"
)
# ============================================================
# COMPLETE
# ============================================================
print("\n")
print("==========================================")
print("TRAINING COMPLETE")
print("==========================================")
print(
"LoRA adapter:",
OUTPUT_DIR
)
print("")
print("Next step:")
print("1. Test LoRA adapter")
print("2. Compare Base vs LoRA")
print("3. Merge adapter")
print("4. Convert to GGUF")
print("5. Deploy in Ollama")
print("")
print("==========================================")