Team Ai
Apppublic

muffin2006/document-classification-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
example_usage.py232 linesDownload Raw Back to root
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