"""User-facing inference for the Proposed model.""" from typing import Dict, List, Optional import numpy as np import pandas as pd import torch from . import config as cfg from .evaluator import load_meta_acsa, format_aspect_summary _EXPECTED_META_KEYS = ( "features_text", "categories_text", "price", "average_rating", "rating_number", ) def _meta_dict_to_df(meta: Dict) -> pd.DataFrame: """Build a one-row DataFrame the MetaEncoder can transform.""" row = {k: meta.get(k) for k in _EXPECTED_META_KEYS} return pd.DataFrame([row]) def _attention_names(out: Dict, attn_array) -> List[str]: names = out.get("meta_token_names") if names is not None: return list(names) arr = np.asarray(attn_array) width = int(arr.shape[-1]) if arr.ndim else int(arr.size) if width == 3: return ["features", "categories", "numeric"] return [f"meta_chunk_{i + 1}" for i in range(width)] def _attention_dict(names: List[str], weights) -> Dict[str, float]: arr = np.asarray(weights, dtype=np.float64).reshape(-1) return {name: float(w) for name, w in zip(names, arr)} def _top_attention_item(attn: Dict[str, float]) -> Dict[str, float]: if not attn: return {"source": "", "weight": 0.0} source = max(attn, key=attn.get) return {"source": source, "weight": float(attn[source])} def _meta_summary(meta: Dict) -> Dict[str, str]: features = str(meta.get("features_text") or meta.get("features") or "")[:420] categories = str(meta.get("categories_text") or meta.get("category") or "")[:220] numeric = [] for key in ("price", "average_rating", "rating_number"): val = meta.get(key) if val is not None: numeric.append(f"{key}={val}") return { "features": features, "categories": categories, "numeric": ", ".join(numeric) if numeric else "not available", } _ASPECT_EVIDENCE_KEYWORDS = { "SIZE": {"size", "fit", "fits", "fitting", "small", "large", "big", "tight", "loose", "xl", "medium", "waist", "length"}, "MATERIAL": {"material", "fabric", "cotton", "polyester", "soft", "scratchy", "thin", "thick", "stretch", "leather", "wool"}, "QUALITY": {"quality", "stitch", "stitching", "seam", "wash", "washed", "durable", "cheap", "broke", "tear", "torn"}, "APPEARANCE": {"look", "looks", "color", "colour", "photo", "picture", "beautiful", "cute", "print", "design"}, "STYLE": {"style", "stylish", "flattering", "casual", "formal", "dress", "shirt", "fashion", "compliments"}, "VALUE": {"price", "worth", "value", "money", "cheap", "expensive", "discount", "penny", "cost"}, } _SENTIMENT_EVIDENCE_KEYWORDS = { "good", "great", "love", "loved", "perfect", "nice", "excellent", "comfortable", "soft", "bad", "poor", "cheap", "terrible", "awful", "small", "large", "tight", "loose", "thin", "worth", "disappointed", "return", "returned", "recommend", "flattering", "beautiful", } _STOPWORDS = { "a", "an", "the", "and", "or", "but", "if", "then", "than", "so", "as", "at", "by", "for", "from", "in", "into", "of", "on", "to", "with", "without", "is", "are", "was", "were", "be", "been", "being", "it", "its", "this", "that", "these", "those", "i", "me", "my", "we", "our", "you", "your", "he", "she", "they", "them", "his", "her", "their", "very", "really", "just", "also", "too", "would", "could", "should", "can", "will", "did", "do", "does", "have", "has", "had", "there", "here", "about", "after", "before", } def _clean_term(term: str) -> str: return str(term).lower().strip(".,!?;:'\"()[]{}<>/\\|`~@#$%^&*_+=") def _is_informative_term(term: str) -> bool: clean = _clean_term(term) if len(clean) < 2 or clean in _STOPWORDS: return False return any(ch.isalpha() for ch in clean) def _text_evidence_by_aspect(review_text: str, top_k: int = 8) -> Dict[str, List[str]]: raw_terms = [_clean_term(t) for t in str(review_text).split()] result = {} for aspect, aspect_terms in _ASPECT_EVIDENCE_KEYWORDS.items(): hits = [] for term in raw_terms: if not _is_informative_term(term): continue if term in aspect_terms or term in _SENTIMENT_EVIDENCE_KEYWORDS: if term not in hits: hits.append(term) if len(hits) >= top_k: break result[aspect] = hits return result def _format_confidence(out: Dict, row_idx: int = 0) -> Dict[str, Dict[str, float]]: probs = torch.softmax(out["logits"][row_idx], dim=-1).detach().cpu().numpy() result = {} for i, aspect in enumerate(cfg.ASPECTS): cls = int(np.argmax(probs[i])) result[aspect] = { "label": cfg.LABEL_NAMES[cls], "confidence": float(probs[i, cls]), "class_probs": { cfg.LABEL_NAMES[j]: float(probs[i, j]) for j in range(len(cfg.LABEL_NAMES)) }, } return result def _attention_insights(attn_payload: Dict) -> Dict: insights = {} if "meta_attention" in attn_payload: insights["top_meta_source"] = _top_attention_item(attn_payload["meta_attention"]) if "meta_attention_by_aspect" in attn_payload: insights["top_meta_source_by_aspect"] = { aspect: _top_attention_item(weights) for aspect, weights in attn_payload["meta_attention_by_aspect"].items() } return insights def _format_attention(out: Dict) -> Dict: """Format legacy 1D attention or new aspect-specific 2D attention.""" if "meta_attn_weights" not in out or out["meta_attn_weights"] is None: return {} aspect_attn = out["meta_attn_weights"].detach().cpu().numpy()[0] names = _attention_names(out, aspect_attn) if aspect_attn.ndim == 1: return {"meta_attention": _attention_dict(names, aspect_attn)} result = { "meta_attention_by_aspect": { aspect: _attention_dict(names, aspect_attn[i]) for i, aspect in enumerate(cfg.ASPECTS) } } if "global_meta_attn_weights" in out and out["global_meta_attn_weights"] is not None: global_attn = out["global_meta_attn_weights"].detach().cpu().numpy()[0] result["meta_attention"] = _attention_dict(names, global_attn) else: result["meta_attention"] = _attention_dict(names, aspect_attn.mean(axis=0)) return result def _slice_model_out(out: Dict, start: int, end: int) -> Dict: sliced = {} for k, v in out.items(): sliced[k] = v[start:end] if torch.is_tensor(v) else v return sliced class AspectPredictor: """Load the Proposed model and expose one-shot and batch predictions.""" def __init__(self, checkpoint_dir=None, device=None): self.model, self.tokenizer, self.meta_encoder, self.device = load_meta_acsa( checkpoint_dir=checkpoint_dir, device=device, ) def predict(self, review_text: str, product_meta: Dict, return_attention: bool = True) -> Dict: enc = self.tokenizer( review_text, max_length=cfg.MAX_LENGTH, truncation=True, padding="max_length", return_tensors="pt", ) input_ids = enc["input_ids"].to(self.device) attn_mask = enc["attention_mask"].to(self.device) meta_df = _meta_dict_to_df(product_meta) meta_vec = torch.from_numpy(self.meta_encoder.transform(meta_df)).float().to(self.device) with torch.no_grad(): out = self.model(input_ids, attn_mask, meta_vec) preds = out["logits"][0].argmax(dim=-1).cpu().numpy() result = { "aspects": format_aspect_summary(preds), "aspect_details": _format_confidence(out, 0), "metadata_summary": _meta_summary(product_meta), } if "overall_logits" in out and out["overall_logits"] is not None: overall_probs = torch.softmax(out["overall_logits"][0], dim=-1).detach().cpu().numpy() overall_cls = int(np.argmax(overall_probs)) result["overall"] = { "label": cfg.OVERALL_LABEL_NAMES[overall_cls], "confidence": float(overall_probs[overall_cls]), "class_probs": { cfg.OVERALL_LABEL_NAMES[j]: float(overall_probs[j]) for j in range(len(cfg.OVERALL_LABEL_NAMES)) }, } if return_attention: attn_payload = _format_attention(out) result.update(attn_payload) result.update(_attention_insights(attn_payload)) return result def predict_batch(self, reviews: List[Dict], batch_size: int = 16) -> List[Dict]: """Each item is {'review_text': str, 'product_meta': {...}}.""" results = [] for i in range(0, len(reviews), batch_size): chunk = reviews[i:i + batch_size] enc = self.tokenizer( [r["review_text"] for r in chunk], max_length=cfg.MAX_LENGTH, truncation=True, padding="max_length", return_tensors="pt", ) input_ids = enc["input_ids"].to(self.device) attn_mask = enc["attention_mask"].to(self.device) meta_df = pd.concat( [_meta_dict_to_df(r["product_meta"]) for r in chunk], ignore_index=True, ) meta_vec = torch.from_numpy(self.meta_encoder.transform(meta_df)).float().to(self.device) with torch.no_grad(): out = self.model(input_ids, attn_mask, meta_vec) preds = out["logits"].argmax(dim=-1).cpu().numpy() for j in range(len(chunk)): item = { "aspects": format_aspect_summary(preds[j]), "aspect_details": _format_confidence(out, j), "metadata_summary": _meta_summary(chunk[j]["product_meta"]), } if "overall_logits" in out and out["overall_logits"] is not None: overall_probs = torch.softmax(out["overall_logits"][j], dim=-1).detach().cpu().numpy() overall_cls = int(np.argmax(overall_probs)) item["overall"] = { "label": cfg.OVERALL_LABEL_NAMES[overall_cls], "confidence": float(overall_probs[overall_cls]), "class_probs": { cfg.OVERALL_LABEL_NAMES[k]: float(overall_probs[k]) for k in range(len(cfg.OVERALL_LABEL_NAMES)) }, } attn_payload = _format_attention(_slice_model_out(out, j, j + 1)) item.update(attn_payload) item.update(_attention_insights(attn_payload)) results.append(item) return results