"""
Add validation cases for torch.jit.ScriptModule APIs on NPU:

PyTorch community lacks sufficient and direct API validations for some APIs, so this file is added.
This file validates torch.jit.ScriptModule.bfloat16, torch.jit.ScriptModule.buffers,
torch.jit.ScriptModule.children, torch.jit.ScriptModule.code
torch.jit.ScriptModule.compile, torch.jit.ScriptModule.cpu
torch.jit.ScriptModule.add_module, torch.jit.ScriptModule.apply
"""

import os
import tempfile
import torch
from torch.testing._internal.common_utils import TestCase, run_tests
import torch.nn as nn
import torch.nn.functional as F

class Model(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv1 = nn.Conv2d(1, 20, 5)
        self.conv2 = nn.Conv2d(20, 20, 5)

        # Register buffer
        self.register_buffer('register_buffer_out1', torch.randn(20))
        self.register_buffer('register_buffer_out2', torch.randn(20, 1, 5, 5))

    def forward(self, x):
        x = F.relu(self.conv1(x))
        return F.relu(self.conv2(x))


class TestJitScriptModuleBfloat16(TestCase):
    def test_bfloat16(self):
        model = Model()
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        self.assertIsInstance(model, nn.Module)

        model_bf16 = model.bfloat16()
        self.assertIsInstance(model_bf16, nn.Module)
        # The validation parameters have been converted to bfloat16
        for param in model_bf16.parameters():
            self.assertEqual(param.dtype, torch.bfloat16)

    def test_jit_bfloat16(self):
        model = Model().bfloat16()
        jit_model = torch.jit.script(model)
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        self.assertIsInstance(jit_model, torch.jit.ScriptModule)
        # Verify the feasibility of JIT model execution
        dummy_input = torch.randn(1, 1, 28, 28, dtype=torch.bfloat16, device="cpu")
        output = jit_model(dummy_input)
        self.assertEqual(output.dtype, torch.bfloat16)


class TestJitScriptModuleBuffers(TestCase):
    def test_buffers(self):
        model = Model()
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        buffers = list(model.buffers())
        self.assertEqual(len(buffers), 2)  # Two registered buffers
        self.assertTrue(all(isinstance(b, torch.Tensor) for b in buffers))

    def test_jit_buffers(self):
        model = torch.jit.script(Model())
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        buffers = list(model.buffers())
        self.assertEqual(len(buffers), 2)
        self.assertTrue(all(isinstance(b, torch.Tensor) for b in buffers))


class TestJitScriptModuleChildren(TestCase):
    def test_children(self):
        model = Model()
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        children = list(model.children())
        self.assertEqual(len(children), 2)  # conv1, conv2
        self.assertTrue(all(isinstance(c, nn.Module) for c in children))

    def test_jit_children(self):
        model = torch.jit.script(Model())
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        children = list(model.children())
        self.assertEqual(len(children), 2)
        self.assertTrue(all(isinstance(c, torch.nn.Module) for c in children))


class TestJitScriptModuleCode(TestCase):
    def test_jit_code(self):
        model = torch.jit.script(Model())
        uninit_param = torch.nn.parameter.UninitializedParameter(device="npu")

        self.assertEqual(uninit_param.device.type, "npu")
        self.assertTrue(hasattr(model, 'code'))
        self.assertIn('forward', str(model.code))


class TestJitScriptModuleCPU(TestCase):
    def _check_model_device(self, model: torch.jit.ScriptModule, device_type: str):
        """Check whether all parameters and buffers of the model are on the specified device."""
        for param in model.parameters():
            self.assertEqual(param.device.type, device_type)
        for buf in model.buffers():
            self.assertEqual(buf.device.type, device_type)

    def _run_infer(self, model: torch.jit.ScriptModule, device):
        x = torch.randn(2, 1, 32, 32).to(device)
        out = model(x)
        self.assertEqual(out.dim(), 4)

    def test_script_module_cpu_from_npu(self):
        # Skip test if NPU is not available
        if not hasattr(torch, "npu") or not torch.npu.is_available():
            self.skipTest("NPU device unavailable, skip this test case")

        # Instantiate model and convert to ScriptModule
        model = Model()
        script_model = torch.jit.script(model)

        # Move ScriptModule to NPU and verify device
        script_model = script_model.npu()
        self._check_model_device(script_model, "npu")

        # Transfer model from NPU back to CPU via cpu() method
        cpu_script_model = script_model.cpu()

        # Assert model type, device and normal inference execution
        self.assertIsInstance(cpu_script_model, torch.jit.ScriptModule)
        self._check_model_device(cpu_script_model, "cpu")
        self._run_infer(cpu_script_model, torch.device("cpu"))


class TestJitScriptModuleCompile(TestCase):
    def _check_model_device(self, model: torch.jit.ScriptModule, device_type: str):
        """Check whether all parameters and buffers of the model are on the specified device."""
        for param in model.parameters():
            self.assertEqual(param.device.type, device_type)
        for buf in model.buffers():
            self.assertEqual(buf.device.type, device_type)

    def _run_infer(self, model, device):
        """Run model inference and verify output tensor shape validity."""
        x = torch.randn(2, 1, 32, 32).to(device)
        out = model(x)
        self.assertEqual(out.dim(), 4)

    def test_script_module_compile(self):
        """Test compilation and inference using torch.compile with NPU backend."""
        if not hasattr(torch, "npu") or not torch.npu.is_available():
            self.skipTest("NPU device is unavailable, skip this test case")
        torch.npu.config.allow_internal_format = False

        # Initialize model and convert to ScriptModule
        model = Model()
        scripted_model = torch.jit.script(model).npu()

        # Move model to NPU and check device placement
        self._check_model_device(scripted_model, "npu")

        # Compile model with NPU backend
        compiled_model = torch.compile(scripted_model, backend="npugraph_ex", fullgraph=True, dynamic=False)

        # run inference validation
        self._run_infer(compiled_model, torch.device("npu"))


class TestScriptModuleAddModule(TestCase):
    class LinearLayer3to2(nn.Module):
        def __init__(self):
            super().__init__()
            self.linear = nn.Linear(3, 2)

        def forward(self, x: torch.Tensor) -> torch.Tensor:
            return self.linear(x)

    class LinearLayer5to3(nn.Module):
        def __init__(self):
            super().__init__()
            self.linear = nn.Linear(5, 3)

        def forward(self, x: torch.Tensor) -> torch.Tensor:
            return self.linear(x)

    class ScriptedModelWithDynamicAdd(torch.jit.ScriptModule):
        def __init__(self):
            super().__init__()
            main_mod = TestScriptModuleAddModule.LinearLayer5to3()
            scripted_main = torch.jit.script(main_mod)
            self.add_module("main", scripted_main)

            sub_mod = TestScriptModuleAddModule.LinearLayer3to2()
            scripted_sub = torch.jit.script(sub_mod)
            self.add_module("sub_module", scripted_sub)

        @torch.jit.script_method
        def forward(self, x):
            return self.main(x)

    def test_add_module_on_scriptmodule_after_script(self):
        module = self.ScriptedModelWithDynamicAdd()
        # NPU
        module.npu()
        self.assertTrue(hasattr(module, "main"))
        self.assertTrue(hasattr(module, "sub_module"))
        self.assertIsInstance(module.main, torch.jit.ScriptModule)
        self.assertIsInstance(module.sub_module, torch.jit.ScriptModule)
        self.assertIn("main", module._modules)
        self.assertIn("sub_module", module._modules)

        main_params = list(module.main.parameters())
        self.assertEqual(len(main_params), 2)
        self.assertEqual(main_params[0].shape, (3, 5))
        self.assertEqual(main_params[0].device.type, "npu")

        sub_params = list(module.sub_module.parameters())
        self.assertEqual(len(sub_params), 2)
        self.assertEqual(sub_params[0].shape, (2, 3))
        self.assertEqual(sub_params[0].device.type, "npu")
        self.assertEqual(sub_params[1].device.type, "npu")

        dummy_hidden = torch.randn(1, 3, device="npu")
        sub_out = module.sub_module(dummy_hidden)
        self.assertEqual(sub_out.shape, (1, 2))
        self.assertEqual(sub_out.device.type, "npu")

        dummy_input = torch.randn(1, 5, device="npu")
        main_out = module(dummy_input)
        self.assertEqual(main_out.shape, (1, 3))
        self.assertEqual(main_out.device.type, "npu")

        with tempfile.NamedTemporaryFile(suffix='.pt', delete=False) as f:
            tmp_path = f.name
        try:
            module.save(tmp_path)
            loaded = torch.jit.load(tmp_path)
            loaded_npu = loaded.to("npu")
            loaded_out = loaded_npu(dummy_input)
            self.assertEqual(loaded_out.shape, (1, 3))
            self.assertEqual(loaded_out.device.type, "npu")
        finally:
            if os.path.exists(tmp_path):
                os.remove(tmp_path)


class TestScriptModuleApply(TestCase):
    def test_apply_modify_parameters(self):
        class SimpleModel(nn.Module):
            def __init__(self):
                super().__init__()
                self.conv = nn.Conv2d(1, 2, 3)
                self.fc = nn.Linear(2 * 26 * 26, 5)

            def forward(self, x):
                x = self.conv(x)
                x = x.view(x.size(0), -1)
                return self.fc(x)

        device = torch.device("npu")
        model = SimpleModel().to(device)
        scripted = torch.jit.script(model)
        self.assertIsInstance(scripted, torch.jit.ScriptModule)

        def zero_params(module):
            if hasattr(module, 'weight') and module.weight is not None:
                module.weight.data.fill_(0.0)
            if hasattr(module, 'bias') and module.bias is not None:
                module.bias.data.fill_(0.0)

        scripted.apply(zero_params)

        for param in scripted.parameters():
            self.assertTrue(torch.all(param == 0))
            self.assertEqual(param.device.type, "npu")

        dummy = torch.randn(1, 1, 28, 28, device=device)
        out = scripted(dummy)
        self.assertTrue(torch.all(out == 0))
        self.assertEqual(out.shape, (1, 5))
        self.assertEqual(out.device.type, "npu")

    def test_apply_recursively_visits_all_modules(self):
        class Leaf(nn.Module):
            def __init__(self):
                super().__init__()
                self.param = nn.Parameter(torch.randn(2, 2))

            def forward(self, x):
                return x

        class Container(nn.Module):
            def __init__(self):
                super().__init__()
                self.leaf1 = Leaf()
                self.leaf2 = Leaf()

            def forward(self, x):
                return x

        device = torch.device("npu")
        model = Container().to(device)
        scripted = torch.jit.script(model)

        visited = set()

        def record(module):
            visited.add(id(module))

        scripted.apply(record)

        expected = {id(scripted), id(scripted.leaf1), id(scripted.leaf2)}
        self.assertEqual(visited, expected)

    def test_apply_returns_self(self):
        class Dummy(nn.Module):
            def forward(self, x):
                return x

        device = torch.device("npu")
        model = Dummy().to(device)
        scripted = torch.jit.script(model)

        def noop(module):
            pass

        ret = scripted.apply(noop)
        self.assertIs(ret, scripted)


class TestJitScriptModuleCodeWithConstants(TestCase):
    def _check_model_device(self, model: torch.jit.ScriptModule, device_type: str):
        """Check whether all parameters and buffers of the model are on the specified device."""
        for param in model.parameters():
            self.assertEqual(param.device.type, device_type)
        for buf in model.buffers():
            self.assertEqual(buf.device.type, device_type)

    def test_script_module_code_with_constants(self):
        @torch.jit.script
        def apply_scale(x, scale=torch.tensor([2.0])):
            return x * scale

        class ModelWithConstants(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.conv = torch.nn.Conv2d(1, 1, kernel_size=3, padding=1)
                self.threshold = 0.5

            def forward(self, x):
                x = self.conv(x)
                x = apply_scale(x)
                if x.mean() > self.threshold:
                    x = x + 1.0
                return x

        # Initialize model and convert to ScriptModule
        model = ModelWithConstants()
        scripted_model = torch.jit.script(model).npu()

        # Check device placement
        self._check_model_device(scripted_model, "npu")

        # Verify the functionality of code_with_constants API
        code_str, constants = scripted_model.code_with_constants

        # 1. Verify the extracted code is a string, contains forward, and matches .code
        self.assertIsInstance(code_str, str)
        self.assertIn("def forward", code_str)
        self.assertEqual(code_str, scripted_model.code)

        # 2. Verify the constants object type
        self.assertIsInstance(constants, torch.jit._script.ConstMap)

        # 3. Verify constants mapping is non-empty and contains the expected tensor
        const_keys = list(constants.const_mapping.keys())
        self.assertGreater(len(const_keys), 0, "ConstMap should not be empty")

        # 4. Assert real CONSTANT.c0 value
        self.assertEqual(constants.c0, torch.tensor([2.0]))

        # 5. Verify the constant is referenced in the code string
        self.assertIn("CONSTANTS.c0", code_str)

if __name__ == "__main__":
    run_tests()