1 parent da60a87 commit ae094c8Copy full SHA for ae094c8
4 files changed
bert_pytorch/__main__.py
@@ -10,7 +10,7 @@
10
def train():
11
parser = argparse.ArgumentParser()
12
13
- parser.add_argument("-d", "--train_dataset", required=True, type=str)
+ parser.add_argument("-c", "--train_dataset", required=True, type=str)
14
parser.add_argument("-t", "--test_dataset", type=str, default=None)
15
parser.add_argument("-v", "--vocab_path", required=True, type=str)
16
parser.add_argument("-o", "--output_path", required=True, type=str)
@@ -23,7 +23,7 @@ def train():
23
parser.add_argument("-b", "--batch_size", type=int, default=64)
24
parser.add_argument("-e", "--epochs", type=int, default=10)
25
parser.add_argument("-w", "--num_workers", type=int, default=5)
26
- parser.add_argument("-c", "--with_cuda", type=bool, default=True)
+ parser.add_argument("--with_cuda", type=bool, default=True)
27
parser.add_argument("--log_freq", type=int, default=10)
28
parser.add_argument("--corpus_lines", type=int, default=None)
29
bert_pytorch/dataset/__init__.py
@@ -1,3 +1,2 @@
1
from .dataset import BERTDataset
2
-from .creator import BERTDatasetCreator
3
from .vocab import WordVocab
bert_pytorch/dataset/dataset.py
@@ -38,7 +38,7 @@ def __getitem__(self, item):
38
output = {"bert_input": bert_input,
39
"bert_label": bert_label,
40
"segment_label": segment_label,
41
- "is_next": self.datas[item]["is_next"]}
+ "is_next": is_next_label}
42
43
return {key: torch.tensor(value) for key, value in output.items()}
44
bert_pytorch/dataset/vocab.py
@@ -167,7 +167,7 @@ def load_vocab(vocab_path: str) -> 'WordVocab':
167
return pickle.load(f)
168
169
170
-if __name__ == "__main__":
+def build():
171
import argparse
172
173
0 commit comments