Download train.py from Aadhitech1/samplellm: direct link, hf CLI and curl.
- Browser
- Download file 13.1 kB
-
https://huggingface.co/Aadhitech1/samplellm/resolve/main/train.py
- Command line
-
hf download hf://Aadhitech1/samplellm/train.py
-
curl -L -o train.py https://huggingface.co/Aadhitech1/samplellm/resolve/main/train.py
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) |