Lightning-AI / Lightning-AI/pytorch-lightning

Frozen dataclass as a batch

Open
#21,577 6 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

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

Description

### Description & Motivation

So I have a frozen immutable dataclass as a batch (this is my way to control complexity), and with `bf16-true` precision I am having this error in the `HalfPrecision` which calls `apply_to_collection` without `allow_frozen` parameter:

```
[rank0]: ValueError: A frozen dataclass was passed to `apply_to_collection` but this is not allowed.
```

I would like to suggest modification of the `_apply_to_collection_slow`, so in the `is_dataclass_instance(data):` you could check for the `to()` method in the dataclass with `hasattr(data, 'to')` so you could call just `data.to(dtype=dtype)` and let the user to handle the situation with the conversion themselves.

So I could add to my batch the new method:

```python
import dataclasses
import typing

@dataclasses.dataclass(frozen=True)
class MyBatch:

input: torch.Tensor

def to(self, device: torch.device, dtype: torch.dtype, non_blocking: bool = False, dataset_idx: int = 0) -> typing.Self:
# I am handling conversion myself
return MyBatch(
input=input.to(device=device, dtype=dtype, non_blocking=non_blocking),
)
```

This feature also allow to get rid of the `transfer_batch_to_device()` callback, as the batch allows to move to the other device itself.

### Pitch

The feature allows to provide more flexibility and keep the code neat.

### Alternatives

- make dataclass mutable (unfrozen) - it is possible, but makes the code fragile
- convert dataclass to dict and back - two additional operations

### Additional context

_No response_

cc @lantiga

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 tracing HalfPrecision and its call to apply_to_collection, then inspect _apply_to_collection_slow for the frozen-dataclass handling. Determine how a dataclass to() method should interact with dtype and device conversion, and verify that frozen batches can be processed without the ValueError while preserving existing behavior for other dataclasses.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning
Issue type
Feature
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
48/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.