diff --git a/monai/networks/nets/dints.py b/monai/networks/nets/dints.py index 88f671152a0..8bdbacc8676 100644 --- a/monai/networks/nets/dints.py +++ b/monai/networks/nets/dints.py @@ -13,6 +13,7 @@ import datetime import warnings +from typing import Any import numpy as np import torch @@ -481,6 +482,19 @@ def __init__( nn.Upsample(scale_factor=2 ** (res_idx != 0), mode=mode, align_corners=True), ) + def __setattr__(self, name: str, value: Any) -> None: + """ + Mirror ``node_a`` into ``_node_a_py`` on assignment. ``forward()`` branches on its + entries, which dynamo cannot constant-fold off a tensor, so a compiled model breaks into + one graph per branch (https://github.com/Project-MONAI/MONAI/issues/9144); the compiled + path reads the mirror instead. Mirroring here rather than in ``__init__()`` keeps code + that replaces ``node_a`` after construction correct. + """ + super().__setattr__(name, value) + if name == "node_a" and value is not None: + # `!= 0`, not an int cast: casting would make a fractional flag such as 0.5 falsy. + object.__setattr__(self, "_node_a_py", (torch.as_tensor(value) != 0).tolist()) + def weight_parameters(self): return [param for name, param in self.named_parameters()] @@ -491,12 +505,20 @@ def forward(self, x: torch.Tensor): Args: x: input tensor. """ + # Branching on `node_a` elements is a data-dependent tensor read that dynamo cannot + # constant-fold, costing 13 graph breaks; hand it the folded snapshot instead. Eager + # keeps the live tensor, so in-place edits still apply. TorchScript cannot type + # `tolist()` and folds `is_scripting()` away, so it indexes the tensor as before. + node_a = self.node_a != 0 + if not torch.jit.is_scripting(): + node_a = self._node_a_py if torch.compiler.is_compiling() else node_a.tolist() + inputs = [] for d in range(self.num_depths): # allow multi-resolution input _mod_w: StemInterface = self.stem_down[str(d)] # type: ignore[assignment] x_out = _mod_w.forward(x) - if self.node_a[0][d]: + if node_a[0][d]: inputs.append(x_out) else: inputs.append(torch.zeros_like(x_out)) @@ -510,7 +532,7 @@ def forward(self, x: torch.Tensor): _mod_up: StemInterface = self.stem_up[str(res_idx)] # type: ignore[assignment] if start: _temp = _mod_up.forward(outputs[res_idx] + _temp) - elif self.node_a[blk_idx + 1][res_idx]: + elif node_a[blk_idx + 1][res_idx]: start = True _temp = _mod_up.forward(outputs[res_idx]) prediction = self.stem_finals(_temp) @@ -629,6 +651,16 @@ def __init__( self._norm_name, ) + def __setattr__(self, name: str, value: Any) -> None: + """ + Mirror ``arch_code_a`` into ``_arch_code_a_py`` on assignment. See + ``DiNTS.__setattr__`` for why. + """ + super().__setattr__(name, value) + if name == "arch_code_a" and value is not None: + # See `DiNTS._node_a_py`. + object.__setattr__(self, "_arch_code_a_py", (torch.as_tensor(value) != 0).tolist()) + def forward(self, x): """This function to be implemented by the architecture instances or search spaces.""" @@ -675,11 +707,16 @@ def forward(self, x: list[torch.Tensor]) -> list[torch.Tensor]: Args: x: input tensor. """ + # See `DiNTS.forward`. + arch_code_a = self.arch_code_a != 0 + if not torch.jit.is_scripting(): + arch_code_a = self._arch_code_a_py if torch.compiler.is_compiling() else arch_code_a.tolist() + # generate path activation probability inputs = x for blk_idx in range(self.num_blocks): outputs = [torch.tensor(0.0, dtype=x[0].dtype, device=x[0].device)] * self.num_depths - for res_idx, activation in enumerate(self.arch_code_a[blk_idx].data): + for res_idx, activation in enumerate(arch_code_a[blk_idx]): if activation: mod: CellInterface = self.cell_tree[str((blk_idx, res_idx))] # type: ignore[assignment] _out = mod.forward(x=inputs[self.arch_code2in[res_idx]], weight=None) diff --git a/tests/networks/nets/test_dints_network.py b/tests/networks/nets/test_dints_network.py index d39c6c06ae0..b554d69eb5d 100644 --- a/tests/networks/nets/test_dints_network.py +++ b/tests/networks/nets/test_dints_network.py @@ -211,6 +211,88 @@ def test_dints_forward_tensor_arch_code(self): self.assertEqual(result.shape, (1, 2, 16, 16, 16)) +class TestDintsTopologyCache(unittest.TestCase): + """When compiling, `forward` branches on Python snapshots of `node_a` / `arch_code_a`. + + These tests pin both halves of that: the snapshot agrees with its source, and eager -- which + still reads the tensors -- honours mutations of the source. + """ + + def _build(self, node_a=None, all_paths_active=False): + num_blocks, num_depths, spatial_dims = 6, 3, 3 + cell = Cell(1, 1, 0, spatial_dims=spatial_dims) + rng = np.random.RandomState(0) + if all_paths_active: + # a random code can leave a block with no active path, handing the next block a + # scalar; tests that actually run `forward` need every path live. + arch_code_a = np.ones((num_blocks, 3 * num_depths - 2), dtype=int) + else: + arch_code_a = rng.randint(0, 2, size=(num_blocks, 3 * num_depths - 2)) + arch_code_a[0, 0] = 1 # keep at least one active path + arch_code_c = rng.randint(len(cell.OPS), size=(num_blocks, 3 * num_depths - 2)) + grid = TopologyInstance( + num_blocks=num_blocks, + num_depths=num_depths, + spatial_dims=spatial_dims, + device="cpu", + arch_code=[arch_code_a, arch_code_c], + ) + net = DiNTS(dints_space=grid, in_channels=1, num_classes=2, spatial_dims=spatial_dims, node_a=node_a) + return net, grid + + def test_cache_matches_source_at_construction(self): + net, grid = self._build() + self.assertEqual(net._node_a_py, (net.node_a != 0).tolist()) + self.assertEqual(grid._arch_code_a_py, (grid.arch_code_a != 0).tolist()) + + def test_cache_preserves_truthiness_not_int_value(self): + """Any non-zero entry is active; truncating to int would make 0.5 and -0.5 falsy.""" + node_a = torch.ones((7, 3)) + node_a[0, 0], node_a[1, 1], node_a[2, 2] = 0.5, -0.5, 0.0 + net, _ = self._build(node_a=node_a) + self.assertEqual(net._node_a_py[0][0], True) + self.assertEqual(net._node_a_py[1][1], True) + self.assertEqual(net._node_a_py[2][2], False) + self.assertEqual(net._node_a_py, (node_a != 0).tolist()) + + def test_cache_resyncs_when_node_a_is_replaced(self): + """Deployment code assigns `node_a` after construction; the snapshot must follow.""" + net, _ = self._build() + replacement = torch.zeros_like(torch.as_tensor(net.node_a)) + replacement[0, 0] = 1 + net.node_a = replacement + self.assertEqual(net._node_a_py, (replacement != 0).tolist()) + + def test_eager_honours_in_place_node_a_edit(self): + """Eager reads the live tensor, so an in-place edit changes the output as before.""" + net, _ = self._build(node_a=torch.ones((7, 3)), all_paths_active=True) + net.eval() + x = torch.randn(1, 1, 32, 32, 32) + with torch.no_grad(): + before = net(x).clone() + net.node_a[0][0] = 0 + with torch.no_grad(): + after = net(x) + self.assertFalse(torch.allclose(before, after)) + + def test_eager_honours_caller_owned_arch_code_array(self): + """`torch.from_numpy` aliases the caller's array; eager must see edits to it.""" + num_blocks, num_depths, spatial_dims = 6, 3, 3 + cell = Cell(1, 1, 0, spatial_dims=spatial_dims) + arch_code_a = np.ones((num_blocks, 3 * num_depths - 2)) + arch_code_c = np.random.RandomState(0).randint(len(cell.OPS), size=(num_blocks, 3 * num_depths - 2)) + grid = TopologyInstance( + num_blocks=num_blocks, + num_depths=num_depths, + spatial_dims=spatial_dims, + device="cpu", + arch_code=[arch_code_a, arch_code_c], + ) + self.assertTrue(bool(grid.arch_code_a[0, 0])) + arch_code_a[0, 0] = 0 # caller edits the array it passed in + self.assertFalse(bool(grid.arch_code_a[0, 0])) + + class TestDintsTS(unittest.TestCase): @parameterized.expand(TEST_CASES_3D + TEST_CASES_2D) def test_script(self, dints_grid_params, dints_params, input_shape, _):