initial commit

This commit is contained in:
2021-06-08 10:55:09 +08:00
commit c4d46626ce
30 changed files with 47900 additions and 0 deletions
+5
View File
@@ -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.
+19
View File
@@ -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
Binary file not shown.
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
+43
View File
@@ -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
+35
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+23860
View File
File diff suppressed because it is too large Load Diff