pytorch / pytorch/pytorch

pytorch while_loop / cond overly strict about memory aliasing

Open
#169,769 5 comments 0 reactions 0 assignees View on GitHub
module: higher order operators module: pt2-dispatcher oncall: pt2 triaged
Dominant language
Python
Stars
103k
Forks
29.5k
PR merge metrics
PR metrics pending

Description

Currently the while_loop will raise a memory aliasing error if you return an input tensor without changes, this requires you to clone all input tensors even if you are not going to modify them. This seems overly strict.

for example
```python
import torch
from torch._higher_order_ops.while_loop import while_loop

def no_op_while_loop(x: torch.Tensor) -> torch.Tensor:
# carried state is just a single tensor
carried = (x,)

def cond_fn(x):
# Make the loop *never* run, to emphasize that this is
# purely a static aliasing restriction, not a runtime issue.
return (x.sum() < 0) # always False for non-negative tensors

def body_fn(x):
# Intentionally return the *same* tensor, without touching it.
# This is logically safe, but violates while_loop's
# "output cannot alias inputs" restriction.
return (x,)

(out,) = while_loop(cond_fn, body_fn, carried)
return out

if __name__ == "__main__":
x = torch.ones(3)
print(no_op_while_loop(x))
```
raises
```
Traceback (most recent call last):
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/higher_order_ops.py", line 76, in graph_break_as_hard_error
return fn(*args, **kwargs)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/higher_order_ops.py", line 1362, in call_function
) = speculate_subgraph(
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/higher_order_ops.py", line 863, in speculate_subgraph
raise ex
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/higher_order_ops.py", line 835, in speculate_subgraph
unimplemented_v2(
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/exc.py", line 528, in unimplemented_v2
raise Unsupported(msg)
torch._dynamo.exc.Unsupported: Encountered aliasing during higher order op tracing
Explanation: Higher order ops do not support aliasing. Found in while_loop
Hint: Consider using the debug context to change user code to avoid aliasing.
Hint: Please open an issue.

Developer debug context: Input-to-output aliasing detected at nodes l_args_2_0_ and l_args_2_0_ in
graph():
%l_args_2_0_ : torch.Tensor [num_users=1] = placeholder[target=l_args_2_0_]
return (l_args_2_0_,)

The above exception was the direct cause of the following exception:

Traceback (most recent call last):
File "/workspaces/data-engine/jobs/model-analysis/notebooks/temp/differs.py", line 25, in
print(no_op_while_loop(x))
File "/workspaces/data-engine/jobs/model-analysis/notebooks/temp/differs.py", line 19, in no_op_while_loop
(out,) = while_loop(cond_fn, body_fn, carried)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_higher_order_ops/while_loop.py", line 176, in while_loop
return torch.compile(
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/eval_frame.py", line 736, in compile_wrapper
return fn(*args, **kwargs)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 1495, in __call__
return self._torchdynamo_orig_callable(
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 629, in __call__
return _compile(
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 1111, in _compile
guarded_code = compile_inner(code, one_graph, hooks, transform)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_utils_internal.py", line 97, in wrapper_function
return function(*args, **kwargs)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 793, in compile_inner
return _compile_inner(code, one_graph, hooks, transform)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 832, in _compile_inner
out_code = transform_code_object(code, transform)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/bytecode_transformation.py", line 1424, in transform_code_object
transformations(instructions, code_options)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 267, in _fn
return fn(*args, **kwargs)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/convert_frame.py", line 753, in transform
tracer.run()
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 3497, in run
super().run()
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 1363, in run
while self.step():
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 1267, in step
self.dispatch_table[inst.opcode](self, inst)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 834, in wrapper
return inner_fn(self, inst)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 2228, in CALL_FUNCTION_EX
self.call_function(fn, argsvars.items, kwargsvars)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/symbolic_convert.py", line 1193, in call_function
self.push(fn.call_function(self, args, kwargs)) # type: ignore[arg-type]
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/lazy.py", line 201, in realize_and_forward
return getattr(self.realize(), name)(*args, **kwargs)
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_dynamo/variables/higher_order_ops.py", line 79, in graph_break_as_hard_error
raise UncapturedHigherOrderOpError(reason + msg) from e
torch._dynamo.exc.UncapturedHigherOrderOpError: while_loop doesn't work unless it is captured completely with torch.compile. Scroll up to find out what causes the graph break.

from user code:
File "/workspaces/data-engine/jobs/model-analysis/.venv/lib/python3.10/site-packages/torch/_higher_order_ops/while_loop.py", line 167, in _while_loop_op_wrapper
return while_loop_op(*args, **kwargs)
```
this would therefore require that I clone all the inputs that are unchanged which is very costly. Please loosen the restrictions.

```
PyTorch version: 2.9.1+cu128
Is debug build: False
CUDA used to build PyTorch: 12.8
ROCM used to build PyTorch: N/A

OS: Debian GNU/Linux 12 (bookworm) (x86_64)
GCC version: (Debian 12.2.0-14+deb12u1) 12.2.0
Clang version: Could not collect
CMake version: Could not collect
Libc version: glibc-2.36

Python version: 3.10.18 (main, Jun 7 2025, 14:54:56) [GCC 12.2.0] (64-bit runtime)
Python platform: Linux-6.14.0-1017-aws-x86_64-with-glibc2.36
Is CUDA available: True
CUDA runtime version: Could not collect
CUDA_MODULE_LOADING set to:
GPU models and configuration: GPU 0: NVIDIA L40S
Nvidia driver version: 580.65.06
cuDNN version: Could not collect
Is XPU available: False
HIP runtime version: N/A
MIOpen runtime version: N/A
Is XNNPACK available: True

CPU:
Architecture: x86_64
CPU op-mode(s): 32-bit, 64-bit
Address sizes: 48 bits physical, 48 bits virtual
Byte Order: Little Endian
CPU(s): 8
On-line CPU(s) list: 0-7
Vendor ID: AuthenticAMD
Model name: AMD EPYC 7R13 Processor
CPU family: 25
Model: 1
Thread(s) per core: 2
Core(s) per socket: 4
Socket(s): 1
Stepping: 1
BogoMIPS: 5299.99
Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr sse sse2 ht syscall nx mmxext fxsr_opt pdpe1gb rdtscp lm constant_tsc rep_good nopl xtopology nonstop_tsc cpuid extd_apicid aperfmperf tsc_known_freq pni pclmulqdq ssse3 fma cx16 pcid sse4_1 sse4_2 x2apic movbe popcnt aes xsave avx f16c rdrand hypervisor lahf_lm cmp_legacy cr8_legacy abm sse4a misalignsse 3dnowprefetch topoext ssbd ibrs ibpb stibp vmmcall fsgsbase bmi1 avx2 smep bmi2 invpcid rdseed adx smap clflushopt clwb sha_ni xsaveopt xsavec xgetbv1 clzero xsaveerptr rdpru wbnoinvd arat npt nrip_save vaes vpclmulqdq rdpid
Hypervisor vendor: KVM
Virtualization type: full
L1d cache: 128 KiB (4 instances)
L1i cache: 128 KiB (4 instances)
L2 cache: 2 MiB (4 instances)
L3 cache: 16 MiB (1 instance)
NUMA node(s): 1
NUMA node0 CPU(s): 0-7
Vulnerability Gather data sampling: Not affected
Vulnerability Ghostwrite: Not affected
Vulnerability Indirect target selection: Not affected
Vulnerability Itlb multihit: Not affected
Vulnerability L1tf: Not affected
Vulnerability Mds: Not affected
Vulnerability Meltdown: Not affected
Vulnerability Mmio stale data: Not affected
Vulnerability Reg file data sampling: Not affected
Vulnerability Retbleed: Not affected
Vulnerability Spec rstack overflow: Mitigation; Safe RET
Vulnerability Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl
Vulnerability Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization
Vulnerability Spectre v2: Mitigation; Retpolines; IBPB conditional; IBRS_FW; STIBP always-on; RSB filling; PBRSB-eIBRS Not affected; BHI Not affected
Vulnerability Srbds: Not affected
Vulnerability Tsa: Mitigation; Clear CPU buffers
Vulnerability Tsx async abort: Not affected
Vulnerability Vmscape: Not affected

Versions of relevant libraries:
[pip3] mypy==1.19.0
[pip3] mypy_extensions==1.1.0
[pip3] numpy==1.26.4
[pip3] nvidia-cublas-cu12==12.8.4.1
[pip3] nvidia-cuda-cupti-cu12==12.8.90
[pip3] nvidia-cuda-nvrtc-cu12==12.8.93
[pip3] nvidia-cuda-runtime-cu12==12.8.90
[pip3] nvidia-cudnn-cu12==9.10.2.21
[pip3] nvidia-cufft-cu12==11.3.3.83
[pip3] nvidia-curand-cu12==10.3.9.90
[pip3] nvidia-cusolver-cu12==11.7.3.90
[pip3] nvidia-cusparse-cu12==12.5.8.93
[pip3] nvidia-cusparselt-cu12==0.7.1
[pip3] nvidia-nccl-cu12==2.27.5
[pip3] nvidia-nvjitlink-cu12==12.8.93
[pip3] nvidia-nvtx-cu12==12.8.90
[pip3] onnx==1.17.0
[pip3] onnxconverter-common==1.16.0
[pip3] onnxoptimizer==0.3.13
[pip3] onnxruntime==1.23.2
[pip3] onnxruntime-gpu==1.23.2
[pip3] torch==2.9.1
[pip3] torchaudio==2.9.1
[pip3] torchvision==0.24.1
[pip3] triton==3.5.1
[conda] Could not collect
```

cc @chauhang @penguinwu @ydwu4 @bdhirsh

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.