import torch
import torch_npu
from torch_npu.testing.testcase import TestCase, run_tests
class TestNpuOneHot(TestCase):
def generate_single_data(self, low, high):
npu_input1 = torch.arange(low, high)
return npu_input1
def cpu_op_exec(self, input1, num_classes):
output = torch.nn.functional.one_hot(input1, num_classes=num_classes)
output = output.to(torch.int32)
output = output.numpy()
return output
def npu_op_exec(self, input1, num_classes):
input1 = input1.to(torch.int32)
input1 = input1.to("npu")
output = torch_npu.npu_one_hot(input1, depth=num_classes).to(torch.int32)
output = output.to("cpu")
output = output.numpy()
return output
def test_npu_one_hot_1(self):
input1 = self.generate_single_data(0, 5)
cpu_output = self.cpu_op_exec(input1, 5)
npu_output = self.npu_op_exec(input1, 5)
self.assertRtolEqual(cpu_output, npu_output)
def test_npu_one_hot_2(self):
input1 = self.generate_single_data(0, 5)
npu_output = self.npu_op_exec(input1, 6)
cpu_output = self.cpu_op_exec(input1, 6)
self.assertRtolEqual(cpu_output, npu_output)
def test_npu_one_hot_3(self):
input1 = self.generate_single_data(0, 10)
cpu_output = self.cpu_op_exec(input1, 10)
npu_output = self.npu_op_exec(input1, 10)
self.assertRtolEqual(cpu_output, npu_output)
def test_npu_one_hot_4(self):
input1 = self.generate_single_data(0, 10)
cpu_output = self.cpu_op_exec(input1, 12)
npu_output = self.npu_op_exec(input1, 12)
self.assertRtolEqual(cpu_output, npu_output)
def test_npu_one_hot_aicpu_int64(self):
input1 = torch.randint(0, 4, size=(4, 64, 64, 64)).npu()
cpu_output = self.cpu_op_exec(input1.cpu(), 4)
npu_output = self.npu_op_exec(input1, 4)
self.assertRtolEqual(cpu_output, npu_output)
if __name__ == "__main__":
run_tests()