MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / BasicLayer

Class BasicLayer

Image_Classification/src/models/focalnet.py:217–312  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

215
216
217class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected