CoolFace
Apppublic

RL-Project/Fetch-Reinforcement_learning_Project

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py105 linesDownload Raw Back to root
1import os2import gradio as gr3import numpy as np4import torch5import imageio6from stable_baselines3 import SAC7from custom_env import create_env8 9# Update your run function to accept a model_name10def run_model_episode(x_start, y_start, x_targ, y_targ, z_targ, model_name, random_coords):11    12    # map the radio‐choice to the actual checkpoint on disk13    model_paths = {14        "Pick & Place (HER)": "App/model/pick_and_place_her.zip",15        "Pick & Place (Dense)":        "App/model/pick_and_place_dense.zip",16        "Push":         "App/model/push.zip",17        "Reach":         "App/model/reach.zip",18    }19    checkpoint_path = model_paths[model_name]20 21    # map the radio‐choice to the actual environment name22    environments = {23        "Pick & Place (HER)": "FetchPickAndPlace-v3",24        "Pick & Place (Dense)":        "FetchPickAndPlaceDense-v3",25        "Push":         "FetchPush-v3",26        "Reach":         "FetchReach-v3",27    }28    environment = environments[model_name]29 30    # Handle environment coordinates31    if(environment == "FetchPush-v3"):32        z_targ = 0.033 34    block_xy=(x_start, y_start),35    goal_xyz=(x_targ, y_targ, z_targ)36 37    if random_coords:38        block_xy = None39        goal_xyz = None40        41    # create the env42    env = create_env(43        render_mode="rgb_array",44        block_xy=block_xy,45        goal_xyz=goal_xyz,46        environment=environment47    )48 49    # load the selected model50    model = SAC.load(checkpoint_path, env=env, verbose=0)51 52    frames = []53    obs, info = env.reset()54    for _ in range(200):55        action, _ = model.predict(obs, deterministic=True)56        obs, reward, done, trunc, info = env.step(action)57        frames.append(env.render())58        if done or trunc:59            obs, info = env.reset()60    env.close()61 62    video_path = "run_video.mp4"63    imageio.mimsave(video_path, frames, fps=30)64    return video_path65 66 67with gr.Blocks() as demo:68    gr.Markdown("## Fetch Robot: Model Demo App")69    gr.Markdown("Enter coordinates, pick a model, then click **Run Model**.")70    gr.Markdown("Coordinates are relative to the center of the table.")71 72    # 1) add a radio (or gr.Dropdown) for model selection73    model_selector = gr.Radio(74        choices=["Pick & Place (HER)", "Pick & Place (Dense)", "Push", "Reach"],75        value="Pick & Place (HER)",76        label="Select a model/environment"77    )78 79    # Randomize coordinates80    randomize = gr.Checkbox(81        label="Use randomized coordinates?",82        value=False83    )84 85    with gr.Row():86        x_start = gr.Number(label="Start X", value=0.0)87        y_start = gr.Number(label="Start Y", value=0.0)88 89    with gr.Row():90        x_targ = gr.Number(label="Target X", value=0.1)91        y_targ = gr.Number(label="Target Y", value=0.1)92        z_targ = gr.Number(label="Target Z", value=0.1)93 94    run_button   = gr.Button("Run Model")95    output_video = gr.Video()96 97    # 2) include the selector as an input to your click callback98    run_button.click(99        fn=run_model_episode,100        inputs=[x_start, y_start, x_targ, y_targ, z_targ, model_selector, randomize],101        outputs=output_video102    )103 104demo.launch(share=True)105