Args: cfg (CfgNode): model (nn.Module): evaluators (list[DatasetEvaluator] or None): if None, will call :meth:`build_evaluator`. Otherwise, must have the same length as ``cfg.DATASETS.TEST``. Returns: d
(cls, cfg, model, evaluators=None)
| 501 | |
| 502 | @classmethod |
| 503 | def test(cls, cfg, model, evaluators=None): |
| 504 | """ |
| 505 | Args: |
| 506 | cfg (CfgNode): |
| 507 | model (nn.Module): |
| 508 | evaluators (list[DatasetEvaluator] or None): if None, will call |
| 509 | :meth:`build_evaluator`. Otherwise, must have the same length as |
| 510 | ``cfg.DATASETS.TEST``. |
| 511 | |
| 512 | Returns: |
| 513 | dict: a dict of result metrics |
| 514 | """ |
| 515 | logger = logging.getLogger(__name__) |
| 516 | if isinstance(evaluators, DatasetEvaluator): |
| 517 | evaluators = [evaluators] |
| 518 | if evaluators is not None: |
| 519 | assert len(cfg.DATASETS.TEST) == len(evaluators), "{} != {}".format( |
| 520 | len(cfg.DATASETS.TEST), len(evaluators) |
| 521 | ) |
| 522 | |
| 523 | results = OrderedDict() |
| 524 | for idx, dataset_name in enumerate(cfg.DATASETS.TEST): |
| 525 | data_loader = cls.build_test_loader(cfg, dataset_name) |
| 526 | # When evaluators are passed in as arguments, |
| 527 | # implicitly assume that evaluators can be created before data_loader. |
| 528 | if evaluators is not None: |
| 529 | evaluator = evaluators[idx] |
| 530 | else: |
| 531 | try: |
| 532 | evaluator = cls.build_evaluator(cfg, dataset_name) |
| 533 | except NotImplementedError: |
| 534 | logger.warn( |
| 535 | "No evaluator found. Use `DefaultTrainer.test(evaluators=)`, " |
| 536 | "or implement its `build_evaluator` method." |
| 537 | ) |
| 538 | results[dataset_name] = {} |
| 539 | continue |
| 540 | results_i = inference_on_dataset(model, data_loader, evaluator) |
| 541 | results[dataset_name] = results_i |
| 542 | if comm.is_main_process(): |
| 543 | assert isinstance( |
| 544 | results_i, dict |
| 545 | ), "Evaluator must return a dict on the main process. Got {} instead.".format( |
| 546 | results_i |
| 547 | ) |
| 548 | logger.info("Evaluation results for {} in csv format:".format(dataset_name)) |
| 549 | print_csv_format(results_i) |
| 550 | |
| 551 | if len(results) == 1: |
| 552 | results = list(results.values())[0] |
| 553 | return results |
| 554 | |
| 555 | @staticmethod |
| 556 | def auto_scale_workers(cfg, num_workers: int): |