Lightning-AI / Lightning-AI/lightning-thunder
no_grad is lost in jitted functions
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 1.5k
- Forks
- 121
- PR merge metrics
- No merged PRs in 30d
Description
## 🐛 Bug
For a function decorated with `torch.no_grad`, the compile data of the jitted version has `is_grad_enabled` set to True when I would expect it to be False.
#### Code sample
```py
import torch
import thunder
@torch.no_grad
def f(x):
print(torch.is_grad_enabled())
return x * 2
jf = thunder.jit(f)
x = torch.ones((2,2))
f(x) # prints False
jf(x) # prints True
thunder.compile_data(jf).is_grad_enabled # True
```
Contributor guide
No contributing guide indexed for this repository
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 by running the provided Python example and compare f(x) with thunder.jit(f)(x), including thunder.compile_data(jf).is_grad_enabled. Then inspect the thunder.jit and compile_data entry points for how torch.no_grad state is carried into compiled functions; done means the jitted call and compile data report gradients disabled.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- compilers, machine-learning
- Issue type
- Bug
- Difficulty
- 3/5
- Estimated time
- 1-2 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 45/100