(
self,
dim: int,
num_heads: int,
mlp_ratio: float = 4.0,
qkv_bias: bool = True,
proj_bias: bool = True,
ffn_bias: bool = True,
drop: float = 0.0,
attn_drop: float = 0.0,
init_values=None,
drop_path: float = 0.0,
act_layer: Callable[..., nn.Module] = nn.GELU,
norm_layer: Callable[..., nn.Module] = nn.LayerNorm,
attn_class: Callable[..., nn.Module] = Attention,
ffn_layer: Callable[..., nn.Module] = Mlp,
qk_norm: bool = False,
fused_attn: bool = True, # use F.scaled_dot_product_attention or not
rope=None,
ttt_mode=True,
ttt_params=None
)
| 20 | |
| 21 | class Block(nn.Module): |
| 22 | def __init__( |
| 23 | self, |
| 24 | dim: int, |
| 25 | num_heads: int, |
| 26 | mlp_ratio: float = 4.0, |
| 27 | qkv_bias: bool = True, |
| 28 | proj_bias: bool = True, |
| 29 | ffn_bias: bool = True, |
| 30 | drop: float = 0.0, |
| 31 | attn_drop: float = 0.0, |
| 32 | init_values=None, |
| 33 | drop_path: float = 0.0, |
| 34 | act_layer: Callable[..., nn.Module] = nn.GELU, |
| 35 | norm_layer: Callable[..., nn.Module] = nn.LayerNorm, |
| 36 | attn_class: Callable[..., nn.Module] = Attention, |
| 37 | ffn_layer: Callable[..., nn.Module] = Mlp, |
| 38 | qk_norm: bool = False, |
| 39 | fused_attn: bool = True, # use F.scaled_dot_product_attention or not |
| 40 | rope=None, |
| 41 | ttt_mode=True, |
| 42 | ttt_params=None |
| 43 | ) -> None: |
| 44 | super().__init__() |
| 45 | |
| 46 | self.ttt_mode = ttt_mode |
| 47 | |
| 48 | self.norm1 = norm_layer(dim) |
| 49 | if not self.ttt_mode: |
| 50 | self.attn = attn_class( |
| 51 | dim, |
| 52 | num_heads=num_heads, |
| 53 | qkv_bias=qkv_bias, |
| 54 | proj_bias=proj_bias, |
| 55 | attn_drop=attn_drop, |
| 56 | proj_drop=drop, |
| 57 | qk_norm=qk_norm, |
| 58 | fused_attn=fused_attn, |
| 59 | rope=rope, |
| 60 | ) |
| 61 | else: |
| 62 | self.ttt_block = FastWeightGluMLPMultihead( |
| 63 | dim, |
| 64 | **ttt_params |
| 65 | ) |
| 66 | self.ls1 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() |
| 67 | |
| 68 | self.drop_path1 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() |
| 69 | self.norm2 = norm_layer(dim) |
| 70 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 71 | self.mlp = ffn_layer( |
| 72 | in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop, bias=ffn_bias |
| 73 | ) |
| 74 | self.ls2 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() |
| 75 | self.drop_path2 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() |
| 76 | |
| 77 | self.sample_drop_ratio = drop_path |
| 78 | |
| 79 | def forward(self, x: Tensor, pos=None, info={}) -> Tensor: |
nothing calls this directly
no test coverage detected