Use rasterized speed
This commit is contained in:
@@ -8,4 +8,5 @@
|
|||||||
#python -m expert --env=NRasterizedRandomAgent --min_timesteps=10000 --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentw36h36mppx2.pkl'
|
#python -m expert --env=NRasterizedRandomAgent --min_timesteps=10000 --env_args='{width:36,height:36,m_per_px:2}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRandomAgentw36h36mppx2.pkl'
|
||||||
#python -m expert --env=NRasterizedRouteRandomAgent --min_timesteps=10000 --env_args='{width:70,height:70,m_per_px:1}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteRandomAgentw70h70mppx1.pkl'
|
#python -m expert --env=NRasterizedRouteRandomAgent --min_timesteps=10000 --env_args='{width:70,height:70,m_per_px:1}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteRandomAgentw70h70mppx1.pkl'
|
||||||
#python -m expert --env=NRasterizedRouteRandomAgentLocation --min_timesteps=100000 --env_args='{width:70,height:70,m_per_px:1}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N100000_NRasterizedRouteRandomAgentLocationw70h70mppx1.pkl'
|
#python -m expert --env=NRasterizedRouteRandomAgentLocation --min_timesteps=100000 --env_args='{width:70,height:70,m_per_px:1}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N100000_NRasterizedRouteRandomAgentLocationw70h70mppx1.pkl'
|
||||||
python -m expert --env=NRasterizedRouteRandomAgentLocation --min_timesteps=100000 --env_args='{width:70,height:70,m_per_px:1,map_color:128}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N100000_NRasterizedRouteRandomAgentLocationw70h70mppx1mapc128.pkl'
|
#python -m expert --env=NRasterizedRouteRandomAgentLocation --min_timesteps=100000 --env_args='{width:70,height:70,m_per_px:1,map_color:128}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N100000_NRasterizedRouteRandomAgentLocationw70h70mppx1mapc128.pkl'
|
||||||
|
python -m expert --env=NRasterizedRouteSpeedRandomAgentLocation --min_timesteps=10000 --env_args='{width:70,height:70,m_per_px:1,map_color:128,mu:0.001}' --policy_args='{mu:0.001}' --path='NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteSpeedRandomAgentLocationw70h70mppx1mapc128mu.001.pkl'
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from imitation.algorithms import adversarial
|
|||||||
import stable_baselines3
|
import stable_baselines3
|
||||||
import torch.utils.data
|
import torch.utils.data
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from intersim.envs import NRasterizedRouteRandomAgentLocation
|
from intersim.envs import NRasterizedRouteSpeedRandomAgentLocation
|
||||||
import itertools
|
import itertools
|
||||||
from torch.distributions import Categorical
|
from torch.distributions import Categorical
|
||||||
import gym
|
import gym
|
||||||
@@ -37,7 +37,7 @@ def train(
|
|||||||
n_disc_updates_per_round=2,
|
n_disc_updates_per_round=2,
|
||||||
n_gen_updates_per_round=10,
|
n_gen_updates_per_round=10,
|
||||||
):
|
):
|
||||||
env = NRasterizedRouteRandomAgentLocation(**env_settings)
|
env = NRasterizedRouteSpeedRandomAgentLocation(**env_settings)
|
||||||
env.discount = discount
|
env.discount = discount
|
||||||
|
|
||||||
tempdir = tempfile.TemporaryDirectory(prefix="quickstart")
|
tempdir = tempfile.TemporaryDirectory(prefix="quickstart")
|
||||||
@@ -45,7 +45,7 @@ def train(
|
|||||||
logger.configure(tempdir_path / "GAIL/")
|
logger.configure(tempdir_path / "GAIL/")
|
||||||
print(f"All Tensorboards and logging are being written inside {tempdir_path}/.")
|
print(f"All Tensorboards and logging are being written inside {tempdir_path}/.")
|
||||||
|
|
||||||
venv = make_vec_env(NRasterizedRouteRandomAgentLocation, n_envs=1, env_kwargs=env_settings)
|
venv = make_vec_env(NRasterizedRouteSpeedRandomAgentLocation, n_envs=1, env_kwargs=env_settings)
|
||||||
discriminator = adversarial.GAIL(
|
discriminator = adversarial.GAIL(
|
||||||
expert_data=expert_data,
|
expert_data=expert_data,
|
||||||
expert_batch_size=expert_batch_size,
|
expert_batch_size=expert_batch_size,
|
||||||
@@ -88,13 +88,13 @@ def video(model_name, env):
|
|||||||
def evaluate():
|
def evaluate():
|
||||||
video(
|
video(
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
env=NRasterizedRouteRandomAgentLocation(**env_settings)
|
env=NRasterizedRouteSpeedRandomAgentLocation(**env_settings)
|
||||||
)
|
)
|
||||||
|
|
||||||
# %%
|
# %%
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
|
|
||||||
with open("data/NormalizedIntersimpleExpertMu.001N100000_NRasterizedRouteRandomAgentLocationw70h70mppx1mapc128.pkl", "rb") as f:
|
with open("data/NormalizedIntersimpleExpertMu.001N10000_NRasterizedRouteSpeedRandomAgentLocationw70h70mppx1mapc128mu.001.pkl", "rb") as f:
|
||||||
trajectories = pickle.load(f)
|
trajectories = pickle.load(f)
|
||||||
transitions = rollout.flatten_trajectories(trajectories)
|
transitions = rollout.flatten_trajectories(trajectories)
|
||||||
train(transitions)
|
train(transitions)
|
||||||
|
|||||||
Reference in New Issue
Block a user