"""
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)
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)
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)
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)
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)
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):
if not hasattr(torch, "npu") or not torch.npu.is_available():
self.skipTest("NPU device unavailable, skip this test case")
model = Model()
script_model = torch.jit.script(model)
script_model = script_model.npu()
self._check_model_device(script_model, "npu")
cpu_script_model = script_model.cpu()
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
model = Model()
scripted_model = torch.jit.script(model).npu()
self._check_model_device(scripted_model, "npu")
compiled_model = torch.compile(scripted_model, backend="npugraph_ex", fullgraph=True, dynamic=False)
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()
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
model = ModelWithConstants()
scripted_model = torch.jit.script(model).npu()
self._check_model_device(scripted_model, "npu")
code_str, constants = scripted_model.code_with_constants
self.assertIsInstance(code_str, str)
self.assertIn("def forward", code_str)
self.assertEqual(code_str, scripted_model.code)
self.assertIsInstance(constants, torch.jit._script.ConstMap)
const_keys = list(constants.const_mapping.keys())
self.assertGreater(len(const_keys), 0, "ConstMap should not be empty")
self.assertEqual(constants.c0, torch.tensor([2.0]))
self.assertIn("CONSTANTS.c0", code_str)
if __name__ == "__main__":
run_tests()