samplellm / test.py
Aadhitech1's picture
Upload 8 files
6082a5d verified
Raw History Blame Contribute Delete
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)