@@ -40,8 +40,7 @@ def __init__(self):
4040
4141 # encoder/ decoder
4242 # Restore the model and construct the encoder and decoder.
43- self ._encoder = Encoder (
44- INPUT_DIM , E_EMB_DIM , E_HID_DIM , E_DROPOUT )
43+ self ._encoder = Encoder (INPUT_DIM , E_EMB_DIM , E_HID_DIM , E_DROPOUT )
4544
4645 self ._decoder = AttentionDecoder (
4746 OUTPUT_DIM , D_EMB_DIM , D_HID_DIM , D_DROPOUT
@@ -81,33 +80,36 @@ def romanize(self, text: str) -> str:
8180 input_tensor = self ._prepare_sequence_in (text ).view (1 , - 1 )
8281 input_length = [len (text ) + 1 ]
8382
84- target_tensor_logits = self ._network (input_tensor ,
85- input_length ,
86- None , 0 )
83+ target_tensor_logits = self ._network (
84+ input_tensor , input_length , None , 0
85+ )
8786
8887 # Seq2seq model returns <END> as the first token,
8988 # As a result, target_tensor_logits.size() is torch.Size([0])
9089 if target_tensor_logits .size (0 ) == 0 :
9190 target = ["<PAD>" ]
9291 else :
9392 target_tensor = (
94- torch .argmax (
95- target_tensor_logits .squeeze (1 ),
96- 1 ).cpu ().detach ().numpy ()
97- )
93+ torch .argmax (target_tensor_logits .squeeze (1 ), 1 )
94+ .cpu ()
95+ .detach ()
96+ .numpy ()
97+ )
9898 target = [self ._ix_to_target_char [t ] for t in target_tensor ]
9999
100100 return "" .join (target )
101101
102102
103103class Encoder (nn .Module ):
104- def __init__ (self , vocabulary_size , embedding_size ,
105- hidden_size , dropout = 0.5 ):
104+ def __init__ (
105+ self , vocabulary_size , embedding_size , hidden_size , dropout = 0.5
106+ ):
106107 """Constructor"""
107108 super (Encoder , self ).__init__ ()
108109 self .hidden_size = hidden_size
109- self .character_embedding = nn .Embedding (vocabulary_size ,
110- embedding_size )
110+ self .character_embedding = nn .Embedding (
111+ vocabulary_size , embedding_size
112+ )
111113 self .rnn = nn .LSTM (
112114 input_size = embedding_size ,
113115 hidden_size = hidden_size // 2 ,
@@ -141,8 +143,7 @@ def forward(self, sequences, sequences_lengths):
141143 sequences , sequences_lengths .copy (), batch_first = True
142144 )
143145
144- sequences_output , self .hidden = self .rnn (sequences_packed ,
145- self .hidden )
146+ sequences_output , self .hidden = self .rnn (sequences_packed , self .hidden )
146147
147148 sequences_output , _ = nn .utils .rnn .pad_packed_sequence (
148149 sequences_output , batch_first = True
@@ -183,22 +184,25 @@ def __init__(self, method, hidden_size):
183184 def forward (self , hidden , encoder_outputs , mask ):
184185 # Calculate energies for each encoder output
185186 if self .method == "dot" :
186- attn_energies = torch .bmm (encoder_outputs ,
187- hidden .transpose (1 , 2 )).squeeze (2 )
187+ attn_energies = torch .bmm (
188+ encoder_outputs , hidden .transpose (1 , 2 )
189+ ).squeeze (2 )
188190 elif self .method == "general" :
189191 attn_energies = self .attn (
190192 encoder_outputs .view (- 1 , encoder_outputs .size (- 1 ))
191193 ) # (batch_size * sequence_len, hidden_size)
192194 attn_energies = torch .bmm (
193- attn_energies .view (
194- * encoder_outputs .size ()), hidden .transpose (1 , 2 )
195- ).squeeze (2 ) # (batch_size, sequence_len)
195+ attn_energies .view (* encoder_outputs .size ()),
196+ hidden .transpose (1 , 2 ),
197+ ).squeeze (
198+ 2
199+ ) # (batch_size, sequence_len)
196200 elif self .method == "concat" :
197201 attn_energies = self .attn (
198- torch .cat ((
199- hidden .expand (* encoder_outputs .size ()),
200- encoder_outputs
201- ), 2 )
202+ torch .cat (
203+ ( hidden .expand (* encoder_outputs .size ()), encoder_outputs ),
204+ 2 ,
205+ )
202206 ) # (batch_size, sequence_len, hidden_size)
203207 attn_energies = torch .bmm (
204208 attn_energies ,
@@ -212,14 +216,16 @@ def forward(self, hidden, encoder_outputs, mask):
212216
213217
214218class AttentionDecoder (nn .Module ):
215- def __init__ (self , vocabulary_size , embedding_size ,
216- hidden_size , dropout = 0.5 ):
219+ def __init__ (
220+ self , vocabulary_size , embedding_size , hidden_size , dropout = 0.5
221+ ):
217222 """Constructor"""
218223 super (AttentionDecoder , self ).__init__ ()
219224 self .vocabulary_size = vocabulary_size
220225 self .hidden_size = hidden_size
221- self .character_embedding = nn .Embedding (vocabulary_size ,
222- embedding_size )
226+ self .character_embedding = nn .Embedding (
227+ vocabulary_size , embedding_size
228+ )
223229 self .rnn = nn .LSTM (
224230 input_size = embedding_size + self .hidden_size ,
225231 hidden_size = hidden_size ,
@@ -262,8 +268,12 @@ def forward(self, input, last_hidden, encoder_outputs, mask):
262268
263269class Seq2Seq (nn .Module ):
264270 def __init__ (
265- self , encoder , decoder , target_start_token ,
266- target_end_token , max_length
271+ self ,
272+ encoder ,
273+ decoder ,
274+ target_start_token ,
275+ target_end_token ,
276+ max_length ,
267277 ):
268278 super ().__init__ ()
269279
@@ -294,22 +304,24 @@ def forward(
294304 max_len = self .max_length
295305 target_vocab_size = self .decoder .vocabulary_size
296306
297- outputs = torch .zeros (max_len ,
298- batch_size ,
299- target_vocab_size ). to ( device )
307+ outputs = torch .zeros (max_len , batch_size , target_vocab_size ). to (
308+ device
309+ )
300310
301311 if target_seq is None :
302312 assert teacher_forcing_ratio == 0 , "Must be zero during inference"
303313 inference = True
304314 else :
305315 inference = False
306316
307- encoder_outputs , encoder_hidden = self .encoder (source_seq ,
308- source_seq_len )
317+ encoder_outputs , encoder_hidden = self .encoder (
318+ source_seq , source_seq_len
319+ )
309320
310321 decoder_input = (
311- torch .tensor ([[start_token ] * batch_size ]).view (batch_size ,
312- 1 ).to (device )
322+ torch .tensor ([[start_token ] * batch_size ])
323+ .view (batch_size , 1 )
324+ .to (device )
313325 )
314326
315327 encoder_hidden_h_t = torch .cat (
@@ -341,6 +353,7 @@ def forward(
341353
342354 return outputs
343355
356+
344357_THAI_TO_ROM = ThaiTransliterator ()
345358
346359
0 commit comments