Spaces:
Runtime error
Runtime error
File size: 4,580 Bytes
bbf97b5 | 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 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | import streamlit as st
import matplotlib.pyplot as plt
from environment import GridWorld
from agent import RLAgent
from dqn_agent import DQNAgent
from utils import (
plot_policy,
plot_heatmap,
plot_visits,
save_agent_walk_gif,
plot_dqn_qvalues
)
from generate_pdf_report import generate_pdf_report
st.set_page_config(layout="wide")
st.title("๐ฎ RL GridWorld: SARSA / Q-Learning / DQN")
# Sidebar
algo = st.sidebar.radio("Select Algorithm", ["SARSA", "Q-Learning", "DQN"])
alpha = st.sidebar.slider("Learning Rate (ฮฑ)", 0.001, 1.0, 0.1)
gamma = st.sidebar.slider("Discount Factor (ฮณ)", 0.0, 1.0, 0.99)
initial_epsilon = st.sidebar.slider("Initial Exploration Rate (ฮต)", 0.0, 1.0, 0.1)
episodes = st.sidebar.number_input("Episodes", 10, 5000, 300)
env = GridWorld()
# Agent selection
if algo == "DQN":
agent = DQNAgent(env.actions, alpha=alpha, gamma=gamma, epsilon=initial_epsilon)
st.sidebar.subheader("๐พ DQN Model Options")
if st.sidebar.button("Save Model"):
agent.save()
st.sidebar.success("Model saved as dqn_model.pth")
if st.sidebar.button("Load Model"):
try:
agent.load()
st.sidebar.success("Model loaded from dqn_model.pth")
except:
st.sidebar.error("Model file not found or invalid.")
else:
agent = RLAgent(env.actions, alpha, gamma, initial_epsilon, method=algo.lower())
# Training
total_rewards = []
all_trajectories = []
best_idx = 0
worst_idx = 0
for ep in range(episodes):
state = env.reset()
action = agent.act(state) if algo == "DQN" else agent.choose_action(state)
done = False
total_reward = 0
trajectory = [state]
while not done:
next_state, reward, done = env.step(action)
if algo == "DQN":
next_action = agent.act(next_state)
agent.remember(state, action, reward, next_state, done)
agent.train_step()
else:
next_action = agent.choose_action(next_state)
agent.update(state, action, reward, next_state, next_action)
state, action = next_state, next_action
trajectory.append(state)
total_reward += reward
total_rewards.append(total_reward)
all_trajectories.append(trajectory)
if env.state == env.goal and total_reward > total_rewards[best_idx]:
best_idx = ep
if total_reward < total_rewards[worst_idx]:
worst_idx = ep
st.success(f"โ
Training completed over {episodes} episodes using **{algo}**.")
# Visualizations
if algo == "DQN":
st.pyplot(plot_dqn_qvalues(agent.model, env.actions))
else:
st.pyplot(plot_policy(agent.Q))
st.pyplot(plot_heatmap(agent.Q))
st.pyplot(plot_visits(agent.visits))
# Reward Graph
st.subheader("๐ Total Reward per Episode")
fig = plt.figure(figsize=(8, 3))
plt.plot(total_rewards)
plt.xlabel("Episode")
plt.ylabel("Total Reward")
plt.grid(True)
st.pyplot(fig)
# Trajectory
st.subheader("๐ถ Agent Path in Last Episode")
trajectory_str = " โ ".join(dict.fromkeys([chr(65 + s[0]*4 + s[1]) for s in all_trajectories[-1]]))
st.code(trajectory_str, language="markdown")
# Q-table viewer
if algo != "DQN":
st.subheader("๐ง Inspect Q-Values")
selected_state = st.selectbox("Choose a state (row, col)", [(i, j) for i in range(4) for j in range(4)])
if selected_state in agent.Q:
st.json(agent.Q[selected_state])
else:
st.info("No Q-values learned yet for this state.")
# ๐๏ธ Walk animation + Loop + Download
st.subheader("๐๏ธ Agent Walk Animation")
episode_option = st.selectbox("Episode to Animate", ["Last Episode", "Best Episode", "Worst Episode"])
loop_forever = st.checkbox("๐ Loop animation infinitely", value=True)
ep_idx = {"Last Episode": episodes-1, "Best Episode": best_idx, "Worst Episode": worst_idx}[episode_option]
if st.button("Generate Walk Animation"):
save_agent_walk_gif(all_trajectories[ep_idx], episode=ep_idx + 1, loop=loop_forever)
st.image("agent_walk.gif", caption=f"Agent Movement - {episode_option}")
with open("agent_walk.gif", "rb") as f:
st.download_button("โฌ๏ธ Download GIF", f, file_name="agent_walk.gif", mime="image/gif")
# ๐ PDF export
st.subheader("๐ Export Report")
if st.button("Generate PDF Report"):
generate_pdf_report(agent, total_rewards, all_trajectories[-1], algo, episodes, gif_path="agent_walk.gif")
st.success("โ
RL_Report.pdf saved in your folder")
|