| 24 | # use attention from torch.nn.MultiHeadAttention |
| 25 | # Block contains a cross-attention layer, a self-attention layer, and a MLP |
| 26 | def __init__( |
| 27 | self, |
| 28 | inner_dim: int, |
| 29 | cond_dim: int, |
| 30 | num_heads: int, |
| 31 | eps: float, |
| 32 | attn_drop: float = 0., |
| 33 | attn_bias: bool = False, |
| 34 | mlp_ratio: float = 4., |
| 35 | mlp_drop: float = 0., |
| 36 | ): |
| 37 | super().__init__() |
| 38 | |
| 39 | self.norm1 = nn.LayerNorm(inner_dim) |
| 40 | self.cross_attn = nn.MultiheadAttention( |
| 41 | embed_dim=inner_dim, num_heads=num_heads, kdim=cond_dim, vdim=cond_dim, |
| 42 | dropout=attn_drop, bias=attn_bias, batch_first=True) |
| 43 | self.norm2 = nn.LayerNorm(inner_dim) |
| 44 | self.self_attn = nn.MultiheadAttention( |
| 45 | embed_dim=inner_dim, num_heads=num_heads, |
| 46 | dropout=attn_drop, bias=attn_bias, batch_first=True) |
| 47 | self.norm3 = nn.LayerNorm(inner_dim) |
| 48 | self.mlp = nn.Sequential( |
| 49 | nn.Linear(inner_dim, int(inner_dim * mlp_ratio)), |
| 50 | nn.GELU(), |
| 51 | nn.Dropout(mlp_drop), |
| 52 | nn.Linear(int(inner_dim * mlp_ratio), inner_dim), |
| 53 | nn.Dropout(mlp_drop), |
| 54 | ) |
| 55 | |
| 56 | def forward(self, x, cond): |
| 57 | # x: [N, L, D] |