(is_ref=True, device=None, reverse_indexing=False)
| 75 | |
| 76 | |
| 77 | def compare_2d(is_ref=True, device=None, reverse_indexing=False): |
| 78 | batch_size = 32 |
| 79 | img_a = [create_test_image_2d(28, 28, 5, rad_max=6, noise_max=1)[0][None] for _ in range(batch_size)] |
| 80 | img_b = [create_test_image_2d(28, 28, 5, rad_max=6, noise_max=1)[0][None] for _ in range(batch_size)] |
| 81 | img_a = np.stack(img_a, axis=0) |
| 82 | img_b = np.stack(img_b, axis=0) |
| 83 | img_a = torch.as_tensor(img_a, device=device) |
| 84 | img_b = torch.as_tensor(img_b, device=device) |
| 85 | model = STNBenchmark(is_ref=is_ref, reverse_indexing=reverse_indexing).to(device) |
| 86 | optimizer = optim.SGD(model.parameters(), lr=0.001) |
| 87 | model.train() |
| 88 | init_loss = None |
| 89 | for _ in range(20): |
| 90 | optimizer.zero_grad() |
| 91 | output_a = model(img_a) |
| 92 | loss = torch.mean((output_a - img_b) ** 2) |
| 93 | if init_loss is None: |
| 94 | init_loss = loss.item() |
| 95 | loss.backward() |
| 96 | optimizer.step() |
| 97 | return model(img_a).detach().cpu().numpy(), loss.item(), init_loss |
| 98 | |
| 99 | |
| 100 | class TestSpatialTransformerCore(DistTestCase): |
no test coverage detected
searching dependent graphs…