muffin2006/document-classification-env
1
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 