| 9 | class RCTCDecoder(nn.Module): |
| 10 | |
| 11 | def __init__(self, |
| 12 | in_channels, |
| 13 | out_channels=6625, |
| 14 | return_feats=False, |
| 15 | **kwargs): |
| 16 | super(RCTCDecoder, self).__init__() |
| 17 | self.char_token = nn.Parameter( |
| 18 | torch.zeros([1, 1, in_channels], dtype=torch.float32), |
| 19 | requires_grad=True, |
| 20 | ) |
| 21 | trunc_normal_(self.char_token, mean=0, std=0.02) |
| 22 | self.fc = nn.Linear( |
| 23 | in_channels, |
| 24 | out_channels, |
| 25 | bias=True, |
| 26 | ) |
| 27 | self.fc_kv = nn.Linear( |
| 28 | in_channels, |
| 29 | 2 * in_channels, |
| 30 | bias=True, |
| 31 | ) |
| 32 | self.w_atten_block = Block(dim=in_channels, |
| 33 | num_heads=in_channels // 32, |
| 34 | mlp_ratio=4.0, |
| 35 | qkv_bias=False) |
| 36 | self.out_channels = out_channels |
| 37 | self.return_feats = return_feats |
| 38 | |
| 39 | def forward(self, x, data=None): |
| 40 | |