import itertools
import torch
import numpy as np

import torch_npu
from torch_npu.testing.testcase import TestCase, run_tests
from torch_npu.testing.common_utils import create_common_tensor, SupportedDevices


class TestLerp(TestCase):

    def cpu_op_exec(self, input1, input2, input3):
        output = torch.lerp(input1, input2, input3)
        output = output.numpy()
        return output

    def cpu_op_exec_fp16(self, input1, input2, input3):
        input1 = input1.to(torch.float32)
        input2 = input2.to(torch.float32)
        input3 = input3.to(torch.float32)
        output = torch.lerp(input1, input2, input3)
        output = output.numpy()
        output = output.astype(np.float16)
        return output

    def npu_op_exec(self, input1, input2, input3):
        output = torch.lerp(input1, input2, input3)
        output = output.to("cpu")
        output = output.numpy()
        return output

    def cpu_op_out_exec(self, input1, input2, input3):
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.numpy()
        return output

    def cpu_op_out_exec_fp16(self, input1, input2, input3):
        input1 = input1.to(torch.float32)
        input2 = input2.to(torch.float32)
        input3 = input3.to(torch.float32)
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.numpy()
        output = output.astype(np.float16)
        return output

    def npu_op_out_exec(self, input1, input2, input3):
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.to("cpu")
        output = output.numpy()
        return output

    def cpu_op_scalar_out_exec(self, input1, input2, input3):
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.numpy()
        return output

    def cpu_op_scalar_exec_fp16(self, input1, input2, input3):
        input1 = input1.to(torch.float32)
        input2 = input2.to(torch.float32)
        output = torch.lerp(input1, input2, input3)
        output = output.numpy()
        output = output.astype(np.float16)
        return output

    def cpu_op_scalar_out_exec_fp16(self, input1, input2, input3):
        input1 = input1.to(torch.float32)
        input2 = input2.to(torch.float32)
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.numpy()
        output = output.astype(np.float16)
        return output

    def npu_op_scalar_out_exec(self, input1, input2, input3):
        output = torch.ones_like(input1)
        torch.lerp(input1, input2, input3, out=output)
        output = output.to("cpu")
        output = output.numpy()
        return output

    def test_lerp_common_shape_format(self):
        shape_format = [
            [[np.float32, -1, (4, 2, 2, 3)]],
            [[np.float32, -1, (2, 2, 3, 4)]],
            [[np.float32, -1, (3, 3, 3)]],
            [[np.float32, -1, (4, 4, 4)]]
        ]
        for item in shape_format:
            cpu_input1, npu_input1 = create_common_tensor(item[0], 1, 100)
            cpu_input2, npu_input2 = create_common_tensor(item[0], 1, 100)
            cpu_input3, npu_input3 = create_common_tensor(item[0], 1, 100)
            cpu_output = self.cpu_op_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            cpu_output1 = self.cpu_op_out_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output1 = self.npu_op_out_exec(npu_input1, npu_input2, npu_input3)
            self.assertRtolEqual(cpu_output, npu_output)
            self.assertRtolEqual(cpu_output1, npu_output1)

    def test_lerp_float16_shape_format(self):
        shape_format = [
            [[np.float16, -1, (100, 4, 5, 5)]],
            [[np.float16, -1, (100, 5, 5, 4)]],
        ]
        for item in shape_format:
            cpu_input1, npu_input1 = create_common_tensor(item[0], 10, 100)
            cpu_input2, npu_input2 = create_common_tensor(item[0], 10, 100)
            cpu_input3, npu_input3 = create_common_tensor(item[0], 10, 100)
            cpu_output = self.cpu_op_exec_fp16(cpu_input1, cpu_input2, cpu_input3)
            npu_output = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            cpu_output1 = self.cpu_op_out_exec_fp16(cpu_input1, cpu_input2, cpu_input3)
            npu_output1 = self.npu_op_out_exec(npu_input1, npu_input2, npu_input3)
            self.assertRtolEqual(cpu_output, npu_output, prec=0.003, prec16=0.003)
            self.assertRtolEqual(cpu_output1, npu_output1, prec=0.003, prec16=0.003)

    def test_lerp_scalar_common_shape_format(self):
        shape_format = [
            [[np.float32, -1, (4, 2, 2, 3)], 1.0],
            [[np.float32, -1, (2, 2, 3, 4)], 2.0],
            [[np.float32, -1, (3, 3, 3)], 1.2],
            [[np.float32, -1, (4, 4, 4)], 1.2]
        ]
        for item in shape_format:
            cpu_input1, npu_input1 = create_common_tensor(item[0], 1, 100)
            cpu_input2, npu_input2 = create_common_tensor(item[0], 1, 100)
            cpu_input3 = item[1]
            npu_input3 = item[1]
            cpu_output = self.cpu_op_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            cpu_output1 = self.cpu_op_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output1 = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            self.assertRtolEqual(cpu_output, npu_output)
            self.assertRtolEqual(cpu_output1, npu_output1)

    def test_lerp_scalar_float16_shape_format(self):
        shape_format = [
            [[np.float16, -1, (100, 4, 5, 5)], 1.2],
            [[np.float16, -1, (100, 5, 5, 4)], 1.2],
        ]
        for item in shape_format:
            cpu_input1, npu_input1 = create_common_tensor(item[0], 10, 100)
            cpu_input2, npu_input2 = create_common_tensor(item[0], 10, 100)
            cpu_input3 = item[1]
            npu_input3 = item[1]
            cpu_output = self.cpu_op_scalar_exec_fp16(cpu_input1, cpu_input2, cpu_input3)
            npu_output = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            cpu_output1 = self.cpu_op_scalar_out_exec_fp16(cpu_input1, cpu_input2, cpu_input3)
            npu_output1 = self.npu_op_scalar_out_exec(npu_input1, npu_input2, npu_input3)
            self.assertRtolEqual(cpu_output, npu_output, prec16=0.02)
            self.assertRtolEqual(cpu_output1, npu_output1, prec16=0.02)

    def test_lerp_broadcast_shape_format(self):
        shape_list = [
            [],
            [5, ],
            [5, 5],
        ]
        for shapes in itertools.product(shape_list, shape_list):
            cpu_input1, npu_input1 = create_common_tensor([np.float32, -1, shapes[0]], 10, 100)
            cpu_input2, npu_input2 = create_common_tensor([np.float32, -1, shapes[1]], 10, 100)
            cpu_input3, npu_input3 = create_common_tensor([np.float32, -1, shapes[0]], 10, 100)
            cpu_output = self.cpu_op_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output = self.npu_op_exec(npu_input1, npu_input2, npu_input3)
            cpu_output1 = self.cpu_op_out_exec(cpu_input1, cpu_input2, cpu_input3)
            npu_output1 = self.npu_op_out_exec(npu_input1, npu_input2, npu_input3)
            self.assertRtolEqual(cpu_output, npu_output)
            self.assertRtolEqual(cpu_output1, npu_output1)

    @SupportedDevices(['Ascend910B'])
    def test_lerp_inplace_shape_format_910b(self):
        cpu_input1, npu_input1 = create_common_tensor([np.float32, -1, [2, 1]], 10, 100)
        cpu_input2, npu_input2 = create_common_tensor([np.float32, -1, [2, 6]], 10, 100)
        cpu_input2.lerp_(cpu_input1, 1)
        npu_input2.lerp_(npu_input1, 1)
        self.assertRtolEqual(cpu_input2, npu_input2.cpu())

        def lerp_inplace(npu_input1, npu_input2):
            npu_input1.lerp_(npu_input2, 1)
        self.assertRaisesRegex(
            Exception, "CheckShape failed", lerp_inplace, npu_input1, npu_input2)

    @SupportedDevices(['Ascend910A'])
    def test_lerp_inplace_shape_format(self):
        cpu_input1, npu_input1 = create_common_tensor([np.float32, -1, [2, 1]], 10, 100)
        cpu_input2, npu_input2 = create_common_tensor([np.float32, -1, [2, 6]], 10, 100)
        cpu_input2.lerp_(cpu_input1, 1)
        npu_input2.lerp_(npu_input1, 1)
        self.assertRtolEqual(cpu_input2, npu_input2.cpu())

        def lerp_inplace(npu_input1, npu_input2):
            npu_input1.lerp_(npu_input2, 1)
        self.assertRaisesRegex(
            Exception, "doesn't match the broadcast shape", lerp_inplace, npu_input1, npu_input2)


    def test_lerp_weight_0d_cpu_tensor(self):
        cpu_a = torch.tensor([1.0, 2.0, 3.0])
        cpu_b = torch.tensor([4.0, 5.0, 6.0])
        cpu_w = torch.tensor(0.5)
        npu_a = cpu_a.npu()
        npu_b = cpu_b.npu()
        cpu_output = torch.lerp(cpu_a, cpu_b, cpu_w)
        npu_output = torch.lerp(npu_a, npu_b, cpu_w)
        self.assertRtolEqual(cpu_output, npu_output.cpu())

    def test_lerp_inplace_weight_0d_cpu_tensor(self):
        cpu_a = torch.tensor([1.0, 2.0, 3.0])
        cpu_b = torch.tensor([4.0, 5.0, 6.0])
        cpu_w = torch.tensor(0.5)
        npu_a = cpu_a.clone().npu()
        npu_b = cpu_b.npu()
        cpu_a.lerp_(cpu_b, cpu_w)
        npu_a.lerp_(npu_b, cpu_w)
        self.assertRtolEqual(cpu_a, npu_a.cpu())


if __name__ == '__main__':
    run_tests()