samplellm / compare_models.py
Aadhitech1's picture
Upload 8 files
6082a5d verified
Raw History Blame Contribute Delete
49.3 kB
# ============================================================
# TELECOM FAST MODEL COMPARISON
# ============================================================
#
# Compare:
#
# 1. Original Qwen/Qwen2.5-1.5B-Instruct
# 2. Fine-tuned QLoRA adapter
#
# Evaluation:
#
# - Fault Severity Accuracy
# - Fault Severity Macro F1
# - Fault Severity Weighted F1
# - Root Cause Token-F1
# - Explanation Token-F1
# - Recommendation Token-F1
# - Four-output Format Compliance
# - Input Tokens
# - Output Tokens
# - Total Tokens
# - Generation Time
# - Tokens / Second
#
# IMPORTANT:
#
# This version does NOT use an LLM judge.
# It is therefore faster and reproducible.
#
# ============================================================
import os
import re
import json
import csv
import argparse
import time
import gc
from collections import Counter
import torch
from datasets import load_dataset
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
BitsAndBytesConfig,
)
from peft import PeftModel
# ============================================================
# CONFIGURATION
# ============================================================
MODEL_NAME = "Qwen/Qwen2.5-1.5B-Instruct"
DEFAULT_ADAPTER = (
"models/telecom-qwen-four-output"
)
DEFAULT_TEST_FILE = (
"sample/data/test_four_output.jsonl"
)
DEFAULT_OUTPUT_DIR = (
"evaluation_results"
)
CACHE_DIR = (
"sample/.cache/huggingface"
)
# Number of examples:
#
# 10 = very quick test
# 50 = quick evaluation
# 100 = better evaluation
# 0 = entire test set
#
DEFAULT_SAMPLES = 50
MAX_INPUT_LENGTH = 1024
MAX_NEW_TOKENS = 220
# ============================================================
# COMMAND LINE
# ============================================================
parser = argparse.ArgumentParser(
description="Fast comparison of original and fine-tuned Qwen models"
)
parser.add_argument(
"--adapter",
default=DEFAULT_ADAPTER,
help="Path to LoRA adapter",
)
parser.add_argument(
"--test",
default=DEFAULT_TEST_FILE,
help="Path to test JSONL",
)
parser.add_argument(
"--samples",
type=int,
default=DEFAULT_SAMPLES,
help="Number of examples. 0 = all examples.",
)
parser.add_argument(
"--output",
default=DEFAULT_OUTPUT_DIR,
help="Output directory",
)
parser.add_argument(
"--debug",
action="store_true",
help="Print detailed prediction for every example",
)
args = parser.parse_args()
# ============================================================
# HEADER
# ============================================================
print("=" * 80)
print("TELECOM FAST MODEL COMPARISON")
print("=" * 80)
print()
print("Base model:")
print(MODEL_NAME)
print()
print("Adapter:")
print(args.adapter)
print()
print("Test file:")
print(args.test)
# ============================================================
# GPU CHECK
# ============================================================
print()
print("=" * 80)
print("GPU")
print("=" * 80)
print()
if not torch.cuda.is_available():
raise RuntimeError(
"CUDA GPU is required for this evaluation."
)
print(
"CUDA available:",
torch.cuda.is_available(),
)
print(
"GPU:",
torch.cuda.get_device_name(0),
)
gpu_memory = (
torch.cuda.get_device_properties(0)
.total_memory
/ (1024 ** 3)
)
print(
"VRAM:",
round(gpu_memory, 2),
"GB",
)
# ============================================================
# DIRECTORY / FILE CHECK
# ============================================================
if not os.path.exists(args.test):
raise FileNotFoundError(
"\nTest file not found:\n"
+ os.path.abspath(args.test)
)
if not os.path.exists(args.adapter):
raise FileNotFoundError(
"\nAdapter not found:\n"
+ os.path.abspath(args.adapter)
)
os.makedirs(
args.output,
exist_ok=True,
)
# ============================================================
# LOAD DATASET
# ============================================================
print()
print("=" * 80)
print("LOADING TEST DATA")
print("=" * 80)
dataset = load_dataset(
"json",
data_files=args.test,
split="train",
cache_dir=CACHE_DIR,
)
print(
"Total examples:",
len(dataset),
)
if args.samples > 0:
number_to_test = min(
args.samples,
len(dataset),
)
dataset = dataset.select(
range(number_to_test)
)
print(
"Testing examples:",
len(dataset),
)
# ============================================================
# DATA VALIDATION
# ============================================================
print()
print("=" * 80)
print("VALIDATING TEST DATA")
print("=" * 80)
for i, example in enumerate(dataset):
if "messages" not in example:
raise ValueError(
f"Example {i} does not contain 'messages'."
)
if not isinstance(
example["messages"],
list,
):
raise ValueError(
f"Example {i}: messages is not a list."
)
roles = [
message.get("role")
for message in example["messages"]
]
if "user" not in roles:
raise ValueError(
f"Example {i} does not contain a user message."
)
if "assistant" not in roles:
raise ValueError(
f"Example {i} does not contain an assistant reference."
)
print(
"Test data validation completed."
)
# ============================================================
# 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 = "left"
# ============================================================
# 4-BIT 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 ORIGINAL MODEL
# ============================================================
print()
print("=" * 80)
print("LOADING ORIGINAL MODEL")
print("=" * 80)
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,
)
base_model.eval()
base_model.config.use_cache = True
print(
"Original model loaded."
)
# ============================================================
# LOAD FINE-TUNED MODEL
# ============================================================
print()
print("=" * 80)
print("LOADING FINE-TUNED MODEL")
print("=" * 80)
fine_tuned_model = PeftModel.from_pretrained(
base_model,
args.adapter,
)
fine_tuned_model.eval()
print(
"Fine-tuned adapter loaded."
)
# ============================================================
# PROMPT CREATION
# ============================================================
def get_prompt(example):
messages = []
for message in example["messages"]:
if message["role"] != "assistant":
messages.append(message)
return tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
# ============================================================
# REFERENCE RESPONSE
# ============================================================
def get_reference(example):
for message in example["messages"]:
if message["role"] == "assistant":
content = message.get(
"content",
"",
)
return str(content)
return ""
# ============================================================
# SAFE SECTION EXTRACTION
# ============================================================
def extract_section(
text,
section_name,
next_sections,
):
"""
Safely extracts text between section headings.
Example:
Probable Root Cause:
Something happened.
Explanation:
Explanation text.
Recommendation:
Restart equipment.
Fault Severity:
1
This implementation deliberately avoids complicated
dynamically-generated lookahead regexes.
"""
if text is None:
return ""
text = str(text).strip()
if not text:
return ""
# --------------------------------------------------------
# Find requested heading
# --------------------------------------------------------
start_pattern = (
r"(?im)^\s*"
+ re.escape(section_name)
+ r"\s*:?\s*"
)
start_match = re.search(
start_pattern,
text,
)
if start_match is None:
return ""
content_start = start_match.end()
remaining = text[
content_start:
]
# --------------------------------------------------------
# Find all possible following section headings
# --------------------------------------------------------
next_positions = []
for section in next_sections:
pattern = (
r"(?im)^\s*"
+ re.escape(section)
+ r"\s*:?\s*"
)
match = re.search(
pattern,
remaining,
)
if match is not None:
next_positions.append(
match.start()
)
# --------------------------------------------------------
# Stop at earliest next heading
# --------------------------------------------------------
if next_positions:
content_end = min(
next_positions
)
content = remaining[
:content_end
]
else:
content = remaining
return content.strip()
# ============================================================
# ROOT CAUSE
# ============================================================
def extract_root_cause(text):
return extract_section(
text,
"Probable Root Cause",
[
"Explanation",
"Recommendation",
"Fault Severity",
],
)
# ============================================================
# EXPLANATION
# ============================================================
def extract_explanation(text):
return extract_section(
text,
"Explanation",
[
"Recommendation",
"Fault Severity",
],
)
# ============================================================
# RECOMMENDATION
# ============================================================
def extract_recommendation(text):
return extract_section(
text,
"Recommendation",
[
"Fault Severity",
],
)
# ============================================================
# SEVERITY
# ============================================================
def extract_severity(text):
if not text:
return None
patterns = [
r"(?im)^\s*Fault\s+Severity\s*[:\-]?\s*([012])\b",
r"(?im)^\s*fault_severity\s*[:=]\s*([012])\b",
r"(?im)^\s*Severity\s*[:\-]?\s*([012])\b",
]
for pattern in patterns:
match = re.search(
pattern,
str(text),
)
if match:
return int(
match.group(1)
)
return None
# ============================================================
# FOUR-OUTPUT FORMAT
# ============================================================
def format_compliant(text):
if not text:
return False
required_sections = [
"Probable Root Cause",
"Explanation",
"Recommendation",
"Fault Severity",
]
lower_text = text.lower()
for section in required_sections:
if section.lower() not in lower_text:
return False
return True
# ============================================================
# TOKEN NORMALIZATION
# ============================================================
def normalize_tokens(text):
if not text:
return []
text = str(text).lower()
text = re.sub(
r"[^\w\s]",
" ",
text,
)
return text.split()
# ============================================================
# TOKEN F1
# ============================================================
def token_f1(
prediction,
reference,
):
pred_tokens = normalize_tokens(
prediction
)
ref_tokens = normalize_tokens(
reference
)
if (
not pred_tokens
and
not ref_tokens
):
return 1.0
if (
not pred_tokens
or
not ref_tokens
):
return 0.0
pred_counts = Counter(
pred_tokens
)
ref_counts = Counter(
ref_tokens
)
common = sum(
(
pred_counts
&
ref_counts
).values()
)
if common == 0:
return 0.0
precision = (
common
/
len(pred_tokens)
)
recall = (
common
/
len(ref_tokens)
)
if precision + recall == 0:
return 0.0
return (
2
*
precision
*
recall
/
(
precision
+
recall
)
)
# ============================================================
# GENERATION
# ============================================================
def generate_response(
model,
prompt,
):
encoded = tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=MAX_INPUT_LENGTH,
)
input_tokens = int(
encoded["input_ids"].shape[1]
)
# --------------------------------------------------------
# Move tensors to model device
# --------------------------------------------------------
encoded = {
key: value.to(model.device)
for key, value in encoded.items()
}
if torch.cuda.is_available():
torch.cuda.synchronize()
start_time = time.perf_counter()
with torch.inference_mode():
output = model.generate(
**encoded,
max_new_tokens=MAX_NEW_TOKENS,
do_sample=False,
num_beams=1,
use_cache=True,
pad_token_id=(
tokenizer.pad_token_id
),
eos_token_id=(
tokenizer.eos_token_id
),
)
if torch.cuda.is_available():
torch.cuda.synchronize()
elapsed = (
time.perf_counter()
-
start_time
)
generated_tokens = output[
0
][
input_tokens:
]
output_token_count = int(
len(generated_tokens)
)
total_tokens = (
input_tokens
+
output_token_count
)
response = tokenizer.decode(
generated_tokens,
skip_special_tokens=True,
).strip()
tokens_per_second = (
output_token_count
/
elapsed
if elapsed > 0
else 0.0
)
return {
"response": response,
"input_tokens": input_tokens,
"output_tokens": output_token_count,
"total_tokens": total_tokens,
"generation_time_seconds": elapsed,
"tokens_per_second": tokens_per_second,
}
# ============================================================
# EVALUATE MODEL
# ============================================================
def evaluate_model(
model,
model_name,
):
print()
print("=" * 80)
print(
"TESTING:",
model_name,
)
print("=" * 80)
results = []
total_examples = len(
dataset
)
for index, example in enumerate(
dataset
):
prompt = get_prompt(
example
)
reference = get_reference(
example
)
generated = generate_response(
model,
prompt,
)
prediction = generated[
"response"
]
# ----------------------------------------------------
# Reference sections
# ----------------------------------------------------
reference_root = (
extract_root_cause(
reference
)
)
reference_explanation = (
extract_explanation(
reference
)
)
reference_recommendation = (
extract_recommendation(
reference
)
)
reference_severity = (
extract_severity(
reference
)
)
# ----------------------------------------------------
# Prediction sections
# ----------------------------------------------------
predicted_root = (
extract_root_cause(
prediction
)
)
predicted_explanation = (
extract_explanation(
prediction
)
)
predicted_recommendation = (
extract_recommendation(
prediction
)
)
predicted_severity = (
extract_severity(
prediction
)
)
# ----------------------------------------------------
# Similarity
# ----------------------------------------------------
root_score = token_f1(
predicted_root,
reference_root,
)
explanation_score = token_f1(
predicted_explanation,
reference_explanation,
)
recommendation_score = token_f1(
predicted_recommendation,
reference_recommendation,
)
# ----------------------------------------------------
# Severity
# ----------------------------------------------------
severity_correct = (
predicted_severity is not None
and
reference_severity is not None
and
predicted_severity
==
reference_severity
)
# ----------------------------------------------------
# Format
# ----------------------------------------------------
format_ok = format_compliant(
prediction
)
# ----------------------------------------------------
# Result
# ----------------------------------------------------
row = {
"index": index + 1,
"model": model_name,
"reference_severity": (
reference_severity
),
"predicted_severity": (
predicted_severity
),
"severity_correct": (
severity_correct
),
"format_compliant": (
format_ok
),
"root_cause_f1": (
root_score
),
"explanation_f1": (
explanation_score
),
"recommendation_f1": (
recommendation_score
),
"input_tokens": (
generated[
"input_tokens"
]
),
"output_tokens": (
generated[
"output_tokens"
]
),
"total_tokens": (
generated[
"total_tokens"
]
),
"generation_time_seconds": (
generated[
"generation_time_seconds"
]
),
"tokens_per_second": (
generated[
"tokens_per_second"
]
),
"reference": reference,
"prediction": prediction,
"reference_root_cause": (
reference_root
),
"predicted_root_cause": (
predicted_root
),
"reference_explanation": (
reference_explanation
),
"predicted_explanation": (
predicted_explanation
),
"reference_recommendation": (
reference_recommendation
),
"predicted_recommendation": (
predicted_recommendation
),
}
results.append(row)
# ----------------------------------------------------
# Debug
# ----------------------------------------------------
if args.debug:
print()
print("-" * 80)
print(
"Example:",
index + 1,
)
print()
print(
"REFERENCE:"
)
print(reference)
print()
print(
"PREDICTION:"
)
print(prediction)
print()
print(
"Severity:",
reference_severity,
"->",
predicted_severity,
)
print(
"Root Cause F1:",
f"{root_score:.4f}",
)
print(
"Explanation F1:",
f"{explanation_score:.4f}",
)
print(
"Recommendation F1:",
f"{recommendation_score:.4f}",
)
print(
"Format:",
format_ok,
)
print(
"Tokens:",
generated[
"input_tokens"
],
"/",
generated[
"output_tokens"
],
"/",
generated[
"total_tokens"
],
)
print(
"Time:",
f"{generated['generation_time_seconds']:.3f}s",
)
else:
if (
(index + 1) % 10 == 0
or
(index + 1)
==
total_examples
):
print(
f"{model_name}: "
f"{index + 1}/"
f"{total_examples}"
)
return results
# ============================================================
# RUN ORIGINAL
# ============================================================
original_results = evaluate_model(
base_model,
"Original",
)
# ============================================================
# RUN FINE-TUNED
# ============================================================
fine_tuned_results = evaluate_model(
fine_tuned_model,
"Fine-tuned",
)
# ============================================================
# AVERAGE
# ============================================================
def average(
rows,
key,
):
if not rows:
return 0.0
values = [
float(row[key])
for row in rows
if isinstance(
row[key],
(int, float),
)
]
if not values:
return 0.0
return (
sum(values)
/
len(values)
)
# ============================================================
# ACCURACY
# ============================================================
def severity_accuracy(
rows,
):
if not rows:
return 0.0
correct = sum(
1
for row in rows
if row[
"severity_correct"
]
)
return (
correct
/
len(rows)
)
# ============================================================
# FORMAT ACCURACY
# ============================================================
def format_accuracy(
rows,
):
if not rows:
return 0.0
correct = sum(
1
for row in rows
if row[
"format_compliant"
]
)
return (
correct
/
len(rows)
)
# ============================================================
# CLASSIFICATION METRICS
# ============================================================
def classification_metrics(
rows,
):
labels = [0, 1, 2]
matrix = {
actual: {
predicted: 0
for predicted in labels
}
for actual in labels
}
for row in rows:
actual = row[
"reference_severity"
]
predicted = row[
"predicted_severity"
]
if (
actual in labels
and
predicted in labels
):
matrix[
actual
][
predicted
] += 1
metrics = {}
for label in labels:
true_positive = matrix[
label
][
label
]
false_positive = sum(
matrix[
actual
][
label
]
for actual in labels
if actual != label
)
false_negative = sum(
matrix[
label
][
predicted
]
for predicted in labels
if predicted != label
)
support = sum(
matrix[
label
].values()
)
precision = (
true_positive
/
(
true_positive
+
false_positive
)
if (
true_positive
+
false_positive
) > 0
else 0.0
)
recall = (
true_positive
/
(
true_positive
+
false_negative
)
if (
true_positive
+
false_negative
) > 0
else 0.0
)
f1 = (
2
*
precision
*
recall
/
(
precision
+
recall
)
if (
precision
+
recall
) > 0
else 0.0
)
metrics[label] = {
"precision": precision,
"recall": recall,
"f1": f1,
"support": support,
}
valid_f1 = [
metrics[label]["f1"]
for label in labels
if metrics[label]["support"] > 0
]
macro_f1 = (
sum(valid_f1)
/
len(valid_f1)
if valid_f1
else 0.0
)
total_support = sum(
metrics[label]["support"]
for label in labels
)
weighted_f1 = (
sum(
metrics[label]["f1"]
*
metrics[label]["support"]
for label in labels
)
/
total_support
if total_support > 0
else 0.0
)
return (
metrics,
matrix,
macro_f1,
weighted_f1,
)
# ============================================================
# SUMMARY VALUES
# ============================================================
original_accuracy = severity_accuracy(
original_results
)
fine_tuned_accuracy = severity_accuracy(
fine_tuned_results
)
original_format = format_accuracy(
original_results
)
fine_tuned_format = format_accuracy(
fine_tuned_results
)
original_root = average(
original_results,
"root_cause_f1",
)
fine_tuned_root = average(
fine_tuned_results,
"root_cause_f1",
)
original_explanation = average(
original_results,
"explanation_f1",
)
fine_tuned_explanation = average(
fine_tuned_results,
"explanation_f1",
)
original_recommendation = average(
original_results,
"recommendation_f1",
)
fine_tuned_recommendation = average(
fine_tuned_results,
"recommendation_f1",
)
original_input = average(
original_results,
"input_tokens",
)
fine_tuned_input = average(
fine_tuned_results,
"input_tokens",
)
original_output = average(
original_results,
"output_tokens",
)
fine_tuned_output = average(
fine_tuned_results,
"output_tokens",
)
original_total = average(
original_results,
"total_tokens",
)
fine_tuned_total = average(
fine_tuned_results,
"total_tokens",
)
original_time = average(
original_results,
"generation_time_seconds",
)
fine_tuned_time = average(
fine_tuned_results,
"generation_time_seconds",
)
original_speed = average(
original_results,
"tokens_per_second",
)
fine_tuned_speed = average(
fine_tuned_results,
"tokens_per_second",
)
(
original_classes,
original_matrix,
original_macro_f1,
original_weighted_f1,
) = classification_metrics(
original_results
)
(
fine_tuned_classes,
fine_tuned_matrix,
fine_tuned_macro_f1,
fine_tuned_weighted_f1,
) = classification_metrics(
fine_tuned_results
)
# ============================================================
# PRINT QUALITY COMPARISON
# ============================================================
print()
print("=" * 80)
print("FAST QUALITY COMPARISON")
print("=" * 80)
print()
print(
f"{'Metric':<32}"
f"{'Original':>15}"
f"{'Fine-tuned':>15}"
f"{'Change':>15}"
)
print("-" * 80)
def print_metric(
name,
original,
fine_tuned,
percent=False,
):
change = (
fine_tuned
-
original
)
if percent:
print(
f"{name:<32}"
f"{original * 100:>14.2f}%"
f"{fine_tuned * 100:>14.2f}%"
f"{change * 100:>+14.2f}%"
)
else:
print(
f"{name:<32}"
f"{original:>15.4f}"
f"{fine_tuned:>15.4f}"
f"{change:>+15.4f}"
)
print_metric(
"Severity Accuracy",
original_accuracy,
fine_tuned_accuracy,
percent=True,
)
print_metric(
"Severity Macro F1",
original_macro_f1,
fine_tuned_macro_f1,
)
print_metric(
"Severity Weighted F1",
original_weighted_f1,
fine_tuned_weighted_f1,
)
print_metric(
"Root Cause Token-F1",
original_root,
fine_tuned_root,
)
print_metric(
"Explanation Token-F1",
original_explanation,
fine_tuned_explanation,
)
print_metric(
"Recommendation Token-F1",
original_recommendation,
fine_tuned_recommendation,
)
print_metric(
"Four-output Format",
original_format,
fine_tuned_format,
percent=True,
)
# ============================================================
# TOKEN USAGE
# ============================================================
print()
print("=" * 80)
print("TOKEN USAGE AND SPEED")
print("=" * 80)
print()
print_metric(
"Average Input Tokens",
original_input,
fine_tuned_input,
)
print_metric(
"Average Output Tokens",
original_output,
fine_tuned_output,
)
print_metric(
"Average Total Tokens",
original_total,
fine_tuned_total,
)
print_metric(
"Average Generation Time",
original_time,
fine_tuned_time,
)
print_metric(
"Average Tokens / Second",
original_speed,
fine_tuned_speed,
)
# ============================================================
# CLASS METRICS
# ============================================================
print()
print("=" * 80)
print("FAULT SEVERITY CLASS METRICS")
print("=" * 80)
for label in [0, 1, 2]:
print()
print(
f"CLASS {label}"
)
print()
print(
"Original:"
)
print(
f" Precision = "
f"{original_classes[label]['precision']:.4f}"
)
print(
f" Recall = "
f"{original_classes[label]['recall']:.4f}"
)
print(
f" F1 = "
f"{original_classes[label]['f1']:.4f}"
)
print(
f" Support = "
f"{original_classes[label]['support']}"
)
print()
print(
"Fine-tuned:"
)
print(
f" Precision = "
f"{fine_tuned_classes[label]['precision']:.4f}"
)
print(
f" Recall = "
f"{fine_tuned_classes[label]['recall']:.4f}"
)
print(
f" F1 = "
f"{fine_tuned_classes[label]['f1']:.4f}"
)
print(
f" Support = "
f"{fine_tuned_classes[label]['support']}"
)
# ============================================================
# CONFUSION MATRIX
# ============================================================
def print_confusion_matrix(
matrix,
title,
):
print()
print(title)
print()
print(
" Pred 0 Pred 1 Pred 2"
)
for actual in [0, 1, 2]:
print(
f"Actual {actual}"
f"{matrix[actual][0]:>12}"
f"{matrix[actual][1]:>10}"
f"{matrix[actual][2]:>10}"
)
print_confusion_matrix(
original_matrix,
"Original confusion matrix:",
)
print_confusion_matrix(
fine_tuned_matrix,
"Fine-tuned confusion matrix:",
)
# ============================================================
# TOKEN RANGES
# ============================================================
def token_range(
rows,
key,
):
values = [
row[key]
for row in rows
]
if not values:
return 0, 0
return (
min(values),
max(values),
)
print()
print("=" * 80)
print("TOKEN RANGE")
print("=" * 80)
for key in [
"input_tokens",
"output_tokens",
"total_tokens",
]:
original_min, original_max = (
token_range(
original_results,
key,
)
)
fine_tuned_min, fine_tuned_max = (
token_range(
fine_tuned_results,
key,
)
)
print()
print(key)
print(
" Original: "
f"{original_min} - "
f"{original_max}"
)
print(
" Fine-tuned: "
f"{fine_tuned_min} - "
f"{fine_tuned_max}"
)
# ============================================================
# SAVE PATHS
# ============================================================
json_path = os.path.join(
args.output,
"fast_comparison_results.json",
)
csv_path = os.path.join(
args.output,
"fast_comparison_predictions.csv",
)
summary_path = os.path.join(
args.output,
"fast_comparison_summary.csv",
)
# ============================================================
# SAVE JSON
# ============================================================
json_output = {
"base_model": MODEL_NAME,
"adapter": os.path.abspath(
args.adapter
),
"test_file": os.path.abspath(
args.test
),
"number_of_examples": len(
dataset
),
"metrics": {
"severity_accuracy": {
"original":
original_accuracy,
"fine_tuned":
fine_tuned_accuracy,
"change":
fine_tuned_accuracy
-
original_accuracy,
},
"severity_macro_f1": {
"original":
original_macro_f1,
"fine_tuned":
fine_tuned_macro_f1,
"change":
fine_tuned_macro_f1
-
original_macro_f1,
},
"severity_weighted_f1": {
"original":
original_weighted_f1,
"fine_tuned":
fine_tuned_weighted_f1,
"change":
fine_tuned_weighted_f1
-
original_weighted_f1,
},
"root_cause_token_f1": {
"original":
original_root,
"fine_tuned":
fine_tuned_root,
"change":
fine_tuned_root
-
original_root,
},
"explanation_token_f1": {
"original":
original_explanation,
"fine_tuned":
fine_tuned_explanation,
"change":
fine_tuned_explanation
-
original_explanation,
},
"recommendation_token_f1": {
"original":
original_recommendation,
"fine_tuned":
fine_tuned_recommendation,
"change":
fine_tuned_recommendation
-
original_recommendation,
},
"format_compliance": {
"original":
original_format,
"fine_tuned":
fine_tuned_format,
"change":
fine_tuned_format
-
original_format,
},
},
"token_usage": {
"input_tokens": {
"original":
original_input,
"fine_tuned":
fine_tuned_input,
},
"output_tokens": {
"original":
original_output,
"fine_tuned":
fine_tuned_output,
},
"total_tokens": {
"original":
original_total,
"fine_tuned":
fine_tuned_total,
},
"generation_time_seconds": {
"original":
original_time,
"fine_tuned":
fine_tuned_time,
},
"tokens_per_second": {
"original":
original_speed,
"fine_tuned":
fine_tuned_speed,
},
},
"class_metrics": {
"original":
original_classes,
"fine_tuned":
fine_tuned_classes,
},
"confusion_matrix": {
"original":
original_matrix,
"fine_tuned":
fine_tuned_matrix,
},
"predictions": {
"original":
original_results,
"fine_tuned":
fine_tuned_results,
},
}
with open(
json_path,
"w",
encoding="utf-8",
) as file:
json.dump(
json_output,
file,
indent=2,
ensure_ascii=False,
)
# ============================================================
# SAVE PREDICTIONS CSV
# ============================================================
combined_rows = []
for original, fine_tuned in zip(
original_results,
fine_tuned_results,
):
combined_rows.append({
"index":
original["index"],
"reference_severity":
original[
"reference_severity"
],
"original_severity":
original[
"predicted_severity"
],
"fine_tuned_severity":
fine_tuned[
"predicted_severity"
],
"original_correct":
original[
"severity_correct"
],
"fine_tuned_correct":
fine_tuned[
"severity_correct"
],
"original_format":
original[
"format_compliant"
],
"fine_tuned_format":
fine_tuned[
"format_compliant"
],
"original_root_f1":
original[
"root_cause_f1"
],
"fine_tuned_root_f1":
fine_tuned[
"root_cause_f1"
],
"original_explanation_f1":
original[
"explanation_f1"
],
"fine_tuned_explanation_f1":
fine_tuned[
"explanation_f1"
],
"original_recommendation_f1":
original[
"recommendation_f1"
],
"fine_tuned_recommendation_f1":
fine_tuned[
"recommendation_f1"
],
"original_input_tokens":
original[
"input_tokens"
],
"fine_tuned_input_tokens":
fine_tuned[
"input_tokens"
],
"original_output_tokens":
original[
"output_tokens"
],
"fine_tuned_output_tokens":
fine_tuned[
"output_tokens"
],
"original_total_tokens":
original[
"total_tokens"
],
"fine_tuned_total_tokens":
fine_tuned[
"total_tokens"
],
"original_time_seconds":
original[
"generation_time_seconds"
],
"fine_tuned_time_seconds":
fine_tuned[
"generation_time_seconds"
],
"original_tokens_per_second":
original[
"tokens_per_second"
],
"fine_tuned_tokens_per_second":
fine_tuned[
"tokens_per_second"
],
"reference":
original[
"reference"
],
"original_prediction":
original[
"prediction"
],
"fine_tuned_prediction":
fine_tuned[
"prediction"
],
})
if combined_rows:
with open(
csv_path,
"w",
newline="",
encoding="utf-8",
) as file:
writer = csv.DictWriter(
file,
fieldnames=list(
combined_rows[0].keys()
),
)
writer.writeheader()
writer.writerows(
combined_rows
)
# ============================================================
# SUMMARY CSV
# ============================================================
summary_rows = [
[
"Severity Accuracy",
original_accuracy,
fine_tuned_accuracy,
fine_tuned_accuracy
-
original_accuracy,
],
[
"Severity Macro F1",
original_macro_f1,
fine_tuned_macro_f1,
fine_tuned_macro_f1
-
original_macro_f1,
],
[
"Severity Weighted F1",
original_weighted_f1,
fine_tuned_weighted_f1,
fine_tuned_weighted_f1
-
original_weighted_f1,
],
[
"Root Cause Token-F1",
original_root,
fine_tuned_root,
fine_tuned_root
-
original_root,
],
[
"Explanation Token-F1",
original_explanation,
fine_tuned_explanation,
fine_tuned_explanation
-
original_explanation,
],
[
"Recommendation Token-F1",
original_recommendation,
fine_tuned_recommendation,
fine_tuned_recommendation
-
original_recommendation,
],
[
"Format Compliance",
original_format,
fine_tuned_format,
fine_tuned_format
-
original_format,
],
[
"Average Input Tokens",
original_input,
fine_tuned_input,
fine_tuned_input
-
original_input,
],
[
"Average Output Tokens",
original_output,
fine_tuned_output,
fine_tuned_output
-
original_output,
],
[
"Average Total Tokens",
original_total,
fine_tuned_total,
fine_tuned_total
-
original_total,
],
[
"Average Generation Time",
original_time,
fine_tuned_time,
fine_tuned_time
-
original_time,
],
[
"Average Tokens / Second",
original_speed,
fine_tuned_speed,
fine_tuned_speed
-
original_speed,
],
]
with open(
summary_path,
"w",
newline="",
encoding="utf-8",
) as file:
writer = csv.writer(
file
)
writer.writerow([
"Metric",
"Original",
"Fine-tuned",
"Change",
])
writer.writerows(
summary_rows
)
# ============================================================
# FINAL MESSAGE
# ============================================================
print()
print("=" * 80)
print("COMPARISON COMPLETE")
print("=" * 80)
print()
print(
"Examples tested:",
len(dataset),
)
print()
print(
"Detailed predictions:"
)
print(
os.path.abspath(
csv_path
)
)
print()
print(
"Summary:"
)
print(
os.path.abspath(
summary_path
)
)
print()
print(
"JSON:"
)
print(
os.path.abspath(
json_path
)
)
print()
print("=" * 80)
print("IMPORTANT")
print("=" * 80)
print(
"This is an objective reference-based evaluation."
)
print()
print(
"Token-F1 measures similarity to the "
"reference response. It does not prove "
"real-world telecom engineering correctness."
)
print()
print(
"Relevance, Coherence, Faithfulness and "
"Helpfulness are intentionally not estimated "
"by the 1.5B model itself."
)
print()
print(
"For those qualities, use engineer validation "
"or a separate stronger evaluation model."
)
print("=" * 80)
# ============================================================
# CLEANUP
# ============================================================
try:
del fine_tuned_model
except Exception:
pass
try:
del base_model
except Exception:
pass
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()