CoolFace
Apppublic

Israelbliz/User-Modeling-Agent

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
test_task_a.py141 linesDownload Raw Back to scripts
1"""Quick end-to-end test of the Task A agent on real data.2 3Picks a user from the training set, picks one of their held-out test4reviews (which is real ground truth we know), generates a predicted5rating + review for that item, and prints both side by side.6 7Usage:8    python -m scripts.test_task_a9    python -m scripts.test_task_a --user <user_id>10    python -m scripts.test_task_a --naija11    python -m scripts.test_task_a --user <user_id> --naija12"""13from __future__ import annotations14 15import argparse16import logging17 18import pandas as pd19 20from core.config import settings21from core.persona import PersonaEngine22from task_a_user_modeling.agent import ImpersonationAgent, ItemInput23 24logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")25 26 27def main():28    ap = argparse.ArgumentParser()29    ap.add_argument("--user", type=str, default=None,30                    help="Specific user_id; else picks a cross-domain user with the most reviews")31    ap.add_argument("--naija", action="store_true",32                    help="Apply Nigerian English style transfer to the generated review")33    args = ap.parse_args()34 35    reviews_path = settings.processed_dir / "reviews.parquet"36    items_path = settings.processed_dir / "items.parquet"37    if not reviews_path.exists() or not items_path.exists():38        raise SystemExit("Run `python data/prepare_data.py` first.")39 40    reviews = pd.read_parquet(reviews_path)41    items = pd.read_parquet(items_path)42    train = reviews[reviews["split"] == "train"]43    test = reviews[reviews["split"] == "test"]44 45    # Pick a user46    if args.user:47        user_id = args.user48    else:49        cross_users = (train.groupby("user_id")50                            .agg(n=("rating", "size"), d=("domain", "nunique"))51                            .reset_index())52        cross_users = cross_users[cross_users["d"] >= 2]53        # Prefer users who also have test reviews54        users_with_test = set(test["user_id"])55        cross_users = cross_users[cross_users["user_id"].isin(users_with_test)]56        if cross_users.empty:57            raise SystemExit("No cross-domain user has test reviews. Try --user <id>")58        user_id = cross_users.nlargest(1, "n").iloc[0]["user_id"]59        print(f"Auto-selected cross-domain user: {user_id}\n")60 61    # Pick a test review for this user62    user_test = test[test["user_id"] == user_id]63    if user_test.empty:64        raise SystemExit(f"User {user_id} has no test reviews — try a different user.")65    test_review = user_test.iloc[0]66    target_item_id = test_review["parent_asin"]67 68    # Look up item metadata69    item_meta = items[items["parent_asin"] == target_item_id]70    if item_meta.empty:71        print(f"WARN: no item metadata for {target_item_id}; using review title only")72        item = ItemInput(73            parent_asin=target_item_id,74            title=str(test_review.get("title", "")),75            description="",76            categories="",77            domain=test_review["domain"],78        )79    else:80        meta = item_meta.iloc[0]81        item = ItemInput(82            parent_asin=target_item_id,83            title=str(meta.get("title", "")),84            description=str(meta.get("description", ""))[:1500],85            categories=str(meta.get("categories", "")),86            domain=test_review["domain"],87            average_rating=float(meta["average_rating"]) if pd.notna(meta.get("average_rating")) else None,88        )89 90    # Build persona (with LLM enrichment)91    print(f"Building persona for {user_id}...")92    engine = PersonaEngine()93    persona = engine.from_dataframe(user_id, train)94    persona = engine.enrich(persona)95 96    # Run the agent97    print(f"\nGenerating review for item: {item.title[:80]}...\n")98    agent = ImpersonationAgent()99    result = agent.run(persona, item, naija_mode=args.naija)100 101    # Print side-by-side comparison with ground truth102    print("=" * 70)103    print("PERSONA SUMMARY")104    print("=" * 70)105    print(f"User: {user_id}")106    print(f"Avg rating: {persona.avg_rating:.2f}  |  Tone: {persona.tone}")107    print(f"Voice: {persona.voice_one_liner}")108 109    print("\n" + "=" * 70)110    print("TARGET ITEM")111    print("=" * 70)112    print(f"Domain: {item.domain}")113    print(f"Title: {item.title}")114    if item.description:115        print(f"Description: {item.description[:300]}...")116 117    print("\n" + "=" * 70)118    print(f"AI-GENERATED PREDICTION  {'(Naija mode)' if args.naija else ''}")119    print("=" * 70)120    print(f"Rating:    {result.rating}★")121    print(f"Reasoning: {result.reasoning}")122    print(f"\nReview:\n{result.review}")123 124    print("\n" + "=" * 70)125    print("GROUND TRUTH (what the user actually wrote)")126    print("=" * 70)127    print(f"Rating:    {test_review['rating']}★")128    print(f"\nReview:\n{test_review['text']}")129 130    print("\n" + "=" * 70)131    print("DELTA")132    print("=" * 70)133    rating_delta = abs(result.rating - float(test_review["rating"]))134    print(f"Rating absolute error: {rating_delta:.1f} stars")135    print(f"Generated review length: {len(result.review.split())} words")136    print(f"Ground truth length:    {len(str(test_review['text']).split())} words")137 138 139if __name__ == "__main__":140    main()141