diff --git a/bat/download-kifu.sh b/bat/download-kifu.sh new file mode 100755 index 0000000..65fd979 --- /dev/null +++ b/bat/download-kifu.sh @@ -0,0 +1,17 @@ +# This shell script is supposed to run in a Datalab container to +# download Shogi kifu from Floodgate. +set -x +mkdir -p ../kifu/zip +cd ../kifu/zip +time wget -c --trust-server-names "https://osdn.net/frs/redir.php?m=jaist&f=shogi-server%2F68500%2Fwdoor2016.7z" +cd .. +which 7z || apt-get update -y && apt-get install -y p7zip-full --allow-unauthenticated +[ -d 2016 ] || time 7z x zip/wdoor2016.7z -aos # -y + +cd ../python-dlshogi +pip install python-shogi statistics tqdm +pip install --no-cache-dir -e . +time python utils/filter_csa.py ../kifu/2016/ # 20 minutes +time python utils/make_kifu_list.py ../kifu/2016/ ../kifu/kifulist +time python pydlshogi/read_kifu.py ../kifu/kifulist_test.txt +time python pydlshogi/read_kifu.py ../kifu/kifulist_train.txt diff --git a/pydlshogi/read_kifu.py b/pydlshogi/read_kifu.py index 1e56b81..7c819f3 100644 --- a/pydlshogi/read_kifu.py +++ b/pydlshogi/read_kifu.py @@ -2,13 +2,44 @@ import shogi.CSA import copy +import argparse +import logging +import os +import numpy as np +import re +from tqdm import tqdm + from pydlshogi.features import * +logging.basicConfig(format='%(asctime)s\t%(levelname)s\t%(message)s', + datefmt='%Y/%m/%d %H:%M:%S', + level=os.environ.get("LOGLEVEL", "DEBUG")) + + +# 保存済みのダンプファイルを読み込む +def load_dump(filename): + logging.info('Loading from %s' % (filename)) + with open(filename, 'rb') as f: + positions = np.load(f) + return positions + + +# ファイルに棋譜データをダンプする +def save_dump(filename, positions): + logging.info('Saving to %s' % (filename)) + np.array(positions).dump(filename) + + # read kifu def read_kifu(kifu_list_file): + logging.info('read kifu start') + dump_fname = re.sub(r'\.[^\.]+$', '', kifu_list_file) + '.pickle' + logging.info('dump_fname %s' % (dump_fname)) + if os.path.exists(dump_fname): return load_dump(dump_fname) + positions = [] with open(kifu_list_file, 'r') as f: - for line in f.readlines(): + for line in tqdm(f.readlines(), ncols=70): filepath = line.rstrip('\r\n') kifu = shogi.CSA.Parser.parse_file(filepath)[0] win_color = shogi.BLACK if kifu['win'] == 'b' else shogi.WHITE @@ -31,4 +62,15 @@ def read_kifu(kifu_list_file): positions.append((piece_bb, occupied, pieces_in_hand, move_label, win)) board.push_usi(move) - return positions \ No newline at end of file + + save_dump(dump_fname, positions) + + logging.info('read kifu end') + return positions + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('kifulist', type=str, help='kifu list') + args = parser.parse_args() + read_kifu(args.kifulist) \ No newline at end of file diff --git a/train_policy.py b/train_policy.py index bf4adfd..1563939 100644 --- a/train_policy.py +++ b/train_policy.py @@ -11,9 +11,6 @@ import argparse import random -import pickle -import os -import re import logging @@ -48,36 +45,11 @@ logging.info('Load optimizer state from {}'.format(args.resume)) serializers.load_npz(args.resume, optimizer) -logging.info('read kifu start') -# 保存済みのpickleファイルがある場合、pickleファイルを読み込む -# train date -train_pickle_filename = re.sub(r'\..*?$', '', args.kifulist_train) + '.pickle' -if os.path.exists(train_pickle_filename): - with open(train_pickle_filename, 'rb') as f: - positions_train = pickle.load(f) - logging.info('load train pickle') -else: - positions_train = read_kifu(args.kifulist_train) +# train data +positions_train = read_kifu(args.kifulist_train) # test data -test_pickle_filename = re.sub(r'\..*?$', '', args.kifulist_test) + '.pickle' -if os.path.exists(test_pickle_filename): - with open(test_pickle_filename, 'rb') as f: - positions_test = pickle.load(f) - logging.info('load test pickle') -else: - positions_test = read_kifu(args.kifulist_test) - -# 保存済みのpickleがない場合、pickleファイルを保存する -if not os.path.exists(train_pickle_filename): - with open(train_pickle_filename, 'wb') as f: - pickle.dump(positions_train, f, pickle.HIGHEST_PROTOCOL) - logging.info('save train pickle') -if not os.path.exists(test_pickle_filename): - with open(test_pickle_filename, 'wb') as f: - pickle.dump(positions_test, f, pickle.HIGHEST_PROTOCOL) - logging.info('save test pickle') -logging.info('read kifu end') +positions_test = read_kifu(args.kifulist_test) logging.info('train position num = {}'.format(len(positions_train))) logging.info('test position num = {}'.format(len(positions_test))) diff --git a/utils/filter_csa.py b/utils/filter_csa.py index b313d9a..692badc 100644 --- a/utils/filter_csa.py +++ b/utils/filter_csa.py @@ -3,6 +3,8 @@ import re import statistics +from tqdm import tqdm + parser = argparse.ArgumentParser() parser.add_argument('dir', type=str) args = parser.parse_args() @@ -16,7 +18,8 @@ def find_all_files(directory): kifu_count = 0 rates = [] -for filepath in find_all_files(args.dir): +n_kifu = len(list(find_all_files(args.dir))) +for filepath in tqdm(find_all_files(args.dir), total=n_kifu, ncols=70): rate = {} move_len = 0 toryo = False