MCPcopy Create free account
hub / github.com/modelscope/modelscope / single_gpu_test

Function single_gpu_test

modelscope/trainers/utils/inference.py:18–77  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

16
17
18def 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

Callers 3

evaluation_loopMethod · 0.90
evaluation_loopMethod · 0.90
evaluateMethod · 0.85

Calls 4

to_deviceFunction · 0.90
evaluate_batchFunction · 0.85
get_metric_valuesFunction · 0.85
updateMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…