A basic Focal Transformer layer for one stage. Args: dim (int): Number of input channels. input_resolution (tuple[int]): Input resolution. depth (int): Number of blocks. window_size (int): Local window size. mlp_ratio (float): Ratio of mlp hidden dim to
| 215 | |
| 216 | |
| 217 | class BasicLayer(nn.Module): |
| 218 | """ A basic Focal Transformer layer for one stage. |
| 219 | |
| 220 | Args: |
| 221 | dim (int): Number of input channels. |
| 222 | input_resolution (tuple[int]): Input resolution. |
| 223 | depth (int): Number of blocks. |
| 224 | window_size (int): Local window size. |
| 225 | mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. |
| 226 | qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True |
| 227 | qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. |
| 228 | drop (float, optional): Dropout rate. Default: 0.0 |
| 229 | drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0 |
| 230 | norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm |
| 231 | downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None |
| 232 | use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False. |
| 233 | focal_level (int): Number of focal levels |
| 234 | focal_window (int): Focal window size at first focal level |
| 235 | use_layerscale (bool): Whether use layerscale |
| 236 | layerscale_value (float): Initial layerscale value |
| 237 | use_postln (bool): Whether use layernorm after modulation |
| 238 | """ |
| 239 | |
| 240 | def __init__(self, dim, out_dim, input_resolution, depth, |
| 241 | mlp_ratio=4., drop=0., drop_path=0., norm_layer=nn.LayerNorm, |
| 242 | downsample=None, use_checkpoint=False, |
| 243 | focal_level=1, focal_window=1, |
| 244 | use_conv_embed=False, |
| 245 | use_layerscale=False, layerscale_value=1e-4, |
| 246 | use_postln=False, |
| 247 | use_postln_in_modulation=False, |
| 248 | normalize_modulator=False): |
| 249 | |
| 250 | super().__init__() |
| 251 | self.dim = dim |
| 252 | self.input_resolution = input_resolution |
| 253 | self.depth = depth |
| 254 | self.use_checkpoint = use_checkpoint |
| 255 | |
| 256 | # build blocks |
| 257 | self.blocks = nn.ModuleList([ |
| 258 | FocalNetBlock( |
| 259 | dim=dim, |
| 260 | input_resolution=input_resolution, |
| 261 | mlp_ratio=mlp_ratio, |
| 262 | drop=drop, |
| 263 | drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, |
| 264 | norm_layer=norm_layer, |
| 265 | focal_level=focal_level, |
| 266 | focal_window=focal_window, |
| 267 | use_layerscale=use_layerscale, |
| 268 | layerscale_value=layerscale_value, |
| 269 | use_postln=use_postln, |
| 270 | use_postln_in_modulation=use_postln_in_modulation, |
| 271 | normalize_modulator=normalize_modulator, |
| 272 | ) |
| 273 | for i in range(depth)]) |
| 274 |