| 37 | |
| 38 | |
| 39 | class DecoderBlock(nn.Module): |
| 40 | def __init__( |
| 41 | self, |
| 42 | dim, |
| 43 | num_heads, |
| 44 | mlp_ratio=4.0, |
| 45 | qkv_bias=False, |
| 46 | proj_bias: bool = False, |
| 47 | ffn_bias: bool = True, |
| 48 | drop: float = 0.0, |
| 49 | attn_drop: float = 0.0, |
| 50 | init_values=None, |
| 51 | drop_path: float = 0.0, |
| 52 | act_layer: Callable[..., nn.Module] = nn.GELU, |
| 53 | norm_layer: Callable[..., nn.Module] = nn.LayerNorm, |
| 54 | self_attn_class: Callable[..., nn.Module] = Attention, |
| 55 | cross_attn_class: Callable[..., nn.Module] = CrossAttention, |
| 56 | ffn_layer: Callable[..., nn.Module] = Mlp, |
| 57 | ): |
| 58 | super().__init__() |
| 59 | self.norm1 = norm_layer(dim) |
| 60 | self.self_attn = self_attn_class( |
| 61 | dim, |
| 62 | num_heads=num_heads, |
| 63 | qkv_bias=qkv_bias, |
| 64 | proj_bias=proj_bias, |
| 65 | attn_drop=attn_drop, |
| 66 | proj_drop=drop, |
| 67 | ) |
| 68 | self.ls1 = ( |
| 69 | LayerScale(dim, init_values=init_values) if init_values else nn.Identity() |
| 70 | ) |
| 71 | self.drop_path1 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() |
| 72 | |
| 73 | self.q_norm2 = norm_layer(dim) |
| 74 | self.kv_norm2 = norm_layer(dim) |
| 75 | self.cross_attn = cross_attn_class( |
| 76 | dim, |
| 77 | num_heads=num_heads, |
| 78 | qkv_bias=qkv_bias, |
| 79 | proj_bias=proj_bias, |
| 80 | attn_drop=attn_drop, |
| 81 | proj_drop=drop, |
| 82 | ) |
| 83 | self.ls2 = ( |
| 84 | LayerScale(dim, init_values=init_values) if init_values else nn.Identity() |
| 85 | ) |
| 86 | self.drop_path2 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() |
| 87 | |
| 88 | self.norm3 = norm_layer(dim) |
| 89 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 90 | self.mlp = ffn_layer( |
| 91 | in_features=dim, |
| 92 | hidden_features=mlp_hidden_dim, |
| 93 | act_layer=act_layer, |
| 94 | drop=drop, |
| 95 | bias=ffn_bias, |
| 96 | ) |
nothing calls this directly
no outgoing calls
no test coverage detected