TanujInsane/document-classification-env
3
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 