CoolFace
Apppublic

Blablablab/audio-classification

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes
simulator_manager.py399 linesDownload Raw Back to simulator
1"""2Simulator manager for orchestrating multiple simulated users.3 4This module provides the SimulatorManager class that manages multiple5SimulatedUser instances, handling parallel execution and result aggregation.6"""7 8import json9import logging10import random11import time12from concurrent.futures import ThreadPoolExecutor, as_completed13from typing import Dict, List, Any, Optional14 15from .config import (16    SimulatorConfig,17    UserConfig,18    CompetenceLevel,19    AnnotationStrategyType,20)21from .user_simulator import SimulatedUser, UserSimulationResult22from .reporting import SimulationReporter23 24logger = logging.getLogger(__name__)25 26 27class SimulatorManager:28    """Orchestrates multiple simulated users.29 30    The SimulatorManager handles:31    - Generating user configurations based on competence distribution32    - Running simulations in parallel or sequentially33    - Aggregating results across all users34    - Exporting results via SimulationReporter35    """36 37    def __init__(38        self,39        config: SimulatorConfig,40        server_url: str,41        gold_standards: Optional[Dict[str, Dict[str, Any]]] = None,42    ):43        """Initialize simulator manager.44 45        Args:46            config: Simulator configuration47            server_url: Base URL of the Potato server48            gold_standards: Optional gold standard answers keyed by instance_id49        """50        self.config = config51        self.server_url = server_url.rstrip("/")52        self.gold_standards = gold_standards or {}53 54        # Load gold standards from file if specified55        if config.gold_standard_file and not gold_standards:56            self.gold_standards = self._load_gold_standards(config.gold_standard_file)57 58        # Generate user configs if not provided59        self.user_configs = self._generate_user_configs()60 61        # Results tracking62        self.results: Dict[str, UserSimulationResult] = {}63        self.reporter = SimulationReporter(config.output_dir)64 65    def _load_gold_standards(self, filepath: str) -> Dict[str, Dict[str, Any]]:66        """Load gold standards from JSON file.67 68        Expected format:69        [70            {"id": "instance_001", "label_field": "value", ...},71            ...72        ]73 74        Args:75            filepath: Path to JSON file76 77        Returns:78            Gold standards dict keyed by instance ID79        """80        try:81            with open(filepath, "r") as f:82                items = json.load(f)83 84            gold_standards = {}85            for item in items:86                item_id = item.pop("id", None)87                if item_id:88                    gold_standards[item_id] = item89 90            logger.info(f"Loaded {len(gold_standards)} gold standards from {filepath}")91            return gold_standards92 93        except Exception as e:94            logger.warning(f"Failed to load gold standards from {filepath}: {e}")95            return {}96 97    def _generate_user_configs(self) -> List[UserConfig]:98        """Generate user configurations based on competence distribution.99 100        If explicit user configs are provided, uses those.101        Otherwise, generates based on user_count and competence_distribution.102 103        Returns:104            List of UserConfig instances105        """106        if self.config.users:107            return self.config.users108 109        users = []110 111        # Get competence distribution112        competence_levels = list(self.config.competence_distribution.keys())113        competence_weights = list(self.config.competence_distribution.values())114 115        # Normalize weights116        total_weight = sum(competence_weights)117        if total_weight > 0:118            competence_weights = [w / total_weight for w in competence_weights]119 120        for i in range(self.config.user_count):121            # Select competence level based on distribution122            competence_str = random.choices(123                competence_levels, weights=competence_weights, k=1124            )[0]125 126            try:127                competence = CompetenceLevel(competence_str)128            except ValueError:129                competence = CompetenceLevel.AVERAGE130 131            users.append(132                UserConfig(133                    user_id=f"sim_user_{i:04d}",134                    competence=competence,135                    strategy=self.config.strategy,136                    timing=self.config.timing,137                    llm_config=self.config.llm_config,138                    biased_config=self.config.biased_config,139                    agent_config=self.config.agent_config,140                )141            )142 143        logger.info(f"Generated {len(users)} user configurations")144        return users145 146    def run_single_user(147        self, user_config: UserConfig, max_annotations: Optional[int] = None148    ) -> UserSimulationResult:149        """Run simulation for a single user.150 151        Args:152            user_config: Configuration for the user153            max_annotations: Maximum annotations for this user154 155        Returns:156            UserSimulationResult with tracking data157        """158        user = SimulatedUser(159            user_config=user_config,160            server_url=self.server_url,161            gold_standards=self.gold_standards,162            simulate_wait=self.config.simulate_wait,163            attention_check_fail_rate=self.config.attention_check_fail_rate,164            respond_fast_rate=self.config.respond_fast_rate,165            interactive_config=self.config.interactive,166        )167 168        result = user.run_simulation(max_annotations)169        self.results[user_config.user_id] = result170 171        return result172 173    def run_parallel(174        self, max_annotations_per_user: Optional[int] = None175    ) -> Dict[str, UserSimulationResult]:176        """Run simulation for all users in parallel.177 178        Args:179            max_annotations_per_user: Maximum annotations per user180 181        Returns:182            Dict mapping user_id to UserSimulationResult183        """184        logger.info(185            f"Starting parallel simulation with {len(self.user_configs)} users "186            f"({self.config.parallel_users} concurrent)"187        )188 189        with ThreadPoolExecutor(max_workers=self.config.parallel_users) as executor:190            futures = {}191 192            for i, user_config in enumerate(self.user_configs):193                # Stagger user starts194                if i > 0 and self.config.delay_between_users > 0:195                    time.sleep(self.config.delay_between_users)196 197                future = executor.submit(198                    self.run_single_user, user_config, max_annotations_per_user199                )200                futures[future] = user_config.user_id201 202            # Wait for completion203            completed = 0204            for future in as_completed(futures):205                user_id = futures[future]206                completed += 1207                try:208                    result = future.result()209                    logger.info(210                        f"[{completed}/{len(futures)}] User {user_id} completed: "211                        f"{len(result.annotations)} annotations"212                    )213                except Exception as e:214                    logger.error(f"User {user_id} failed: {e}")215 216        logger.info(f"Parallel simulation completed: {len(self.results)} users")217        return self.results218 219    def run_sequential(220        self, max_annotations_per_user: Optional[int] = None221    ) -> Dict[str, UserSimulationResult]:222        """Run simulation for all users sequentially.223 224        Args:225            max_annotations_per_user: Maximum annotations per user226 227        Returns:228            Dict mapping user_id to UserSimulationResult229        """230        logger.info(231            f"Starting sequential simulation with {len(self.user_configs)} users"232        )233 234        for i, user_config in enumerate(self.user_configs):235            result = self.run_single_user(user_config, max_annotations_per_user)236            logger.info(237                f"[{i+1}/{len(self.user_configs)}] User {user_config.user_id} "238                f"completed: {len(result.annotations)} annotations"239            )240 241        logger.info(f"Sequential simulation completed: {len(self.results)} users")242        return self.results243 244    def get_summary(self) -> Dict[str, Any]:245        """Get summary statistics for all users.246 247        Returns:248            Summary dictionary with aggregate statistics249        """250        if not self.results:251            return {"error": "No results available"}252 253        total_annotations = sum(len(r.annotations) for r in self.results.values())254        total_time = sum(r.total_time for r in self.results.values())255 256        total_attention_passed = sum(257            r.attention_checks_passed for r in self.results.values()258        )259        total_attention_failed = sum(260            r.attention_checks_failed for r in self.results.values()261        )262        total_gold_correct = sum(263            r.gold_standard_correct for r in self.results.values()264        )265        total_gold_incorrect = sum(266            r.gold_standard_incorrect for r in self.results.values()267        )268 269        blocked_users = sum(1 for r in self.results.values() if r.was_blocked)270        users_with_errors = sum(1 for r in self.results.values() if r.errors)271 272        # Calculate response time statistics273        all_response_times = [274            record.response_time275            for result in self.results.values()276            for record in result.annotations277        ]278 279        response_time_stats = {}280        if all_response_times:281            response_time_stats = {282                "min": min(all_response_times),283                "max": max(all_response_times),284                "mean": sum(all_response_times) / len(all_response_times),285            }286 287        # Competence level distribution in results288        competence_distribution = {}289        for user_id in self.results:290            for config in self.user_configs:291                if config.user_id == user_id:292                    level = config.competence.value293                    competence_distribution[level] = (294                        competence_distribution.get(level, 0) + 1295                    )296                    break297 298        return {299            "user_count": len(self.results),300            "total_annotations": total_annotations,301            "total_time_seconds": total_time,302            "average_annotations_per_user": (303                total_annotations / len(self.results) if self.results else 0304            ),305            "average_time_per_user": (306                total_time / len(self.results) if self.results else 0307            ),308            "attention_checks": {309                "passed": total_attention_passed,310                "failed": total_attention_failed,311                "pass_rate": (312                    total_attention_passed313                    / (total_attention_passed + total_attention_failed)314                    if (total_attention_passed + total_attention_failed) > 0315                    else None316                ),317            },318            "gold_standards": {319                "correct": total_gold_correct,320                "incorrect": total_gold_incorrect,321                "accuracy": (322                    total_gold_correct / (total_gold_correct + total_gold_incorrect)323                    if (total_gold_correct + total_gold_incorrect) > 0324                    else None325                ),326            },327            "blocked_users": blocked_users,328            "users_with_errors": users_with_errors,329            "response_time_stats": response_time_stats,330            "competence_distribution": competence_distribution,331            "per_user": {332                user_id: {333                    "annotations": len(r.annotations),334                    "total_time": r.total_time,335                    "attention_passed": r.attention_checks_passed,336                    "attention_failed": r.attention_checks_failed,337                    "gold_correct": r.gold_standard_correct,338                    "gold_incorrect": r.gold_standard_incorrect,339                    "was_blocked": r.was_blocked,340                    "errors": len(r.errors),341                }342                for user_id, r in self.results.items()343            },344        }345 346    def export_results(self) -> str:347        """Export all results using the reporter.348 349        Returns:350            Path to the output directory351        """352        self.reporter.export_results(self.results, self.get_summary())353        return self.config.output_dir354 355    def print_summary(self) -> None:356        """Print a summary of results to stdout."""357        summary = self.get_summary()358 359        print("\n" + "=" * 60)360        print("SIMULATION SUMMARY")361        print("=" * 60)362 363        print(f"\nUsers: {summary['user_count']}")364        print(f"Total annotations: {summary['total_annotations']}")365        print(f"Total time: {summary['total_time_seconds']:.1f}s")366        print(367            f"Avg annotations/user: {summary['average_annotations_per_user']:.1f}"368        )369        print(f"Avg time/user: {summary['average_time_per_user']:.1f}s")370 371        ac = summary["attention_checks"]372        if ac["passed"] or ac["failed"]:373            print(f"\nAttention Checks:")374            print(f"  Passed: {ac['passed']}")375            print(f"  Failed: {ac['failed']}")376            if ac["pass_rate"] is not None:377                print(f"  Pass rate: {ac['pass_rate']:.1%}")378 379        gs = summary["gold_standards"]380        if gs["correct"] or gs["incorrect"]:381            print(f"\nGold Standards:")382            print(f"  Correct: {gs['correct']}")383            print(f"  Incorrect: {gs['incorrect']}")384            if gs["accuracy"] is not None:385                print(f"  Accuracy: {gs['accuracy']:.1%}")386 387        if summary["blocked_users"]:388            print(f"\nBlocked users: {summary['blocked_users']}")389 390        if summary["users_with_errors"]:391            print(f"Users with errors: {summary['users_with_errors']}")392 393        if summary["competence_distribution"]:394            print(f"\nCompetence distribution:")395            for level, count in summary["competence_distribution"].items():396                print(f"  {level}: {count}")397 398        print("\n" + "=" * 60)399