"""User-facing inference for the Proposed model. Typical use: from src.inference import AspectPredictor p = AspectPredictor() result = p.predict( review_text="The size runs small but the fabric feels great.", product_meta={ "features_text": "100% cotton tee, slim fit", "categories_text": "Clothing > Men > T-Shirts", "price": 19.99, "average_rating": 4.3, "rating_number": 217, }, ) # result == { # "aspects": {"SIZE": "Negative", "MATERIAL": "Positive", ...}, # "meta_attention": {"meta_chunk_1": 0.31, ...}, # } """ from typing import Dict, List, Optional import numpy as np import pandas as pd import torch from src import config as cfg from src.evaluator import load_meta_acsa, format_aspect_summary from src.explainer import META_TOKEN_NAMES # Keys this module expects in `product_meta` dict _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]) class AspectPredictor: """Lightweight wrapper that loads the Proposed model on init and exposes one-shot and batch predictions.""" def __init__(self, checkpoint_dir=None, device=None): if checkpoint_dir is None: from pathlib import Path root_path = Path(__file__).resolve().parent.parent checkpoint_dir = str(root_path / "checkpoints" / "meta_acsa") self.model, self.tokenizer, self.meta_encoder, self.device = load_meta_acsa( checkpoint_dir=checkpoint_dir, device=device, ) # ----------------------------------------------------------- single 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() # (num_aspects,) meta_attn = out["meta_attn_weights"][0].cpu().numpy() # (T,) result = {"aspects": format_aspect_summary(preds)} if return_attention: result["meta_attention"] = { name: float(w) for name, w in zip(META_TOKEN_NAMES, meta_attn) } return result # ----------------------------------------------------------- batch def predict_batch(self, reviews: List[Dict], batch_size: int = 16) -> List[Dict]: """Each `reviews` 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() # (B, num_aspects) meta_attn = out["meta_attn_weights"].cpu().numpy() for j in range(len(chunk)): results.append({ "aspects": format_aspect_summary(preds[j]), "meta_attention": {name: float(w) for name, w in zip(META_TOKEN_NAMES, meta_attn[j])}, }) return results