amazon-acsa-dashboard / scripts /06_evaluate.py
Lavender960825's picture
Upload 50 files
8f1213e verified
Raw
History Blame Contribute Delete
8.52 kB
"""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()