CoolFace
Apppublic

sneedium/captcha_pixelplanet

sourceHugging Facebsdupdated 4y agoView on Hugging Face
1likes
model_language.py68 linesDownload Raw Back to modules
1import logging2import torch.nn as nn3from fastai.vision import *4 5from modules.model import _default_tfmer_cfg6from modules.model import Model7from modules.transformer import (PositionalEncoding, 8                                 TransformerDecoder,9                                 TransformerDecoderLayer)10 11 12class BCNLanguage(Model):13    def __init__(self, config):14        super().__init__(config)15        d_model = ifnone(config.model_language_d_model, _default_tfmer_cfg['d_model'])16        nhead = ifnone(config.model_language_nhead, _default_tfmer_cfg['nhead'])17        d_inner = ifnone(config.model_language_d_inner, _default_tfmer_cfg['d_inner'])18        dropout = ifnone(config.model_language_dropout, _default_tfmer_cfg['dropout'])19        activation = ifnone(config.model_language_activation, _default_tfmer_cfg['activation'])20        num_layers = ifnone(config.model_language_num_layers, 4)21        self.d_model = d_model22        self.detach = ifnone(config.model_language_detach, True)23        self.use_self_attn = ifnone(config.model_language_use_self_attn, False)24        self.loss_weight = ifnone(config.model_language_loss_weight, 1.0)25        self.max_length = config.dataset_max_length + 1  # additional stop token26        self.debug = ifnone(config.global_debug, False)27 28        self.proj = nn.Linear(self.charset.num_classes, d_model, False)29        self.token_encoder = PositionalEncoding(d_model, max_len=self.max_length)30        self.pos_encoder = PositionalEncoding(d_model, dropout=0, max_len=self.max_length)31        decoder_layer = TransformerDecoderLayer(d_model, nhead, d_inner, dropout, 32                activation, self_attn=self.use_self_attn, debug=self.debug)33        self.model = TransformerDecoder(decoder_layer, num_layers)34 35        self.cls = nn.Linear(d_model, self.charset.num_classes)36 37        if config.model_language_checkpoint is not None:38            logging.info(f'Read language model from {config.model_language_checkpoint}.')39            self.load(config.model_language_checkpoint)40 41    def forward(self, tokens, lengths):42        """43        Args:44            tokens: (N, T, C) where T is length, N is batch size and C is classes number45            lengths: (N,)46        """47        if self.detach: tokens = tokens.detach()48        embed = self.proj(tokens)  # (N, T, E)49        embed = embed.permute(1, 0, 2)  # (T, N, E)50        embed = self.token_encoder(embed)  # (T, N, E)51        padding_mask = self._get_padding_mask(lengths, self.max_length)52 53        zeros = embed.new_zeros(*embed.shape)54        qeury = self.pos_encoder(zeros)55        location_mask = self._get_location_mask(self.max_length, tokens.device)56        output = self.model(qeury, embed,57                tgt_key_padding_mask=padding_mask,58                memory_mask=location_mask,59                memory_key_padding_mask=padding_mask)  # (T, N, E)60        output = output.permute(1, 0, 2)  # (N, T, E)61 62        logits = self.cls(output)  # (N, T, C)63        pt_lengths = self._get_length(logits)64 65        res =  {'feature': output, 'logits': logits, 'pt_lengths': pt_lengths,66                'loss_weight':self.loss_weight, 'name': 'language'}67        return res68