RL-Project/Fetch-Reinforcement_learning_Project
0
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 