Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions bat/download-kifu.sh
Original file line number Diff line number Diff line change
@@ -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
46 changes: 44 additions & 2 deletions pydlshogi/read_kifu.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

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)
34 changes: 3 additions & 31 deletions train_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,9 +11,6 @@

import argparse
import random
import pickle
import os
import re

import logging

Expand Down Expand Up @@ -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)))
Expand Down
5 changes: 4 additions & 1 deletion utils/filter_csa.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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
Expand Down