Test model in EpochBasedTrainer with a single gpu. Args: trainer (modelscope.trainers.EpochBasedTrainer): Trainer to be tested. data_loader (nn.Dataloader): Pytorch data loader. device (str | torch.device): The target device for the data. metric_classes (List): L
(trainer,
data_loader,
device,
metric_classes=None,
vis_closure=None,
data_loader_iters=None)
| 16 | |
| 17 | |
| 18 | def single_gpu_test(trainer, |
| 19 | data_loader, |
| 20 | device, |
| 21 | metric_classes=None, |
| 22 | vis_closure=None, |
| 23 | data_loader_iters=None): |
| 24 | """Test model in EpochBasedTrainer with a single gpu. |
| 25 | |
| 26 | Args: |
| 27 | trainer (modelscope.trainers.EpochBasedTrainer): Trainer to be tested. |
| 28 | data_loader (nn.Dataloader): Pytorch data loader. |
| 29 | device (str | torch.device): The target device for the data. |
| 30 | metric_classes (List): List of Metric class that uses to collect metrics. |
| 31 | vis_closure (Callable): Collect data for TensorboardHook. |
| 32 | data_loader_iters (int): Used when dataset has no attribute __len__ or only load part of dataset. |
| 33 | |
| 34 | Returns: |
| 35 | list: The prediction results. |
| 36 | """ |
| 37 | dataset = data_loader.dataset |
| 38 | progress_with_iters = False |
| 39 | if data_loader_iters is None: |
| 40 | try: |
| 41 | data_len = len(dataset) |
| 42 | except Exception as e: |
| 43 | logging.error(e) |
| 44 | raise ValueError( |
| 45 | 'Please implement ``__len__`` method for your dataset, or provide ``data_loader_iters``' |
| 46 | ) |
| 47 | desc = 'Total test samples' |
| 48 | else: |
| 49 | progress_with_iters = True |
| 50 | data_len = data_loader_iters |
| 51 | desc = 'Test iterations' |
| 52 | |
| 53 | with tqdm(total=data_len, desc=desc) as pbar: |
| 54 | for i, data in enumerate(data_loader): |
| 55 | data = to_device(data, device) |
| 56 | evaluate_batch(trainer, data, metric_classes, vis_closure) |
| 57 | |
| 58 | if progress_with_iters: |
| 59 | batch_size = 1 # iteration count |
| 60 | else: |
| 61 | if isinstance(data, Mapping): |
| 62 | if 'nsentences' in data: |
| 63 | batch_size = data['nsentences'] |
| 64 | else: |
| 65 | try: |
| 66 | batch_size = len(next(iter(data.values()))) |
| 67 | except Exception: |
| 68 | batch_size = data_loader.batch_size |
| 69 | else: |
| 70 | batch_size = len(data) |
| 71 | for _ in range(batch_size): |
| 72 | pbar.update() |
| 73 | |
| 74 | if progress_with_iters and (i + 1) >= data_len: |
| 75 | break |
no test coverage detected
searching dependent graphs…