(self, d_model, n_head, window_size=None, drop_path_rate=0.0)
| 300 | |
| 301 | class ConditionalResAttBlock(nn.Module): |
| 302 | def __init__(self, d_model, n_head, window_size=None, drop_path_rate=0.0): |
| 303 | super().__init__() |
| 304 | self.window_size = window_size |
| 305 | |
| 306 | self.self_attn = MultiHeadAttention(d_model, d_model, d_model, d_model, n_head) |
| 307 | self.self_attn_ln = LayerNorm(d_model) |
| 308 | |
| 309 | self.cross_attn = MultiHeadAttention(d_model, d_model, d_model, d_model, n_head) |
| 310 | self.cross_attn_ln = LayerNorm(d_model) |
| 311 | |
| 312 | self.mlp = nn.Sequential(OrderedDict([ |
| 313 | ("c_fc", nn.Linear(d_model, d_model * 4, bias=False)), |
| 314 | ("silu", nn.SiLU(inplace=True)), |
| 315 | ("c_proj", nn.Linear(d_model * 4, d_model, bias=False)) |
| 316 | ])) |
| 317 | self.mlp_ln = LayerNorm(d_model) |
| 318 | |
| 319 | self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity() |
| 320 | |
| 321 | def window_attention(self, x, attn_layer, index): |
| 322 | attn_mask = None |
nothing calls this directly
no test coverage detected