MCPcopy Create free account
hub / github.com/deepbrainai-research/float / Attention

Class Attention

models/float/FMT.py:39–87  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37
38
39class Attention(nn.Module):
40 def __init__(
41 self,
42 dim: int,
43 num_heads: int = 8,
44 qkv_bias: bool = False,
45 qk_norm: bool = False,
46 attn_drop: float = 0.,
47 proj_drop: float = 0.,
48 norm_layer: nn.Module = nn.LayerNorm,
49 ) -> None:
50
51 super().__init__()
52 assert dim % num_heads == 0, 'dim should be divisible by num_heads'
53 self.num_heads = num_heads
54 self.head_dim = dim // num_heads
55 self.scale = self.head_dim ** -0.5
56 self.fused_attn = use_fused_attn()
57
58 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
59 self.q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
60 self.k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
61 self.attn_drop = nn.Dropout(attn_drop)
62 self.proj = nn.Linear(dim, dim)
63 self.proj_drop = nn.Dropout(proj_drop)
64
65 def forward(self, x: torch.Tensor, mask: torch.Tensor = None) -> torch.Tensor:
66 B, N, C = x.shape
67 qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
68 q, k, v = qkv.unbind(0)
69 q, k = self.q_norm(q), self.k_norm(k)
70
71 if self.fused_attn:
72 x = F.scaled_dot_product_attention(
73 q, k, v,
74 attn_mask = ~mask,
75 dropout_p=self.attn_drop.p if self.training else 0.,
76 )
77 else:
78 q = q * self.scale
79 attn = q @ k.transpose(-2, -1)
80 attn = attn.softmax(dim=-1)
81 attn = self.attn_drop(attn)
82 x = attn @ v
83
84 x = x.transpose(1, 2).reshape(B, N, C)
85 x = self.proj(x)
86 x = self.proj_drop(x)
87 return x
88
89class TimestepEmbedder(nn.Module):
90 """

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected