mirror of
https://github.com/wassname/Castor.git
synced 2026-09-09 11:13:20 +08:00
Update data path (#127)
This commit is contained in:
+4
-4
@@ -31,12 +31,12 @@ class DatasetFactory(object):
|
||||
@staticmethod
|
||||
def get_dataset(dataset_name, word_vectors_dir, word_vectors_file, batch_size, device, castor_dir="./", utils_trecqa="utils/trec_eval-9.0.5/trec_eval"):
|
||||
if dataset_name == 'sick':
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'sick/')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'datasets', 'sick/')
|
||||
train_loader, dev_loader, test_loader = SICK.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
embedding = nn.Embedding.from_pretrained(SICK.TEXT_FIELD.vocab.vectors)
|
||||
return SICK, embedding, train_loader, test_loader, dev_loader
|
||||
elif dataset_name == 'msrvid':
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'msrvid/')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'datasets', 'msrvid/')
|
||||
dev_loader = None
|
||||
train_loader, test_loader = MSRVID.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
embedding = nn.Embedding.from_pretrained(MSRVID.TEXT_FIELD.vocab.vectors)
|
||||
@@ -44,14 +44,14 @@ class DatasetFactory(object):
|
||||
elif dataset_name == 'trecqa':
|
||||
if not os.path.exists(os.path.join(castor_dir, utils_trecqa)):
|
||||
raise FileNotFoundError('TrecQA requires the trec_eval tool to run. Please run get_trec_eval.sh inside Castor/utils (as working directory) before continuing.')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'TrecQA/')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'datasets', 'TrecQA/')
|
||||
train_loader, dev_loader, test_loader = TRECQA.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
embedding = nn.Embedding.from_pretrained(TRECQA.TEXT_FIELD.vocab.vectors)
|
||||
return TRECQA, embedding, train_loader, test_loader, dev_loader
|
||||
elif dataset_name == 'wikiqa':
|
||||
if not os.path.exists(os.path.join(castor_dir, utils_trecqa)):
|
||||
raise FileNotFoundError('TrecQA requires the trec_eval tool to run. Please run get_trec_eval.sh inside Castor/utils (as working directory) before continuing.')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'WikiQA/')
|
||||
dataset_root = os.path.join(castor_dir, os.pardir, 'Castor-data', 'datasets', 'WikiQA/')
|
||||
train_loader, dev_loader, test_loader = WikiQA.iters(dataset_root, word_vectors_file, word_vectors_dir, batch_size, device=device, unk_init=UnknownWordVecCache.unk)
|
||||
embedding = nn.Embedding.from_pretrained(WikiQA.TEXT_FIELD.vocab.vectors)
|
||||
return WikiQA, embedding, train_loader, test_loader, dev_loader
|
||||
|
||||
Reference in New Issue
Block a user