huggingface / huggingface/diffusers
【BUG】Attention.head_to_batch_dim has bug in terms of tensor permutation
- Lingua principale
- Python
- Stelle
- 34.5k
- Fork
- 7.3k
- Merge medio
- 3g 3h
- PR unite (30g)
- 91
Descrizione
### Describe the bug
https://github.com/huggingface/diffusers/blob/1826a1e7d31df48d345a20028b3ace48f09a4e60/src/diffusers/models/attention_processor.py#L613
when `out_dim==4`, the ourpout shape is mismatch to the function's comment ``[batch_size, seq_len, heads, dim // heads]`
here is original function
```
def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor:
r"""
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is
the number of heads initialized while constructing the `Attention` class.
Args:
tensor (`torch.Tensor`): The tensor to reshape.
out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is
reshaped to `[batch_size * heads, seq_len, dim // heads]`.
Returns:
`torch.Tensor`: The reshaped tensor.
"""
head_size = self.heads
if tensor.ndim == 3:
batch_size, seq_len, dim = tensor.shape
extra_dim = 1
else:
batch_size, extra_dim, seq_len, dim = tensor.shape
tensor = tensor.reshape(batch_size, seq_len * extra_dim, head_size, dim // head_size)
tensor = tensor.permute(0, 2, 1, 3)
if out_dim == 3:
tensor = tensor.reshape(batch_size * head_size, seq_len * extra_dim, dim // head_size)
return tensor
```
and at [Line 633](https://github.com/huggingface/diffusers/blob/1826a1e7d31df48d345a20028b3ace48f09a4e60/src/diffusers/models/attention_processor.py#L633), `tensor = tensor.permute(0, 2, 1, 3)` the tensor permutes again
The correction should be moving Line633 to Line635.5 i.e.,
```
...
tensor = tensor.reshape(batch_size, seq_len * extra_dim, head_size, dim // head_size)
if out_dim == 3:
tensor = tensor.permute(0, 2, 1, 3)
tensor = tensor.reshape(batch_size * head_size, seq_len * extra_dim, dim // head_size)
return tensor
```
### Reproduction
just inside Attention, run
```
self.head_to_batch_dim(query,out_dim=4)
```
### Logs
_No response_
### System Info
This is irrelevant to the system & environment info
### Who can help?
_No response_
Guida per i contributori
Apri la guida per i contributori
Direzione di ricerca
Inizia in src/diffusers/models/attention_processor.py, nell’implementazione collegata di head_to_batch_dim, e analizza i cambiamenti nella forma del tensore attorno alla permutazione. Riproduci il problema chiamando Attention.head_to_batch_dim(query, out_dim=4), quindi verifica che la forma restituita corrisponda a quella documentata, mentre il comportamento di out_dim=3 rimanga corretto.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Valutazione
- Stack tecnologico
- python, pytorch
- Ambito
- machine-learning
- Tipo di issue
- Bug
- Difficoltà
- 3/5
- Tempo stimato
- 1-2 giorni
- Stato di attività
- Ferma
- Chiarezza
- Specificata chiaramente
- Idoneità per principianti
- 45/100