Lightning-AI / Lightning-AI/lightning-thunder

thunderfx: Applying thunderfx should return a nn.Module like object.

Open
#2,757 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

thunderfx ux
Dominant language
Python
Stars
1.5k
Forks
121
PR merge metrics
No merged PRs in 30d

Description

This should allow easily applying `thunderfx` to only a few submodules.

Repro
```python
import thunder
import torch
from typing import Callable

model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 10)
)

cm = torch.compile(model)
print(type(cm)) #
# print(cm.forward)

from thunder.dynamo import thunderfx

tm = thunderfx(model)
print(type(tm))
# Not available
# print(tm.forward) # Errors
```

Repro to replace submodules
```python

import thunder
import torch
from typing import Callable

# The logic is based on https://github.com/pytorch/ao/blob/b34c1037/torchao/quantization/quant_api.py#L230
def _replace_with_custom_fn_if_matches_filter_with_name(
model,
replacement_fn: Callable[[torch.nn.Module, str], torch.nn.Module],
filter_fn: Callable[[torch.nn.Module, str], bool],
cur_fqn="",
) -> None:
"""
Recursively replaces each child module in `model` with the result of `replacement_fn(child)`

replacement_fn (Callable[[torch.nn.Module, str], torch.nn.Module]): The function to replace matching modules.
filter_fn (Callable[[torch.nn.Module, str], bool]): The function to filter matching modules.
cur_fqn (str): The current fully qualified name of the module.

Returns:
None
"""
if filter_fn(model, cur_fqn[:-1]):
model = replacement_fn(model, cur_fqn[:-1])
return model
else:
named_children_list = list(model.named_children())
for name, child in named_children_list:
new_child = _replace_with_custom_fn_if_matches_filter_with_name(
child,
replacement_fn,
filter_fn,
f"{cur_fqn}{name}.",
)
if new_child is not child:
setattr(model, name, new_child)
return model

# works
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 10)
)
_replace_with_custom_fn_if_matches_filter_with_name(model, replacement_fn=lambda module, name: torch.compile(module), filter_fn=lambda module, name: isinstance(module, torch.nn.ReLU))

# works
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 10)
)

from thunder.dynamo import ThunderCompiler
_replace_with_custom_fn_if_matches_filter_with_name(model, replacement_fn=lambda module, name: torch.compile(module, backend=ThunderCompiler()), filter_fn=lambda module, name: isinstance(module, torch.nn.ReLU))

# doesn't work
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 10)
)
_replace_with_custom_fn_if_matches_filter_with_name(model, replacement_fn=lambda module, name: thunderfx(module), filter_fn=lambda module, name: isinstance(module, torch.nn.ReLU))

```

Alternative:
`torch.compile(backend=ThunderCompiler())` works.

cc @borda

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 at thunder.dynamo.thunderfx and compare its result with the torch.compile and ThunderCompiler examples in the issue. Reproduce the failing thunderfx submodule replacement case, then verify that the returned object exposes a forward-like interface and can replace selected modules as shown.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
compilers
Issue type
Feature
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.