import gymnasium as gym import matplotlib.pyplot as plt import numpy as np from agent.ppo import PPOAgent def main(): env = gym.make('Pendulum-v1') state_dim = env.observation_space.shape[0] action_dim = env.action_space.shape[0] action_bound = float(env.action_space.high[0]) agent = PPOAgent(state_dim=state_dim, action_dim=action_dim, action_bound=action_bound) num_episodes = 200 batch_size = 2000 episode_rewards = [] state, _ = env.reset() memory = [] current_ep_reward = 0 episodes_completed = 0 print("开始训练 PPO 智能体...") step_count = 0 while episodes_completed < num_episodes: action = agent.get_action(state) # 交互 next_state, reward, terminated, truncated, _ = env.step(action) # 【极其关键的修复】:只有真正死亡 (terminated) 才清零未来价值 # 绝对不能把时间截断 (truncated) 算作 mask=0 mask = 0.0 if terminated else 1.0 done = terminated or truncated memory.append([state, action, reward, next_state, mask]) state = next_state current_ep_reward += reward step_count += 1 if done: episode_rewards.append(current_ep_reward) episodes_completed += 1 state, _ = env.reset() current_ep_reward = 0 if episodes_completed % 10 == 0: avg_reward = np.mean(episode_rewards[-10:]) print(f"Episode: {episodes_completed}, 平均奖励 (最近10轮): {avg_reward:.2f}") if step_count >= batch_size: agent.update(memory) memory.clear() step_count = 0 env.close() plt.figure(figsize=(10, 5)) plt.plot(episode_rewards) plt.title('PPO Learning Curve on Pendulum-v1') plt.xlabel('Episode') plt.ylabel('Total Reward') plt.grid(True) plt.savefig('ppo_learning_curve_final.png') plt.show() if __name__ == '__main__': main()