from pysc2.agents import base_agent
from pysc2.env import sc2_env
from pysc2.lib import actions, features, units
from absl import app
import time
from datetime import datetime
import numpy as np
SCREEN_SIZE = 84
MINIMAP_SIZE = 64
# A3C example from: https://github.com/greentfrapp/pysc2-RLagents/blob/master/Agents/PySC2_A3C_AtariNet.py
"""
Use the following command to launch Tensorboard:
tensorboard --logdir=worker_0:'./train_0',worker_1:'./train_1',worker_2:'./train_2',worker_3:'./train_3'
"""
# helper functions
# copies one set of variables to another
# used to set worker network parameters to those of the global network
def update_target_graph(from_scope, to_scope):
from_vars = tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES, from_scope)
to_vars = tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES, to_scope)
op_holder = []
for from_var, to_var, in zip(from_vars, to_vars):
op_holder.append(to_var.assign(from_var))
return op_holder
def process_observation(observation, action_spec):
# reward
reward = observation.reward
# features
features = observation.observation
# these have spatial data that we should process with a conv net
spatial_features = ['feature_minimap','feature_screen']
# they're variable features because they change size depending on what you select
variable_features = ['cargo','multi_select','build_queue']
available_actions = ['available_actions']
# shapes of some features depend on the state
# tf requires fixed-input shapes, set max size then pad the input if it falls short
max_no = {'available_actions':len(action_spec.functions), 'cargo': 500, 'multi_select': 500, 'build_queue': 10}
non_spatial_stack = []
for feature_label in observation.observation:
if feature_label not in spatial_features + variable_features + available_actions:
non_spatial_stack = np.concatenate((non_spatial_stack, feature_label.reshape(-1)))
class minigame_agent(base_agent.BaseAgent):
def __init__(self):
super(minigame_agent, self).__init__()
def step(self, obs):
super(minigame_agent, self).step(obs)
# single_select will always be a tuple(7,) i.e. invarient dims
#raveled_single_select = np.ravel(obs.observation.single_select[0]).tolist()
#single_select = raveled_single_select
# multi_select will be a tuple(n,7) *if* something's selected.
# If nothing is selected multi_select will return NoneType
# DefeatRoaches has at most 9 marines you can select...
# since there's gonna be a maximum of 200 units, why not map onto a bigger vector of zeros?
#raveled_multi_select = np.ravel(obs.observation.multi_select).tolist()
#multi_select = [0] * 7 * 9
#for i, item in enumerate(raveled_multi_select):
# multi_select[i] = item
# feature_screen should always be a tuple(17, SCREEN_SIZE, SCREEN_SIZE)
#feature_screen = obs.observation.feature_screen.tolist()
# feature_minimap should always be a tuple(7, SCREEN_SIZE, SCREEN_SIZE)
#feature_minimap = obs.observation.feature_minimap.tolist()
# player should always be a tuple(11,)
#raveled_player = np.ravel(obs.observation.player).tolist()
#player = raveled_player
# feature_units will be a tuple(n,26) where n is the number of units present on the map
# since there's a maximum of 200 units why not map onto a bigger vector of zeros?
# in DefeatRoaches, there are a maximum of 9+5 units
#raveled_feature_units = np.ravel(obs.observation.feature_units).tolist()
#feature_units = [0] * 26 * 14
#for i, item in enumerate(raveled_feature_units):
# feature_units[i] = item
# Since actions are sequenced, it would be nice to have them as a one-hot implementation; 523 total actions
#raveled_available_actions = np.ravel(obs.observation.available_actions).tolist()
#available_actions = [0] * 523
#for i, item in enumerate(raveled_available_actions):
# available_actions[item] = 1
# what happens if we just work off the nonspatial list first?
#nonspatial_list = single_select + multi_select + player + feature_units + available_actions
#process_observation(self.obs_spec, self.action_spec)
# put a break point here to see what's going on
return actions.FUNCTIONS.no_op()
def main(unused_argv):
agent = minigame_agent()
try:
while True:
try:
with sc2_env.SC2Env(
map_name="DefeatRoaches",
agent_interface_format=features.AgentInterfaceFormat(
feature_dimensions=features.Dimensions(screen=SCREEN_SIZE,
minimap=MINIMAP_SIZE),
use_feature_units=True),
step_mul=1,
game_steps_per_episode=0,
visualize=True,
) as env:
agent.setup(env.observation_spec(), env.action_spec())
timesteps = env.reset()
agent.reset()
while True:
step_actions = [agent.step(timesteps[0])]
timesteps = env.step(step_actions)
except KeyboardInterrupt:
raise
except Exception as e:
with open('log.txt', 'a') as myfile:
myfile.write('At %s recorded exception: %s.\n' %
(datetime.now(), repr(e)))
except KeyboardInterrupt:
pass
if __name__ == "__main__":
app.run(main)
Comments