Team Ai
Apppublic

TanujInsane/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
grading.py269 linesDownload Raw Back to root
1"""2Grade agents on document classification.3Uses hierarchical partial credit - related categories get partial score.4Handles tool actions (metadata requests, escalation) gracefully.5"""6 7import numpy as np8import json9from typing import Dict, List, Callable10from environment import DocumentClassificationEnv11from constants import CATEGORY_HIERARCHY12 13 14# Hierarchy map - categories that are "close" get partial credit15# CATEGORY_HIERARCHY is now in constants.py16 17PARTIAL_CREDIT = 0.4  # Score for picking a related category18TASK_ID_MAP = {19    "easy": 1,20    "medium": 2,21    "hard": 3,22}23 24 25def hierarchical_accuracy(true_label: str, pred_label: str) -> float:26    """27    Returns:28      1.0 if exact match29      0.4 if same parent category (e.g. Billing vs Billing-Dispute)30      0.0 otherwise (including TOOL_METADATA and ESCALATED)31    """32    if pred_label in ("TOOL_METADATA", "ESCALATED"):33        return 0.034    if true_label == pred_label:35        return 1.036    related = CATEGORY_HIERARCHY.get(true_label, [])37    if pred_label in related:38        return PARTIAL_CREDIT39    return 0.040 41 42class AgentGrader:43    """Grade an agent on a specific task difficulty.44 45    Handles the full action space including tool actions:46    - Standard classification actions (0..N-1)47    - Request Metadata tool (action N)48    - Escalate to Human (action N+1)49    """50 51    def __init__(self, task_difficulty: str, num_episodes: int = 3):52        self.task_difficulty = task_difficulty53        self.num_episodes = num_episodes54        self.env = DocumentClassificationEnv(task_difficulty=task_difficulty, seed=42)55 56    def grade(self, agent_fn: Callable) -> Dict:57        """58        Run agent through multiple episodes and return score 0.0-1.0.59 60        Args:61            agent_fn: callable(observation) -> action (int)62 63        Returns:64            dict with score and detailed metrics65        """66        all_scores = []67        all_accuracies = []68        all_partial_credits = []69        all_tool_usage_rates = []70        all_escalation_rates = []71        all_agent_times = []72        73        all_ep_true = []74        all_ep_pred = []75 76        import time77 78        for ep in range(self.num_episodes):79            obs, _ = self.env.reset(seed=ep * 777)80            ep_true = []81            ep_pred = []82            ep_rewards = []83            tool_requests = 084            escalations = 085            classifications = 086            terminated = False87 88            while not terminated:89                t0 = time.time()90                action = agent_fn(obs)91                t1 = time.time()92                all_agent_times.append((t1 - t0) * 1000)93                obs, reward, terminated, _, info = self.env.step(action)94 95                pred_cat = info.get("predicted_category", "")96                true_cat = info.get("true_category", "")97 98                if pred_cat == "TOOL_METADATA":99                    tool_requests += 1100                    # Don't count tool usage as a classification attempt101                    continue102                elif pred_cat == "ESCALATED":103                    escalations += 1104                    ep_true.append(true_cat)105                    ep_pred.append(pred_cat)106                    ep_rewards.append(reward)107                    classifications += 1108                else:109                    ep_true.append(true_cat)110                    ep_pred.append(pred_cat)111                    ep_rewards.append(reward)112                    classifications += 1113 114            if len(ep_true) == 0:115                all_accuracies.append(0.0)116                all_scores.append(0.0)117                all_partial_credits.append(0.0)118                all_tool_usage_rates.append(0.0)119                all_escalation_rates.append(0.0)120                continue121 122            # Exact accuracy (only count real classifications, not escalations)123            real_classifications = [(t, p) for t, p in zip(ep_true, ep_pred) if p != "ESCALATED"]124            if real_classifications:125                exact = sum(t == p for t, p in real_classifications) / len(real_classifications)126            else:127                exact = 0.0128 129            # Hierarchical accuracy (with partial credit)130            hier_scores = [131                hierarchical_accuracy(t, p) for t, p in zip(ep_true, ep_pred)132            ]133            hier_acc = np.mean(hier_scores) if hier_scores else 0.0134 135            # Partial credit ratio136            partial = sum(1 for s in hier_scores if 0 < s < 1) / len(hier_scores) if hier_scores else 0.0137 138            # Tool usage metrics139            total_actions = classifications + tool_requests140            tool_rate = tool_requests / total_actions if total_actions > 0 else 0.0141            esc_rate = escalations / classifications if classifications > 0 else 0.0142 143            all_accuracies.append(exact)144            all_scores.append(hier_acc)145            all_partial_credits.append(partial)146            all_tool_usage_rates.append(tool_rate)147            all_escalation_rates.append(esc_rate)148            all_ep_true.append(ep_true)149            all_ep_pred.append(ep_pred)150 151        avg_accuracy = float(np.mean(all_accuracies))152        avg_hier = float(np.mean(all_scores))153        avg_partial = float(np.mean(all_partial_credits))154        avg_tool_rate = float(np.mean(all_tool_usage_rates))155        avg_esc_rate = float(np.mean(all_escalation_rates))156        avg_processing_time = float(np.mean(all_agent_times)) if all_agent_times else 100.0157 158        # Generate Confusion Matrix159        from collections import defaultdict160        cm = defaultdict(lambda: defaultdict(int))161        for ep_t, ep_p in zip(all_ep_true, all_ep_pred):162            for t_label, p_label in zip(ep_t, ep_p):163                cm[t_label][p_label] += 1164        165        cm_dict = {k: dict(v) for k, v in cm.items()}166        try:167            out_path = f"/tmp/confusion_matrix_{self.task_difficulty}.json"168            with open(out_path, "w") as f:169                json.dump(cm_dict, f, indent=2)170        except Exception:171            pass172 173        # Compute final score based on difficulty174        metrics = {175            "accuracy": avg_accuracy,176            "hierarchical_accuracy": avg_hier,177            "partial_credit_rate": avg_partial,178            "average_processing_time_ms": avg_processing_time,179            "tool_usage_rate": avg_tool_rate,180            "escalation_rate": avg_esc_rate,181        }182 183        score_raw = self._compute_score(metrics)184        score_clamped = float(min(1.0, max(0.0, score_raw)))185 186        import uuid187        task_id = TASK_ID_MAP.get(self.task_difficulty, 0)188        episode_id = str(uuid.uuid4())189        scenario_id = f"document_classification_{self.task_difficulty}"190 191        breakdown = {192            "ACCURACY": avg_accuracy,193            "HIERARCHICAL_ACCURACY": avg_hier,194            "PARTIAL_CREDIT_RATE": avg_partial,195            "AVERAGE_PROCESSING_TIME_MS": avg_processing_time,196            "TOOL_USAGE_RATE": avg_tool_rate,197            "ESCALATION_RATE": avg_esc_rate,198        }199 200        return {201            "score": score_clamped,202            "task_id": task_id,203            "episode_id": episode_id,204            "scenario_id": scenario_id,205            "breakdown": breakdown,206            "grader_version": "1.0.0",207            "accuracy": avg_accuracy,208            "hierarchical_accuracy": avg_hier,209            "partial_credit_rate": avg_partial,210            "tool_usage_rate": avg_tool_rate,211            "escalation_rate": avg_esc_rate,212            "task_difficulty": self.task_difficulty,213            "num_episodes": self.num_episodes,214        }215 216    def _compute_score(self, metrics: Dict) -> float:217        accuracy = metrics["accuracy"]218        hier_accuracy = metrics["hierarchical_accuracy"]219        avg_time = metrics.get("average_processing_time_ms", 100)220 221        if self.task_difficulty == "easy":222            time_bonus = 0.1 if avg_time < 100 else 0.0223            # Easy: mostly exact accuracy224            score = 0.85 * accuracy + 0.05 * hier_accuracy + 0.1 * time_bonus225 226        elif self.task_difficulty == "medium":227            time_bonus = 0.0228            if avg_time < 200:229                time_bonus = 0.15230            elif avg_time < 500:231                time_bonus = 0.10232            # Medium: mix of exact + hierarchical233            score = 0.70 * accuracy + 0.15 * hier_accuracy + 0.15 * time_bonus234 235        else:  # hard236            time_bonus = 0.0237            if avg_time < 100:238                time_bonus = 0.25239            elif avg_time < 300:240                time_bonus = 0.15241            elif avg_time < 500:242                time_bonus = 0.05243            # Hard: hierarchical accuracy matters more (22 categories)244            score = 0.50 * accuracy + 0.30 * hier_accuracy + 0.20 * time_bonus245 246        return min(1.0, max(0.0, score))247 248 249def grade_baseline(task_difficulty: str) -> float:250    """Quick helper to grade the TF-IDF baseline and return score."""251    from baseline_inference import load_or_train252 253    model = load_or_train(task_difficulty)254    grader = AgentGrader(task_difficulty)255 256    def agent_fn(obs):257        pred = model.predict([obs["content"]])[0]258        return int(pred)259 260    result = grader.grade(agent_fn)261    return result["score"]262 263 264if __name__ == "__main__":265    print("Grading baseline agent with hierarchical partial credit...\n")266    for difficulty in ["easy", "medium", "hard"]:267        score = grade_baseline(difficulty)268        print(f"  {difficulty.upper()} Score: {score:.4f}")269