🐛 [Bug] acc tracer doesn't handle torch.max(tensor).values correctly
Open
Nobody has claimed this yet.
bug
story: Dynamo Frontend & Partitioning
- Dominant language
- Python
- Stars
- 3k
- Forks
- 410
- Avg merge
- 3d 18h
- Merged PRs (30d)
- 78
Description
Bug Description
acc tracer doesn't handle torch.max(tensor).values correctly
To Reproduce
import torch
import torch.nn as nn
import torch_tensorrt.fx.tracer.acc_tracer.acc_tracer as acc_tracer
device = torch.device("cuda")
class MyModule(nn.Module):
def forward(self, x):
a = torch.max(torch.abs(x), dim=1).values
return a
# create an instance of the module
module = MyModule().to(device)
input_data = torch.randn(10, 10, device=device)
acc_tracer.trace(module, [input_data])
Error:
AssertionError:Expected torch.Tensor type for <class 'torch.return_types.max'>
If I switch to fx tracer, it works pretty fine.
Expected behavior
It runs okay.
Environment
trunk
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the provided MyModule reproduction and the acc_tracer.trace entry point, then inspect how torch.max(..., dim=1).values is handled. The work is done when the reproduction traces successfully without the torch.return_types.max assertion and the behavior is covered by an appropriate regression test.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- compilers, machine-learning
- Issue type
- Bug
- Difficulty
- 3/5
- Estimated time
- 1-2 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 35/100