Skip to content

Commit 13937e8

Browse files
committed
Format code
1 parent c5a73df commit 13937e8

3 files changed

Lines changed: 103 additions & 77 deletions

File tree

pythainlp/tokenize/newmm.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
# -*- coding: utf-8 -*-
22
"""
3-
Dictionary-based Thai Word Segmentation
4-
using maximal matching algorithm and Thai Character Cluster (TCC).
3+
Dictionary-based maximal matching word segmentation, constrained with
4+
Thai Character Cluster (TCC) boundaries.
55
6-
The code is based on the notebooks created by Korakot Chaovavanich.
6+
The code is based on the notebooks created by Korakot Chaovavanich,
7+
with heuristic graph size limit added to avoid exponential wait time.
78
89
:See Also:
910
* \

pythainlp/transliterate/thai2rom.py

Lines changed: 50 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -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

103103
class 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

214218
class 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

263269
class 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

pythainlp/transliterate/thaig2p.py

Lines changed: 49 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -41,8 +41,7 @@ def __init__(self):
4141

4242
# encoder/ decoder
4343
# Restore the model and construct the encoder and decoder.
44-
self._encoder = Encoder(
45-
INPUT_DIM, E_EMB_DIM, E_HID_DIM, E_DROPOUT)
44+
self._encoder = Encoder(INPUT_DIM, E_EMB_DIM, E_HID_DIM, E_DROPOUT)
4645

4746
self._decoder = AttentionDecoder(
4847
OUTPUT_DIM, D_EMB_DIM, D_HID_DIM, D_DROPOUT
@@ -82,33 +81,36 @@ def g2p(self, text: str) -> str:
8281
input_tensor = self._prepare_sequence_in(text).view(1, -1)
8382
input_length = [len(text) + 1]
8483

85-
target_tensor_logits = self._network(input_tensor,
86-
input_length,
87-
None, 0)
84+
target_tensor_logits = self._network(
85+
input_tensor, input_length, None, 0
86+
)
8887

8988
# Seq2seq model returns <END> as the first token,
9089
# As a result, target_tensor_logits.size() is torch.Size([0])
9190
if target_tensor_logits.size(0) == 0:
9291
target = ["<PAD>"]
9392
else:
9493
target_tensor = (
95-
torch.argmax(
96-
target_tensor_logits.squeeze(1),
97-
1).cpu().detach().numpy()
98-
)
94+
torch.argmax(target_tensor_logits.squeeze(1), 1)
95+
.cpu()
96+
.detach()
97+
.numpy()
98+
)
9999
target = [self._ix_to_target_char[t] for t in target_tensor]
100100

101101
return "".join(target)
102102

103103

104104
class Encoder(nn.Module):
105-
def __init__(self, vocabulary_size, embedding_size,
106-
hidden_size, dropout=0.5):
105+
def __init__(
106+
self, vocabulary_size, embedding_size, hidden_size, dropout=0.5
107+
):
107108
"""Constructor"""
108109
super(Encoder, self).__init__()
109110
self.hidden_size = hidden_size
110-
self.character_embedding = nn.Embedding(vocabulary_size,
111-
embedding_size)
111+
self.character_embedding = nn.Embedding(
112+
vocabulary_size, embedding_size
113+
)
112114
self.rnn = nn.LSTM(
113115
input_size=embedding_size,
114116
hidden_size=hidden_size // 2,
@@ -142,8 +144,7 @@ def forward(self, sequences, sequences_lengths):
142144
sequences, sequences_lengths.copy(), batch_first=True
143145
)
144146

145-
sequences_output, self.hidden = self.rnn(sequences_packed,
146-
self.hidden)
147+
sequences_output, self.hidden = self.rnn(sequences_packed, self.hidden)
147148

148149
sequences_output, _ = nn.utils.rnn.pad_packed_sequence(
149150
sequences_output, batch_first=True
@@ -184,22 +185,25 @@ def __init__(self, method, hidden_size):
184185
def forward(self, hidden, encoder_outputs, mask):
185186
# Calculate energies for each encoder output
186187
if self.method == "dot":
187-
attn_energies = torch.bmm(encoder_outputs,
188-
hidden.transpose(1, 2)).squeeze(2)
188+
attn_energies = torch.bmm(
189+
encoder_outputs, hidden.transpose(1, 2)
190+
).squeeze(2)
189191
elif self.method == "general":
190192
attn_energies = self.attn(
191193
encoder_outputs.view(-1, encoder_outputs.size(-1))
192194
) # (batch_size * sequence_len, hidden_size)
193195
attn_energies = torch.bmm(
194-
attn_energies.view(
195-
*encoder_outputs.size()), hidden.transpose(1, 2)
196-
).squeeze(2) # (batch_size, sequence_len)
196+
attn_energies.view(*encoder_outputs.size()),
197+
hidden.transpose(1, 2),
198+
).squeeze(
199+
2
200+
) # (batch_size, sequence_len)
197201
elif self.method == "concat":
198202
attn_energies = self.attn(
199-
torch.cat((
200-
hidden.expand(*encoder_outputs.size()),
201-
encoder_outputs
202-
), 2)
203+
torch.cat(
204+
(hidden.expand(*encoder_outputs.size()), encoder_outputs),
205+
2,
206+
)
203207
) # (batch_size, sequence_len, hidden_size)
204208
attn_energies = torch.bmm(
205209
attn_energies,
@@ -213,14 +217,16 @@ def forward(self, hidden, encoder_outputs, mask):
213217

214218

215219
class AttentionDecoder(nn.Module):
216-
def __init__(self, vocabulary_size, embedding_size,
217-
hidden_size, dropout=0.5):
220+
def __init__(
221+
self, vocabulary_size, embedding_size, hidden_size, dropout=0.5
222+
):
218223
"""Constructor"""
219224
super(AttentionDecoder, self).__init__()
220225
self.vocabulary_size = vocabulary_size
221226
self.hidden_size = hidden_size
222-
self.character_embedding = nn.Embedding(vocabulary_size,
223-
embedding_size)
227+
self.character_embedding = nn.Embedding(
228+
vocabulary_size, embedding_size
229+
)
224230
self.rnn = nn.LSTM(
225231
input_size=embedding_size + self.hidden_size,
226232
hidden_size=hidden_size,
@@ -263,8 +269,12 @@ def forward(self, input, last_hidden, encoder_outputs, mask):
263269

264270
class Seq2Seq(nn.Module):
265271
def __init__(
266-
self, encoder, decoder, target_start_token,
267-
target_end_token, max_length
272+
self,
273+
encoder,
274+
decoder,
275+
target_start_token,
276+
target_end_token,
277+
max_length,
268278
):
269279
super().__init__()
270280

@@ -295,22 +305,24 @@ def forward(
295305
max_len = self.max_length
296306
target_vocab_size = self.decoder.vocabulary_size
297307

298-
outputs = torch.zeros(max_len,
299-
batch_size,
300-
target_vocab_size).to(device)
308+
outputs = torch.zeros(max_len, batch_size, target_vocab_size).to(
309+
device
310+
)
301311

302312
if target_seq is None:
303313
assert teacher_forcing_ratio == 0, "Must be zero during inference"
304314
inference = True
305315
else:
306316
inference = False
307317

308-
encoder_outputs, encoder_hidden = self.encoder(source_seq,
309-
source_seq_len)
318+
encoder_outputs, encoder_hidden = self.encoder(
319+
source_seq, source_seq_len
320+
)
310321

311322
decoder_input = (
312-
torch.tensor([[start_token] * batch_size]).view(batch_size,
313-
1).to(device)
323+
torch.tensor([[start_token] * batch_size])
324+
.view(batch_size, 1)
325+
.to(device)
314326
)
315327

316328
encoder_hidden_h_t = torch.cat(

0 commit comments

Comments
 (0)