initial commit
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
## Train
|
||||
`onmt_train -config seq2seq.yaml -continue`
|
||||
## Generate
|
||||
`onmt_translate -model model/model_step_100000.pt -src corpus/dailydialog/dialogues_test_src.txt -output output/pred_100000.txt -verbose
|
||||
`
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,19 @@
|
||||
# This is a sample Python script.
|
||||
|
||||
# Press Shift+F10 to execute it or replace it with your code.
|
||||
# Press Double Shift to search everywhere for classes, files, tool windows, actions, and settings.
|
||||
|
||||
|
||||
# onmt_translate -model model/model_step_100000.pt -src corpus/dailydialog/dialogues_test_src.txt -output output/pred_100000.txt -verbose
|
||||
# onmt_train -config seq2seq.yaml -continue
|
||||
def print_hi(name):
|
||||
# Use a breakpoint in the code line below to debug your script.
|
||||
print(f'Hi, {name}') # Press Ctrl+F8 to toggle the breakpoint.
|
||||
|
||||
|
||||
# Press the green button in the gutter to run the script.
|
||||
if __name__ == '__main__':
|
||||
print_hi('PyCharm')
|
||||
|
||||
# See PyCharm help at https://www.jetbrains.com/help/pycharm/
|
||||
# asdfas
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,43 @@
|
||||
# seq2seq.yaml
|
||||
|
||||
## Where the samples will be written
|
||||
save_data: saved_data/
|
||||
## Where the vocab(s) will be written
|
||||
src_vocab: vocab/dialog.vocab.src
|
||||
tgt_vocab: vocab/dialog.vocab.tgt
|
||||
# Prevent overwriting existing files in the folder
|
||||
overwrite: False
|
||||
|
||||
# Corpus opts:
|
||||
data:
|
||||
corpus_1:
|
||||
path_src: corpus/dailydialog/dialogues_train_src.txt
|
||||
path_tgt: corpus/dailydialog/dialogues_train_tgt.txt
|
||||
valid:
|
||||
path_src: corpus/dailydialog/dialogues_validation_src.txt
|
||||
path_tgt: corpus/dailydialog/dialogues_validation_tgt.txt
|
||||
|
||||
both_embeddings: embeddings/glove.6B/glove.6B.300d.txt
|
||||
embeddings_type: "GloVe"
|
||||
word_vec_size: 300
|
||||
|
||||
global_attention: "mlp"
|
||||
|
||||
encoder_type: brnn
|
||||
decoder_type: rnn
|
||||
|
||||
hidden_size: 500
|
||||
optim: sgd
|
||||
switchout_temperature: 0.01
|
||||
|
||||
# Train on a single GPU
|
||||
world_size: 1
|
||||
gpu_ranks: [ 0 ]
|
||||
|
||||
# Where to save the checkpoints
|
||||
save_model: model
|
||||
save_checkpoint_steps: 1000
|
||||
train_steps: 100000
|
||||
valid_steps: 1000
|
||||
|
||||
train_from: model_step_78000.pt
|
||||
@@ -0,0 +1,35 @@
|
||||
daily_dialog = [
|
||||
['corpus/dailydialog/train/dialogues_train.txt', 'corpus/dailydialog/dialogues_train.txt'],
|
||||
['corpus/dailydialog/test/dialogues_test.txt', 'corpus/dailydialog/dialogues_test.txt'],
|
||||
['corpus/dailydialog/validation/dialogues_validation.txt', 'corpus/dailydialog/dialogues_validation.txt']
|
||||
]
|
||||
|
||||
|
||||
def process_daily_dialog(file_name: str, output_name: str):
|
||||
lines = []
|
||||
conversation_list = []
|
||||
with open(file_name, 'r', encoding='utf-8') as f:
|
||||
lines = f.readlines()
|
||||
for line in lines:
|
||||
utterances = line[:-1].split('__eou__')
|
||||
conversation_list.append([u for u in utterances if u != ''])
|
||||
|
||||
src_list = []
|
||||
tgt_list = []
|
||||
for conversation in conversation_list:
|
||||
if len(conversation) == 1:
|
||||
continue
|
||||
for i in range(len(conversation) - 1):
|
||||
src_list.append(conversation[i].strip(' ').rstrip(' '))
|
||||
tgt_list.append(conversation[i + 1].strip(' ').rstrip(' '))
|
||||
with open(output_name.split('.')[0] + '_src.txt', 'w', encoding='utf-8') as f:
|
||||
for src in src_list:
|
||||
f.write(src + '\n')
|
||||
with open(output_name.split('.')[0] + '_tgt.txt', 'w', encoding='utf-8') as f:
|
||||
for tgt in tgt_list:
|
||||
f.write(tgt + '\n')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
for f in daily_dialog:
|
||||
process_daily_dialog(f[0], f[1])
|
||||
+23866
File diff suppressed because it is too large
Load Diff
+23860
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user