samplellm / train.py
Aadhitech1's picture
Upload 8 files
6082a5d verified
Raw History Blame Contribute Delete
13.1 kB
# ============================================================
# TELECOM QLoRA FINE-TUNING
# FOUR-OUTPUT TELECOM DIAGNOSIS
#
# Model:
# Qwen/Qwen2.5-1.5B-Instruct
#
# Output:
# Probable Root Cause
# Explanation
# Recommendation
# Fault Severity
# ============================================================
import os
import re
import torch
from datasets import load_dataset
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
BitsAndBytesConfig,
)
from peft import (
LoraConfig,
prepare_model_for_kbit_training,
)
from trl import (
SFTTrainer,
SFTConfig,
)
# ============================================================
# CONFIGURATION
# ============================================================
MODEL_NAME = "Qwen/Qwen2.5-1.5B-Instruct"
TRAIN_FILE = "sample/data/train_four_output.jsonl"
VALIDATION_FILE = "sample/data/validation_four_output.jsonl"
OUTPUT_DIR = "models/telecom-qwen-four-output"
CACHE_DIR = "sample/.cache/huggingface"
EPOCHS = 3
MAX_LENGTH = 1024
TRAIN_BATCH_SIZE = 1
EVAL_BATCH_SIZE = 1
GRADIENT_ACCUMULATION = 8
LEARNING_RATE = 2e-4
SEED = 42
# ============================================================
# REQUIRED OUTPUT FORMAT
# ============================================================
REQUIRED_SECTIONS = [
"Probable Root Cause:",
"Explanation:",
"Recommendation:",
"Fault Severity:",
]
# ============================================================
# HELPER FUNCTIONS
# ============================================================
def get_assistant_message(example):
for message in example["messages"]:
if message["role"] == "assistant":
return message["content"]
raise ValueError(
"Assistant message not found."
)
def extract_fault_severity(text):
matches = re.findall(
r"fault[_\s-]*severity\s*[:=]\s*([012])",
text,
re.IGNORECASE,
)
if not matches:
return None
return int(matches[-1])
def verify_example(example):
# --------------------------------------------------------
# Check messages
# --------------------------------------------------------
if "messages" not in example:
raise ValueError(
"Dataset example does not contain 'messages'."
)
roles = [
message["role"]
for message in example["messages"]
]
required_roles = {
"system",
"user",
"assistant",
}
if not required_roles.issubset(set(roles)):
raise ValueError(
"Example must contain system, user and assistant messages."
)
# --------------------------------------------------------
# Check assistant response
# --------------------------------------------------------
assistant = get_assistant_message(example)
for section in REQUIRED_SECTIONS:
if section not in assistant:
raise ValueError(
f"Missing required section: {section}"
)
# --------------------------------------------------------
# Check severity
# --------------------------------------------------------
severity = extract_fault_severity(
assistant
)
if severity not in [0, 1, 2]:
raise ValueError(
f"Invalid fault severity: {severity}"
)
return example
def print_distribution(
dataset,
name,
):
counts = {
0: 0,
1: 0,
2: 0,
}
for example in dataset:
assistant = get_assistant_message(
example
)
severity = extract_fault_severity(
assistant
)
if severity in counts:
counts[severity] += 1
total = len(dataset)
print()
print(
f"{name} fault_severity distribution:"
)
for severity in [0, 1, 2]:
percentage = (
counts[severity] / total * 100
)
print(
f" {severity}: "
f"{counts[severity]} "
f"({percentage:.2f}%)"
)
# ============================================================
# START
# ============================================================
print("=" * 80)
print(
"TELECOM QLoRA FINE-TUNING"
)
print(
"FOUR-OUTPUT TELECOM DIAGNOSIS"
)
print("=" * 80)
# ============================================================
# GPU CHECK
# ============================================================
print()
print(
"CUDA available:",
torch.cuda.is_available()
)
if not torch.cuda.is_available():
raise RuntimeError(
"CUDA GPU is required for this training configuration."
)
gpu_name = torch.cuda.get_device_name(0)
gpu_memory = (
torch.cuda.get_device_properties(0)
.total_memory
/ (1024 ** 3)
)
print(
"GPU:",
gpu_name
)
print(
"VRAM:",
round(gpu_memory, 2),
"GB"
)
# ============================================================
# DIRECTORIES
# ============================================================
os.makedirs(
CACHE_DIR,
exist_ok=True
)
os.makedirs(
OUTPUT_DIR,
exist_ok=True
)
# ============================================================
# CHECK DATA FILES
# ============================================================
if not os.path.exists(TRAIN_FILE):
raise FileNotFoundError(
f"Training dataset not found:\n{TRAIN_FILE}"
)
if not os.path.exists(
VALIDATION_FILE
):
raise FileNotFoundError(
f"Validation dataset not found:\n"
f"{VALIDATION_FILE}"
)
# ============================================================
# LOAD DATA
# ============================================================
print()
print("=" * 80)
print(
"LOADING DATASETS"
)
print("=" * 80)
train_dataset = load_dataset(
"json",
data_files=TRAIN_FILE,
split="train",
cache_dir=CACHE_DIR,
)
validation_dataset = load_dataset(
"json",
data_files=VALIDATION_FILE,
split="train",
cache_dir=CACHE_DIR,
)
print(
"Training examples:",
len(train_dataset)
)
print(
"Validation examples:",
len(validation_dataset)
)
# ============================================================
# VERIFY DATA
# ============================================================
print()
print(
"Verifying four-output format..."
)
train_dataset = train_dataset.map(
verify_example
)
validation_dataset = validation_dataset.map(
verify_example
)
print(
"Dataset verification completed."
)
# ============================================================
# PRINT LABEL DISTRIBUTION
# ============================================================
print_distribution(
train_dataset,
"Training"
)
print_distribution(
validation_dataset,
"Validation"
)
# ============================================================
# TOKENIZER
# ============================================================
print()
print("=" * 80)
print(
"LOADING TOKENIZER"
)
print("=" * 80)
tokenizer = AutoTokenizer.from_pretrained(
MODEL_NAME,
trust_remote_code=True,
cache_dir=CACHE_DIR,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
# ============================================================
# 4-BIT QUANTIZATION
# ============================================================
print()
print("=" * 80)
print(
"CONFIGURING 4-BIT QUANTIZATION"
)
print("=" * 80)
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
)
# ============================================================
# LOAD BASE MODEL
# ============================================================
print()
print("=" * 80)
print(
"LOADING BASE MODEL"
)
print("=" * 80)
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
quantization_config=bnb_config,
device_map="auto",
dtype=torch.float16,
trust_remote_code=True,
cache_dir=CACHE_DIR,
)
print(
"Base model loaded."
)
# ============================================================
# PREPARE QLORA
# ============================================================
print()
print(
"Preparing model for QLoRA..."
)
model = prepare_model_for_kbit_training(
model
)
model.config.use_cache = False
# ============================================================
# LORA CONFIGURATION
# ============================================================
print()
print("=" * 80)
print(
"CONFIGURING LoRA"
)
print("=" * 80)
lora_config = LoraConfig(
r=8,
lora_alpha=16,
lora_dropout=0.05,
target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
],
bias="none",
task_type="CAUSAL_LM",
)
# ============================================================
# CHAT FORMAT
# ============================================================
def formatting_func(example):
return tokenizer.apply_chat_template(
example["messages"],
tokenize=False,
add_generation_prompt=False,
)
# ============================================================
# TRAINING CONFIGURATION
# ============================================================
print()
print("=" * 80)
print(
"CONFIGURING TRAINING"
)
print("=" * 80)
# IMPORTANT:
# max_length belongs to SFTConfig, not SFTTrainer.
training_args = SFTConfig(
output_dir=OUTPUT_DIR,
seed=SEED,
data_seed=SEED,
per_device_train_batch_size=TRAIN_BATCH_SIZE,
per_device_eval_batch_size=EVAL_BATCH_SIZE,
gradient_accumulation_steps=GRADIENT_ACCUMULATION,
num_train_epochs=EPOCHS,
learning_rate=LEARNING_RATE,
fp16=False,
bf16=False,
gradient_checkpointing=True,
gradient_checkpointing_kwargs={
"use_reentrant": False
},
optim="paged_adamw_8bit",
logging_steps=10,
eval_strategy="epoch",
save_strategy="epoch",
save_total_limit=2,
load_best_model_at_end=True,
metric_for_best_model="eval_loss",
greater_is_better=False,
report_to="none",
remove_unused_columns=False,
max_length=MAX_LENGTH,
packing=False,
)
# ============================================================
# CREATE SFT TRAINER
# ============================================================
print()
print("=" * 80)
print(
"CREATING SFT TRAINER"
)
print("=" * 80)
trainer = SFTTrainer(
model=model,
train_dataset=train_dataset,
eval_dataset=validation_dataset,
peft_config=lora_config,
formatting_func=formatting_func,
processing_class=tokenizer,
args=training_args,
)
# ============================================================
# TRAIN
# ============================================================
print()
print("=" * 80)
print(
"STARTING TRAINING"
)
print("=" * 80)
print(
"Training examples:",
len(train_dataset)
)
print(
"Validation examples:",
len(validation_dataset)
)
print(
"Epochs:",
EPOCHS
)
print(
"Batch size:",
TRAIN_BATCH_SIZE
)
print(
"Gradient accumulation:",
GRADIENT_ACCUMULATION
)
print(
"Effective batch size:",
TRAIN_BATCH_SIZE
* GRADIENT_ACCUMULATION
)
print(
"Learning rate:",
LEARNING_RATE
)
print(
"Maximum sequence length:",
MAX_LENGTH
)
print("=" * 80)
try:
result = trainer.train()
except KeyboardInterrupt:
print()
print(
"Training interrupted."
)
print(
"Saving current adapter..."
)
trainer.save_model(
OUTPUT_DIR
)
tokenizer.save_pretrained(
OUTPUT_DIR
)
trainer.save_state()
raise
# ============================================================
# SAVE MODEL
# ============================================================
print()
print("=" * 80)
print(
"SAVING MODEL"
)
print("=" * 80)
trainer.save_model(
OUTPUT_DIR
)
tokenizer.save_pretrained(
OUTPUT_DIR
)
trainer.save_state()
# ============================================================
# RESULTS
# ============================================================
print()
print("=" * 80)
print(
"TRAINING COMPLETE"
)
print("=" * 80)
print(
"Adapter saved to:"
)
print(
os.path.abspath(
OUTPUT_DIR
)
)
if result is not None:
print()
print(
"Training metrics:"
)
for key, value in result.metrics.items():
print(
f"{key}: {value}"
)
# ============================================================
# EXPECTED OUTPUT
# ============================================================
print()
print("=" * 80)
print(
"EXPECTED MODEL RESPONSE FORMAT"
)
print("=" * 80)
print(
"""Probable Root Cause:
<probable root cause based only on supplied evidence>
Explanation:
<technical explanation>
Recommendation:
1. <recommended action>
2. <recommended action>
3. <recommended action>
Fault Severity:
fault_severity=0|1|2
"""
)
print("=" * 80)