Skip to content
Merged
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
43 changes: 40 additions & 3 deletions monai/networks/nets/dints.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@

import datetime
import warnings
from typing import Any

import numpy as np
import torch
Expand Down Expand Up @@ -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:
Comment thread
ericspod marked this conversation as resolved.
"""
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())
Comment thread
coderabbitai[bot] marked this conversation as resolved.

def weight_parameters(self):
return [param for name, param in self.named_parameters()]

Expand All @@ -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))
Expand All @@ -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)
Expand Down Expand Up @@ -629,6 +651,16 @@ def __init__(
self._norm_name,
)

def __setattr__(self, name: str, value: Any) -> None:
Comment thread
ericspod marked this conversation as resolved.
"""
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."""

Expand Down Expand Up @@ -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)
Expand Down
82 changes: 82 additions & 0 deletions tests/networks/nets/test_dints_network.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, _):
Expand Down
Loading