Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4import unittest5import torch6from torch.autograd import gradcheck7 8from tensormask.layers.swap_align2nat import SwapAlign2Nat9 10 11class SwapAlign2NatTest(unittest.TestCase):12 @unittest.skipIf(not torch.cuda.is_available(), "CUDA not available")13 def test_swap_align2nat_gradcheck_cuda(self):14 dtype = torch.float6415 device = torch.device("cuda")16 m = SwapAlign2Nat(2).to(dtype=dtype, device=device)17 x = torch.rand(2, 4, 10, 10, dtype=dtype, device=device, requires_grad=True)18 19 self.assertTrue(gradcheck(m, x), "gradcheck failed for SwapAlign2Nat CUDA")20 21 def _swap_align2nat(self, tensor, lambda_val):22 """23 The basic setup for testing Swap_Align24 """25 op = SwapAlign2Nat(lambda_val, pad_val=0.0)26 input = torch.from_numpy(tensor[None, :, :, :].astype("float32"))27 output = op.forward(input.cuda()).cpu().numpy()28 return output[0]29 30 31if __name__ == "__main__":32 unittest.main()33 