forked from pedropintoo/SmartTLS
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest.py
More file actions
52 lines (40 loc) Β· 1.68 KB
/
Copy pathtest.py
File metadata and controls
52 lines (40 loc) Β· 1.68 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
import gymnasium as gym
import sys
from stable_baselines3 import PPO
from stable_baselines3.common.env_checker import check_env
from marl_tls.env import TLSEnv
import optparse
def get_options():
optParser = optparse.OptionParser()
optParser.add_option("--load_model", action="store", type="string", default="data/trained_model_ppo", help="file to load the model")
optParser.add_option("--simulation", action="store", type="string", default="cross/cross", help="path to the simulation")
optParser.add_option("--traffic_scale", action="store", type="string", default="1", help="Scale Traffic")
optParser.add_option("--render_mode", action="store", type="string", default="human", help="Render Mode")
options, args = optParser.parse_args()
return options
def run(vec_env, model, end):
obs = vec_env.reset()
step = 0
while True:
actions, _states = model.predict(obs)
obs, rewards, dones, infos = vec_env.step(actions)
step += 1
if step >= end - 1: # end-1 because vec_env.reset() is called inside step() and starts a new simulation
break
if __name__ == "__main__":
options = get_options()
load_model = options.load_model
simulation_path = options.simulation
traffic_scale = options.traffic_scale
render_mode = options.render_mode
model = PPO.load(load_model)
end = 2250
vec_env = TLSEnv.get_vec_env(
TLSEnv,
render_mode=render_mode if render_mode == "human" else None,
simulation_path=simulation_path,
traffic_scale=traffic_scale,
end=end,
) # new environment with human visualization
run(vec_env, model,end)
vec_env.close()