from pysc2.agents import base_agent from pysc2.env import sc2_env from pysc2.lib import actions, features, units from absl import app import random import numpy as np import pandas as pd import time # from Steven Brown: https://itnext.io/build-a-zerg-bot-with-pysc2-2-0-295375d2f58e # from Steven Brown: https://chatbotslife.com/building-a-smart-pysc2-agent-cdc269cb095d _PLAYER_SELF = 1 _NOT_QUEUED = [0] _QUEUED = [1] ACTION_DO_NOTHING = 'donothing' ACTION_SELECT_SCV = 'selectscv' ACTION_BUILD_SUPPLY_DEPOT = 'buildsupplydepot' ACTION_BUILD_BARRACKS = 'buildbarracks' ACTION_SELECT_BARRACKS = 'selectbarracks' ACTION_BUILD_MARINE = 'buildmarine' ACTION_SELECT_ARMY = 'selectarmy' ACTION_ATTACK = 'attack' smart_actions = [ ACTION_DO_NOTHING, ACTION_SELECT_SCV, ACTION_BUILD_SUPPLY_DEPOT, ACTION_BUILD_BARRACKS, ACTION_SELECT_BARRACKS, ACTION_BUILD_MARINE, ACTION_SELECT_ARMY, ACTION_ATTACK, ] KILL_UNIT_REWARD = 0.2 KILL_BUILDING_REWARD = 0.5 class QLearningTable: def __init__(self, actions, learning_rate = 0.01, reward_decay = 0.9, e_greedy = 0.9): self.actions = actions self.lr = learning_rate self.gamma = reward_decay self.epsilon = e_greedy self.q_table = pd.DataFrame(columns=self.actions, dtype=np.float64) def choose_action(self, observation): self.check_state_exists(observation) if np.random.uniform() < self.epsilon: # choose best action state_action = self.q_table.ix[observation, :] # some actions have the same value state_action = state_action.reindex(np.random.permutation(state_action.index)) action = state_action.idxmax() else: # choose random action action = np.random.choice(self.actions) return action def learn(self, s, a, r, s_): self.check_state_exists(s_) self.check_state_exists(s) q_predict = self.q_table.ix[s, a] q_target = r + self.gamma * self.q_table.ix[s_, :].max() # current_state = [ # supply_depot_count, # barracks_count, # supply_limit, # army_supply, # ] # update self.q_table.ix[s, a] += self.lr * (q_target - q_predict) def check_state_exists(self, state): if state not in self.q_table.index: #append new state to q table self.q_table = self.q_table.append(pd.Series([0] * len(self.actions), index=self.q_table.columns, name=state)) class SmartAgent(base_agent.BaseAgent): def __init__(self): super(SmartAgent, self).__init__() self.qlearn = QLearningTable(actions=list(range(len(smart_actions)))) self.previous_killed_unit_score = 0 self.previous_killed_building_score = 0 self.previous_action = None self.previous_state = None def can_do(self, obs, action): return action in obs.observation.available_actions def transformLocation(self, x, x_distance, y, y_distance): if not self.base_top_left: return [x - x_distance, y - y_distance] return [x + x_distance, y + y_distance] def unit_type_is_selected(self, obs, unit_type): if (len(obs.observation.single_select) > 0 and obs.observation.single_select[0].unit_type == unit_type): return True if (len(obs.observation.multi_select) > 0 and obs.observation.multi_select[0].unit_type == unit_type): return True return False def get_units_by_type(self, obs, unit_type): return [unit for unit in obs.observation.feature_units if unit.unit_type == unit_type] def step(self, obs): super(SmartAgent, self).step(obs) # Figure out where the player is: top left or bottom right player_y, player_x = (obs.observation.feature_minimap.player_relative == _PLAYER_SELF).nonzero() self.base_top_left = 1 if player_y.any() and player_y.mean() <= 31 else 0 # prefetch values for current state supply_depot_count = len(self.get_units_by_type(obs, units.Terran.SupplyDepot)) barracks_count = len(self.get_units_by_type(obs, units.Terran.Barracks)) supply_limit = obs.observation.player.food_cap army_supply = obs.observation.player.food_army current_state = [ supply_depot_count, barracks_count, supply_limit, army_supply, ] killed_unit_score = obs.observation.score_cumulative.killed_value_units killed_building_score = obs.observation.score_cumulative.killed_value_structures if self.previous_action is not None: reward = 0 if killed_unit_score > self.previous_killed_unit_score: reward += KILL_UNIT_REWARD if killed_building_score > self.previous_killed_building_score: reward += KILL_BUILDING_REWARD self.qlearn.learn(str(self.previous_state), self.previous_action, reward, str(current_state)) rl_action = self.qlearn.choose_action(str(current_state)) smart_action = smart_actions[rl_action] self.previous_killed_unit_score = killed_unit_score self.previous_killed_building_score = killed_building_score self.previous_state = current_state self.previous_action = rl_action print(smart_action) # no_op if smart_action == ACTION_DO_NOTHING: return actions.FUNCTIONS.no_op() # select an SCV elif smart_action == ACTION_SELECT_SCV: scvs = self.get_units_by_type(obs, units.Terran.SCV) #time.sleep(0.5) if len(scvs) > 0: #time.sleep(0.5) scv = random.choice(scvs) #print(scv.x, scv.y) if (scv.x > 0) & (scv.x < 83) & (scv.y > 0) & (scv.y < 83): return actions.FUNCTIONS.select_point("select_all_type", (scv.x, scv.y)) else: return actions.FUNCTIONS.no_op() # build a supply depot elif smart_action == ACTION_BUILD_SUPPLY_DEPOT: if self.can_do(obs, actions.FUNCTIONS.Build_SupplyDepot_screen.id): x = random.randint(0, 83) y = random.randint(0, 83) return actions.FUNCTIONS.Build_SupplyDepot_screen("now", (x, y)) elif smart_action == ACTION_BUILD_BARRACKS: if self.can_do(obs, actions.FUNCTIONS.Build_Barracks_screen.id): x = random.randint(0, 83) y = random.randint(0, 83) return actions.FUNCTIONS.Build_Barracks_screen("now", (x, y)) elif smart_action == ACTION_SELECT_BARRACKS: barracks = self.get_units_by_type(obs, units.Terran.Barracks) if len(barracks) > 0: barrack = random.choice(barracks) return actions.FUNCTIONS.select_point("select_all_type", (barrack.x, barrack.y)) elif smart_action == ACTION_BUILD_MARINE: if self.can_do(obs, actions.FUNCTIONS.Train_Marine_quick.id): return actions.FUNCTIONS.Train_Marine_quick("now") elif smart_action == ACTION_SELECT_ARMY: if self.can_do(obs, actions.FUNCTIONS.select_army.id): return actions.FUNCTIONS.select_army("select") elif smart_action == ACTION_ATTACK: if self.can_do(obs, actions.FUNCTIONS.Attack_minimap.id): if self.base_top_left: return actions.FUNCTIONS.Attack_minimap("now", (39, 45)) return actions.FUNCTIONS.Attack_minimap("now", (21, 24)) return actions.FUNCTIONS.no_op() def main(unused_argv): agent = SmartAgent() try: while True: with sc2_env.SC2Env( map_name="Simple64", players=[sc2_env.Agent(sc2_env.Race.terran), sc2_env.Bot(sc2_env.Race.random, sc2_env.Difficulty.very_easy)], agent_interface_format=features.AgentInterfaceFormat( feature_dimensions=features.Dimensions(screen=84, minimap=64), use_feature_units=True,), step_mul=16, 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])] if timesteps[0].last(): break timesteps = env.step(step_actions) except KeyboardInterrupt: pass if __name__ == "__main__": app.run(main)