MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / run_test

Function run_test

tests/integration/test_integration_determinism.py:28–75  ·  view source on GitHub ↗
(batch_size=64, train_steps=200, device="cuda:0")

Source from the content-addressed store, hash-verified

26
27
28def run_test(batch_size=64, train_steps=200, device="cuda:0"):
29 class _TestBatch(Dataset):
30 def __init__(self, transforms):
31 self.transforms = transforms
32
33 def __getitem__(self, _unused_id):
34 im, seg = create_test_image_2d(128, 128, noise_max=1, num_objs=4, num_seg_classes=1)
35 seed = np.random.randint(2147483647)
36 self.transforms.set_random_state(seed=seed)
37 im = self.transforms(im)
38 self.transforms.set_random_state(seed=seed)
39 seg = self.transforms(seg)
40 return im, seg
41
42 def __len__(self):
43 return train_steps
44
45 net = UNet(
46 spatial_dims=2, in_channels=1, out_channels=1, channels=(4, 8, 16, 32), strides=(2, 2, 2), num_res_units=2
47 ).to(device)
48
49 loss = DiceLoss(sigmoid=True)
50 opt = torch.optim.Adam(net.parameters(), 1e-2)
51 train_transforms = Compose(
52 [
53 EnsureChannelFirst(channel_dim="no_channel"),
54 ScaleIntensity(),
55 RandSpatialCrop((96, 96), random_size=False),
56 RandRotate90(),
57 ]
58 )
59
60 src = DataLoader(_TestBatch(train_transforms), batch_size=batch_size, shuffle=True)
61
62 net.train()
63 epoch_loss = 0
64 step = 0
65 for img, seg in src:
66 step += 1
67 opt.zero_grad()
68 output = net(img.to(device))
69 step_loss = loss(output, seg.to(device))
70 step_loss.backward()
71 opt.step()
72 epoch_loss += step_loss.item()
73 epoch_loss /= step
74
75 return epoch_loss, step
76
77
78class TestDeterminism(DistTestCase):

Callers 1

test_trainingMethod · 0.70

Calls 12

UNetClass · 0.90
DiceLossClass · 0.90
ComposeClass · 0.90
EnsureChannelFirstClass · 0.90
ScaleIntensityClass · 0.90
RandSpatialCropClass · 0.90
RandRotate90Class · 0.90
DataLoaderClass · 0.90
_TestBatchClass · 0.70
trainMethod · 0.45
backwardMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…