Download test.py from Aadhitech1/samplellm: direct link, hf CLI and curl.
- Browser
- Download file 20.2 kB
-
https://huggingface.co/Aadhitech1/samplellm/resolve/main/test.py
- Command line
-
hf download hf://Aadhitech1/samplellm/test.py
-
curl -L -o test.py https://huggingface.co/Aadhitech1/samplellm/resolve/main/test.py
20.2 kB
| # ============================================================ | |
| # TELECOM QLoRA TEST / EVALUATION | |
| # FOUR-OUTPUT TELECOM DIAGNOSIS | |
| # ============================================================ | |
| import os | |
| import re | |
| import json | |
| import argparse | |
| from collections import Counter | |
| import torch | |
| from datasets import load_dataset | |
| from transformers import ( | |
| AutoTokenizer, | |
| AutoModelForCausalLM, | |
| BitsAndBytesConfig, | |
| ) | |
| from peft import PeftModel | |
| from sklearn.metrics import ( | |
| accuracy_score, | |
| precision_recall_fscore_support, | |
| confusion_matrix, | |
| ) | |
| # ============================================================ | |
| # CONFIGURATION | |
| # ============================================================ | |
| MODEL_NAME = "C:\\Users\\aadhi\\Downloads\\telecom_four_output_complete_workspace\\models\\telecom-qwen-four-output" | |
| CACHE_DIR = "sample/.cache/huggingface" | |
| DEFAULT_TEST_FILE = ( | |
| "sample/data/test_four_output.jsonl" | |
| ) | |
| # ============================================================ | |
| # REQUIRED OUTPUT SECTIONS | |
| # ============================================================ | |
| REQUIRED_SECTIONS = [ | |
| "Probable Root Cause:", | |
| "Explanation:", | |
| "Recommendation:", | |
| "Fault Severity:", | |
| ] | |
| # ============================================================ | |
| # COMMAND-LINE ARGUMENTS | |
| # ============================================================ | |
| parser = argparse.ArgumentParser( | |
| description=( | |
| "Evaluate telecom QLoRA four-output model." | |
| ) | |
| ) | |
| parser.add_argument( | |
| "--adapter", | |
| required=True, | |
| help=( | |
| "Path to trained LoRA adapter." | |
| ), | |
| ) | |
| parser.add_argument( | |
| "--test", | |
| default=DEFAULT_TEST_FILE, | |
| help=( | |
| "Path to test JSONL dataset." | |
| ), | |
| ) | |
| parser.add_argument( | |
| "--samples", | |
| type=int, | |
| default=0, | |
| help=( | |
| "Number of test samples. " | |
| "0 = all samples." | |
| ), | |
| ) | |
| parser.add_argument( | |
| "--debug", | |
| action="store_true", | |
| help=( | |
| "Print generated responses." | |
| ), | |
| ) | |
| args = parser.parse_args() | |
| # ============================================================ | |
| # HELPERS | |
| # ============================================================ | |
| 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 get_assistant_message(example): | |
| for message in example["messages"]: | |
| if message["role"] == "assistant": | |
| return message["content"] | |
| return "" | |
| def extract_section( | |
| text, | |
| section_name, | |
| next_sections, | |
| ): | |
| if section_name not in text: | |
| return "" | |
| content = text.split( | |
| section_name, | |
| 1 | |
| )[1] | |
| end_positions = [] | |
| for section in next_sections: | |
| position = content.find( | |
| section | |
| ) | |
| if position >= 0: | |
| end_positions.append( | |
| position | |
| ) | |
| if end_positions: | |
| content = content[ | |
| :min(end_positions) | |
| ] | |
| return content.strip() | |
| def tokenize_for_overlap(text): | |
| return set( | |
| re.findall( | |
| r"[a-z0-9_]+", | |
| text.lower() | |
| ) | |
| ) | |
| def token_f1( | |
| reference, | |
| prediction, | |
| ): | |
| reference_tokens = ( | |
| tokenize_for_overlap( | |
| reference | |
| ) | |
| ) | |
| prediction_tokens = ( | |
| tokenize_for_overlap( | |
| prediction | |
| ) | |
| ) | |
| if not reference_tokens: | |
| return 0.0 | |
| if not prediction_tokens: | |
| return 0.0 | |
| intersection = ( | |
| reference_tokens | |
| & | |
| prediction_tokens | |
| ) | |
| precision = ( | |
| len(intersection) | |
| / | |
| len(prediction_tokens) | |
| ) | |
| recall = ( | |
| len(intersection) | |
| / | |
| len(reference_tokens) | |
| ) | |
| if precision + recall == 0: | |
| return 0.0 | |
| return ( | |
| 2 | |
| * precision | |
| * recall | |
| / | |
| (precision + recall) | |
| ) | |
| # ============================================================ | |
| # GPU CHECK | |
| # ============================================================ | |
| print("=" * 80) | |
| print( | |
| "TELECOM FOUR-OUTPUT MODEL EVALUATION" | |
| ) | |
| print("=" * 80) | |
| print() | |
| print( | |
| "CUDA available:", | |
| torch.cuda.is_available() | |
| ) | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError( | |
| "CUDA GPU is required." | |
| ) | |
| print( | |
| "GPU:", | |
| torch.cuda.get_device_name(0) | |
| ) | |
| print( | |
| "VRAM:", | |
| round( | |
| torch.cuda.get_device_properties( | |
| 0 | |
| ).total_memory | |
| / (1024 ** 3), | |
| 2, | |
| ), | |
| "GB", | |
| ) | |
| # ============================================================ | |
| # CHECK FILES | |
| # ============================================================ | |
| if not os.path.exists( | |
| args.adapter | |
| ): | |
| raise FileNotFoundError( | |
| "Adapter directory not found:\n" | |
| f"{args.adapter}" | |
| ) | |
| if not os.path.exists( | |
| args.test | |
| ): | |
| raise FileNotFoundError( | |
| "Test dataset not found:\n" | |
| f"{args.test}" | |
| ) | |
| # ============================================================ | |
| # LOAD TEST DATA | |
| # ============================================================ | |
| print() | |
| print( | |
| "Loading test dataset..." | |
| ) | |
| test_dataset = load_dataset( | |
| "json", | |
| data_files=args.test, | |
| split="train", | |
| cache_dir=CACHE_DIR, | |
| ) | |
| print( | |
| "Total test examples:", | |
| len(test_dataset) | |
| ) | |
| if args.samples > 0: | |
| number = min( | |
| args.samples, | |
| len(test_dataset) | |
| ) | |
| test_dataset = ( | |
| test_dataset.select( | |
| range(number) | |
| ) | |
| ) | |
| print( | |
| "Examples to evaluate:", | |
| len(test_dataset) | |
| ) | |
| # ============================================================ | |
| # LOAD TOKENIZER | |
| # ============================================================ | |
| print() | |
| print( | |
| "Loading tokenizer..." | |
| ) | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| args.adapter, | |
| trust_remote_code=True, | |
| cache_dir=CACHE_DIR, | |
| ) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = ( | |
| tokenizer.eos_token | |
| ) | |
| # ============================================================ | |
| # QUANTIZATION | |
| # ============================================================ | |
| 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( | |
| "Loading base Qwen model..." | |
| ) | |
| base_model = ( | |
| AutoModelForCausalLM.from_pretrained( | |
| MODEL_NAME, | |
| quantization_config=bnb_config, | |
| device_map="auto", | |
| dtype=torch.float16, | |
| trust_remote_code=True, | |
| cache_dir=CACHE_DIR, | |
| ) | |
| ) | |
| # ============================================================ | |
| # LOAD LORA ADAPTER | |
| # ============================================================ | |
| print() | |
| print( | |
| "Loading LoRA adapter..." | |
| ) | |
| model = PeftModel.from_pretrained( | |
| base_model, | |
| args.adapter, | |
| ) | |
| model.eval() | |
| print( | |
| "Model ready." | |
| ) | |
| # ============================================================ | |
| # EVALUATION STORAGE | |
| # ============================================================ | |
| predictions = [] | |
| y_true = [] | |
| y_pred = [] | |
| root_scores = [] | |
| explanation_scores = [] | |
| recommendation_scores = [] | |
| # ============================================================ | |
| # GENERATE RESPONSES | |
| # ============================================================ | |
| print() | |
| print("=" * 80) | |
| print( | |
| "GENERATING TEST RESPONSES" | |
| ) | |
| print("=" * 80) | |
| for index, example in enumerate( | |
| test_dataset | |
| ): | |
| # -------------------------------------------------------- | |
| # Get messages | |
| # -------------------------------------------------------- | |
| messages = example[ | |
| "messages" | |
| ] | |
| # -------------------------------------------------------- | |
| # Remove assistant answer | |
| # -------------------------------------------------------- | |
| prompt_messages = [ | |
| message | |
| for message in messages | |
| if message["role"] | |
| != "assistant" | |
| ] | |
| # -------------------------------------------------------- | |
| # Reference answer | |
| # -------------------------------------------------------- | |
| reference = ( | |
| get_assistant_message( | |
| example | |
| ) | |
| ) | |
| # -------------------------------------------------------- | |
| # Build prompt | |
| # -------------------------------------------------------- | |
| prompt = ( | |
| tokenizer.apply_chat_template( | |
| prompt_messages, | |
| tokenize=False, | |
| add_generation_prompt=True, | |
| ) | |
| ) | |
| # -------------------------------------------------------- | |
| # Tokenize | |
| # -------------------------------------------------------- | |
| inputs = tokenizer( | |
| prompt, | |
| return_tensors="pt", | |
| ) | |
| inputs = { | |
| key: value.to( | |
| model.device | |
| ) | |
| for key, value | |
| in inputs.items() | |
| } | |
| # -------------------------------------------------------- | |
| # Generate | |
| # -------------------------------------------------------- | |
| with torch.no_grad(): | |
| output = model.generate( | |
| **inputs, | |
| max_new_tokens=350, | |
| do_sample=False, | |
| pad_token_id=( | |
| tokenizer.pad_token_id | |
| ), | |
| eos_token_id=( | |
| tokenizer.eos_token_id | |
| ), | |
| ) | |
| # -------------------------------------------------------- | |
| # Decode ONLY generated tokens | |
| # -------------------------------------------------------- | |
| generated_tokens = output[ | |
| 0 | |
| ][ | |
| inputs["input_ids"].shape[1]: | |
| ] | |
| generated = tokenizer.decode( | |
| generated_tokens, | |
| skip_special_tokens=True, | |
| ).strip() | |
| # -------------------------------------------------------- | |
| # Severity | |
| # -------------------------------------------------------- | |
| true_label = ( | |
| extract_fault_severity( | |
| reference | |
| ) | |
| ) | |
| predicted_label = ( | |
| extract_fault_severity( | |
| generated | |
| ) | |
| ) | |
| y_true.append( | |
| true_label | |
| ) | |
| y_pred.append( | |
| predicted_label | |
| ) | |
| # -------------------------------------------------------- | |
| # Root Cause | |
| # -------------------------------------------------------- | |
| root_reference = extract_section( | |
| reference, | |
| "Probable Root Cause:", | |
| [ | |
| "Explanation:", | |
| "Recommendation:", | |
| "Fault Severity:", | |
| ], | |
| ) | |
| root_prediction = extract_section( | |
| generated, | |
| "Probable Root Cause:", | |
| [ | |
| "Explanation:", | |
| "Recommendation:", | |
| "Fault Severity:", | |
| ], | |
| ) | |
| root_f1 = token_f1( | |
| root_reference, | |
| root_prediction, | |
| ) | |
| root_scores.append( | |
| root_f1 | |
| ) | |
| # -------------------------------------------------------- | |
| # Explanation | |
| # -------------------------------------------------------- | |
| explanation_reference = extract_section( | |
| reference, | |
| "Explanation:", | |
| [ | |
| "Recommendation:", | |
| "Fault Severity:", | |
| ], | |
| ) | |
| explanation_prediction = extract_section( | |
| generated, | |
| "Explanation:", | |
| [ | |
| "Recommendation:", | |
| "Fault Severity:", | |
| ], | |
| ) | |
| explanation_f1 = token_f1( | |
| explanation_reference, | |
| explanation_prediction, | |
| ) | |
| explanation_scores.append( | |
| explanation_f1 | |
| ) | |
| # -------------------------------------------------------- | |
| # Recommendation | |
| # -------------------------------------------------------- | |
| recommendation_reference = extract_section( | |
| reference, | |
| "Recommendation:", | |
| [ | |
| "Fault Severity:", | |
| ], | |
| ) | |
| recommendation_prediction = extract_section( | |
| generated, | |
| "Recommendation:", | |
| [ | |
| "Fault Severity:", | |
| ], | |
| ) | |
| recommendation_f1 = token_f1( | |
| recommendation_reference, | |
| recommendation_prediction, | |
| ) | |
| recommendation_scores.append( | |
| recommendation_f1 | |
| ) | |
| # -------------------------------------------------------- | |
| # Format check | |
| # -------------------------------------------------------- | |
| all_sections_present = all( | |
| section in generated | |
| for section | |
| in REQUIRED_SECTIONS | |
| ) | |
| # -------------------------------------------------------- | |
| # Save result | |
| # -------------------------------------------------------- | |
| predictions.append({ | |
| "index": index, | |
| "true_fault_severity": ( | |
| true_label | |
| ), | |
| "predicted_fault_severity": ( | |
| predicted_label | |
| ), | |
| "severity_correct": ( | |
| true_label | |
| == | |
| predicted_label | |
| ), | |
| "all_sections_present": ( | |
| all_sections_present | |
| ), | |
| "root_cause_token_f1": ( | |
| root_f1 | |
| ), | |
| "explanation_token_f1": ( | |
| explanation_f1 | |
| ), | |
| "recommendation_token_f1": ( | |
| recommendation_f1 | |
| ), | |
| "model_response": ( | |
| generated | |
| ), | |
| "reference_response": ( | |
| reference | |
| ), | |
| }) | |
| # -------------------------------------------------------- | |
| # Debug output | |
| # -------------------------------------------------------- | |
| if args.debug: | |
| print() | |
| print("-" * 80) | |
| print( | |
| "Example:", | |
| index + 1 | |
| ) | |
| print( | |
| "Expected severity:", | |
| true_label | |
| ) | |
| print( | |
| "Predicted severity:", | |
| predicted_label | |
| ) | |
| print() | |
| print( | |
| "MODEL RESPONSE:" | |
| ) | |
| print() | |
| print( | |
| generated | |
| ) | |
| # -------------------------------------------------------- | |
| # Progress | |
| # -------------------------------------------------------- | |
| if not args.debug: | |
| if ( | |
| (index + 1) % 10 == 0 | |
| or | |
| (index + 1) | |
| == len(test_dataset) | |
| ): | |
| print( | |
| f"Processed " | |
| f"{index + 1}/" | |
| f"{len(test_dataset)}" | |
| ) | |
| # ============================================================ | |
| # SAVE PREDICTIONS | |
| # ============================================================ | |
| prediction_file = os.path.join( | |
| args.adapter, | |
| "test_four_output_predictions.jsonl", | |
| ) | |
| with open( | |
| prediction_file, | |
| "w", | |
| encoding="utf-8", | |
| ) as file: | |
| for row in predictions: | |
| file.write( | |
| json.dumps( | |
| row, | |
| ensure_ascii=False, | |
| ) | |
| + "\n" | |
| ) | |
| # ============================================================ | |
| # VALID SEVERITY PREDICTIONS | |
| # ============================================================ | |
| valid_pairs = [ | |
| ( | |
| true, | |
| pred | |
| ) | |
| for true, pred | |
| in zip( | |
| y_true, | |
| y_pred | |
| ) | |
| if true in [0, 1, 2] | |
| and pred in [0, 1, 2] | |
| ] | |
| if not valid_pairs: | |
| print() | |
| print( | |
| "ERROR: Model did not produce" | |
| " valid fault_severity values." | |
| ) | |
| print( | |
| "Check the generated responses." | |
| ) | |
| raise RuntimeError( | |
| "No valid severity predictions." | |
| ) | |
| valid_true = [ | |
| pair[0] | |
| for pair in valid_pairs | |
| ] | |
| valid_pred = [ | |
| pair[1] | |
| for pair in valid_pairs | |
| ] | |
| # ============================================================ | |
| # SEVERITY METRICS | |
| # ============================================================ | |
| accuracy = accuracy_score( | |
| valid_true, | |
| valid_pred, | |
| ) | |
| precision, recall, f1, support = ( | |
| precision_recall_fscore_support( | |
| valid_true, | |
| valid_pred, | |
| labels=[ | |
| 0, | |
| 1, | |
| 2 | |
| ], | |
| zero_division=0, | |
| ) | |
| ) | |
| matrix = confusion_matrix( | |
| valid_true, | |
| valid_pred, | |
| labels=[ | |
| 0, | |
| 1, | |
| 2 | |
| ], | |
| ) | |
| # ============================================================ | |
| # FORMAT COMPLIANCE | |
| # ============================================================ | |
| format_compliance = ( | |
| sum( | |
| row[ | |
| "all_sections_present" | |
| ] | |
| for row | |
| in predictions | |
| ) | |
| / | |
| len(predictions) | |
| ) | |
| # ============================================================ | |
| # AVERAGE TEXT SCORES | |
| # ============================================================ | |
| average_root_f1 = ( | |
| sum(root_scores) | |
| / | |
| len(root_scores) | |
| ) | |
| average_explanation_f1 = ( | |
| sum(explanation_scores) | |
| / | |
| len(explanation_scores) | |
| ) | |
| average_recommendation_f1 = ( | |
| sum(recommendation_scores) | |
| / | |
| len(recommendation_scores) | |
| ) | |
| # ============================================================ | |
| # RESULTS | |
| # ============================================================ | |
| print() | |
| print("=" * 80) | |
| print( | |
| "FOUR-OUTPUT EVALUATION RESULTS" | |
| ) | |
| print("=" * 80) | |
| print() | |
| print( | |
| f"Fault Severity Accuracy: " | |
| f"{accuracy:.4f}" | |
| ) | |
| print( | |
| f"Fault Severity Accuracy: " | |
| f"{accuracy * 100:.2f}%" | |
| ) | |
| print() | |
| print( | |
| f"Four-section format compliance: " | |
| f"{format_compliance * 100:.2f}%" | |
| ) | |
| print() | |
| print( | |
| f"Average Root Cause token-F1: " | |
| f"{average_root_f1:.4f}" | |
| ) | |
| print( | |
| f"Average Explanation token-F1: " | |
| f"{average_explanation_f1:.4f}" | |
| ) | |
| print( | |
| f"Average Recommendation token-F1: " | |
| f"{average_recommendation_f1:.4f}" | |
| ) | |
| # ============================================================ | |
| # PER-CLASS METRICS | |
| # ============================================================ | |
| print() | |
| print( | |
| "Per-class fault severity metrics:" | |
| ) | |
| for index, severity in enumerate( | |
| [0, 1, 2] | |
| ): | |
| print( | |
| f"Class {severity}: " | |
| f"precision={precision[index]:.4f}, " | |
| f"recall={recall[index]:.4f}, " | |
| f"f1={f1[index]:.4f}, " | |
| f"support={support[index]}" | |
| ) | |
| # ============================================================ | |
| # CONFUSION MATRIX | |
| # ============================================================ | |
| print() | |
| print( | |
| "Confusion matrix:" | |
| ) | |
| print() | |
| print( | |
| " Predicted" | |
| ) | |
| print( | |
| " 0 1 2" | |
| ) | |
| for index, row in enumerate( | |
| matrix | |
| ): | |
| print( | |
| f"Actual {index} " | |
| f"{row[0]:6d} " | |
| f"{row[1]:6d} " | |
| f"{row[2]:6d}" | |
| ) | |
| # ============================================================ | |
| # MAJORITY BASELINE | |
| # ============================================================ | |
| counter = Counter( | |
| valid_true | |
| ) | |
| majority_class = ( | |
| counter.most_common(1)[0][0] | |
| ) | |
| majority_accuracy = ( | |
| sum( | |
| value | |
| == majority_class | |
| for value | |
| in valid_true | |
| ) | |
| / | |
| len(valid_true) | |
| ) | |
| print() | |
| print( | |
| "Majority-class baseline:" | |
| ) | |
| print( | |
| "Majority class:", | |
| majority_class | |
| ) | |
| print( | |
| f"Baseline accuracy: " | |
| f"{majority_accuracy:.4f}" | |
| ) | |
| # ============================================================ | |
| # OUTPUT FILE | |
| # ============================================================ | |
| print() | |
| print( | |
| "Detailed predictions saved to:" | |
| ) | |
| print( | |
| os.path.abspath( | |
| prediction_file | |
| ) | |
| ) | |
| # ============================================================ | |
| # WARNING | |
| # ============================================================ | |
| print() | |
| print("=" * 80) | |
| print( | |
| "IMPORTANT" | |
| ) | |
| print("=" * 80) | |
| print( | |
| "Root Cause, Explanation and Recommendation " | |
| "token-F1 are similarity metrics against the " | |
| "structured reference responses. They are NOT " | |
| "proof of real-world telecom engineering correctness." | |
| ) | |
| print() | |
| print( | |
| "For production deployment, replace synthetic " | |
| "root-cause/recommendation targets with " | |
| "engineer-validated annotations." | |
| ) | |
| print("=" * 80) |