Lightning-AI / Lightning-AI/lightning-thunder
Excessive caching for `torch.dtype` objects
Open
@t-vi is already working on this.
Since May 13, 2024.
bug
jit
- Dominant language
- Python
- Stars
- 1.5k
- Forks
- 121
- PR merge metrics
- No merged PRs in 30d
Description
🐛 Bug
It seems that when passing a torch.dtype as argument of a jitted function, Thunder caches the value and the resulting rerun of the function outputs wrong values. For example:
To Reproduce
The following code:
import torch
import thunder
def to_dtype(a, dtype):
return a.to(dtype)
jit_dtype = thunder.jit(to_dtype)
a = torch.randn((1, 2), dtype=torch.float32)
print(jit_dtype(a, torch.bfloat16))
print(jit_dtype(a, torch.float16))
print(jit_dtype(a, torch.float64))
Will print 3 times a tensor of dtype bf16.
Code sample
Expected behavior
I expect the dtype to change every time, not just once.
Additional context
The issue can be solved by temporarily disabling cacheing (i.e. cache=NO_CACHING argument to thunder.jit).
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.
Assessment
This issue has not been assessed yet.