Files
tugo/data_processing/parallel_processor.py
2023-05-30 17:16:48 +08:00

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