238 lines
9.7 KiB
Python
238 lines
9.7 KiB
Python
# from __future__ import print_function
|
|
# from __future__ import absolute_import
|
|
import os
|
|
import glob
|
|
import os.path
|
|
import tarfile
|
|
import gzip
|
|
import shutil
|
|
import numpy as np
|
|
import multiprocessing
|
|
from os import sys
|
|
import torch
|
|
|
|
from tugo.gosgf import Sgf_game
|
|
from game_logic.goboard_fast import Board, GameState, Move
|
|
from game_logic.gotypes import Player, Point
|
|
from tugo.data_processing.index_processor import KGSIndex
|
|
from tugo.data_processing.sampling import Sampler
|
|
from tugo.data_processing.generator import DataGenerator
|
|
from tugo.encoders.base import get_encoder_by_name
|
|
from torch.utils.data import TensorDataset
|
|
|
|
|
|
def worker(jobinfo):
|
|
try:
|
|
clazz, encoder, zip_file, data_file_name, game_list = jobinfo
|
|
clazz(encoder=encoder).process_zip(zip_file, data_file_name, game_list)
|
|
except (KeyboardInterrupt, SystemExit):
|
|
raise Exception('>>> Exiting child process.')
|
|
|
|
|
|
class GoDataProcessor:
|
|
def __init__(self, encoder='simple', data_directory='train_data'):
|
|
self.encoder_string = encoder
|
|
self.encoder = get_encoder_by_name(encoder, 19)
|
|
self.data_dir = data_directory
|
|
|
|
def load_go_data(self, data_type='train', num_samples=1000,
|
|
use_generator=False):
|
|
index = KGSIndex(data_directory=self.data_dir)
|
|
index.download_files()
|
|
|
|
sampler = Sampler(data_dir=self.data_dir)
|
|
data = sampler.draw_data(data_type, num_samples)
|
|
|
|
self.map_to_workers(data_type, data)
|
|
if use_generator:
|
|
generator = DataGenerator(self.data_dir, data)
|
|
return generator
|
|
else:
|
|
features_and_labels = self.consolidate_games(data_type, data)
|
|
return features_and_labels
|
|
|
|
def unzip_data(self, zip_file_name):
|
|
this_gz = gzip.open(self.data_dir + '/' + zip_file_name)
|
|
|
|
tar_file = zip_file_name[0:-3]
|
|
this_tar = open(self.data_dir + '/' + tar_file, 'wb')
|
|
|
|
shutil.copyfileobj(this_gz, this_tar)
|
|
this_tar.close()
|
|
return tar_file
|
|
|
|
def process_zip(self, zip_file_name, data_file_name, game_list):
|
|
print(">>> Processing zip file:", zip_file_name, "Data file name:", data_file_name)
|
|
tar_file = self.unzip_data(zip_file_name)
|
|
zip_file = tarfile.open(self.data_dir + '/' + tar_file)
|
|
name_list = zip_file.getnames()
|
|
total_examples = self.num_total_examples(zip_file, game_list, name_list)
|
|
print(">>> Total examples:", total_examples)
|
|
|
|
shape = self.encoder.shape()
|
|
feature_shape = np.insert(shape, 0, np.asarray([total_examples]))
|
|
features = np.zeros(feature_shape)
|
|
labels = np.zeros((total_examples,))
|
|
|
|
counter = 0
|
|
for index in game_list:
|
|
name = name_list[index + 1]
|
|
if not name.endswith('.sgf'):
|
|
raise ValueError(name + ' is not a valid sgf')
|
|
sgf_content = zip_file.extractfile(name).read()
|
|
sgf = Sgf_game.from_string(sgf_content)
|
|
if sgf.get_handicap() is not None and sgf.get_handicap() != 0:
|
|
print('<func process_zip>', 'get_handicap:', sgf.get_handicap(), ' => skipping handicaped game ...')
|
|
continue
|
|
|
|
game_state, first_move_done = self.get_handicap(sgf)
|
|
|
|
for item in sgf.main_sequence_iter():
|
|
color, move_tuple = item.get_move()
|
|
point = None
|
|
if color is not None:
|
|
if move_tuple is not None:
|
|
row, col = move_tuple
|
|
point = Point(row + 1, col + 1)
|
|
move = Move.play(point)
|
|
else:
|
|
move = Move.pass_turn()
|
|
if first_move_done and point is not None:
|
|
features[counter] = self.encoder.encode(game_state)
|
|
labels[counter] = self.encoder.encode_point(point)
|
|
counter += 1
|
|
game_state = game_state.apply_move(move)
|
|
first_move_done = True
|
|
|
|
feature_file_base = self.data_dir + '/' + data_file_name + '_features_%d'
|
|
label_file_base = self.data_dir + '/' + data_file_name + '_labels_%d'
|
|
|
|
chunk = 0 # Due to files with large content, split up after chunksize
|
|
chunksize = 1024
|
|
while features.shape[0] > 0:
|
|
feature_file = feature_file_base % chunk
|
|
label_file = label_file_base % chunk
|
|
chunk += 1
|
|
# current_features, features = features[:chunksize], features[chunksize:]
|
|
# current_labels, labels = labels[:chunksize], labels[chunksize:]
|
|
current_chunksize = min(chunksize, features.shape[0])
|
|
current_features, features = features[:current_chunksize], features[current_chunksize:]
|
|
current_labels, labels = labels[:current_chunksize], labels[current_chunksize:]
|
|
np.save(feature_file, current_features)
|
|
# print(">>>>>> save feature_file:", feature_file)
|
|
np.save(label_file, current_labels)
|
|
|
|
def consolidate_games(self, name, samples):
|
|
files_needed = set(file_name for file_name, index in samples)
|
|
file_names = []
|
|
for zip_file_name in files_needed:
|
|
file_name = zip_file_name.replace('.tar.gz', '') + name
|
|
file_names.append(file_name)
|
|
|
|
feature_list = []
|
|
label_list = []
|
|
for file_name in file_names:
|
|
file_prefix = file_name.replace('.tar.gz', '')
|
|
base = self.data_dir + '/' + file_prefix + '_features_*.npy'
|
|
# print("base:", base)
|
|
# print(glob.glob(base))
|
|
for feature_file in glob.glob(base):
|
|
# print(f"Processing feature file: {feature_file}")
|
|
label_file = feature_file.replace('features', 'labels')
|
|
x = np.load(feature_file)
|
|
y = np.load(label_file)
|
|
x = x.astype('float32')
|
|
# y = torch.nn.functional.one_hot(torch.tensor(y.astype(int)), 19 * 19)
|
|
y = torch.tensor(y.astype(int))
|
|
# print("Feature shape:", x.shape)
|
|
# print("Label shape:", y.shape)
|
|
feature_list.append(x)
|
|
label_list.append(y)
|
|
|
|
assert feature_list, "No features found. Please check the data files."
|
|
|
|
features = torch.tensor(np.concatenate(feature_list, axis=0)).float()
|
|
labels = torch.cat(label_list, axis=0)
|
|
|
|
feature_file = self.data_dir + '/' + name + '_feature.pt'
|
|
label_file = self.data_dir + '/' + name + '_label.pt'
|
|
|
|
torch.save(features, feature_file)
|
|
torch.save(labels, label_file)
|
|
|
|
dataset = TensorDataset(torch.tensor(features), labels)
|
|
return dataset
|
|
|
|
|
|
|
|
@staticmethod
|
|
def get_handicap(sgf): # Get handicap stones
|
|
go_board = Board(19, 19)
|
|
first_move_done = False
|
|
move = None
|
|
game_state = GameState.new_game(19)
|
|
if sgf.get_handicap() is not None and sgf.get_handicap() != 0:
|
|
# print('get_handicap2:', sgf.get_handicap())
|
|
for setup in sgf.get_root().get_setup_stones():
|
|
# print("------- setup:", setup, type(setup))
|
|
for move in setup:
|
|
# print("------- move:", move, type(move))
|
|
row, col = move
|
|
point = Point(row + 1, col + 1)
|
|
move = Move.play(point)
|
|
go_board.place_stone(Player.black, Point(row + 1, col + 1)) # black gets handicap
|
|
first_move_done = True
|
|
game_state = GameState(go_board, Player.white, None, move)
|
|
return game_state, first_move_done
|
|
|
|
def map_to_workers(self, data_type, samples):
|
|
zip_names = set()
|
|
indices_by_zip_name = {}
|
|
for filename, index in samples:
|
|
zip_names.add(filename)
|
|
if filename not in indices_by_zip_name:
|
|
indices_by_zip_name[filename] = []
|
|
indices_by_zip_name[filename].append(index)
|
|
|
|
zips_to_process = []
|
|
for zip_name in zip_names:
|
|
base_name = zip_name.replace('.tar.gz', '')
|
|
data_file_name = base_name + data_type
|
|
if not os.path.isfile(self.data_dir + '/' + data_file_name):
|
|
zips_to_process.append((self.__class__, self.encoder_string, zip_name,
|
|
data_file_name, indices_by_zip_name[zip_name]))
|
|
cores = multiprocessing.cpu_count() # Determine number of CPU cores and split work load among them
|
|
pool = multiprocessing.Pool(processes=cores)
|
|
p = pool.map_async(worker, zips_to_process)
|
|
try:
|
|
_ = p.get()
|
|
except KeyboardInterrupt: # Caught keyboard interrupt, terminating workers
|
|
pool.terminate()
|
|
pool.join()
|
|
sys.exit(-1)
|
|
|
|
def num_total_examples(self, zip_file, game_list, name_list):
|
|
total_examples = 0
|
|
for index in game_list:
|
|
name = name_list[index + 1]
|
|
if name.endswith('.sgf'):
|
|
sgf_content = zip_file.extractfile(name).read()
|
|
sgf = Sgf_game.from_string(sgf_content)
|
|
if sgf.get_handicap() is not None and sgf.get_handicap() != 0:
|
|
print('get_handicap:', sgf.get_handicap(), ' => skipping handicaped game ...')
|
|
continue
|
|
game_state, first_move_done = self.get_handicap(sgf)
|
|
|
|
num_moves = 0
|
|
for item in sgf.main_sequence_iter():
|
|
color, move = item.get_move()
|
|
if color is not None:
|
|
if first_move_done:
|
|
num_moves += 1
|
|
first_move_done = True
|
|
total_examples = total_examples + num_moves
|
|
else:
|
|
raise ValueError(name + ' is not a valid sgf')
|
|
return total_examples
|
|
|