diff --git a/monai/networks/nets/dints.py b/monai/networks/nets/dints.py index ad4b350d30e..9d917392198 100644 --- a/monai/networks/nets/dints.py +++ b/monai/networks/nets/dints.py @@ -349,7 +349,7 @@ class DiNTS(nn.Module): use_downsample: use downsample in the stem. If ``False``, the search space will be in resolution [1, 1/2, 1/4, 1/8], if ``True``, the search space will be in resolution [1/2, 1/4, 1/8, 1/16]. - node_a: node activation numpy matrix. Its shape is `(num_depths, num_blocks + 1)`. + node_a: node activation matrix (numpy array or tensor). Its shape is `(num_blocks + 1, num_depths)`. +1 for multi-resolution inputs. In model searching stage, ``node_a`` can be None. In deployment stage, ``node_a`` cannot be None. """ @@ -522,7 +522,8 @@ class TopologyConstruction(nn.Module): The base class for `TopologyInstance` and `TopologySearch`. Args: - arch_code: `[arch_code_a, arch_code_c]`, numpy arrays. The architecture codes defining the model. + arch_code: `[arch_code_a, arch_code_c]`, numpy arrays, torch tensors or nested lists (anything + accepted by ``torch.as_tensor``). The architecture codes defining the model. For example, for a ``num_depths=4, num_blocks=12`` search space: - `arch_code_a` is a 12x10 (10 paths) binary matrix representing if a path is activated. @@ -608,8 +609,12 @@ def __init__( arch_code_a = torch.ones((self.num_blocks, len(self.arch_code2out))).to(self.device) arch_code_c = torch.ones((self.num_blocks, len(self.arch_code2out), self.num_cell_ops)).to(self.device) else: - arch_code_a = torch.from_numpy(arch_code[0]).to(self.device) - arch_code_c = F.one_hot(torch.from_numpy(arch_code[1]).to(torch.int64), self.num_cell_ops).to(self.device) + # accept numpy arrays (legacy search checkpoints), torch tensors (checkpoints that are + # loadable with ``torch.load(weights_only=True)``) or nested lists (JSON/YAML configs) + arch_code_a = torch.as_tensor(arch_code[0], device=self.device) + arch_code_c = F.one_hot( + torch.as_tensor(arch_code[1], device=self.device).to(torch.int64), self.num_cell_ops + ) self.arch_code_a = arch_code_a self.arch_code_c = arch_code_c diff --git a/tests/networks/nets/test_dints_network.py b/tests/networks/nets/test_dints_network.py index 80ade00db70..fa7be5c8918 100644 --- a/tests/networks/nets/test_dints_network.py +++ b/tests/networks/nets/test_dints_network.py @@ -11,6 +11,8 @@ from __future__ import annotations +import os +import tempfile import unittest import numpy as np @@ -152,6 +154,49 @@ def test_dints_search(self, dints_grid_params, dints_params, input_shape, expect self.assertEqual(result.shape, expected_shape) self.assertTrue(isinstance(net.weight_parameters(), list)) + @parameterized.expand(TEST_CASES_3D + TEST_CASES_2D) + def test_dints_arch_code_types(self, dints_grid_params, _dints_params, _input_shape, _expected_shape): + """arch_code may be given as numpy arrays, torch tensors or nested lists (issue #9025).""" + dints_grid_params = {k: v for k, v in dints_grid_params.items() if k != "arch_code"} # shared across tests + num_blocks = dints_grid_params["num_blocks"] + num_depths = dints_grid_params["num_depths"] + num_cell_ops = len(Cell(1, 1, 0, spatial_dims=dints_grid_params["spatial_dims"]).OPS) + rng = np.random.RandomState(0) + arch_code_a = rng.randint(2, size=(num_blocks, 3 * num_depths - 2)) + arch_code_a[:, 0] = 1 # keep at least one active path per block + arch_code_c = rng.randint(num_cell_ops, size=(num_blocks, 3 * num_depths - 2)) + + ref = TopologyInstance(arch_code=[arch_code_a, arch_code_c], **dints_grid_params) + for variant in ( + [torch.as_tensor(arch_code_a), torch.as_tensor(arch_code_c)], + [arch_code_a.tolist(), arch_code_c.tolist()], + ): + grid = TopologyInstance(arch_code=variant, **dints_grid_params) + torch.testing.assert_close(grid.arch_code_a, ref.arch_code_a, check_dtype=False) + torch.testing.assert_close(grid.arch_code_c, ref.arch_code_c, check_dtype=False) + self.assertEqual(set(grid.cell_tree.keys()), set(ref.cell_tree.keys())) + + def test_dints_search_code_weights_only_roundtrip(self): + """A search checkpoint saved with tensors loads under ``torch.load(weights_only=True)`` + and can be used to build ``TopologyInstance`` / ``DiNTS`` (issue #9025).""" + grid_params = {"channel_mul": 0.2, "num_blocks": 2, "num_depths": 2, "device": "cpu", "use_downsample": True} + node_a, arch_code_a, arch_code_c, arch_code_a_max = TopologySearch(**grid_params).decode() + with tempfile.TemporaryDirectory() as tempdir: + path = os.path.join(tempdir, "search_code.pt") + torch.save( + { + "node_a": torch.as_tensor(node_a), + "arch_code_a": torch.as_tensor(arch_code_a), + "arch_code_c": torch.as_tensor(arch_code_c), + "arch_code_a_max": torch.as_tensor(arch_code_a_max), + }, + path, + ) + ckpt = torch.load(path, map_location="cpu", weights_only=True) + grid = TopologyInstance(arch_code=[ckpt["arch_code_a"], ckpt["arch_code_c"]], **grid_params) + net = DiNTS(dints_space=grid, in_channels=1, num_classes=2, use_downsample=True, node_a=ckpt["node_a"]) + self.assertEqual(net(torch.randn(1, 1, 16, 16, 16)).shape, (1, 2, 16, 16, 16)) + class TestDintsTS(unittest.TestCase): @parameterized.expand(TEST_CASES_3D + TEST_CASES_2D)