Blablablab/audio-classification
0
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 