# ============================================================ # 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: Explanation: Recommendation: 1. 2. 3. Fault Severity: fault_severity=0|1|2 """ ) print("=" * 80)