`torch.compile` on a `custom_op`'s fake when using `torch.func.vjp` crashes
- Dominant language
- Python
- Stars
- 103k
- Forks
- 29.5k
- PR merge metrics
- PR metrics pending
Description
### 🐛 Describe the bug
Here's a crazy bug for you. When you use `torch.func.vjp` in the fake for a `custom_op` and then try to `torch.compile` it, the compiler crashes.
A repro is as follows:
```
import torch
@torch.library.custom_op("reproducer::simple_op", mutates_args=())
def simple_op(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""Simple op that does a matrix multiplication."""
return torch.mm(x, weight)
@simple_op.register_fake
def _(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
"""Fake implementation that calls vjp."""
def func(w: torch.Tensor) -> torch.Tensor:
return torch.mm(x, w)
output, vjp_fn = torch.func.vjp(func, weight)
return output
def test_eager_mode():
"""Test in eager mode."""
print("=== Eager mode ===")
x = torch.randn(4, 8)
weight = torch.randn(8, 16)
result = torch.ops.reproducer.simple_op(x, weight)
print(f"Result shape: {result.shape}")
def test_compile_mode():
"""Test with torch.compile."""
print("\n=== Compile mode ===")
device = "cuda" if torch.cuda.is_available() else "cpu"
x = torch.randn(4, 8, device=device)
weight = torch.randn(8, 16, device=device)
@torch.compile(fullgraph=True, backend="inductor")
def forward(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor:
return torch.ops.reproducer.simple_op(x, weight)
result = forward(x, weight)
print(f"Result shape: {result.shape}")
if __name__ == "__main__":
test_eager_mode()
test_compile_mode()
```
### Error logs
I originally managed to produce 2 errors by messing with my "real" use case that exposed the bug. One is repro'd in the above code:
```
(environment) ryan@ryan-dev-box:~/src/environment$ TORCHDYNAMO_VERBOSE=1 python tools/vjp_fake_mode_reproducer.py
=== Eager mode ===
Result shape: torch.Size([4, 16])
=== Compile mode ===
Traceback (most recent call last):
File "/home/ryan/src/environment/tools/vjp_fake_mode_reproducer.py", line 53, in
test_compile_mode()
File "/home/ryan/src/environment/tools/vjp_fake_mode_reproducer.py", line 47, in test_compile_mode
result = forward(x, weight)
^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py", line 736, in compile_wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 1495, in __call__
return self._torchdynamo_orig_callable(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 629, in __call__
return _compile(
^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 1111, in _compile
guarded_code = compile_inner(code, one_graph, hooks, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_utils_internal.py", line 97, in wrapper_function
return function(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 793, in compile_inner
return _compile_inner(code, one_graph, hooks, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 832, in _compile_inner
out_code = transform_code_object(code, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/bytecode_transformation.py", line 1424, in transform_code_object
transformations(instructions, code_options)
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 267, in _fn
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py", line 753, in transform
tracer.run()
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 3497, in run
super().run()
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 1363, in run
while self.step():
^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 1267, in step
self.dispatch_table[inst.opcode](self, inst)
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 834, in wrapper
return inner_fn(self, inst)
^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 2910, in CALL
self._call(inst)
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py", line 2904, in _call
self.call_function(fn, args, kwargs)
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/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 "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/variables/lazy.py", line 201, in realize_and_forward
return getattr(self.realize(), name)(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/variables/torch.py", line 1338, in call_function
tensor_variable = wrap_fx_proxy(
^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/variables/builder.py", line 2559, in wrap_fx_proxy
return wrap_fx_proxy_cls(target_cls=TensorVariable, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/variables/builder.py", line 2625, in wrap_fx_proxy_cls
return _wrap_fx_proxy(
^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/variables/builder.py", line 2723, in _wrap_fx_proxy
example_value = get_fake_value(proxy.node, tx, allow_non_graph_fake=True)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 3355, in get_fake_value
raise TorchRuntimeError(str(e)).with_traceback(e.__traceback__) from None
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 3253, in get_fake_value
ret_val = wrap_fake_exception(
^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 2753, in wrap_fake_exception
return fn()
^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 3254, in
lambda: run_node(tx.output, node, args, kwargs, nnmodule)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 3462, in run_node
raise RuntimeError(make_error_message(e)).with_traceback(
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/utils.py", line 3421, in run_node
return node.target(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_ops.py", line 1243, in __call__
return self._op(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py", line 111, in autograd_impl
result = forward_no_grad(*args, Metadata(keyset, keyword_only_args))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py", line 40, in forward_no_grad
result = op.redispatch(keyset & _C._after_autograd_keyset, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_ops.py", line 836, in redispatch
return self._handle.redispatch_boxed(keyset, *args, **kwargs) # type: ignore[return-value]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/utils/_stats.py", line 28, in wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 1352, in __torch_dispatch__
return self.dispatch(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2058, in dispatch
return self._cached_dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 1487, in _cached_dispatch_impl
output = self._dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2658, in _dispatch_impl
result = maybe_fake_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/utils.py", line 32, in __call__
return self.func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/library.py", line 1436, in inner
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/custom_ops.py", line 632, in fake_impl
return self._abstract_fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/src/environment/tools/vjp_fake_mode_reproducer.py", line 23, in _
output, vjp_fn = torch.func.vjp(func, weight)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/eager_transforms.py", line 300, in vjp
return _vjp_with_argnums(func, *primals, has_aux=has_aux)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/vmap.py", line 48, in fn
return f(*args, **kwargs)
^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/eager_transforms.py", line 358, in _vjp_with_argnums
primals_out = func(*primals)
^^^^^^^^^^^^^^
File "/home/ryan/src/environment/tools/vjp_fake_mode_reproducer.py", line 21, in func
return torch.mm(x, w)
^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/utils/_stats.py", line 28, in wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 1352, in __torch_dispatch__
return self.dispatch(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2058, in dispatch
return self._cached_dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 1457, in _cached_dispatch_impl
return self._dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2352, in _dispatch_impl
(flat_args, flat_arg_fake_tensors) = self.validate_and_convert_non_fake_tensors(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2803, in validate_and_convert_non_fake_tensors
validated_args = [validate(a) for a in flat_args]
^^^^^^^^^^^
File "/home/ryan/anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py", line 2791, in validate
raise AssertionError(
torch._dynamo.exc.TorchRuntimeError: Dynamo failed to run FX node with fake tensors: call_function reproducer.simple_op(*(FakeTensor(..., device='cuda:0', size=(4, 8)), FakeTensor(..., device='cuda:0', size=(8, 16))), **{}): got AssertionError("Please convert all Tensors to FakeTensors first or instantiate FakeTensorMode with 'allow_non_fake_inputs'. Found in aten.mm.default(FakeTensor(..., device='cuda:0', size=(4, 8)), GradTrackingTensor(lvl=1, value=\n FakeTensor(..., device='cuda:0', size=(8, 16))\n))")
from user code:
File "/home/ryan/src/environment/tools/vjp_fake_mode_reproducer.py", line 45, in forward
return torch.ops.reproducer.simple_op(x, weight)
```
The other error message I didn't find a specific repro for is as shown here:
```
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py:749: in compile_wrapper
raise e.remove_dynamo_frames() from None # see TORCHDYNAMO_VERBOSE=1
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py:736: in compile_wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:1495: in __call__
return self._torchdynamo_orig_callable(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:629: in __call__
return _compile(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:1111: in _compile
guarded_code = compile_inner(code, one_graph, hooks, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_utils_internal.py:97: in wrapper_function
return function(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:793: in compile_inner
return _compile_inner(code, one_graph, hooks, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:832: in _compile_inner
out_code = transform_code_object(code, transform)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/bytecode_transformation.py:1424: in transform_code_object
transformations(instructions, code_options)
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:267: in _fn
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/convert_frame.py:753: in transform
tracer.run()
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py:3497: in run
super().run()
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py:1363: in run
while self.step():
^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py:1267: in step
self.dispatch_table[inst.opcode](self, inst)
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py:3672: in RETURN_VALUE
self._return(inst)
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/symbolic_convert.py:3653: in _return
all_stack_locals_metadata = self.output.compile_subgraph(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/output_graph.py:1422: in compile_subgraph
self.compile_and_call_fx_graph(tx, pass2.graph_output_vars(), root)
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/output_graph.py:1696: in compile_and_call_fx_graph
compiled_fn = self.call_user_compiler(gm, self.example_inputs())
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/output_graph.py:1811: in call_user_compiler
return self._call_user_compiler(gm, example_inputs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/output_graph.py:1871: in _call_user_compiler
raise BackendCompilerFailed(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/output_graph.py:1846: in _call_user_compiler
compiled_fn = compiler_fn(gm, example_inputs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/repro/after_dynamo.py:150: in __call__
compiled_gm = compiler_fn(gm, example_inputs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/__init__.py:2380: in __call__
return compile_fx(model_, inputs_, config_patches=self.config)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_inductor/compile_fx.py:2418: in compile_fx
return aot_autograd(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/backends/common.py:109: in __call__
cg = aot_module_simplified(gm, example_inputs, **self.kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/aot_autograd.py:1199: in aot_module_simplified
compiled_fn = AOTAutogradCache.load(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/autograd_cache.py:1140: in load
compiled_fn = dispatch_and_compile()
^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/aot_autograd.py:1184: in dispatch_and_compile
compiled_fn, _ = create_aot_dispatcher_function(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/aot_autograd.py:576: in create_aot_dispatcher_function
return _create_aot_dispatcher_function(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/aot_autograd.py:836: in _create_aot_dispatcher_function
compiled_fn, fw_metadata = compiler_fn(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/jit_compile_runtime_wrappers.py:1262: in aot_dispatch_autograd
fx_g, joint_inputs, maybe_subclass_meta = aot_dispatch_autograd_graph(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/dispatch_and_compile_graph.py:318: in aot_dispatch_autograd_graph
fx_g = _create_graph(joint_fn_to_trace, updated_joint_inputs, aot_config=aot_config)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/dispatch_and_compile_graph.py:55: in _create_graph
fx_g = make_fx(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:2318: in wrapped
return make_fx_tracer.trace(f, *args)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:2250: in trace
return self._trace_inner(f, *args)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:2221: in _trace_inner
t = dispatch_trace(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_compile.py:53: in inner
return disable_fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py:929: in _fn
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1254: in dispatch_trace
graph = tracer.trace(root, concrete_args) # type: ignore[arg-type]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py:929: in _fn
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:850: in trace
(self.create_arg(fn(*args)),),
^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:703: in flatten_fn
tree_out = root_fn(*tree_args)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1312: in wrapped
out = f(*tensors) # type:ignore[call-arg]
^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/traced_function_transforms.py:720: in inner_fn
outs = fn(*args)
^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/traced_function_transforms.py:671: in joint_helper
return _functionalized_f_helper(primals, tangents)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/traced_function_transforms.py:419: in _functionalized_f_helper
f_outs = fn(*f_args)
^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/traced_function_transforms.py:286: in inner_fn_with_anomaly
return inner_fn(*args)
^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/_aot_autograd/traced_function_transforms.py:271: in inner_fn
backward_out = torch.autograd.grad(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/autograd/__init__.py:452: in grad
return handle_torch_function(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/overrides.py:1725: in handle_torch_function
result = mode.__torch_function__(public_api, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1360: in __torch_function__
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/autograd/__init__.py:503: in grad
result = _engine_run_backward(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/autograd/graph.py:829: in _engine_run_backward
return Variable._execution_engine.run_backward( # Calls into the C++ engine to run the backward pass
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/autograd/function.py:311: in apply
return user_fn(self, *args)
^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py:181: in new_backward
grad_inputs = orig_backward(ctx, *grads)
^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py:77: in backward
result = info._backward_fn(ctx, *grads)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
environment/impl/nn/parallel/parallel_execute.py:398: in parallel_execute_backward
backward_result = torch.ops.environment.parallel_execute(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_ops.py:1243: in __call__
return self._op(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py:111: in autograd_impl
result = forward_no_grad(*args, Metadata(keyset, keyword_only_args))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/autograd.py:40: in forward_no_grad
result = op.redispatch(keyset & _C._after_autograd_keyset, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_ops.py:836: in redispatch
return self._handle.redispatch_boxed(keyset, *args, **kwargs) # type: ignore[return-value]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/functional_tensor.py:511: in __torch_dispatch__
outs_unwrapped = func._op_dk(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/utils/_stats.py:28: in wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1462: in __torch_dispatch__
return proxy_call(self, func, self.pre_dispatch, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:914: in proxy_call
out = func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_ops.py:829: in __call__
return self._op(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/utils/_stats.py:28: in wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:1352: in __torch_dispatch__
return self.dispatch(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2058: in dispatch
return self._cached_dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:1487: in _cached_dispatch_impl
output = self._dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2658: in _dispatch_impl
result = maybe_fake_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/utils.py:32: in __call__
return self.func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/library.py:1431: in inner
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_library/custom_ops.py:632: in fake_impl
return self._abstract_fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
environment/impl/nn/parallel/parallel_execute.py:188: in parallel_execute_fake
return _parallel_execute_impl(
environment/impl/nn/parallel/parallel_execute.py:57: in _parallel_execute_impl
result = module(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:825: in module_call_wrapper
return self.call_module(mod, forward, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1048: in call_module
return forward(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:818: in forward
return _orig_module_call(mod, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1773: in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1784: in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
environment/impl/nn/parallel/parallel_execution_registrar.py:269: in forward
return self.__forward_impl(
environment/impl/nn/parallel/parallel_execution_registrar.py:322: in __forward_impl
raw_results[i] = registrar.get(self.__derivative_order)(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:825: in module_call_wrapper
return self.call_module(mod, forward, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1048: in call_module
return forward(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:818: in forward
return _orig_module_call(mod, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1773: in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1784: in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
environment/impl/nn/parallel/checkpoint_registrar.py:104: in forward
_, vjp_fn = torch.func.vjp(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/eager_transforms.py:300: in vjp
return _vjp_with_argnums(func, *primals, has_aux=has_aux)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/vmap.py:48: in fn
return f(*args, **kwargs)
^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/eager_transforms.py:358: in _vjp_with_argnums
primals_out = func(*primals)
^^^^^^^^^^^^^^
environment/impl/nn/parallel/checkpoint_registrar.py:93: in functional_wrapper
return self.__prior_module(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:825: in module_call_wrapper
return self.call_module(mod, forward, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1048: in call_module
return forward(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:818: in forward
return _orig_module_call(mod, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1773: in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1784: in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
environment/impl/nn/parallel/checkpoint_registrar.py:50: in forward
return torch.func.functional_call(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_functorch/functional_call.py:148: in functional_call
return nn.utils.stateless._functional_call(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/utils/stateless.py:282: in _functional_call
return module(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:825: in module_call_wrapper
return self.call_module(mod, forward, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1048: in call_module
return forward(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:818: in forward
return _orig_module_call(mod, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1773: in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1784: in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/container.py:244: in forward
input = module(input)
^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:825: in module_call_wrapper
return self.call_module(mod, forward, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/experimental/proxy_tensor.py:1048: in call_module
return forward(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/fx/_symbolic_trace.py:818: in forward
return _orig_module_call(mod, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1773: in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/module.py:1784: in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/nn/modules/linear.py:125: in forward
return F.linear(input, self.weight, self.bias)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/utils/_stats.py:28: in wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:1352: in __torch_dispatch__
return self.dispatch(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2058: in dispatch
return self._cached_dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:1457: in _cached_dispatch_impl
return self._dispatch_impl(func, types, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2352: in _dispatch_impl
(flat_args, flat_arg_fake_tensors) = self.validate_and_convert_non_fake_tensors(
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2803: in validate_and_convert_non_fake_tensors
validated_args = [validate(a) for a in flat_args]
^^^^^^^^^^^
_ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _
x = GradTrackingTensor(lvl=-2, value=
FakeTensor(..., device='cuda:0', size=(4, 8))
)
def validate(x: T) -> Union[T, FakeTensor]:
if not isinstance(x, Tensor):
return x
nonlocal flat_arg_fake_tensors
if not self.is_our_fake(x):
if hasattr(func, "tags") and torch.Tag.inplace_view in func.tags:
args, kwargs = pytree.tree_unflatten(flat_args, args_spec)
raise AssertionError(
f"Can't call metadata mutating ops on non-Fake Tensor inputs. Found in {render_call(func, args, kwargs)}"
)
allow_non_fake_inputs = (
self.allow_non_fake_inputs
if fake_tensor_tls.allow_non_fake_inputs_override is None
else fake_tensor_tls.allow_non_fake_inputs_override
)
if not allow_non_fake_inputs:
if isinstance(x, FakeTensor) and x.fake_mode is not self:
raise AssertionError("Mixing fake modes NYI")
args, kwargs = pytree.tree_unflatten(flat_args, args_spec)
> raise AssertionError(
f"Please convert all Tensors to FakeTensors first or instantiate FakeTensorMode "
f"with 'allow_non_fake_inputs'. Found in {render_call(func, args, kwargs)}"
)
E torch._dynamo.exc.BackendCompilerFailed: backend='inductor' raised:
E AssertionError: Please convert all Tensors to FakeTensors first or instantiate FakeTensorMode with 'allow_non_fake_inputs'. Found in aten.linear.default(GradTrackingTensor(lvl=1, value=
E FakeTensor(..., device='cuda:0', size=(4, 8))
E ), GradTrackingTensor(lvl=1, value=
E Parameter(FakeTensor(..., device='cuda:0', size=(16, 8), requires_grad=True))
E ), GradTrackingTensor(lvl=1, value=
E Parameter(FakeTensor(..., device='cuda:0', size=(16,), requires_grad=True))
E ))
../../anaconda3/envs/environment/lib/python3.12/site-packages/torch/_subclasses/fake_tensor.py:2791: BackendCompilerFailed
```
### Versions
(output re-used from previously filed bug, as environment hasnt changed)
```
python collect_env.py
--2025-11-21 18:09:41-- https://raw.githubusercontent.com/pytorch/pytorch/main/torch/utils/collect_env.py
Resolving raw.githubusercontent.com (raw.githubusercontent.com)... 185.199.109.133, 185.199.110.133, 185.199.108.133, ...
Connecting to raw.githubusercontent.com (raw.githubusercontent.com)|185.199.109.133|:443... connected.
HTTP request sent, awaiting response... 200 OK
Length: 30662 (30K) [text/plain]
Saving to: ‘collect_env.py.1’
collect_env.py.1 100%[===================>] 29.94K --.-KB/s in 0s
2025-11-21 18:09:41 (123 MB/s) - ‘collect_env.py.1’ saved [30662/30662]
Collecting environment information...
PyTorch version: 2.8.0+cu129
Is debug build: False
CUDA used to build PyTorch: 12.9
ROCM used to build PyTorch: N/A
OS: Ubuntu 25.10 (x86_64)
GCC version: (Ubuntu 15.2.0-4ubuntu4) 15.2.0
Clang version: Could not collect
CMake version: Could not collect
Libc version: glibc-2.42
Python version: 3.12.0 | packaged by Anaconda, Inc. | (main, Oct 2 2023, 17:29:18) [GCC 11.2.0] (64-bit runtime)
Python platform: Linux-6.17.0-6-generic-x86_64-with-glibc2.42
Is CUDA available: True
CUDA runtime version: 12.9.86
CUDA_MODULE_LOADING set to: LAZY
GPU models and configuration: GPU 0: NVIDIA GeForce RTX 5060 Ti
Nvidia driver version: 580.95.05
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): 12
On-line CPU(s) list: 0-11
Vendor ID: AuthenticAMD
Model name: AMD Ryzen 5 7600X 6-Core Processor
CPU family: 25
Model: 97
Thread(s) per core: 2
Core(s) per socket: 6
Socket(s): 1
Stepping: 2
Frequency boost: enabled
CPU(s) scaling MHz: 70%
CPU max MHz: 5457.1050
CPU min MHz: 427.3640
BogoMIPS: 9381.80
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 amd_lbr_v2 nopl xtopology nonstop_tsc cpuid extd_apicid aperfmperf rapl pni pclmulqdq monitor ssse3 fma cx16 sse4_1 sse4_2 movbe popcnt aes xsave avx f16c rdrand lahf_lm cmp_legacy svm extapic cr8_legacy abm sse4a misalignsse 3dnowprefetch osvw ibs skinit wdt tce topoext perfctr_core perfctr_nb bpext perfctr_llc mwaitx cpuid_fault cpb cat_l3 cdp_l3 hw_pstate ssbd mba perfmon_v2 ibrs ibpb stibp ibrs_enhanced vmmcall fsgsbase bmi1 avx2 smep bmi2 erms invpcid cqm rdt_a avx512f avx512dq rdseed adx smap avx512ifma clflushopt clwb avx512cd sha_ni avx512bw avx512vl xsaveopt xsavec xgetbv1 xsaves cqm_llc cqm_occup_llc cqm_mbm_total cqm_mbm_local user_shstk avx512_bf16 clzero irperf xsaveerptr rdpru wbnoinvd cppc arat npt lbrv svm_lock nrip_save tsc_scale vmcb_clean flushbyasid decodeassists pausefilter pfthreshold avic vgif x2avic v_spec_ctrl vnmi avx512vbmi umip pku ospke avx512_vbmi2 gfni vaes vpclmulqdq avx512_vnni avx512_bitalg avx512_vpopcntdq rdpid overflow_recov succor smca fsrm flush_l1d amd_lbr_pmc_freeze
Virtualization: AMD-V
L1d cache: 192 KiB (6 instances)
L1i cache: 192 KiB (6 instances)
L2 cache: 6 MiB (6 instances)
L3 cache: 32 MiB (1 instance)
NUMA node(s): 1
NUMA node0 CPU(s): 0-11
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 Old microcode: 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; Enhanced / Automatic IBRS; IBPB conditional; STIBP always-on; 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: Mitigation; IBPB before exit to userspace
Versions of relevant libraries:
[pip3] cuequivariance-ops-torch-cu12==0.7.0
[pip3] cuequivariance-torch==0.7.0
[pip3] mypy==1.18.1
[pip3] mypy_extensions==1.1.0
[pip3] numpy==2.2.6
[pip3] nvidia-cublas-cu12==12.9.1.4
[pip3] nvidia-cuda-cupti-cu12==12.9.79
[pip3] nvidia-cuda-nvrtc-cu12==12.9.86
[pip3] nvidia-cuda-runtime-cu12==12.9.79
[pip3] nvidia-cudnn-cu12==9.10.2.21
[pip3] nvidia-cufft-cu12==11.4.1.4
[pip3] nvidia-curand-cu12==10.3.10.19
[pip3] nvidia-cusolver-cu12==11.7.5.82
[pip3] nvidia-cusparse-cu12==12.5.10.65
[pip3] nvidia-cusparselt-cu12==0.7.1
[pip3] nvidia-nccl-cu12==2.27.3
[pip3] nvidia-nvjitlink-cu12==12.9.86
[pip3] nvidia-nvtx-cu12==12.9.79
[pip3] torch==2.8.0+cu129
[pip3] torch_cluster==1.6.3+pt28cu129
[pip3] torch_geometric==2.5.2
[pip3] torch_scatter==2.1.2+pt28cu129
[pip3] torchaudio==2.8.0+cu129
[pip3] torchvision==0.23.0+cu129
[pip3] triton==3.4.0
[conda] cuequivariance-ops-torch-cu12 0.7.0 pypi_0 pypi
[conda] cuequivariance-torch 0.7.0 pypi_0 pypi
[conda] numpy 2.2.6 pypi_0 pypi
[conda] nvidia-cublas-cu12 12.9.1.4 pypi_0 pypi
[conda] nvidia-cuda-cupti-cu12 12.9.79 pypi_0 pypi
[conda] nvidia-cuda-nvrtc-cu12 12.9.86 pypi_0 pypi
[conda] nvidia-cuda-runtime-cu12 12.9.79 pypi_0 pypi
[conda] nvidia-cudnn-cu12 9.10.2.21 pypi_0 pypi
[conda] nvidia-cufft-cu12 11.4.1.4 pypi_0 pypi
[conda] nvidia-curand-cu12 10.3.10.19 pypi_0 pypi
[conda] nvidia-cusolver-cu12 11.7.5.82 pypi_0 pypi
[conda] nvidia-cusparse-cu12 12.5.10.65 pypi_0 pypi
[conda] nvidia-cusparselt-cu12 0.7.1 pypi_0 pypi
[conda] nvidia-nccl-cu12 2.27.3 pypi_0 pypi
[conda] nvidia-nvjitlink-cu12 12.9.86 pypi_0 pypi
[conda] nvidia-nvtx-cu12 12.9.79 pypi_0 pypi
[conda] torch 2.8.0+cu129 pypi_0 pypi
[conda] torch-cluster 1.6.3+pt28cu129 pypi_0 pypi
[conda] torch-geometric 2.5.2 pypi_0 pypi
[conda] torch-scatter 2.1.2+pt28cu129 pypi_0 pypi
[conda] torchaudio 2.8.0+cu129 pypi_0 pypi
[conda] torchvision 0.23.0+cu129 pypi_0 pypi
[conda] triton 3.4.0 pypi_0 pypi
```
cc @chauhang @penguinwu @Chillee @samdow @kshitij12345 @bdhirsh @bobrenjc93
Contributor guide
Assessment
This issue has not been assessed yet.