NightPrince commited on
Commit
6d26def
·
verified ·
1 Parent(s): 49f1e7f

Add catt/transformer.py

Browse files
Files changed (1) hide show
  1. catt/transformer.py +559 -0
catt/transformer.py ADDED
@@ -0,0 +1,559 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ @author : Hyunwoong
3
+ @when : 2019-12-18
4
+ @homepage : https://github.com/gusdnd852
5
+ """
6
+
7
+ import math
8
+ import torch
9
+ import torch.nn as nn
10
+
11
+
12
+ class EncoderLayer(nn.Module):
13
+
14
+ def __init__(self, d_model, ffn_hidden, n_head, drop_prob):
15
+ super(EncoderLayer, self).__init__()
16
+ self.attention = MultiHeadAttention(d_model=d_model, n_head=n_head)
17
+ self.norm1 = LayerNorm(d_model=d_model)
18
+ self.dropout1 = nn.Dropout(p=drop_prob)
19
+
20
+ self.ffn = PositionwiseFeedForward(d_model=d_model, hidden=ffn_hidden, drop_prob=drop_prob)
21
+ self.norm2 = LayerNorm(d_model=d_model)
22
+ self.dropout2 = nn.Dropout(p=drop_prob)
23
+
24
+ def forward(self, x, s_mask):
25
+ # 1. compute self attention
26
+ _x = x
27
+ x = self.attention(q=x, k=x, v=x, mask=s_mask)
28
+
29
+ # 2. add and norm
30
+ x = self.dropout1(x)
31
+ x = self.norm1(x + _x)
32
+
33
+ # 3. positionwise feed forward network
34
+ _x = x
35
+ x = self.ffn(x)
36
+
37
+ # 4. add and norm
38
+ x = self.dropout2(x)
39
+ x = self.norm2(x + _x)
40
+ return x
41
+
42
+
43
+ class DecoderLayer(nn.Module):
44
+
45
+ def __init__(self, d_model, ffn_hidden, n_head, drop_prob):
46
+ super(DecoderLayer, self).__init__()
47
+ self.self_attention = MultiHeadAttention(d_model=d_model, n_head=n_head)
48
+ self.norm1 = LayerNorm(d_model=d_model)
49
+ self.dropout1 = nn.Dropout(p=drop_prob)
50
+
51
+ self.enc_dec_attention = MultiHeadAttention(d_model=d_model, n_head=n_head)
52
+ self.norm2 = LayerNorm(d_model=d_model)
53
+ self.dropout2 = nn.Dropout(p=drop_prob)
54
+
55
+ self.ffn = PositionwiseFeedForward(d_model=d_model, hidden=ffn_hidden, drop_prob=drop_prob)
56
+ self.norm3 = LayerNorm(d_model=d_model)
57
+ self.dropout3 = nn.Dropout(p=drop_prob)
58
+
59
+ def forward(self, dec, enc, t_mask, s_mask):
60
+ # 1. compute self attention
61
+ _x = dec
62
+ x = self.self_attention(q=dec, k=dec, v=dec, mask=t_mask)
63
+
64
+ # 2. add and norm
65
+ x = self.dropout1(x)
66
+ x = self.norm1(x + _x)
67
+
68
+ if enc is not None:
69
+ # 3. compute encoder - decoder attention
70
+ _x = x
71
+ x = self.enc_dec_attention(q=x, k=enc, v=enc, mask=s_mask)
72
+
73
+ # 4. add and norm
74
+ x = self.dropout2(x)
75
+ x = self.norm2(x + _x)
76
+
77
+ # 5. positionwise feed forward network
78
+ _x = x
79
+ x = self.ffn(x)
80
+
81
+ # 6. add and norm
82
+ x = self.dropout3(x)
83
+ x = self.norm3(x + _x)
84
+ return x
85
+
86
+
87
+ class ScaleDotProductAttention(nn.Module):
88
+ """
89
+ compute scale dot product attention
90
+
91
+ Query : given sentence that we focused on (decoder)
92
+ Key : every sentence to check relationship with Qeury(encoder)
93
+ Value : every sentence same with Key (encoder)
94
+ """
95
+
96
+ def __init__(self):
97
+ super(ScaleDotProductAttention, self).__init__()
98
+ self.softmax = nn.Softmax(dim=-1)
99
+
100
+ def forward(self, q, k, v, mask=None, e=1e-12):
101
+ # input is 4 dimension tensor
102
+ # [batch_size, head, length, d_tensor]
103
+ batch_size, head, length, d_tensor = k.size()
104
+
105
+ # 1. dot product Query with Key^T to compute similarity
106
+ k_t = k.transpose(2, 3) # transpose
107
+ score = (q @ k_t) / math.sqrt(d_tensor) # scaled dot product
108
+
109
+ # 2. apply masking (opt)
110
+ if mask is not None:
111
+ score = score.masked_fill(mask == 0, -10000)
112
+
113
+ # 3. pass them softmax to make [0, 1] range
114
+ score = self.softmax(score)
115
+
116
+ # 4. multiply with Value
117
+ v = score @ v
118
+
119
+ return v, score
120
+
121
+
122
+ class PositionwiseFeedForward(nn.Module):
123
+
124
+ def __init__(self, d_model, hidden, drop_prob=0.1):
125
+ super(PositionwiseFeedForward, self).__init__()
126
+ self.linear1 = nn.Linear(d_model, hidden)
127
+ self.linear2 = nn.Linear(hidden, d_model)
128
+ self.relu = nn.ReLU()
129
+ self.dropout = nn.Dropout(p=drop_prob)
130
+
131
+ def forward(self, x):
132
+ x = self.linear1(x)
133
+ x = self.relu(x)
134
+ x = self.dropout(x)
135
+ x = self.linear2(x)
136
+ return x
137
+
138
+
139
+ class MultiHeadAttention(nn.Module):
140
+
141
+ def __init__(self, d_model, n_head):
142
+ super(MultiHeadAttention, self).__init__()
143
+ self.n_head = n_head
144
+ self.attention = ScaleDotProductAttention()
145
+ self.w_q = nn.Linear(d_model, d_model, bias=False)
146
+ self.w_k = nn.Linear(d_model, d_model, bias=False)
147
+ self.w_v = nn.Linear(d_model, d_model, bias=False)
148
+ self.w_concat = nn.Linear(d_model, d_model, bias=False)
149
+
150
+ def forward(self, q, k, v, mask=None):
151
+ # 1. dot product with weight matrices
152
+ q, k, v = self.w_q(q), self.w_k(k), self.w_v(v)
153
+
154
+ # 2. split tensor by number of heads
155
+ q, k, v = self.split(q), self.split(k), self.split(v)
156
+
157
+ # 3. do scale dot product to compute similarity
158
+ out, attention = self.attention(q, k, v, mask=mask)
159
+
160
+ # 4. concat and pass to linear layer
161
+ out = self.concat(out)
162
+ out = self.w_concat(out)
163
+
164
+ # 5. visualize attention map
165
+ # TODO : we should implement visualization
166
+
167
+ return out
168
+
169
+ def split(self, tensor):
170
+ """
171
+ split tensor by number of head
172
+
173
+ :param tensor: [batch_size, length, d_model]
174
+ :return: [batch_size, head, length, d_tensor]
175
+ """
176
+ batch_size, length, d_model = tensor.size()
177
+
178
+ d_tensor = d_model // self.n_head
179
+ tensor = tensor.view(batch_size, length, self.n_head, d_tensor).transpose(1, 2)
180
+ # it is similar with group convolution (split by number of heads)
181
+
182
+ return tensor
183
+
184
+ def concat(self, tensor):
185
+ """
186
+ inverse function of self.split(tensor : torch.Tensor)
187
+
188
+ :param tensor: [batch_size, head, length, d_tensor]
189
+ :return: [batch_size, length, d_model]
190
+ """
191
+ batch_size, head, length, d_tensor = tensor.size()
192
+ d_model = head * d_tensor
193
+
194
+ tensor = tensor.transpose(1, 2).contiguous().view(batch_size, length, d_model)
195
+ return tensor
196
+
197
+
198
+ class LayerNorm(nn.Module):
199
+ def __init__(self, d_model, eps=1e-12):
200
+ super(LayerNorm, self).__init__()
201
+ self.gamma = nn.Parameter(torch.ones(d_model))
202
+ self.beta = nn.Parameter(torch.zeros(d_model))
203
+ self.eps = eps
204
+
205
+ def forward(self, x):
206
+ mean = x.mean(-1, keepdim=True)
207
+ var = x.var(-1, unbiased=False, keepdim=True)
208
+ # '-1' means last dimension.
209
+
210
+ out = (x - mean) / torch.sqrt(var + self.eps)
211
+ out = self.gamma * out + self.beta
212
+ return out
213
+
214
+
215
+ class TransformerEmbedding(nn.Module):
216
+ """
217
+ token embedding + positional encoding (sinusoid)
218
+ positional encoding can give positional information to network
219
+ """
220
+
221
+ def __init__(self, vocab_size, d_model, max_len, drop_prob, padding_idx, learnable_pos_emb=True):
222
+ """
223
+ class for word embedding that included positional information
224
+
225
+ :param vocab_size: size of vocabulary
226
+ :param d_model: dimensions of model
227
+ """
228
+ super(TransformerEmbedding, self).__init__()
229
+ self.tok_emb = TokenEmbedding(vocab_size, d_model, padding_idx)
230
+ if learnable_pos_emb:
231
+ self.pos_emb = LearnablePositionalEncoding(d_model, max_len)
232
+ else:
233
+ self.pos_emb = SinusoidalPositionalEncoding(d_model, max_len)
234
+ self.drop_out = nn.Dropout(p=drop_prob)
235
+
236
+ def forward(self, x):
237
+ tok_emb = self.tok_emb(x)
238
+ pos_emb = self.pos_emb(x).to(tok_emb.device)
239
+ return self.drop_out(tok_emb + pos_emb)
240
+
241
+
242
+ class TokenEmbedding(nn.Embedding):
243
+ """
244
+ Token Embedding using torch.nn
245
+ they will dense representation of word using weighted matrix
246
+ """
247
+
248
+ def __init__(self, vocab_size, d_model, padding_idx):
249
+ """
250
+ class for token embedding that included positional information
251
+
252
+ :param vocab_size: size of vocabulary
253
+ :param d_model: dimensions of model
254
+ """
255
+ super(TokenEmbedding, self).__init__(vocab_size, d_model, padding_idx=padding_idx)
256
+
257
+
258
+ class SinusoidalPositionalEncoding(nn.Module):
259
+ """
260
+ compute sinusoid encoding.
261
+ """
262
+
263
+ def __init__(self, d_model, max_len):
264
+ """
265
+ constructor of sinusoid encoding class
266
+
267
+ :param d_model: dimension of model
268
+ :param max_len: max sequence length
269
+
270
+ """
271
+ super(SinusoidalPositionalEncoding, self).__init__()
272
+
273
+ # same size with input matrix (for adding with input matrix)
274
+ self.encoding = torch.zeros(max_len, d_model)
275
+ self.encoding.requires_grad = False # we don't need to compute gradient
276
+
277
+ pos = torch.arange(0, max_len)
278
+ pos = pos.float().unsqueeze(dim=1)
279
+ # 1D => 2D unsqueeze to represent word's position
280
+
281
+ _2i = torch.arange(0, d_model, step=2).float()
282
+ # 'i' means index of d_model (e.g. embedding size = 50, 'i' = [0,50])
283
+ # "step=2" means 'i' multiplied with two (same with 2 * i)
284
+
285
+ self.encoding[:, 0::2] = torch.sin(pos / (10000 ** (_2i / d_model)))
286
+ self.encoding[:, 1::2] = torch.cos(pos / (10000 ** (_2i / d_model)))
287
+ # compute positional encoding to consider positional information of words
288
+
289
+ def forward(self, x):
290
+ # self.encoding
291
+ # [max_len = 512, d_model = 512]
292
+
293
+ batch_size, seq_len = x.size()
294
+ # [batch_size = 128, seq_len = 30]
295
+
296
+ return self.encoding[:seq_len, :]
297
+ # [seq_len = 30, d_model = 512]
298
+ # it will add with tok_emb : [128, 30, 512]
299
+
300
+
301
+ class LearnablePositionalEncoding(nn.Module):
302
+ """
303
+ compute sinusoid encoding.
304
+ """
305
+
306
+ def __init__(self, d_model, max_seq_len):
307
+ """
308
+ constructor of learnable positonal encoding class
309
+
310
+ :param d_model: dimension of model
311
+ :param max_seq_len: max sequence length
312
+
313
+ """
314
+ super(LearnablePositionalEncoding, self).__init__()
315
+ self.max_seq_len = max_seq_len
316
+ self.wpe = nn.Embedding(max_seq_len, d_model)
317
+
318
+ def forward(self, x):
319
+ # self.encoding
320
+ # [max_len = 512, d_model = 512]
321
+ device = x.device
322
+ batch_size, seq_len = x.size()
323
+ assert seq_len <= self.max_seq_len, f"Cannot forward sequence of length {seq_len}, max_seq_len is {self.max_seq_len}"
324
+ pos = torch.arange(0, seq_len, dtype=torch.long, device=device) # shape (seq_len)
325
+ pos_emb = self.wpe(pos) # position embeddings of shape (seq_len, d_model)
326
+
327
+ return pos_emb
328
+ # [seq_len = 30, d_model = 512]
329
+ # it will add with tok_emb : [128, 30, 512]
330
+
331
+
332
+ class Encoder(nn.Module):
333
+
334
+ def __init__(self, enc_voc_size, max_len, d_model, ffn_hidden, n_head, n_layers, drop_prob, padding_idx, learnable_pos_emb=True):
335
+ super().__init__()
336
+ self.emb = TransformerEmbedding(d_model=d_model,
337
+ max_len=max_len,
338
+ vocab_size=enc_voc_size,
339
+ drop_prob=drop_prob,
340
+ padding_idx=padding_idx,
341
+ learnable_pos_emb=learnable_pos_emb
342
+ )
343
+
344
+ self.layers = nn.ModuleList([EncoderLayer(d_model=d_model,
345
+ ffn_hidden=ffn_hidden,
346
+ n_head=n_head,
347
+ drop_prob=drop_prob)
348
+ for _ in range(n_layers)])
349
+
350
+ def forward(self, x, s_mask):
351
+ x = self.emb(x)
352
+
353
+ for layer in self.layers:
354
+ x = layer(x, s_mask)
355
+
356
+ return x
357
+
358
+ class Decoder(nn.Module):
359
+ def __init__(self, dec_voc_size, max_len, d_model, ffn_hidden, n_head, n_layers, drop_prob, padding_idx, learnable_pos_emb=True):
360
+ super().__init__()
361
+ self.emb = TransformerEmbedding(d_model=d_model,
362
+ drop_prob=drop_prob,
363
+ max_len=max_len,
364
+ vocab_size=dec_voc_size,
365
+ padding_idx=padding_idx,
366
+ learnable_pos_emb=learnable_pos_emb
367
+ )
368
+
369
+ self.layers = nn.ModuleList([DecoderLayer(d_model=d_model,
370
+ ffn_hidden=ffn_hidden,
371
+ n_head=n_head,
372
+ drop_prob=drop_prob)
373
+ for _ in range(n_layers)])
374
+
375
+ self.linear = nn.Linear(d_model, dec_voc_size)
376
+
377
+ def forward(self, trg, enc_src, trg_mask, src_mask):
378
+ trg = self.emb(trg)
379
+
380
+ for layer in self.layers:
381
+ trg = layer(trg, enc_src, trg_mask, src_mask)
382
+
383
+ # pass to LM head
384
+ output = self.linear(trg)
385
+ return output
386
+
387
+ class Transformer(nn.Module):
388
+
389
+ def __init__(self, src_pad_idx, trg_pad_idx, enc_voc_size, dec_voc_size, d_model, n_head, max_len,
390
+ ffn_hidden, n_layers, drop_prob, learnable_pos_emb=True):
391
+ super().__init__()
392
+ self.src_pad_idx = src_pad_idx
393
+ self.trg_pad_idx = trg_pad_idx
394
+ self.encoder = Encoder(d_model=d_model,
395
+ n_head=n_head,
396
+ max_len=max_len,
397
+ ffn_hidden=ffn_hidden,
398
+ enc_voc_size=enc_voc_size,
399
+ drop_prob=drop_prob,
400
+ n_layers=n_layers,
401
+ padding_idx=src_pad_idx,
402
+ learnable_pos_emb=learnable_pos_emb)
403
+
404
+ self.decoder = Decoder(d_model=d_model,
405
+ n_head=n_head,
406
+ max_len=max_len,
407
+ ffn_hidden=ffn_hidden,
408
+ dec_voc_size=dec_voc_size,
409
+ drop_prob=drop_prob,
410
+ n_layers=n_layers,
411
+ padding_idx=trg_pad_idx,
412
+ learnable_pos_emb=learnable_pos_emb)
413
+
414
+ def get_device(self):
415
+ return next(self.parameters()).device
416
+
417
+ def forward(self, src, trg):
418
+ device = self.get_device()
419
+ src_mask = self.make_pad_mask(src, src, self.src_pad_idx, self.src_pad_idx).to(device)
420
+ src_trg_mask = self.make_pad_mask(trg, src, self.trg_pad_idx, self.src_pad_idx).to(device)
421
+ trg_mask = self.make_pad_mask(trg, trg, self.trg_pad_idx, self.trg_pad_idx).to(device) * \
422
+ self.make_no_peak_mask(trg, trg).to(device)
423
+
424
+ #print(src_mask)
425
+ #print('-'*100)
426
+ #print(trg_mask)
427
+ enc_src = self.encoder(src, src_mask)
428
+ output = self.decoder(trg, enc_src, trg_mask, src_trg_mask)
429
+ return output
430
+
431
+ def make_pad_mask(self, q, k, q_pad_idx, k_pad_idx):
432
+ len_q, len_k = q.size(1), k.size(1)
433
+
434
+ # batch_size x 1 x 1 x len_k
435
+ k = k.ne(k_pad_idx).unsqueeze(1).unsqueeze(2)
436
+ # batch_size x 1 x len_q x len_k
437
+ k = k.repeat(1, 1, len_q, 1)
438
+
439
+ # batch_size x 1 x len_q x 1
440
+ q = q.ne(q_pad_idx).unsqueeze(1).unsqueeze(3)
441
+ # batch_size x 1 x len_q x len_k
442
+ q = q.repeat(1, 1, 1, len_k)
443
+
444
+ mask = k & q
445
+ return mask
446
+
447
+ def make_no_peak_mask(self, q, k):
448
+ len_q, len_k = q.size(1), k.size(1)
449
+
450
+ # len_q x len_k
451
+ mask = torch.tril(torch.ones(len_q, len_k)).type(torch.BoolTensor)
452
+
453
+ return mask
454
+
455
+
456
+ def make_pad_mask(x, pad_idx):
457
+ q = k = x
458
+ q_pad_idx = k_pad_idx = pad_idx
459
+ len_q, len_k = q.size(1), k.size(1)
460
+
461
+ # batch_size x 1 x 1 x len_k
462
+ k = k.ne(k_pad_idx).unsqueeze(1).unsqueeze(2)
463
+ # batch_size x 1 x len_q x len_k
464
+ k = k.repeat(1, 1, len_q, 1)
465
+
466
+ # batch_size x 1 x len_q x 1
467
+ q = q.ne(q_pad_idx).unsqueeze(1).unsqueeze(3)
468
+ # batch_size x 1 x len_q x len_k
469
+ q = q.repeat(1, 1, 1, len_k)
470
+
471
+ mask = k & q
472
+ return mask
473
+
474
+
475
+ from torch.nn.utils.rnn import pad_sequence
476
+ # x_list is a list of tensors of shape TxH where T is the seqlen and H is the feats dim
477
+ def pad_seq_v2(sequences, batch_first=True, padding_value=0.0, prepadding=True):
478
+ lens = [i.shape[0]for i in sequences]
479
+ padded_sequences = pad_sequence(sequences, batch_first=True, padding_value=padding_value) # NxTxH
480
+ if prepadding:
481
+ for i in range(len(lens)):
482
+ padded_sequences[i] = padded_sequences[i].roll(-lens[i])
483
+ if not batch_first:
484
+ padded_sequences = padded_sequences.transpose(0, 1) # TxNxH
485
+ return padded_sequences
486
+
487
+
488
+
489
+ if __name__ == '__main__':
490
+ import torch
491
+ import random
492
+ import numpy as np
493
+
494
+ rand_seed = 10
495
+
496
+ device = 'cpu'
497
+
498
+ # model parameter setting
499
+ batch_size = 128
500
+ max_len = 256
501
+ d_model = 512
502
+ n_layers = 3
503
+ n_heads = 16
504
+ ffn_hidden = 2048
505
+ drop_prob = 0.1
506
+
507
+ # optimizer parameter setting
508
+ init_lr = 1e-5
509
+ factor = 0.9
510
+ adam_eps = 5e-9
511
+ patience = 10
512
+ warmup = 100
513
+ epoch = 1000
514
+ clip = 1.0
515
+ weight_decay = 5e-4
516
+ inf = float('inf')
517
+
518
+ src_pad_idx = 2
519
+ trg_pad_idx = 3
520
+
521
+ enc_voc_size = 37
522
+ dec_voc_size = 15
523
+ model = Transformer(src_pad_idx=src_pad_idx,
524
+ trg_pad_idx=trg_pad_idx,
525
+ d_model=d_model,
526
+ enc_voc_size=enc_voc_size,
527
+ dec_voc_size=dec_voc_size,
528
+ max_len=max_len,
529
+ ffn_hidden=ffn_hidden,
530
+ n_head=n_heads,
531
+ n_layers=n_layers,
532
+ drop_prob=drop_prob
533
+ ).to(device)
534
+
535
+ random.seed(rand_seed)
536
+ # Set the seed to 0 for reproducible results
537
+ np.random.seed(rand_seed)
538
+ torch.manual_seed(rand_seed)
539
+
540
+ x_list = [
541
+ torch.tensor([[1, 1]]).transpose(0, 1), # 2
542
+ torch.tensor([[1, 1, 1, 1, 1, 1, 1]]).transpose(0, 1), # 7
543
+ torch.tensor([[1, 1, 1]]).transpose(0, 1) # 3
544
+ ]
545
+
546
+
547
+ src_pad_idx = model.src_pad_idx
548
+ trg_pad_idx = model.trg_pad_idx
549
+
550
+ src = pad_seq_v2(x_list, padding_value=src_pad_idx, prepadding=False).squeeze(2)
551
+ trg = pad_seq_v2(x_list, padding_value=trg_pad_idx, prepadding=False).squeeze(2)
552
+ out = model(src, trg)
553
+
554
+
555
+
556
+
557
+
558
+
559
+