CoolFace
Apppublic

AdityParbat/disaster-response-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
train.py61 linesDownload Raw Back to root
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