(args: argparse.Namespace, config_file_list: list[str], use_caption=True)
| 280 | |
| 281 | @beartype |
| 282 | def test(args: argparse.Namespace, config_file_list: list[str], use_caption=True) -> None: |
| 283 | scores = {} |
| 284 | max_steps = args.max_steps |
| 285 | |
| 286 | early_stop_thresholds = { |
| 287 | "parsing_failure": args.parsing_failure_th, |
| 288 | "repeating_action": args.repeating_action_failure_th, |
| 289 | } |
| 290 | |
| 291 | if ( |
| 292 | args.observation_type |
| 293 | in [ |
| 294 | "accessibility_tree_with_captioner", |
| 295 | "image_som", |
| 296 | ] |
| 297 | and use_caption |
| 298 | ): |
| 299 | device = torch.device("cuda") if torch.cuda.is_available() else "cpu" |
| 300 | dtype = torch.float16 if torch.cuda.is_available() else torch.float32 |
| 301 | # caption_image_fn = image_utils.get_captioning_fn(device, dtype, args.captioning_model) |
| 302 | caption_image_fn = get_captioning_model(args.captioning_model) |
| 303 | else: |
| 304 | caption_image_fn = None |
| 305 | |
| 306 | # Load a (possibly different) captioning model for running VQA evals. |
| 307 | if caption_image_fn and args.eval_captioning_model == args.captioning_model: |
| 308 | eval_caption_image_fn = caption_image_fn |
| 309 | else: |
| 310 | eval_caption_image_fn = image_utils.get_captioning_fn( |
| 311 | args.eval_captioning_model_device, |
| 312 | torch.float16 |
| 313 | if (torch.cuda.is_available() and args.eval_captioning_model_device == "cuda") |
| 314 | else torch.float32, |
| 315 | args.eval_captioning_model, |
| 316 | ) |
| 317 | |
| 318 | agent = construct_agent( |
| 319 | args, |
| 320 | captioning_fn=caption_image_fn if args.observation_type == "accessibility_tree_with_captioner" else None, |
| 321 | ) # NOTE: captioning_fn here is used for captioning input images. |
| 322 | |
| 323 | env = ScriptBrowserEnv( |
| 324 | headless=not args.render, |
| 325 | slow_mo=args.slow_mo, |
| 326 | observation_type=args.observation_type, |
| 327 | current_viewport_only=args.current_viewport_only, |
| 328 | viewport_size={ |
| 329 | "width": args.viewport_width, |
| 330 | "height": args.viewport_height, |
| 331 | }, |
| 332 | save_trace_enabled=args.save_trace_enabled, |
| 333 | sleep_after_execution=args.sleep_after_execution, |
| 334 | # NOTE: captioning_fn here is used for LLM + captioning baselines. |
| 335 | # This can be different from the captioning model used for evals. |
| 336 | captioning_fn=caption_image_fn, |
| 337 | ) |
| 338 | |
| 339 | for config_file in config_file_list: |
no test coverage detected