Lightning-AI / Lightning-AI/lightning-thunder

Excessive caching for `torch.dtype` objects

Open
#406 0 comments 0 reactions 1 assignee View on GitHub

@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

  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.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.