(dist=False)
| 44 | |
| 45 | |
| 46 | def test_func(dist=False): |
| 47 | dummy_model = DummyModel() |
| 48 | dataset = dummy_dataset.to_torch_dataset() |
| 49 | |
| 50 | dummy_loader = DataLoader( |
| 51 | dataset, |
| 52 | batch_size=2, |
| 53 | ) |
| 54 | |
| 55 | metric_class = SequenceClassificationMetric() |
| 56 | |
| 57 | if dist: |
| 58 | init_dist(launcher='pytorch') |
| 59 | |
| 60 | rank, world_size = get_dist_info() |
| 61 | device = torch.device(f'cuda:{rank}') |
| 62 | dummy_model.cuda() |
| 63 | |
| 64 | if world_size > 1: |
| 65 | from torch.nn.parallel.distributed import DistributedDataParallel |
| 66 | dummy_model = DistributedDataParallel( |
| 67 | dummy_model, device_ids=[torch.cuda.current_device()]) |
| 68 | test_func = multi_gpu_test |
| 69 | else: |
| 70 | test_func = single_gpu_test |
| 71 | |
| 72 | dummy_trainer = DummyTrainer(dummy_model) |
| 73 | |
| 74 | metric_results = test_func( |
| 75 | dummy_trainer, |
| 76 | dummy_loader, |
| 77 | device=device, |
| 78 | metric_classes=[metric_class]) |
| 79 | |
| 80 | return metric_results |
| 81 | |
| 82 | |
| 83 | @unittest.skipIf(not torch.cuda.is_available(), 'cuda unittest') |
no test coverage detected
searching dependent graphs…