From 81e00df7c5a0dc7212c83ec3793861b8ed6b2a01 Mon Sep 17 00:00:00 2001 From: "Patel, Nilaykumar K" Date: Thu, 1 Oct 2026 02:23:04 -0400 Subject: [PATCH 1/3] fix(dints): avoid torch.compile graph breaks on topology branches DiNTS.forward and TopologyInstance.forward branch on elements of the node_a / arch_code_a tensors. Each tensor read is a graph break under torch.compile: torch._dynamo.explain on a 6-block/3-depth DiNTS reports 14 graphs with 13 breaks. Mirror the activation flags into plain Python lists that dynamo can constant-fold, and keep the mirrors in sync through __setattr__ so that deployment code replacing node_a after construction stays correct. The same model then compiles to 1 graph with 0 breaks. The mirrors use '!= 0' rather than an int cast so that the cached flags reproduce the truthiness of the original tensor test exactly; casting would make fractional activations such as 0.5 falsy. Signed-off-by: Patel, Nilaykumar K --- monai/networks/nets/dints.py | 26 ++++++++++-- tests/networks/nets/test_dints_network.py | 49 +++++++++++++++++++++++ 2 files changed, 72 insertions(+), 3 deletions(-) diff --git a/monai/networks/nets/dints.py b/monai/networks/nets/dints.py index 88f671152a0..54a7480356f 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,17 @@ def __init__( nn.Upsample(scale_factor=2 ** (res_idx != 0), mode=mode, align_corners=True), ) + def __setattr__(self, name: str, value: Any) -> None: + super().__setattr__(name, value) + if name == "node_a" and value is not None: + # `forward` branches on these flags. Reading them from a tensor forces a graph + # break under torch.compile, so mirror them into a plain Python list that can be + # constant-folded. `!= 0` reproduces the truthiness of the original tensor test + # exactly; casting to int would make fractional values such as 0.5 falsy. + # This tracks rebinding only -- mutating `node_a` in place leaves the mirror + # stale, so replace the attribute rather than editing it. + 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()] @@ -496,7 +508,7 @@ def forward(self, x: torch.Tensor): # 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 self._node_a_py[0][d]: inputs.append(x_out) else: inputs.append(torch.zeros_like(x_out)) @@ -510,7 +522,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 self._node_a_py[blk_idx + 1][res_idx]: start = True _temp = _mod_up.forward(outputs[res_idx]) prediction = self.stem_finals(_temp) @@ -629,6 +641,14 @@ def __init__( self._norm_name, ) + def __setattr__(self, name: str, value: Any) -> None: + super().__setattr__(name, value) + if name == "arch_code_a" and value is not None: + # Mirrored for the same reason as `DiNTS._node_a_py`: `TopologyInstance.forward` + # branches on these flags, and reading them from a tensor breaks the + # torch.compile graph. Tracks rebinding only, not in-place mutation. + 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.""" @@ -679,7 +699,7 @@ def forward(self, x: list[torch.Tensor]) -> list[torch.Tensor]: 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(self._arch_code_a_py[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..017857c45b6 100644 --- a/tests/networks/nets/test_dints_network.py +++ b/tests/networks/nets/test_dints_network.py @@ -211,6 +211,55 @@ def test_dints_forward_tensor_arch_code(self): self.assertEqual(result.shape, (1, 2, 16, 16, 16)) +class TestDintsTopologyCache(unittest.TestCase): + """`forward` branches on cached Python copies of `node_a` / `arch_code_a`. + + The caches exist so torch.compile can constant-fold the branches instead of breaking the + graph on a tensor read. A cache that disagreed with its source would silently change which + cells are executed, so these tests pin the equivalence. + """ + + def _build(self, node_a=None): + num_blocks, num_depths, spatial_dims = 6, 3, 3 + cell = Cell(1, 1, 0, spatial_dims=spatial_dims) + rng = np.random.RandomState(0) + 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 cache 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()) + + class TestDintsTS(unittest.TestCase): @parameterized.expand(TEST_CASES_3D + TEST_CASES_2D) def test_script(self, dints_grid_params, dints_params, input_shape, _): From a02bcb46e901226f070f2ad2a97ae93a2ed8b665 Mon Sep 17 00:00:00 2001 From: "Patel, Nilaykumar K" Date: Thu, 1 Oct 2026 04:42:31 -0400 Subject: [PATCH 2/3] fix(dints): keep eager topology branches responsive to in-place edits `tolist()` snapshots, so branching on it unconditionally broke in-place edits: `net.node_a[0, 0] = 0` left the stem input active, and because `torch.from_numpy` aliases the caller's array, editing the supplied `arch_code_a` had the same defect. Measured, the forward output changed on `dev` but not on the previous commit -- a silent regression. Gate the snapshot on `torch.compiler.is_compiling()`, which dynamo folds at trace time: the compiled graph still traces to 1 graph with 0 breaks, while eager reads the live tensor and behaves exactly as before. TorchScript cannot type `Tensor.tolist()`, so it gets the tensor behind a `torch.jit.is_scripting()` guard that is pruned at script time. Add regression tests for both mutation paths. Signed-off-by: Patel, Nilaykumar K --- monai/networks/nets/dints.py | 30 ++++++++------ tests/networks/nets/test_dints_network.py | 49 +++++++++++++++++++---- 2 files changed, 59 insertions(+), 20 deletions(-) diff --git a/monai/networks/nets/dints.py b/monai/networks/nets/dints.py index 54a7480356f..77cb38a6c52 100644 --- a/monai/networks/nets/dints.py +++ b/monai/networks/nets/dints.py @@ -485,12 +485,7 @@ def __init__( def __setattr__(self, name: str, value: Any) -> None: super().__setattr__(name, value) if name == "node_a" and value is not None: - # `forward` branches on these flags. Reading them from a tensor forces a graph - # break under torch.compile, so mirror them into a plain Python list that can be - # constant-folded. `!= 0` reproduces the truthiness of the original tensor test - # exactly; casting to int would make fractional values such as 0.5 falsy. - # This tracks rebinding only -- mutating `node_a` in place leaves the mirror - # stale, so replace the attribute rather than editing it. + # `!= 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): @@ -503,12 +498,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_py[0][d]: + if node_a[0][d]: inputs.append(x_out) else: inputs.append(torch.zeros_like(x_out)) @@ -522,7 +525,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_py[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) @@ -644,9 +647,7 @@ def __init__( def __setattr__(self, name: str, value: Any) -> None: super().__setattr__(name, value) if name == "arch_code_a" and value is not None: - # Mirrored for the same reason as `DiNTS._node_a_py`: `TopologyInstance.forward` - # branches on these flags, and reading them from a tensor breaks the - # torch.compile graph. Tracks rebinding only, not in-place mutation. + # See `DiNTS._node_a_py`. object.__setattr__(self, "_arch_code_a_py", (torch.as_tensor(value) != 0).tolist()) def forward(self, x): @@ -695,11 +696,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_py[blk_idx]): + 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 017857c45b6..b554d69eb5d 100644 --- a/tests/networks/nets/test_dints_network.py +++ b/tests/networks/nets/test_dints_network.py @@ -212,19 +212,23 @@ def test_dints_forward_tensor_arch_code(self): class TestDintsTopologyCache(unittest.TestCase): - """`forward` branches on cached Python copies of `node_a` / `arch_code_a`. + """When compiling, `forward` branches on Python snapshots of `node_a` / `arch_code_a`. - The caches exist so torch.compile can constant-fold the branches instead of breaking the - graph on a tensor read. A cache that disagreed with its source would silently change which - cells are executed, so these tests pin the equivalence. + 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): + 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) - 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 + 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, @@ -252,13 +256,42 @@ def test_cache_preserves_truthiness_not_int_value(self): 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 cache must follow.""" + """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) From f63b713f4cecfe904f686a6cf140149110481746 Mon Sep 17 00:00:00 2001 From: "Patel, Nilaykumar K" Date: Mon, 5 Oct 2026 01:05:45 -0400 Subject: [PATCH 3/3] docs(dints): explain why the __setattr__ overrides exist Both overrides mirror a topology tensor into a plain Python list and are non-obvious without the issue for context. Document the compile rationale on DiNTS.__setattr__ and point TopologyConstruction.__setattr__ at it. Addresses review feedback on #9145. Signed-off-by: Patel, Nilaykumar K --- monai/networks/nets/dints.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/monai/networks/nets/dints.py b/monai/networks/nets/dints.py index 77cb38a6c52..8bdbacc8676 100644 --- a/monai/networks/nets/dints.py +++ b/monai/networks/nets/dints.py @@ -483,6 +483,13 @@ def __init__( ) 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. @@ -645,6 +652,10 @@ def __init__( ) 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`.