huggingface / huggingface/diffusers

【BUG】Attention.head_to_batch_dim has bug in terms of tensor permutation

Aperta
#10,303 3 commenti 0 reazioni 0 assegnatari Vedi su GitHub
bug stale
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

Ricevi le nuove issue nella tua casella

Un breve riepilogo di issue GitHub adatte ai principianti.