dmlc--dgl
9f32554296
* change the signature of node/edge filter * upd filter * Support multi-dimension node feature in SPMV * push transformer * remove some experimental settings * stable version * hotfix * upd tutorial * upd README * merge * remove redundency * remove tqdm * several changes * Refactor * Refactor * tutorial train * fixed a bug * fixed perf issue * upd * change dir * move un-related to contrib * tutuorial code * remove redundency * upd * upd * upd * upd * improve viz * universal done * halt norm * fixed a bug * add draw graph * fixed several bugs * remove dependency on core * upd format of README * trigger * trigger * upd viz * trigger * add transformer tutorial * fix tutorial * fix readme * small fix on tutorials * url fix in readme * fixed func link * upd
93 行
4.0 KiB
Python
93 行
4.0 KiB
Python
import numpy as np
|
|
import torch as th
|
|
import os
|
|
from dgl.data.utils import *
|
|
|
|
_urls = {
|
|
'wmt': 'https://s3.us-east-2.amazonaws.com/dgl.ai/dataset/wmt16_en_de.tar.gz',
|
|
'scripts': 'https://s3.us-east-2.amazonaws.com/dgl.ai/dataset/transformer_scripts.zip',
|
|
}
|
|
|
|
def prepare_dataset(dataset_name):
|
|
"download and generate datasets"
|
|
script_dir = os.path.join('scripts')
|
|
if not os.path.exists(script_dir):
|
|
download(_urls['scripts'], path='scripts.zip')
|
|
extract_archive('scripts.zip', 'scripts')
|
|
|
|
directory = os.path.join('data', dataset_name)
|
|
if not os.path.exists(directory):
|
|
os.makedirs(directory)
|
|
else:
|
|
return
|
|
if dataset_name == 'multi30k':
|
|
os.system('bash scripts/prepare-multi30k.sh')
|
|
elif dataset_name == 'wmt14':
|
|
download(_urls['wmt'], path='wmt16_en_de.tar.gz')
|
|
os.system('bash scripts/prepare-wmt14.sh')
|
|
elif dataset_name == 'copy' or dataset_name == 'tiny_copy':
|
|
train_size = 9000
|
|
valid_size = 1000
|
|
test_size = 1000
|
|
char_list = [chr(i) for i in range(ord('a'), ord('z') + 1)]
|
|
with open(os.path.join(directory, 'train.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'train.out'), 'w') as f_out:
|
|
for i, l in zip(range(train_size), np.random.normal(15, 3, train_size).astype(int)):
|
|
l = max(l, 1)
|
|
line = ' '.join(np.random.choice(char_list, l)) + '\n'
|
|
f_in.write(line)
|
|
f_out.write(line)
|
|
|
|
with open(os.path.join(directory, 'valid.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'valid.out'), 'w') as f_out:
|
|
for i, l in zip(range(valid_size), np.random.normal(15, 3, valid_size).astype(int)):
|
|
l = max(l, 1)
|
|
line = ' '.join(np.random.choice(char_list, l)) + '\n'
|
|
f_in.write(line)
|
|
f_out.write(line)
|
|
|
|
with open(os.path.join(directory, 'test.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'test.out'), 'w') as f_out:
|
|
for i, l in zip(range(test_size), np.random.normal(15, 3, test_size).astype(int)):
|
|
l = max(l, 1)
|
|
line = ' '.join(np.random.choice(char_list, l)) + '\n'
|
|
f_in.write(line)
|
|
f_out.write(line)
|
|
|
|
with open(os.path.join(directory, 'vocab.txt'), 'w') as f:
|
|
for c in char_list:
|
|
f.write(c + '\n')
|
|
|
|
elif dataset_name == 'sort' or dataset_name == 'tiny_sort':
|
|
train_size = 9000
|
|
valid_size = 1000
|
|
test_size = 1000
|
|
char_list = [chr(i) for i in range(ord('a'), ord('z') + 1)]
|
|
with open(os.path.join(directory, 'train.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'train.out'), 'w') as f_out:
|
|
for i, l in zip(range(train_size), np.random.normal(15, 3, train_size).astype(int)):
|
|
l = max(l, 1)
|
|
seq = np.random.choice(char_list, l)
|
|
f_in.write(' '.join(seq) + '\n')
|
|
f_out.write(' '.join(np.sort(seq)) + '\n')
|
|
|
|
with open(os.path.join(directory, 'valid.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'valid.out'), 'w') as f_out:
|
|
for i, l in zip(range(valid_size), np.random.normal(15, 3, valid_size).astype(int)):
|
|
l = max(l, 1)
|
|
seq = np.random.choice(char_list, l)
|
|
f_in.write(' '.join(seq) + '\n')
|
|
f_out.write(' '.join(np.sort(seq)) + '\n')
|
|
|
|
with open(os.path.join(directory, 'test.in'), 'w') as f_in,\
|
|
open(os.path.join(directory, 'test.out'), 'w') as f_out:
|
|
for i, l in zip(range(test_size), np.random.normal(15, 3, test_size).astype(int)):
|
|
l = max(l, 1)
|
|
seq = np.random.choice(char_list, l)
|
|
f_in.write(' '.join(seq) + '\n')
|
|
f_out.write(' '.join(np.sort(seq)) + '\n')
|
|
|
|
with open(os.path.join(directory, 'vocab.txt'), 'w') as f:
|
|
for c in char_list:
|
|
f.write(c + '\n')
|