MCPcopy Create free account
hub / github.com/FunAudioLLM/Fun-Audio-Chat / eager_attention_forward

Function eager_attention_forward

funaudiochat/modeling_funaudiochat.py:77–97  ·  view source on GitHub ↗
(
    module: nn.Module,
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    attention_mask: Optional[torch.Tensor],
    scaling: float,
    dropout: float = 0.0,
    **kwargs,
)

Source from the content-addressed store, hash-verified

75
76
77def 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
100class FunAudioChatAudioAttention(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected