Lightning-AI / Lightning-AI/lightning-thunder

no_grad is lost in jitted functions

Open
#1,486 2 comments 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

autograd
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

  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.

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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.