Team Ai
Apppublic

muffin2006/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
grading.py180 linesDownload Raw Back to root
1"""2Grade agents on document classification.3Uses hierarchical partial credit - related categories get partial score.4"""5 6import numpy as np7from typing import Dict, List8from environment import DocumentClassificationEnv9 10 11# Hierarchy map - categories that are "close" get partial credit12CATEGORY_HIERARCHY = {13    "Billing": ["Billing-Dispute", "Billing-Refund"],14    "Billing-Dispute": ["Billing", "Billing-Refund"],15    "Billing-Refund": ["Billing", "Billing-Dispute"],16    "Support": ["Support-Urgent", "Support-Normal"],17    "Support-Urgent": ["Support", "Support-Normal"],18    "Support-Normal": ["Support", "Support-Urgent"],19    "Technical": ["Technical-Bug", "Technical-Feature"],20    "Technical-Bug": ["Technical", "Technical-Feature"],21    "Technical-Feature": ["Technical", "Technical-Bug"],22    "HR": ["HR-Payroll", "HR-Benefits", "HR-Complaint"],23    "HR-Payroll": ["HR", "HR-Benefits"],24    "HR-Benefits": ["HR", "HR-Payroll"],25    "HR-Complaint": ["HR"],26    "Legal": ["Legal-Contract", "Legal-Compliance"],27    "Legal-Contract": ["Legal", "Legal-Compliance"],28    "Legal-Compliance": ["Legal", "Legal-Contract"],29    "Executive": ["Executive-Strategic"],30    "Executive-Strategic": ["Executive"],31}32 33PARTIAL_CREDIT = 0.4  # Score for picking a related category34 35 36def hierarchical_accuracy(true_label: str, pred_label: str) -> float:37    """38    Returns:39      1.0 if exact match40      0.4 if same parent category (e.g. Billing vs Billing-Dispute)41      0.0 otherwise42    """43    if true_label == pred_label:44        return 1.045    related = CATEGORY_HIERARCHY.get(true_label, [])46    if pred_label in related:47        return PARTIAL_CREDIT48    return 0.049 50 51class AgentGrader:52    """Grade an agent on a specific task difficulty"""53 54    def __init__(self, task_difficulty: str, num_episodes: int = 3):55        self.task_difficulty = task_difficulty56        self.num_episodes = num_episodes57        self.env = DocumentClassificationEnv(task_difficulty=task_difficulty, seed=42)58 59    def grade(self, agent_fn) -> Dict:60        """61        Run agent through multiple episodes and return score 0.0-1.0.62 63        Args:64            agent_fn: callable(observation) -> action (int)65 66        Returns:67            dict with score and detailed metrics68        """69        all_scores = []70        all_accuracies = []71        all_partial_credits = []72 73        for ep in range(self.num_episodes):74            obs, _ = self.env.reset(seed=ep * 777)75            ep_true = []76            ep_pred = []77            ep_rewards = []78            terminated = False79 80            while not terminated:81                action = agent_fn(obs)82                obs, reward, terminated, _, info = self.env.step(action)83                ep_true.append(info.get("true_category", ""))84                pred_cat = info.get("predicted_category", "")85                ep_pred.append(pred_cat)86                ep_rewards.append(reward)87 88            # Exact accuracy89            exact = sum(t == p for t, p in zip(ep_true, ep_pred)) / len(ep_true)90 91            # Hierarchical accuracy (with partial credit)92            hier_scores = [93                hierarchical_accuracy(t, p) for t, p in zip(ep_true, ep_pred)94            ]95            hier_acc = np.mean(hier_scores)96 97            # Partial credit ratio98            partial = sum(1 for s in hier_scores if 0 < s < 1) / len(hier_scores)99 100            all_accuracies.append(exact)101            all_scores.append(hier_acc)102            all_partial_credits.append(partial)103 104        avg_accuracy = np.mean(all_accuracies)105        avg_hier = np.mean(all_scores)106        avg_partial = np.mean(all_partial_credits)107 108        # Compute final score based on difficulty109        metrics = {110            "accuracy": avg_accuracy,111            "hierarchical_accuracy": avg_hier,112            "partial_credit_rate": avg_partial,113            "average_processing_time_ms": 50,  # placeholder114        }115 116        score = self._compute_score(metrics)117 118        return {119            "score": score,120            "accuracy": avg_accuracy,121            "hierarchical_accuracy": avg_hier,122            "partial_credit_rate": avg_partial,123            "task_difficulty": self.task_difficulty,124            "num_episodes": self.num_episodes,125        }126 127    def _compute_score(self, metrics: Dict) -> float:128        accuracy = metrics["accuracy"]129        hier_accuracy = metrics["hierarchical_accuracy"]130        avg_time = metrics.get("average_processing_time_ms", 100)131 132        if self.task_difficulty == "easy":133            time_bonus = 0.1 if avg_time < 100 else 0.0134            # Easy: mostly exact accuracy135            score = 0.85 * accuracy + 0.05 * hier_accuracy + 0.1 * time_bonus136 137        elif self.task_difficulty == "medium":138            time_bonus = 0.0139            if avg_time < 200:140                time_bonus = 0.15141            elif avg_time < 500:142                time_bonus = 0.10143            # Medium: mix of exact + hierarchical144            score = 0.70 * accuracy + 0.15 * hier_accuracy + 0.15 * time_bonus145 146        else:  # hard / extreme147            time_bonus = 0.0148            if avg_time < 100:149                time_bonus = 0.25150            elif avg_time < 300:151                time_bonus = 0.15152            elif avg_time < 500:153                time_bonus = 0.05154            # Hard: hierarchical accuracy matters more (22+ categories)155            score = 0.50 * accuracy + 0.30 * hier_accuracy + 0.20 * time_bonus156 157        return min(1.0, max(0.0, score))158 159 160def grade_baseline(task_difficulty: str) -> float:161    """Quick helper to grade and return score"""162    from baseline_inference import load_or_train163 164    model = load_or_train(task_difficulty)165    grader = AgentGrader(task_difficulty)166 167    def agent_fn(obs):168        pred = model.predict([obs["content"]])[0]169        return int(pred)170 171    result = grader.grade(agent_fn)172    return result["score"]173 174 175if __name__ == "__main__":176    print("Grading baseline agent with hierarchical partial credit...\n")177    for difficulty in ["easy", "medium", "hard"]:178        score = grade_baseline(difficulty)179        print(f"  {difficulty.upper()} Score: {score:.4f}")180