Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions monai/networks/nets/dints.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
45 changes: 45 additions & 0 deletions tests/networks/nets/test_dints_network.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@

from __future__ import annotations

import os
import tempfile
import unittest

import numpy as np
Expand Down Expand Up @@ -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)
Expand Down