Blablablab/audio-classification
0
1"""2Command-line interface for the user simulator.3 4Usage:5 python -m potato.simulator --server http://localhost:8000 --users 106 python -m potato.simulator --config simulator-config.yaml --server http://localhost:80007"""8 9import argparse10import logging11import sys12import os13 14from .config import (15 SimulatorConfig,16 TimingConfig,17 LLMStrategyConfig,18 BiasedStrategyConfig,19 AnnotationStrategyType,20)21from .simulator_manager import SimulatorManager22 23 24def setup_logging(verbose: bool = False) -> None:25 """Configure logging for the CLI.26 27 Args:28 verbose: If True, enable debug logging29 """30 level = logging.DEBUG if verbose else logging.INFO31 logging.basicConfig(32 level=level,33 format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",34 datefmt="%H:%M:%S",35 )36 37 # Suppress noisy loggers38 if not verbose:39 logging.getLogger("urllib3").setLevel(logging.WARNING)40 logging.getLogger("requests").setLevel(logging.WARNING)41 42 43def parse_args() -> argparse.Namespace:44 """Parse command-line arguments.45 46 Returns:47 Parsed arguments namespace48 """49 parser = argparse.ArgumentParser(50 description="User Simulator for Potato Annotation Platform",51 formatter_class=argparse.RawDescriptionHelpFormatter,52 epilog="""53Examples:54 # Basic random simulation55 python -m potato.simulator --server http://localhost:8000 --users 1056 57 # With configuration file58 python -m potato.simulator --config simulator.yaml --server http://localhost:800059 60 # LLM-powered simulation with Ollama61 python -m potato.simulator --server http://localhost:8000 --users 5 \\62 --strategy llm --llm-endpoint ollama --llm-model llama3.263 64 # Biased simulation65 python -m potato.simulator --server http://localhost:8000 --users 20 \\66 --strategy biased --bias-weights positive=0.6,negative=0.3,neutral=0.167 68 # Fast scalability test69 python -m potato.simulator --server http://localhost:8000 --users 100 \\70 --parallel 20 --max-annotations 5 --fast-mode71""",72 )73 74 # Required arguments75 parser.add_argument(76 "--server",77 "-s",78 required=True,79 help="Potato server URL (e.g., http://localhost:8000)",80 )81 82 # Configuration file (alternative to CLI args)83 parser.add_argument(84 "--config",85 "-c",86 help="Path to YAML configuration file",87 )88 89 # User configuration90 parser.add_argument(91 "--users",92 "-u",93 type=int,94 default=10,95 help="Number of simulated users (default: 10)",96 )97 parser.add_argument(98 "--competence",99 help="Competence distribution as comma-separated key=value pairs "100 "(e.g., good=0.5,average=0.3,poor=0.2)",101 )102 103 # Strategy configuration104 parser.add_argument(105 "--strategy",106 choices=["random", "biased", "llm", "pattern", "gold_standard"],107 default="random",108 help="Annotation strategy (default: random)",109 )110 111 # LLM configuration112 parser.add_argument(113 "--llm-endpoint",114 choices=["openai", "anthropic", "ollama", "gemini", "huggingface", "vllm"],115 help="LLM endpoint type (for --strategy llm)",116 )117 parser.add_argument(118 "--llm-model",119 help="LLM model name (for --strategy llm)",120 )121 parser.add_argument(122 "--llm-api-key",123 help="LLM API key (or set via environment variable)",124 )125 parser.add_argument(126 "--llm-base-url",127 help="LLM base URL (for local endpoints like Ollama)",128 )129 130 # Biased strategy configuration131 parser.add_argument(132 "--bias-weights",133 help="Label bias weights as comma-separated key=value pairs "134 "(e.g., positive=0.6,negative=0.3,neutral=0.1)",135 )136 137 # Execution configuration138 parser.add_argument(139 "--parallel",140 "-p",141 type=int,142 default=5,143 help="Maximum concurrent users (default: 5)",144 )145 parser.add_argument(146 "--max-annotations",147 "-m",148 type=int,149 help="Maximum annotations per user (default: unlimited)",150 )151 parser.add_argument(152 "--sequential",153 action="store_true",154 help="Run users sequentially instead of in parallel",155 )156 157 # Timing configuration158 parser.add_argument(159 "--fast-mode",160 action="store_true",161 help="Disable waiting between annotations (for testing)",162 )163 parser.add_argument(164 "--timing-min",165 type=float,166 default=2.0,167 help="Minimum annotation time in seconds (default: 2.0)",168 )169 parser.add_argument(170 "--timing-max",171 type=float,172 default=30.0,173 help="Maximum annotation time in seconds (default: 30.0)",174 )175 176 # Quality control testing177 parser.add_argument(178 "--attention-fail-rate",179 type=float,180 default=0.0,181 help="Rate at which to fail attention checks (0-1, default: 0)",182 )183 parser.add_argument(184 "--fast-response-rate",185 type=float,186 default=0.0,187 help="Rate of suspiciously fast responses (0-1, default: 0)",188 )189 190 # Gold standards191 parser.add_argument(192 "--gold-file",193 help="Path to JSON file with gold standard answers",194 )195 196 # Output configuration197 parser.add_argument(198 "--output-dir",199 "-o",200 default="simulator_output",201 help="Output directory for results (default: simulator_output)",202 )203 parser.add_argument(204 "--no-export",205 action="store_true",206 help="Don't export results to files",207 )208 209 # Other options210 parser.add_argument(211 "--verbose",212 "-v",213 action="store_true",214 help="Enable verbose logging",215 )216 217 return parser.parse_args()218 219 220def parse_key_value_pairs(s: str) -> dict:221 """Parse comma-separated key=value pairs.222 223 Args:224 s: String like "key1=val1,key2=val2"225 226 Returns:227 Dictionary of parsed pairs228 """229 result = {}230 if not s:231 return result232 233 for pair in s.split(","):234 if "=" in pair:235 key, value = pair.split("=", 1)236 # Try to convert to float237 try:238 result[key.strip()] = float(value.strip())239 except ValueError:240 result[key.strip()] = value.strip()241 242 return result243 244 245def build_config_from_args(args: argparse.Namespace) -> SimulatorConfig:246 """Build SimulatorConfig from CLI arguments.247 248 Args:249 args: Parsed arguments250 251 Returns:252 SimulatorConfig instance253 """254 # If config file provided, use it as base255 if args.config:256 config = SimulatorConfig.from_yaml(args.config)257 else:258 config = SimulatorConfig()259 260 # Override with CLI arguments261 config.user_count = args.users262 config.parallel_users = args.parallel263 config.simulate_wait = not args.fast_mode264 config.attention_check_fail_rate = args.attention_fail_rate265 config.respond_fast_rate = args.fast_response_rate266 config.output_dir = args.output_dir267 268 # Parse competence distribution269 if args.competence:270 config.competence_distribution = parse_key_value_pairs(args.competence)271 272 # Parse strategy273 try:274 config.strategy = AnnotationStrategyType(args.strategy)275 except ValueError:276 config.strategy = AnnotationStrategyType.RANDOM277 278 # LLM configuration279 if args.strategy == "llm" and args.llm_endpoint:280 api_key = args.llm_api_key281 if not api_key:282 # Try common environment variables283 env_vars = {284 "openai": "OPENAI_API_KEY",285 "anthropic": "ANTHROPIC_API_KEY",286 "huggingface": "HF_TOKEN",287 "gemini": "GOOGLE_API_KEY",288 }289 env_var = env_vars.get(args.llm_endpoint)290 if env_var:291 api_key = os.environ.get(env_var)292 293 config.llm_config = LLMStrategyConfig(294 endpoint_type=args.llm_endpoint,295 model=args.llm_model,296 api_key=api_key,297 base_url=args.llm_base_url,298 )299 300 # Biased configuration301 if args.strategy == "biased" and args.bias_weights:302 config.biased_config = BiasedStrategyConfig(303 label_weights=parse_key_value_pairs(args.bias_weights)304 )305 306 # Timing configuration307 config.timing = TimingConfig(308 annotation_time_min=args.timing_min,309 annotation_time_max=args.timing_max,310 )311 312 # Gold standards file313 if args.gold_file:314 config.gold_standard_file = args.gold_file315 316 return config317 318 319def main() -> int:320 """Main entry point for CLI.321 322 Returns:323 Exit code (0 for success, 1 for error)324 """325 args = parse_args()326 setup_logging(args.verbose)327 328 logger = logging.getLogger(__name__)329 330 try:331 # Build configuration332 config = build_config_from_args(args)333 334 logger.info(f"Starting simulator with {config.user_count} users")335 logger.info(f"Server: {args.server}")336 logger.info(f"Strategy: {config.strategy.value}")337 338 # Create manager339 manager = SimulatorManager(config, args.server)340 341 # Run simulation342 if args.sequential:343 results = manager.run_sequential(args.max_annotations)344 else:345 results = manager.run_parallel(args.max_annotations)346 347 # Print summary348 manager.print_summary()349 350 # Export results351 if not args.no_export:352 manager.export_results()353 354 return 0355 356 except KeyboardInterrupt:357 logger.info("Simulation interrupted by user")358 return 1359 360 except Exception as e:361 logger.error(f"Simulation failed: {e}")362 if args.verbose:363 import traceback364 traceback.print_exc()365 return 1366 367 368if __name__ == "__main__":369 sys.exit(main())370 