MCPcopy Create free account
hub / github.com/buaacxf/VIPTR / SVTRNet

Class SVTRNet

modules/SVTR.py:388–592  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

386
387
388class SVTRNet(nn.Module):
389 def __init__(
390 self,
391 img_size=[32, 100],
392 in_channels=3,
393 embed_dim=[64, 128, 256],
394 depth=[3, 6, 3],
395 num_heads=[2, 4, 8],
396 mixer=['Local'] * 6 + ['Global'] * 6, # Local atten, Global atten, Conv
397 local_mixer=[[7, 11], [7, 11], [7, 11]],
398 patch_merging='Conv', # Conv, Pool, None
399 mlp_ratio=4,
400 qkv_bias=True,
401 qk_scale=None,
402 drop_rate=0.,
403 last_drop=0.1,
404 attn_drop_rate=0.,
405 drop_path_rate=0.1,
406 norm_layer='nn.LayerNorm',
407 sub_norm='nn.LayerNorm',
408 epsilon=1e-6,
409 out_channels=192,
410 out_char_num=25,
411 block_unit='Block',
412 act='nn.GELU',
413 last_stage=True,
414 sub_num=2,
415 prenorm=False,
416 use_lenhead=False,
417 local_rank=-1,
418 **kwargs):
419 super().__init__()
420 self.img_size = img_size
421 self.embed_dim = embed_dim
422 self.out_channels = out_channels
423 self.prenorm = prenorm
424 patch_merging = None if patch_merging != 'Conv' and patch_merging != 'Pool' else patch_merging
425 self.patch_embed = PatchEmbed(
426 img_size=img_size,
427 in_channels=in_channels,
428 embed_dim=embed_dim[0],
429 sub_num=sub_num)
430 num_patches = self.patch_embed.num_patches
431 self.HW = [img_size[0] // (2 ** sub_num), img_size[1] // (2 ** sub_num)]
432 self.pos_embed = nn.Parameter(torch.zeros([1, num_patches, embed_dim[0]])).to(local_rank)
433 self.pos_drop = nn.Dropout(p=drop_rate)
434 Block_unit = eval(block_unit)
435
436 dpr = np.linspace(0, drop_path_rate, sum(depth))
437 self.blocks1 = nn.ModuleList([
438 Block_unit(
439 dim=embed_dim[0],
440 num_heads=num_heads[0],
441 local_rank=local_rank,
442 mixer=mixer[0:depth[0]][i],
443 HW=self.HW,
444 local_mixer=local_mixer[0],
445 mlp_ratio=mlp_ratio,

Callers 2

__init__Method · 0.90
SVTR.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected