(self,
in_channels,
out_channels,
nhead=None,
num_decoder_layers=6,
max_len=25,
attention_dropout_rate=0.0,
residual_dropout_rate=0.1,
scale_embedding=True,
parallel_decoding=False,
autoregressive_decoding=False,
sampler_step=5,
low_confidence_decoding=False,
random_mask_decoding=False,
semi_autoregressive_decoding=False,
cloze_mask_decoding=False,
rec_loss_weight=1.0,
reflect_loss_weight=1.0,
sample_k=0,
temperature=1.0)
| 26 | """ |
| 27 | |
| 28 | def __init__(self, |
| 29 | in_channels, |
| 30 | out_channels, |
| 31 | nhead=None, |
| 32 | num_decoder_layers=6, |
| 33 | max_len=25, |
| 34 | attention_dropout_rate=0.0, |
| 35 | residual_dropout_rate=0.1, |
| 36 | scale_embedding=True, |
| 37 | parallel_decoding=False, |
| 38 | autoregressive_decoding=False, |
| 39 | sampler_step=5, |
| 40 | low_confidence_decoding=False, |
| 41 | random_mask_decoding=False, |
| 42 | semi_autoregressive_decoding=False, |
| 43 | cloze_mask_decoding=False, |
| 44 | rec_loss_weight=1.0, |
| 45 | reflect_loss_weight=1.0, |
| 46 | sample_k=0, |
| 47 | temperature=1.0): |
| 48 | super(MDiffDecoder, self).__init__() |
| 49 | self.out_channels = out_channels |
| 50 | self.ignore_index = out_channels - 1 |
| 51 | self.mask_token_id = out_channels - 2 |
| 52 | self.eos = 0 |
| 53 | self.max_len = max_len |
| 54 | d_model = in_channels |
| 55 | dim_feedforward = d_model * 4 |
| 56 | self.pd = parallel_decoding |
| 57 | self.ar = autoregressive_decoding |
| 58 | self.sampler_step = sampler_step |
| 59 | self.lc = low_confidence_decoding |
| 60 | self.rm = random_mask_decoding |
| 61 | self.semiar = semi_autoregressive_decoding |
| 62 | self.cm = cloze_mask_decoding |
| 63 | self.rec_loss_weight = rec_loss_weight |
| 64 | self.reflect_loss_weight = reflect_loss_weight |
| 65 | self.temperature = temperature |
| 66 | self.sample_k = sample_k |
| 67 | nhead = nhead if nhead is not None else d_model // 32 |
| 68 | self.embedding = Embeddings( |
| 69 | d_model=d_model, |
| 70 | vocab=self.out_channels, |
| 71 | padding_idx=0, |
| 72 | scale_embedding=scale_embedding, |
| 73 | ) |
| 74 | self.pos_embed = nn.Parameter(torch.zeros( |
| 75 | [1, self.max_len + 1, d_model], dtype=torch.float32), |
| 76 | requires_grad=True) |
| 77 | nn.init.trunc_normal_(self.pos_embed, std=0.02) |
| 78 | self.positional_encoding = PositionalEncoding( |
| 79 | dropout=residual_dropout_rate, dim=d_model) |
| 80 | |
| 81 | self.decoder = nn.ModuleList([ |
| 82 | TransformerBlock( |
| 83 | d_model, |
| 84 | nhead, |
| 85 | dim_feedforward, |
no test coverage detected