Lavender960825's picture
Update app.py
13b1de0 verified
Raw History Blame Contribute Delete
19.7 kB
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 "<div class='empty'>No report JSON found.</div>"
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 = ["<div class='metric-bars'>"]
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"<div class='bar-row'><div class='bar-label'>{label}</div>"
f"<div class='bar-wrap'><span>F1 {f1:.4f}</span>"
f"<div class='bar-track'><div class='bar f1' style='width:{100*f1/max_f1:.1f}%'></div></div>"
f"<span>Acc {acc:.4f}</span>"
f"<div class='bar-track'><div class='bar acc' style='width:{100*acc/max_acc:.1f}%'></div></div>"
"</div></div>"
)
parts.append("</div>")
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"<div class='overall-card'>Overall sentiment: <b>{html.escape(overall)}</b></div>"]
parts.append("<div class='aspect-grid'>")
for aspect in cfg.ASPECTS:
label = result["aspects"].get(aspect, "Not_Mentioned")
color = LABEL_COLORS.get(label, "#64748b")
parts.append(
f"<div class='aspect-card' style='border-top-color:{color}'>"
f"<div class='aspect-name'>{aspect}</div>"
f"<div class='aspect-pred' style='color:{color}'>{label}</div>"
f"<div class='aspect-desc'>{ASPECT_DESCRIPTIONS[aspect]}</div>"
"</div>"
)
parts.append("</div>")
return "\n".join(parts)
def _meta_attention_html(result):
attn = result.get("meta_attention", {})
if not attn:
return "<div class='empty'>No metadata attention returned.</div>"
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 = ["<div class='attention-box'>"]
for key, value in attn.items():
value = float(value)
parts.append(
f"<div class='attn-row'><span>{names.get(key, key)}</span>"
f"<div class='attn-track'><div class='attn-fill' style='width:{100*value/max_v:.1f}%'></div></div>"
f"<b>{value:.3f}</b></div>"
)
parts.append("</div>")
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"<div class='diagnosis'><h3>Enterprise Diagnostic Summary</h3>",
f"<p><b>Risk level:</b> {risk}</p>",
f"<p><b>Detected drivers:</b> {', '.join(mentioned) if mentioned else 'No aspect was strongly mentioned.'}</p>",
]
if negatives:
parts.append("<p><b>Immediate action:</b> prioritize review of " + ", ".join(negatives) + " related product claims, PDP copy, sizing guidance, QA, and return reasons.</p>")
if positives:
parts.append("<p><b>Reusable strengths:</b> emphasize " + ", ".join(positives) + " in merchandising, ad copy, and search filters.</p>")
parts.append("<p><b>Interpretation note:</b> metadata attention shows how much the model relied on product attributes in addition to the review text.</p></div>")
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 "<div class='status ok'>Model artifacts found. Proposed model inference is enabled.</div>"
items = "".join(f"<li>{html.escape(x)}</li>" for x in missing)
return f"<div class='status warn'>Model artifacts are missing in this workspace. Upload these files to enable Tabs 3 and 4:<ul>{items}</ul></div>"
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)