muffin2006/document-classification-env
1
1"""2Example usage of the Document Classification OpenEnv Environment3Demonstrates various ways to interact with the environment4"""5 6import numpy as np7from environment import DocumentClassificationEnv8from grading import BaselineAgent, AgentGrader9 10 11def example_basic_usage():12 """Example 1: Basic environment usage"""13 print("=" * 70)14 print("Example 1: Basic Environment Usage")15 print("=" * 70)16 17 # Create environment18 env = DocumentClassificationEnv(task_difficulty="easy")19 print(f"✓ Created environment: {env.task_difficulty}")20 21 # Reset and get initial observation22 obs, info = env.reset()23 print(f"✓ Reset environment")24 print(f" Initial document: {obs['document_id']}")25 print(f" Content preview: {obs['content'][:60]}...")26 27 # Take a few steps28 total_reward = 029 for step in range(5):30 action = env.action_space.sample() # Random action31 obs, reward, done, truncated, info = env.step(action)32 total_reward += reward33 34 print(f"\nStep {step + 1}:")35 print(f" Action: {env.CATEGORY_MAPS['easy'][action]}")36 print(f" Reward: {reward:.3f}")37 print(f" Episode Accuracy: {info['episode_accuracy']:.3f}")38 39 if done:40 print(f" Episode Complete!")41 break42 43 print(f"\nTotal Reward: {total_reward:.3f}")44 45 46def example_state_inspection():47 """Example 2: Inspecting environment state"""48 print("\n" + "=" * 70)49 print("Example 2: Environment State Inspection")50 print("=" * 70)51 52 env = DocumentClassificationEnv(task_difficulty="medium")53 obs, _ = env.reset()54 55 # Take a few steps56 for _ in range(3):57 action = env.action_space.sample()58 obs, reward, done, _, info = env.step(action)59 60 # Get and inspect state61 state = env.state()62 print("Environment State:")63 print(f" Current document index: {state['current_document_index']}")64 print(f" Total documents: {state['total_documents']}")65 print(f" Episode reward: {state['episode_reward_total']:.3f}")66 print(f" Accuracy so far: {state['episode_accuracy']:.3f}")67 print(f" Avg processing time: {state['average_processing_time_ms']:.2f}ms")68 69 print(f"\nCurrent Observation:")70 print(f" Document ID: {state['current_observation']['document_id']}")71 print(f" Word count: {state['current_observation']['word_count'][0]}")72 print(f" Has urgency: {bool(state['current_observation']['has_urgency_markers'][0])}")73 74 75def example_baseline_agent():76 """Example 3: Using the baseline agent"""77 print("\n" + "=" * 70)78 print("Example 3: Baseline Agent Performance")79 print("=" * 70)80 81 for difficulty in ["easy", "medium", "hard"]:82 print(f"\n{difficulty.upper()} Task:")83 84 env = DocumentClassificationEnv(task_difficulty=difficulty)85 agent = BaselineAgent(difficulty)86 87 obs, _ = env.reset()88 89 total_reward = 090 correct = 091 total = 092 93 for step in range(20): # Run 20 steps94 action = agent.decide(obs)95 obs, reward, done, truncated, info = env.step(action)96 97 total_reward += reward98 if info['is_correct']:99 correct += 1100 total += 1101 102 if done:103 break104 105 accuracy = correct / total if total > 0 else 0106 print(f" Accuracy: {accuracy:.3f} ({correct}/{total})")107 print(f" Total Reward: {total_reward:.3f}")108 print(f" Avg Reward: {total_reward / total:.3f}")109 110 111def example_agent_grading():112 """Example 4: Grading an agent"""113 print("\n" + "=" * 70)114 print("Example 4: Agent Grading System")115 print("=" * 70)116 117 difficulty = "easy"118 119 # Create grader120 grader = AgentGrader(difficulty)121 122 # Create a simple custom agent123 def custom_agent(obs):124 """A simple custom agent that prefers 'Support' category"""125 # 70% of the time, choose Support (index 2 in easy)126 if np.random.random() < 0.7:127 return 2 # Support128 else:129 return np.random.randint(0, 5) # Random action130 131 # Grade the custom agent132 score, metrics = grader.grade_agent(custom_agent, verbose=True)133 134 print(f"Final Score: {score:.4f}")135 136 137def example_difficulty_comparison():138 """Example 5: Comparing difficulty levels"""139 print("\n" + "=" * 70)140 print("Example 5: Difficulty Level Comparison")141 print("=" * 70)142 143 difficulties = ["easy", "medium", "hard"]144 145 print(f"\n{'Difficulty':<10} {'Categories':<12} {'Documents':<12} {'Time Limit':<15} {'Features':<10}")146 print("-" * 60)147 148 for diff in difficulties:149 env = DocumentClassificationEnv(task_difficulty=diff)150 config = env.task_config151 152 time_limit = config.get("time_limit", "None")153 time_str = f"{time_limit}s" if time_limit else "None"154 155 print(f"{diff:<10} {config['num_categories']:<12} {config['num_documents']:<12} {time_str:<15} {config['feature_dim']:<10}")156 157 158def example_seed_reproducibility():159 """Example 6: Demonstrating reproducibility with seeds"""160 print("\n" + "=" * 70)161 print("Example 6: Seed-based Reproducibility")162 print("=" * 70)163 164 # Create two environments with the same seed165 env1 = DocumentClassificationEnv(task_difficulty="easy", seed=42)166 env2 = DocumentClassificationEnv(task_difficulty="easy", seed=42)167 168 obs1, _ = env1.reset(seed=42)169 obs2, _ = env2.reset(seed=42)170 171 print("Comparing two environments with seed=42:")172 print(f" Env1 doc ID: {obs1['document_id']}")173 print(f" Env2 doc ID: {obs2['document_id']}")174 print(f" Match: {obs1['document_id'] == obs2['document_id']}")175 176 print(f"\n Env1 content: {obs1['content'][:40]}...")177 print(f" Env2 content: {obs2['content'][:40]}...")178 print(f" Match: {obs1['content'] == obs2['content']}")179 180 print(f"\n Feature vectors equal: {np.allclose(obs1['features'], obs2['features'])}")181 182 183def example_episode_summary():184 """Example 7: Getting episode summary"""185 print("\n" + "=" * 70)186 print("Example 7: Episode Summary")187 print("=" * 70)188 189 env = DocumentClassificationEnv(task_difficulty="easy")190 obs, _ = env.reset()191 192 # Run a full episode193 done = False194 while not done:195 action = env.action_space.sample()196 obs, reward, done, truncated, info = env.step(action)197 198 # Get summary from info199 summary = info.get("episode_summary", {})200 201 print("Episode Summary:")202 for key, value in summary.items():203 if isinstance(value, float):204 print(f" {key}: {value:.4f}")205 else:206 print(f" {key}: {value}")207 208 209def main():210 """Run all examples"""211 try:212 example_basic_usage()213 example_state_inspection()214 example_baseline_agent()215 example_agent_grading()216 example_difficulty_comparison()217 example_seed_reproducibility()218 example_episode_summary()219 220 print("\n" + "=" * 70)221 print("✓ All examples completed successfully!")222 print("=" * 70)223 224 except Exception as e:225 print(f"\n✗ Error running examples: {e}")226 import traceback227 traceback.print_exc()228 229 230if __name__ == "__main__":231 main()232 