Hacktoberfest 2026: the issues maintainers tagged for October, open and beginner-friendly. Browse Hacktoberfest issues

DiNTS forward() breaks the torch.compile graph on every topology branch

Open
#9,144 0 comments 0 reactions 0 assignees View on GitHub

Maintainers usually reply within 3 days

Nobody has claimed this yet.

Assessment

Difficulty
4/5
Estimated time
3-5 days
Newbie friendliness
68/100
Issue type
Bug
Clarity
Clearly specified
Activity status
Active
Tech stack
python, pytorch

Research direction

Start in monai/networks/nets/dints.py, focusing on DiNTS.forward and TopologyInstance.forward and their node_a/arch_code_a control flow. Run the provided TorchDynamo reproducer first, then verify that the topology behavior is preserved and compilation reaches the stated 1 graph with 0 breaks for the sample configuration.

Written by the indexing model from the issue text.

Description

Describe the bug

DiNTS.forward and TopologyInstance.forward pick which cells to run by indexing the node_a / arch_code_a tensors inside Python control flow:

# monai/networks/nets/dints.py
if self.node_a[0][d]: ...
elif self.node_a[blk_idx + 1][res_idx]: ...
for res_idx, activation in enumerate(self.arch_code_a[blk_idx].data): ...

Each is a data-dependent tensor read, so TorchDynamo cannot constant-fold the branch and splits the graph instead: a 6-block / 3-depth DiNTS compiles into 14 graphs with 13 breaks. These flags are fixed once a model is deployed, so the breaks are pure overhead — they block fusion across the network and scale with num_blocks, so deeper searched architectures fragment further.

To Reproduce

  1. Install MONAI from dev, with a torch 2.x providing TorchDynamo.

  2. Run:

    import torch, torch._dynamo as dyn
    from monai.networks.nets.dints import DiNTS, TopologyInstance
    
    grid = dict(channel_mul=0.2, num_blocks=6, num_depths=3,
                use_downsample=True, spatial_dims=3, device="cpu")
    m = DiNTS(dints_space=TopologyInstance(**grid), in_channels=1,
              num_classes=2, spatial_dims=3, use_downsample=True).eval()
    
    expl = dyn.explain(m)(torch.randn(1, 1, 32, 32, 32))
    print(expl.graph_count, expl.graph_break_count)
    
  3. Observed: 14 13. (expl.break_reasons is not populated on this torch version, so only the counts are available.)

Expected behavior

The topology flags are constant for the lifetime of a deployed model, so the branches should be constant-foldable and the network should compile to 1 0 — one graph, no breaks.
Screenshots

N/A — the reproducer output above is the full evidence.

Environment

MONAI version: 1.6.1rc0+8.g860506514
MONAI rev id : 860506514fb1fba41db1578e0eeacfa952583a13   (dev)
Pytorch version: 2.12.1+rocm7.1
Numpy version: 2.5.2

Reproduced on CPU. The graph breaks are introduced at Dynamo trace time, so they are independent of
the accelerator and of the torch build — this is not platform-specific.

python -c "import monai; monai.config.print_debug_info()"
================================
Printing MONAI config...
================================
MONAI version: 1.6.1rc0+8.g860506514
Numpy version: 2.5.2
Pytorch version: 2.12.1+rocm7.1
MONAI flags: HAS_EXT=False, USE_COMPILED=False, USE_META_DICT=False
MONAI rev id: 860506514fb1fba41db1578e0eeacfa952583a13
MONAI __file__: /home/user/upstream_monai/upstream-fork-monai/monai/__init__.py

Optional dependencies:
Pytorch Ignite version: 0.5.5
ITK version: NOT INSTALLED or UNKNOWN VERSION.
Nibabel version: 5.4.2
scikit-image version: 0.26.0
scipy version: 1.18.1
Pillow version: 12.3.0
Tensorboard version: NOT INSTALLED or UNKNOWN VERSION.
gdown version: NOT INSTALLED or UNKNOWN VERSION.
TorchVision version: 0.27.1+rocm7.1
tqdm version: 4.70.1
lmdb version: NOT INSTALLED or UNKNOWN VERSION.
psutil version: NOT INSTALLED or UNKNOWN VERSION.
pandas version: NOT INSTALLED or UNKNOWN VERSION.
einops version: 0.8.2
transformers version: NOT INSTALLED or UNKNOWN VERSION.
mlflow version: NOT INSTALLED or UNKNOWN VERSION.
pynrrd version: NOT INSTALLED or UNKNOWN VERSION.
clearml version: NOT INSTALLED or UNKNOWN VERSION.

================================
Printing system config...
================================
`psutil` required for `print_system_info`

================================
Printing GPU config...
================================
Num GPUs: 1
Has CUDA: True
CUDA version: None
cuDNN enabled: True
NVIDIA_TF32_OVERRIDE: None
TORCH_ALLOW_TF32_CUBLAS_OVERRIDE: None
cuDNN version: 3005001
Current device: 0
Library compiled for CUDA architectures: ['gfx900', 'gfx906', 'gfx908', 'gfx90a', 'gfx942', 'gfx1030', 'gfx1100', 'gfx1101', 'gfx1102', 'gfx1200', 'gfx1201', 'gfx950', 'gfx1150', 'gfx1151']
GPU 0 Name: AMD Instinct MI300X
GPU 0 Is integrated: False
GPU 0 Is multi GPU board: False
GPU 0 Multi processor count: 304
GPU 0 Total memory (GB): 192.0
GPU 0 CUDA capability (maj.min): 9.4
Dominant language
Python
Stars
8.7k
Forks
1.6k
Avg merge
3d 10h
Merged PRs (30d)
17

Getting set up

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

More from Project-MONAI/MONAI

All issues in Project-MONAI/MONAI

Similar issues

More Python issues

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.