(batch_size=64, train_steps=200, device="cuda:0")
| 26 | |
| 27 | |
| 28 | def 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 | |
| 78 | class TestDeterminism(DistTestCase): |
no test coverage detected
searching dependent graphs…