import html import json import re from functools import lru_cache from pathlib import Path import gradio as gr import pandas as pd from src import config as cfg from src.aspect_dict import get_aspect_dict from src.evaluator import aspect_to_overall_sentiment from src.inference import AspectPredictor ROOT = Path(__file__).resolve().parent REPORT_DIR = ROOT / "reports" CHECKPOINT_DIR = ROOT / "checkpoints" / "meta_acsa" META_ENCODER_PATH = ROOT / "data" / "meta_encoder.pkl" ASPECT_DESCRIPTIONS = { "SIZE": "Size, fit, length, width, whether it runs large or small", "MATERIAL": "Fabric, feel, thickness, breathability, material comfort", "QUALITY": "Workmanship, durability, loose threads, zippers, buttons, post-wash performance", "APPEARANCE": "Color, pattern, consistency with pictures, visual presentation", "STYLE": "Style, cut, fashionability, whether it flatters the body shape", "VALUE": "Price, value for money, worthiness of purchase, return/exchange tendency", } LABEL_COLORS = { "Positive": "#15803d", "Negative": "#b91c1c", "Not_Mentioned": "#64748b", } POSITIVE_CUES = { "love", "loved", "great", "good", "perfect", "excellent", "nice", "soft", "comfortable", "beautiful", "pretty", "flattering", "worth", "recommend", "amazing", "stylish", "quality", "durable", "cute", "fits", "fit", } NEGATIVE_CUES = { "bad", "poor", "terrible", "awful", "cheap", "scratchy", "itchy", "thin", "small", "large", "tight", "loose", "broke", "broken", "ripped", "tore", "overpriced", "return", "returning", "waste", "disappointed", "uncomfortable", "shrunk", "defective", "wrong", } def _load_json(path: Path) -> dict: if not path.exists(): return {} with path.open("r", encoding="utf-8") as f: return json.load(f) def _fmt_pct(value): try: return f"{float(value) * 100:.2f}%" except Exception: return "-" def _metric_table(rows): return pd.DataFrame(rows) def baseline_tables(): data = _load_json(REPORT_DIR / "evaluation_comparison.json") overall = data.get("overall_3class_comparison", {}) rows = [] for name, metrics in overall.items(): rows.append({ "Model": name.replace("__", " / ").replace("_", " "), "Macro-F1": round(metrics.get("macro_f1", 0), 4), "Accuracy": round(metrics.get("accuracy", 0), 4), "Macro-F1 (%)": _fmt_pct(metrics.get("macro_f1", 0)), "Accuracy (%)": _fmt_pct(metrics.get("accuracy", 0)), }) aspect_rows = [] proposed = data.get("proposed_per_aspect", {}).get("per_aspect", {}) no_meta = data.get("acsa_no_meta_per_aspect", {}).get("per_aspect", {}) for aspect in cfg.ASPECTS: p = proposed.get(aspect, {}) b = no_meta.get(aspect, {}) aspect_rows.append({ "Aspect": aspect, "Proposed F1": round(p.get("macro_f1", 0), 4), "No-Meta F1": round(b.get("macro_f1", 0), 4), "F1 Gain": round(p.get("macro_f1", 0) - b.get("macro_f1", 0), 4), "Proposed Acc": round(p.get("accuracy", 0), 4), "No-Meta Acc": round(b.get("accuracy", 0), 4), "Acc Gain": round(p.get("accuracy", 0) - b.get("accuracy", 0), 4), }) return _metric_table(rows), _metric_table(aspect_rows), _bars_html(rows, "Model") def ablation_tables(): data = _load_json(REPORT_DIR / "ablation_summary.json") rows = [] for name, block in data.items(): rows.append({ "Variant": name, "Mean Macro-F1": round(block.get("mean_macro_f1", 0), 4), "Mean Accuracy": round(block.get("mean_accuracy", 0), 4), "Macro-F1 (%)": _fmt_pct(block.get("mean_macro_f1", 0)), "Accuracy (%)": _fmt_pct(block.get("mean_accuracy", 0)), }) aspect_rows = [] proposed = data.get("Proposed", {}).get("per_aspect", {}) for aspect in cfg.ASPECTS: row = {"Aspect": aspect} for variant, block in data.items(): m = block.get("per_aspect", {}).get(aspect, {}) row[f"{variant} F1"] = round(m.get("macro_f1", 0), 4) row[f"{variant} Acc"] = round(m.get("accuracy", 0), 4) if variant != "Proposed": row[f"{variant} F1 Delta"] = round( m.get("macro_f1", 0) - proposed.get(aspect, {}).get("macro_f1", 0), 4 ) aspect_rows.append(row) return _metric_table(rows), _metric_table(aspect_rows), _bars_html(rows, "Variant") def _bars_html(rows, label_key): if not rows: return "
No report JSON found.
" max_f1 = max(float(r.get("Macro-F1", r.get("Mean Macro-F1", 0))) for r in rows) or 1 max_acc = max(float(r.get("Accuracy", r.get("Mean Accuracy", 0))) for r in rows) or 1 parts = ["
"] for row in rows: label = html.escape(str(row[label_key])) f1 = float(row.get("Macro-F1", row.get("Mean Macro-F1", 0))) acc = float(row.get("Accuracy", row.get("Mean Accuracy", 0))) parts.append( f"
{label}
" f"
F1 {f1:.4f}" f"
" f"Acc {acc:.4f}" f"
" "
" ) parts.append("
") return "\n".join(parts) def _checkpoint_status(): missing = [] if not (CHECKPOINT_DIR / "best.pt").exists(): missing.append("checkpoints/meta_acsa/best.pt") if not (CHECKPOINT_DIR / "tokenizer").exists(): missing.append("checkpoints/meta_acsa/tokenizer/") if not META_ENCODER_PATH.exists(): missing.append("data/meta_encoder.pkl") return missing @lru_cache(maxsize=1) def get_predictor(): missing = _checkpoint_status() if missing: raise FileNotFoundError("Missing model artifacts: " + ", ".join(missing)) return AspectPredictor(checkpoint_dir=CHECKPOINT_DIR) def _product_meta(features_text, categories_text, price, average_rating, rating_number): return { "features_text": features_text or "", "categories_text": categories_text or "", "price": float(price or 0), "average_rating": float(average_rating or 0), "rating_number": int(rating_number or 0), } def predict(review_text, features_text, categories_text, price, average_rating, rating_number): if not str(review_text or "").strip(): raise gr.Error("Please enter a review.") predictor = get_predictor() result = predictor.predict( review_text=review_text, product_meta=_product_meta(features_text, categories_text, price, average_rating, rating_number), ) aspect_codes = [ {"Not_Mentioned": 0, "Positive": 1, "Negative": 2}[result["aspects"][a]] for a in cfg.ASPECTS ] overall_code = int(aspect_to_overall_sentiment(aspect_codes)[0]) overall = cfg.OVERALL_LABEL_NAMES[overall_code] return result, overall def _prediction_table(result): rows = [] for aspect in cfg.ASPECTS: label = result["aspects"].get(aspect, "Not_Mentioned") rows.append({ "Aspect": aspect, "Prediction": label, "Business Meaning": ASPECT_DESCRIPTIONS[aspect], }) return pd.DataFrame(rows) def _prediction_cards(result, overall): parts = [f"
Overall sentiment: {html.escape(overall)}
"] parts.append("
") for aspect in cfg.ASPECTS: label = result["aspects"].get(aspect, "Not_Mentioned") color = LABEL_COLORS.get(label, "#64748b") parts.append( f"
" f"
{aspect}
" f"
{label}
" f"
{ASPECT_DESCRIPTIONS[aspect]}
" "
" ) parts.append("
") return "\n".join(parts) def _meta_attention_html(result): attn = result.get("meta_attention", {}) if not attn: return "
No metadata attention returned.
" names = { "meta_chunk_1": "Meta chunk 1", "meta_chunk_2": "Meta chunk 2", "meta_chunk_3": "Meta chunk 3", "meta_chunk_4": "Meta chunk 4", } max_v = max(float(v) for v in attn.values()) or 1 parts = ["
"] for key, value in attn.items(): value = float(value) parts.append( f"
{names.get(key, key)}" f"
" f"{value:.3f}
" ) parts.append("
") return "\n".join(parts) def _diagnosis_html(result, overall): negatives = [a for a, label in result["aspects"].items() if label == "Negative"] positives = [a for a, label in result["aspects"].items() if label == "Positive"] mentioned = negatives + positives risk = "High" if len(negatives) >= 2 or overall == "Negative" else "Medium" if negatives else "Low" parts = [ f"

Enterprise Diagnostic Summary

", f"

Risk level: {risk}

", f"

Detected drivers: {', '.join(mentioned) if mentioned else 'No aspect was strongly mentioned.'}

", ] if negatives: parts.append("

Immediate action: prioritize review of " + ", ".join(negatives) + " related product claims, PDP copy, sizing guidance, QA, and return reasons.

") if positives: parts.append("

Reusable strengths: emphasize " + ", ".join(positives) + " in merchandising, ad copy, and search filters.

") parts.append("

Interpretation note: metadata attention shows how much the model relied on product attributes in addition to the review text.

") return "\n".join(parts) def _extract_keyword_evidence(review_text, result): text = str(review_text or "") lower = text.lower() aspect_dict = get_aspect_dict() rows = [] for aspect in cfg.ASPECTS: label = result["aspects"].get(aspect, "Not_Mentioned") matches = [] for kw in sorted(aspect_dict[aspect], key=len, reverse=True): if re.search(r"\b" + re.escape(kw.lower()) + r"\b", lower): matches.append(kw) windows = [] for kw in matches[:6]: m = re.search(r"\b" + re.escape(kw.lower()) + r"\b", lower) if not m: continue start = max(0, m.start() - 45) end = min(len(text), m.end() + 45) windows.append(text[start:end].strip()) cue_words = sorted({w for w in re.findall(r"[a-z][a-z\-]+", lower) if w in POSITIVE_CUES or w in NEGATIVE_CUES}) rows.append({ "Aspect": aspect, "Prediction": label, "Matched Aspect Keywords": ", ".join(matches[:8]) or "-", "Nearby Review Evidence": " | ".join(windows[:3]) or "-", "Sentiment Cue Words": ", ".join(cue_words[:10]) or "-", }) return pd.DataFrame(rows) def enterprise_explain(review_text, features_text, categories_text, price, average_rating, rating_number): try: result, overall = predict(review_text, features_text, categories_text, price, average_rating, rating_number) except FileNotFoundError as e: raise gr.Error(str(e)) return ( _prediction_table(result), _prediction_cards(result, overall), _meta_attention_html(result), _diagnosis_html(result, overall), ) def consumer_keywords(review_text, features_text, categories_text, price, average_rating, rating_number): try: result, overall = predict(review_text, features_text, categories_text, price, average_rating, rating_number) except FileNotFoundError as e: raise gr.Error(str(e)) evidence = _extract_keyword_evidence(review_text, result) negatives = [a for a, label in result["aspects"].items() if label == "Negative"] positives = [a for a, label in result["aspects"].items() if label == "Positive"] if negatives: advice = f"Shopping reference: pay attention to {', '.join(negatives)}. These are the likely pain points behind the bad-review signal." elif positives: advice = f"Shopping reference: the review is mainly positive on {', '.join(positives)}. These are the strongest purchase-supporting signals." else: advice = "Shopping reference: no clear aspect-level good/bad signal was detected from this review." return _prediction_cards(result, overall), evidence, advice def model_status_html(): missing = _checkpoint_status() if not missing: return "
Model artifacts found. Proposed model inference is enabled.
" items = "".join(f"
  • {html.escape(x)}
  • " for x in missing) return f"
    Model artifacts are missing in this workspace. Upload these files to enable Tabs 3 and 4:
    " CSS = """ .metric-bars, .attention-box { display: flex; flex-direction: column; gap: 10px; } .bar-row { display: grid; grid-template-columns: minmax(220px, 1.2fr) 2fr; gap: 14px; align-items: center; padding: 10px; border: 1px solid #e5e7eb; border-radius: 8px; background: #fff; } .bar-label { font-weight: 650; color: #111827; } .bar-wrap { display: grid; grid-template-columns: 86px 1fr 86px 1fr; gap: 8px; align-items: center; font-size: 12px; } .bar-track, .attn-track { height: 10px; background: #e5e7eb; border-radius: 999px; overflow: hidden; } .bar { height: 100%; border-radius: 999px; } .bar.f1 { background: #2563eb; } .bar.acc { background: #059669; } .overall-card { padding: 14px 16px; border: 1px solid #d1d5db; border-radius: 8px; background: #f9fafb; margin-bottom: 12px; font-size: 18px; } .aspect-grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(180px, 1fr)); gap: 10px; } .aspect-card { border: 1px solid #e5e7eb; border-top: 4px solid #64748b; border-radius: 8px; padding: 12px; background: white; } .aspect-name { font-weight: 700; color: #111827; } .aspect-pred { font-size: 18px; font-weight: 750; margin-top: 4px; } .aspect-desc { color: #4b5563; font-size: 13px; margin-top: 6px; } .attn-row { display: grid; grid-template-columns: 120px 1fr 58px; gap: 10px; align-items: center; padding: 8px 10px; border: 1px solid #e5e7eb; border-radius: 8px; } .attn-fill { height: 100%; background: #7c3aed; border-radius: 999px; } .diagnosis { border: 1px solid #d1d5db; border-radius: 8px; padding: 14px 16px; background: #ffffff; } .diagnosis h3 { margin-top: 0; } .status { padding: 12px 14px; border-radius: 8px; margin-bottom: 12px; } .status.ok { background: #ecfdf5; border: 1px solid #86efac; color: #166534; } .status.warn { background: #fff7ed; border: 1px solid #fdba74; color: #9a3412; } .empty { color: #64748b; padding: 12px; } """ with gr.Blocks(css=CSS, title="ACSA-Clothing v4 Demo") as demo: gr.Markdown("# ACSA-Clothing v4: BERT + Metadata Cross-Attention") gr.HTML(model_status_html()) with gr.Tabs(): with gr.Tab("1. Baseline F1 / Accuracy"): gr.Markdown("Overall sentiment baselines and per-aspect Proposed vs no-metadata comparison.") baseline_overall, baseline_aspects, baseline_chart = baseline_tables() gr.HTML(baseline_chart) gr.Dataframe(baseline_overall, label="Overall 3-class comparison", interactive=False) gr.Dataframe(baseline_aspects, label="Per-aspect comparison", interactive=False) with gr.Tab("2. Ablation F1 / Accuracy"): gr.Markdown("A1 removes text metadata, A2 removes numeric metadata, A3 replaces Cross-Attention with concat fusion.") ablation_overall, ablation_aspects, ablation_chart = ablation_tables() gr.HTML(ablation_chart) gr.Dataframe(ablation_overall, label="Ablation summary", interactive=False) gr.Dataframe(ablation_aspects, label="Per-aspect ablation details", interactive=False) with gr.Tab("3. Proposed Enterprise Diagnosis"): gr.Markdown("Fine-grained explanation for enterprise diagnostics: aspect predictions, metadata attention, and operational recommendations.") with gr.Row(): with gr.Column(): ent_review = gr.Textbox( label="Review text", lines=5, value="The size runs small but the fabric feels great. Color looks exactly like the picture.", ) ent_features = gr.Textbox(label="features_text", value="100% cotton tee, slim fit, machine washable") ent_categories = gr.Textbox(label="categories_text", value="Clothing > Women > T-Shirts") with gr.Column(): ent_price = gr.Number(label="price", value=19.99) ent_avg = gr.Number(label="average_rating", value=4.3) ent_rnum = gr.Number(label="rating_number", value=217) ent_btn = gr.Button("Run enterprise diagnosis", variant="primary") ent_table = gr.Dataframe(label="Aspect-level predictions", interactive=False) ent_cards = gr.HTML() ent_attn = gr.HTML(label="Metadata attention") ent_diag = gr.HTML(label="Diagnostic summary") ent_btn.click( enterprise_explain, inputs=[ent_review, ent_features, ent_categories, ent_price, ent_avg, ent_rnum], outputs=[ent_table, ent_cards, ent_attn, ent_diag], ) with gr.Tab("4. Proposed Consumer Keywords"): gr.Markdown("Identify the good/bad-review keywords behind the model result for consumer shopping reference.") with gr.Row(): with gr.Column(): con_review = gr.Textbox( label="Review text", lines=5, value="I love the soft fabric and the style is cute, but it is too tight and not worth the price.", ) con_features = gr.Textbox(label="features_text", value="Soft stretch fabric, fitted silhouette") con_categories = gr.Textbox(label="categories_text", value="Clothing > Women > Dresses") with gr.Column(): con_price = gr.Number(label="price", value=49.99) con_avg = gr.Number(label="average_rating", value=3.8) con_rnum = gr.Number(label="rating_number", value=84) con_btn = gr.Button("Find shopping keywords", variant="primary") con_cards = gr.HTML() con_evidence = gr.Dataframe(label="Keyword evidence", interactive=False) con_advice = gr.Textbox(label="Consumer shopping reference", interactive=False) con_btn.click( consumer_keywords, inputs=[con_review, con_features, con_categories, con_price, con_avg, con_rnum], outputs=[con_cards, con_evidence, con_advice], ) if __name__ == "__main__": demo.launch(server_name="0.0.0.0", server_port=7860)