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()