Lightning-AI / Lightning-AI/pytorch-lightning

Refactor throughput dtype inference into the Precision plugin API

Open
#21,535 1 comment 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

refactor
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Outline & Motivation

While looking into throughput monitoring, I noticed that dtype inference currently relies on explicit isinstance(...) checks against various precision plugin classes.

In both Fabric and PyTorch, the helper _plugin_to_compute_dtype(...):

- Imports multiple precision plugin implementations
- Performs chained type checks
- Accesses internal attributes like _desired_input_dtype and _desired_dtype
- Returns hard-coded dtype mappings per plugin

Although this works, it tightly couples throughput utilities to specific precision plugin implementations. As a result:

- Adding new precision plugins may require modifying throughput logic
- Internal plugin attributes are accessed outside the plugin itself
- Custom precision plugins cannot integrate seamlessly without touching throughput utilities

It feels like the responsibility for exposing the compute dtype should live within the precision plugin itself rather than being inferred externally.

### Pitch

I propose introducing a small API on the Precision base class:

```
def compute_dtype(self) -> torch.dtype:
return torch.float32
```

Each built-in precision plugin would override this method to return its effective compute dtype (e.g., half precision returning torch.float16, transformer engine returning torch.int8, etc.).

Throughput utilities would then simply call:

`plugin.compute_dtype()`

This would allow us to:

- Remove all isinstance(...) checks
- Eliminate plugin-specific imports inside throughput utilities
- Stop accessing internal precision attributes from outside the plugin

The refactor would be internal and fully backward compatible since the base implementation defaults to torch.float32.

### Additional context

This change aligns well with Lightning’s plugin-based architecture philosophy — where each plugin encapsulates its own behavior and metadata.

It would also:

- Reduce maintenance overhead when adding new precision backends
- Automatically support custom precision plugins in throughput tooling
- Improve separation of concerns

If this direction makes sense, I’d be happy to open a PR implementing the change across Fabric and PyTorch precision plugins with tests.

Would love to hear your thoughts.

Contributor guide

Open the contributing guide

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 locating _plugin_to_compute_dtype(...) in the Fabric and PyTorch throughput utilities, then read the Precision base class and the built-in precision plugin implementations it currently checks. Done means dtype inference is delegated through the proposed plugin API, plugin-specific checks and internal attribute access are removed, and the affected tests cover the built-in behaviors.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, performance
Issue type
Refactor
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.