(
module: nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_mask: Optional[torch.Tensor],
scaling: float,
dropout: float = 0.0,
**kwargs,
)
| 75 | |
| 76 | |
| 77 | def eager_attention_forward( |
| 78 | module: nn.Module, |
| 79 | query: torch.Tensor, |
| 80 | key: torch.Tensor, |
| 81 | value: torch.Tensor, |
| 82 | attention_mask: Optional[torch.Tensor], |
| 83 | scaling: float, |
| 84 | dropout: float = 0.0, |
| 85 | **kwargs, |
| 86 | ): |
| 87 | key_states = key.transpose(-1, -2) |
| 88 | attn_weights = torch.matmul(query, key_states) * scaling |
| 89 | |
| 90 | if attention_mask is not None: |
| 91 | attn_weights = attn_weights + attention_mask |
| 92 | |
| 93 | attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) |
| 94 | attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) |
| 95 | attn_output = torch.matmul(attn_weights, value) |
| 96 | |
| 97 | return attn_output, attn_weights |
| 98 | |
| 99 | |
| 100 | class FunAudioChatAudioAttention(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected