122 lines
5.0 KiB
Python
122 lines
5.0 KiB
Python
# ==============================================================================
|
|
# Author: Hongru Liu
|
|
# Affiliation: School of Power and Energy, Northwestern Polytechnical University
|
|
# Version: 1.0
|
|
# Contact: hongruliu@mail.nwpu.edu.cn
|
|
# ==============================================================================
|
|
|
|
import gymnasium as gym
|
|
import torch
|
|
import numpy as np
|
|
import matplotlib.pyplot as plt
|
|
import yaml
|
|
import os
|
|
from algorithms.sac import SAC
|
|
|
|
def evaluate_and_plot(model_path):
|
|
"""
|
|
加载训练好的模型,在环境中运行一个回合,并绘制类似阶跃响应的状态轨迹图
|
|
"""
|
|
# 1. 读取配置和初始化环境
|
|
config_path = os.path.join(os.path.dirname(__file__), 'configs', 'pendulum_config.yaml')
|
|
with open(config_path, "r", encoding="utf-8") as f:
|
|
config = yaml.safe_load(f)
|
|
|
|
env = gym.make('Pendulum-v1')
|
|
if not isinstance(env.observation_space, gym.spaces.Box) or env.observation_space.shape is None:
|
|
raise TypeError("Pendulum-v1 observation_space must be a Box with a valid shape.")
|
|
if not isinstance(env.action_space, gym.spaces.Box) or env.action_space.shape is None:
|
|
raise TypeError("Pendulum-v1 action_space must be a Box with a valid shape.")
|
|
|
|
state_dim = env.observation_space.shape[0]
|
|
action_dim = env.action_space.shape[0]
|
|
max_action = float(env.action_space.high[0])
|
|
|
|
# 2. 实例化算法并加载权重
|
|
agent = SAC(state_dim, action_dim, max_action, config)
|
|
|
|
if os.path.exists(model_path):
|
|
# 仅加载 Actor 的权重,因为测试阶段不需要 Critic 参与评估
|
|
agent.actor.load_state_dict(torch.load(model_path, map_location=agent.device))
|
|
print(f"[*] Successfully loaded model weights from: {model_path}")
|
|
else:
|
|
print(f"[!] Model file not found: {model_path}")
|
|
return
|
|
|
|
# 3. 运行单个 Episode,收集状态和动作数据
|
|
state, _ = env.reset()
|
|
|
|
# 用于记录作图的数据列表
|
|
angles = []
|
|
velocities = []
|
|
actions = []
|
|
time_steps = []
|
|
|
|
episode_reward: float = 0.0
|
|
for step in range(config['max_steps']):
|
|
# 设置 evaluate=True,让网络输出确定的均值动作,关闭随机探索
|
|
action = agent.select_action(state, evaluate=True)
|
|
|
|
# 记录当前步的数据
|
|
# Pendulum 的状态是 [cos(theta), sin(theta), theta_dot]
|
|
cos_theta, sin_theta, theta_dot = state
|
|
# 将 cos 和 sin 转换为实际角度 (弧度),范围 [-pi, pi]
|
|
angle = np.arctan2(sin_theta, cos_theta)
|
|
|
|
angles.append(angle)
|
|
velocities.append(theta_dot)
|
|
actions.append(action[0])
|
|
time_steps.append(step)
|
|
|
|
# 与环境交互
|
|
next_state, reward, terminated, truncated, _ = env.step(action)
|
|
state = next_state
|
|
episode_reward += float(reward)
|
|
|
|
if terminated or truncated:
|
|
break
|
|
|
|
print(f"[*] Evaluation completed. Total Reward: {episode_reward:.2f}")
|
|
|
|
# 4. 绘制状态轨迹图 (类阶跃响应)
|
|
fig, axs = plt.subplots(3, 1, figsize=(10, 10), facecolor='white', sharex=True)
|
|
|
|
# 子图 1: 角度响应 (Pendulum Angle)
|
|
axs[0].set_facecolor('white')
|
|
axs[0].plot(time_steps, angles, linewidth=2.5, color='#d62728', label='Angle (rad)')
|
|
axs[0].axhline(y=0, color='black', linestyle='--', linewidth=1.5, alpha=0.5) # 目标线
|
|
axs[0].set_ylabel('Angle [rad]', fontsize=14, fontweight='bold')
|
|
axs[0].set_title('Pendulum Step Response Simulation', fontsize=16, fontweight='bold')
|
|
axs[0].tick_params(axis='y', labelsize=12)
|
|
axs[0].grid(True, linestyle='--', alpha=0.7)
|
|
axs[0].legend(fontsize=12, loc='upper right')
|
|
|
|
# 子图 2: 角速度响应 (Angular Velocity)
|
|
axs[1].set_facecolor('white')
|
|
axs[1].plot(time_steps, velocities, linewidth=2.5, color='#2ca02c', label='Velocity (rad/s)')
|
|
axs[1].axhline(y=0, color='black', linestyle='--', linewidth=1.5, alpha=0.5) # 目标线
|
|
axs[1].set_ylabel('Velocity [rad/s]', fontsize=14, fontweight='bold')
|
|
axs[1].tick_params(axis='y', labelsize=12)
|
|
axs[1].grid(True, linestyle='--', alpha=0.7)
|
|
axs[1].legend(fontsize=12, loc='upper right')
|
|
|
|
# 子图 3: 控制输入 (Control Torque)
|
|
axs[2].set_facecolor('white')
|
|
axs[2].plot(time_steps, actions, linewidth=2.5, color='#9467bd', label='Torque (N.m)')
|
|
axs[2].set_xlabel('Time Step', fontsize=14, fontweight='bold')
|
|
axs[2].set_ylabel('Control Input', fontsize=14, fontweight='bold')
|
|
axs[2].tick_params(axis='both', which='major', labelsize=12)
|
|
axs[2].grid(True, linestyle='--', alpha=0.7)
|
|
axs[2].legend(fontsize=12, loc='upper right')
|
|
|
|
plt.tight_layout()
|
|
save_path = "step_response.png"
|
|
plt.savefig(save_path, dpi=300, facecolor=fig.get_facecolor(), edgecolor='none')
|
|
plt.close()
|
|
print(f"[*] Step response plot saved to: {save_path}")
|
|
|
|
if __name__ == "__main__":
|
|
# 替换为你实际训练出来的最后一次模型文件名称
|
|
# 例如:sac_actor_pendulum_ep200.pth
|
|
evaluate_and_plot("sac_actor_pendulum_ep200.pth")
|