File size: 8,523 Bytes
8f1213e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 | """Script 06: Evaluate everything on the held-out test set and produce comparison report.
Outputs:
reports/evaluation_comparison.json β overall (3-class) comparison table
across all 4 models
reports/per_aspect_proposed.json β per-aspect detail for Proposed
reports/per_aspect_acsa_no_meta.json β per-aspect detail for Baseline 3
reports/aspect_distribution.png β Pos/Neg share per aspect
reports/category_aspect_*.png β drilldown heatmaps
reports/category_aspect_aggregation.csv
"""
import json
import sys
from pathlib import Path
import numpy as np
import pandas as pd
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from src.utils import setup_logging
from src import config as cfg
from src.meta_encoder import MetaEncoder
from src.evaluator import (
load_meta_acsa, load_acsa, load_bert_overall,
predict_per_aspect, predict_overall, predict_overall_from_proposed,
evaluate_per_aspect, overall_metrics,
aspect_to_overall_sentiment,
aggregate_aspect_distribution_by_category,
)
from src.explainer import plot_aspect_distribution, plot_category_aspect_heatmap
def _safe_load_tfidf_metrics():
p = cfg.CHECKPOINT_DIR / "tfidf_baseline" / "metrics.json"
if p.exists():
with open(p) as f:
return json.load(f)
return {}
def _eval_proposed_or_skip(test_df):
ckpt = cfg.CHECKPOINT_DIR / "meta_acsa" / "best.pt"
if not ckpt.exists():
print("[warn] Proposed model checkpoint not found; skipping.")
return None, None, None, None
print("\n>>> Evaluating Proposed (BERT + Meta Cross-Attention) ...")
model, tok, enc, device = load_meta_acsa()
preds, labels = predict_per_aspect(model, tok, test_df, device, meta_encoder=enc)
detail = evaluate_per_aspect(preds, labels)
return preds, labels, detail, (model, tok, enc, device)
def _eval_acsa_no_meta_or_skip(test_df):
ckpt = cfg.CHECKPOINT_DIR / "acsa" / "best.pt"
if not ckpt.exists():
print("[warn] BERT-ACSA (no meta) checkpoint not found; skipping.")
return None, None, None
print("\n>>> Evaluating BERT-ACSA (no meta) ...")
model, tok, device = load_acsa()
preds, labels = predict_per_aspect(model, tok, test_df, device, meta_encoder=None)
detail = evaluate_per_aspect(preds, labels)
return preds, labels, detail
def _eval_bert_overall_or_skip(test_df):
ckpt = cfg.CHECKPOINT_DIR / "bert_overall" / "best.pt"
if not ckpt.exists():
print("[warn] BERT-overall checkpoint not found; skipping.")
return None
print("\n>>> Evaluating BERT-overall ...")
model, tok, device = load_bert_overall()
preds, labels = predict_overall(model, tok, test_df, device)
return overall_metrics(labels, preds)
def _print_per_aspect(detail, header):
if detail is None:
return
print(f"\n=== {header} per-aspect ===")
for aspect, m in detail["per_aspect"].items():
print(f" {aspect}: F1={m['macro_f1']:.4f} Acc={m['accuracy']:.4f}")
print(f"Mean Macro-F1: {detail['overall']['mean_macro_f1']:.4f} "
f"Mean Acc: {detail['overall']['mean_accuracy']:.4f}")
def main():
setup_logging()
test_df = pd.read_parquet(cfg.TEST_PATH).reset_index(drop=True)
# Run per-aspect models
prop_preds, prop_labels, prop_detail, prop_ctx = _eval_proposed_or_skip(test_df)
base3_preds, _, base3_detail = _eval_acsa_no_meta_or_skip(test_df)
_print_per_aspect(prop_detail, "Proposed (BERT + Meta Fusion)")
_print_per_aspect(base3_detail, "Baseline 3 (BERT-ACSA, no meta)")
# Save per-aspect detailed JSON
if prop_detail:
with open(cfg.REPORT_DIR / "per_aspect_proposed.json", "w") as f:
json.dump(prop_detail, f, indent=2)
if base3_detail:
with open(cfg.REPORT_DIR / "per_aspect_acsa_no_meta.json", "w") as f:
json.dump(base3_detail, f, indent=2)
# ---- Overall 3-class comparison table ----
print("\n" + "=" * 60)
print("OVERALL (3-class) MODEL COMPARISON")
print("=" * 60)
tfidf = _safe_load_tfidf_metrics()
bert_overall = _eval_bert_overall_or_skip(test_df)
y_true_overall = test_df["overall_label"].astype(int).values
# Proposed: use overall_head directly (joint-trained)
proposed_overall_head = None
proposed_overall_agg = None
if prop_ctx is not None:
model, tok, enc, device = prop_ctx
print("\n>>> Evaluating Proposed overall_head (joint-trained) ...")
head_preds, head_labels = predict_overall_from_proposed(
model, tok, test_df, device, meta_encoder=enc)
proposed_overall_head = overall_metrics(head_labels, head_preds)
# Also compute voting-based aggregation for reference
agg = aspect_to_overall_sentiment(np.array(prop_preds).T)
proposed_overall_agg = overall_metrics(y_true_overall, agg)
base3_overall = None
if base3_preds is not None:
agg = aspect_to_overall_sentiment(np.array(base3_preds).T)
base3_overall = overall_metrics(y_true_overall, agg)
comparison = {
"Baseline_1_TFIDF_LogReg": {
"macro_f1": tfidf.get("test_macro_f1"),
"accuracy": tfidf.get("test_accuracy"),
},
"Baseline_2_BERT_overall_3class": {
"macro_f1": (bert_overall or {}).get("test_macro_f1"),
"accuracy": (bert_overall or {}).get("test_accuracy"),
},
"Baseline_3_BERT_ACSA_no_meta__aggregated_to_overall": (
{"macro_f1": base3_overall["test_macro_f1"],
"accuracy": base3_overall["test_accuracy"]}
if base3_overall else None
),
"Proposed_BERT_Meta_Fusion__overall_head": (
{"macro_f1": proposed_overall_head["test_macro_f1"],
"accuracy": proposed_overall_head["test_accuracy"]}
if proposed_overall_head else None
),
"Proposed_BERT_Meta_Fusion__aggregated_to_overall_(reference)": (
{"macro_f1": proposed_overall_agg["test_macro_f1"],
"accuracy": proposed_overall_agg["test_accuracy"]}
if proposed_overall_agg else None
),
}
print(json.dumps(comparison, indent=2))
# ---- Save full report ----
report = {
"overall_3class_comparison": comparison,
"proposed_per_aspect": prop_detail,
"acsa_no_meta_per_aspect": base3_detail,
"note": (
"The Proposed model jointly trains per-aspect heads and an overall "
"sentiment head on the shared fused representation. The overall_head "
"result is the primary overall metric. The aggregated_to_overall "
"result (voting from per-aspect predictions) is included for reference. "
"Baseline 3 vs Proposed isolates the marginal value of "
"metadata cross-attention fusion."
),
}
out = cfg.REPORT_DIR / "evaluation_comparison.json"
with open(out, "w") as f:
json.dump(report, f, indent=2)
print(f"\nFull report -> {out}")
# ---- Drilldown plots (use Proposed predictions if available, else Baseline 3) ----
aspect_preds_for_drill = prop_preds if prop_preds is not None else base3_preds
source = "proposed" if prop_preds is not None else "acsa_no_meta"
if aspect_preds_for_drill is None:
print("[info] No per-aspect predictions available; skipping drilldown plots.")
return
df_with_preds = test_df.copy().reset_index(drop=True)
arr = np.array(aspect_preds_for_drill).T
for i, a in enumerate(cfg.ASPECTS):
df_with_preds[f"pred_{a}"] = arr[:, i]
plot_aspect_distribution(
df_with_preds, output_path=cfg.REPORT_DIR / f"aspect_distribution.png",
)
agg = aggregate_aspect_distribution_by_category(test_df, aspect_preds_for_drill)
if agg is not None and len(agg) > 0:
agg.to_csv(cfg.REPORT_DIR / "category_aspect_aggregation.csv", index=False)
plot_category_aspect_heatmap(agg, "negative_share",
cfg.REPORT_DIR / "category_aspect_negative_heatmap.png")
plot_category_aspect_heatmap(agg, "positive_share",
cfg.REPORT_DIR / "category_aspect_positive_heatmap.png")
print(f"Saved category x aspect drilldown ({source}) to {cfg.REPORT_DIR}/")
else:
print("[info] Drilldown aggregation empty; check category column or row counts.")
if __name__ == "__main__":
main()
|