entropy you silly nugget
This commit is contained in:
@@ -39,7 +39,7 @@ env_config = {
|
||||
# The order in which components reset, step, and generate obs follows their listed order below.
|
||||
'components': [
|
||||
# (1) Building houses
|
||||
('Craft', {'skill_dist': "pareto", 'commodities': ["Gem"],'max_skill_amount_benefit':1.5}),
|
||||
('Craft', {'skill_dist': "pareto", 'commodities': ["Gem"],'max_skill_amount_benefit':2}),
|
||||
# (2) Trading collectible resources
|
||||
('ContinuousDoubleAuction', {'max_num_orders': 10}),
|
||||
# (3) Movement and resource collection
|
||||
@@ -52,7 +52,7 @@ env_config = {
|
||||
# ===== SCENARIO CLASS ARGUMENTS =====
|
||||
# (optional) kwargs that are added by the Scenario class (i.e. not defined in BaseEnvironment)
|
||||
|
||||
'starting_agent_coin': 10,
|
||||
'starting_agent_coin': 20,
|
||||
'fixed_four_skill_and_loc': True,
|
||||
|
||||
# ===== STANDARD ARGUMENTS ======
|
||||
@@ -60,6 +60,7 @@ env_config = {
|
||||
'agent_composition': {"BasicMobileAgent": 20,"TradingAgent":5}, # Number of non-planner agents (must be > 1)
|
||||
'world_size': [5, 5], # [Height, Width] of the env world
|
||||
'episode_length': 256, # Number of timesteps per episode
|
||||
'isoelastic_eta':0.001,
|
||||
'allow_observation_scaling': True,
|
||||
'dense_log_frequency': 100,
|
||||
'world_dense_log_frequency':1,
|
||||
@@ -94,7 +95,7 @@ eval_env_config = {
|
||||
# The order in which components reset, step, and generate obs follows their listed order below.
|
||||
'components': [
|
||||
# (1) Building houses
|
||||
('Craft', {'skill_dist': "pareto", 'commodities': ["Gem"],'max_skill_amount_benefit':1.5}),
|
||||
('Craft', {'skill_dist': "pareto", 'commodities': ["Gem"],'max_skill_amount_benefit':2}),
|
||||
# (2) Trading collectible resources
|
||||
('ContinuousDoubleAuction', {'max_num_orders': 10}),
|
||||
# (3) Movement and resource collection
|
||||
@@ -107,7 +108,7 @@ eval_env_config = {
|
||||
# ===== SCENARIO CLASS ARGUMENTS =====
|
||||
# (optional) kwargs that are added by the Scenario class (i.e. not defined in BaseEnvironment)
|
||||
|
||||
'starting_agent_coin': 10,
|
||||
'starting_agent_coin': 20,
|
||||
'fixed_four_skill_and_loc': True,
|
||||
|
||||
# ===== STANDARD ARGUMENTS ======
|
||||
@@ -116,6 +117,7 @@ eval_env_config = {
|
||||
'world_size': [1, 1], # [Height, Width] of the env world
|
||||
'episode_length': 256, # Number of timesteps per episode
|
||||
'allow_observation_scaling': True,
|
||||
'isoelastic_eta':0.001,
|
||||
'dense_log_frequency': 1,
|
||||
'world_dense_log_frequency':1,
|
||||
'energy_cost':0,
|
||||
@@ -135,7 +137,7 @@ eval_env_config = {
|
||||
'flatten_masks': True,
|
||||
}
|
||||
|
||||
num_frames=5
|
||||
num_frames=1
|
||||
|
||||
class TensorboardCallback(BaseCallback):
|
||||
"""
|
||||
@@ -161,6 +163,23 @@ class TensorboardCallback(BaseCallback):
|
||||
|
||||
return True
|
||||
|
||||
min_at_target_basic=0.5
|
||||
min_lr_basic=5e-6
|
||||
start_lr_basic=9e-4
|
||||
|
||||
min_at_target_trade=0.5
|
||||
min_lr_trade=5e-6
|
||||
start_lr_trade=9e-4
|
||||
|
||||
def learning_rate_adj_basic(x) -> float:
|
||||
diff=start_lr_basic-min_lr_basic
|
||||
lr=min_lr_basic+x*diff
|
||||
return lr
|
||||
|
||||
def learning_rate_adj_trade(x) -> float:
|
||||
diff=start_lr_trade-min_lr_trade
|
||||
lr=min_lr_basic+x*diff
|
||||
return lr
|
||||
|
||||
def printMarket(market):
|
||||
for i in range(len(market)):
|
||||
@@ -273,37 +292,15 @@ runname="run_{}".format(run_number)
|
||||
model_db=[None,None] # object for storing model
|
||||
|
||||
|
||||
model = MaskablePPO("MlpPolicy",n_steps=int(env_config['episode_length']*2),ent_coef=0.1, vf_coef=0.5 ,gamma=0.99, learning_rate=1e-5,env=stackenv_basic, seed=300,verbose=1,device="cuda",tensorboard_log="./log")
|
||||
model_trade=MaskablePPO("MlpPolicy",n_steps=int(env_config['episode_length']*2),ent_coef=0.1, vf_coef=0.5 ,gamma=0.99, learning_rate=1e-5,env=stackenv_traid, seed=300,verbose=1,device="cuda",tensorboard_log="./log")
|
||||
model = MaskablePPO("MlpPolicy",n_steps=int(env_config['episode_length']*2),ent_coef=0.1, vf_coef=0.5 ,gamma=0.99, learning_rate=learning_rate_adj_basic,env=stackenv_basic, seed=445,verbose=1,device="cuda",tensorboard_log="./log")
|
||||
model_trade=MaskablePPO("MlpPolicy",n_steps=int(env_config['episode_length']*2),ent_coef=0.1, vf_coef=0.5 ,gamma=0.99, learning_rate=learning_rate_adj_trade,env=stackenv_traid, seed=445,verbose=1,device="cuda",tensorboard_log="./log")
|
||||
# Setup complete
|
||||
|
||||
n_agents=econ.n_agents
|
||||
|
||||
total_required_for_episode_basic=len(mobileRecieverEconWrapper.agnet_idx)*env_config['episode_length']
|
||||
total_required_for_episode_traid=len(tradeRecieverEconWrapper.agnet_idx)*env_config['episode_length']
|
||||
|
||||
print("this is run {}".format(runname))
|
||||
# Load models
|
||||
model.load("basic.ai")
|
||||
model_trade.load("trade.ai")
|
||||
|
||||
while True:
|
||||
|
||||
|
||||
#Train
|
||||
runname="run_{}_{}".format(run_number,"basic")
|
||||
|
||||
thread_model=Thread(target=train,args=(model,total_required_for_episode_basic*50,econ,True,runname,model_db,0))
|
||||
runname="run_{}_{}".format(run_number,"trader")
|
||||
thread_model_traid=Thread(target=train,args=(model_trade,total_required_for_episode_traid*50,econ,False,runname,model_db,1))
|
||||
|
||||
thread_model.start()
|
||||
thread_model_traid.start()
|
||||
thread_model.join()
|
||||
thread_model_traid.join()
|
||||
#normenv.save("temp-normalizer.ai")
|
||||
model=model_db[0]
|
||||
model_trade=model_db[1]
|
||||
model.save("basic.ai")
|
||||
model_trade.save("trade.ai")
|
||||
|
||||
## Run Eval
|
||||
print("### EVAL ###")
|
||||
obs_basic=stackenv_basic_eval.reset()
|
||||
obs_trade=stackenv_traid_eval.reset()
|
||||
@@ -329,8 +326,8 @@ while True:
|
||||
craft=econ_eval.get_component("Craft")
|
||||
# trades=market.get_dense_log()
|
||||
build=craft.get_dense_log()
|
||||
met=econ.previous_episode_metrics
|
||||
printReplay(econ_eval,0)
|
||||
met=econ_eval.previous_episode_metrics
|
||||
printReplay(econ_eval,21)
|
||||
# printMarket(trades)
|
||||
# printBuilds(builds=build)
|
||||
print("social/productivity: {}".format(met["social/productivity"]))
|
||||
|
||||
Reference in New Issue
Block a user