AdityParbat/disaster-response-env
0
1import os2import argparse3from stable_baselines3 import PPO4from stable_baselines3.common.env_util import make_vec_env5from stable_baselines3.common.callbacks import CheckpointCallback, EvalCallback6from gym_wrapper import DisasterResponseGymEnv7 8def train():9 parser = argparse.ArgumentParser()10 parser.add_argument("--timesteps", type=int, default=20000)11 parser.add_argument("--task", type=str, default="citywide_crisis_management")12 args = parser.parse_args()13 14 print(f"Starting PPO training on {args.task} for {args.timesteps} steps...")15 16 # 1. Create and Wrap Environment17 # We use a lambda to ensure the env is created correctly in the vec_env18 env = DisasterResponseGymEnv(task_id=args.task)19 20 # 2. Initialize PPO Model21 # MultiInputPolicy is required for spaces.Dict observations22 model = PPO(23 policy="MultiInputPolicy",24 env=env,25 verbose=1,26 learning_rate=3e-4,27 n_steps=2048,28 batch_size=64,29 n_epochs=10,30 gamma=0.99,31 gae_lambda=0.95,32 clip_range=0.2,33 tensorboard_log="./logs/ppo_disaster_tensorboard/"34 )35 36 # 3. Setup Callpoints37 checkpoint_callback = CheckpointCallback(38 save_freq=5000,39 save_path="./models/",40 name_prefix="ppo_disaster_model"41 )42 43 # 4. Train44 model.learn(45 total_timesteps=args.timesteps,46 callback=checkpoint_callback,47 progress_bar=False48 )49 50 # 5. Save Final Model51 model_path = os.path.join("models", f"ppo_disaster_final_{args.task}")52 model.save(model_path)53 print(f"Training complete. Model saved to {model_path}")54 55if __name__ == "__main__":56 # Ensure models and logs directories exist57 os.makedirs("models", exist_ok=True)58 os.makedirs("logs", exist_ok=True)59 60 train()61 