def compute_policy_gradient(policy_network, trajectories, gamma=0.99): policy_gradients = [] for trajectory in trajectories: states, actions, rewards = trajectory discounted_rewards = compute_discounted_rewards(rewards, gamma) for t, (state, action, reward) in enumerate(zip(states, actions, discounted_rewards)): state_tensor = torch.tensor(state, dtype=torch.float32) action_tensor = torch.tensor(action, dtype=torch.int64) reward_tensor = torch.tensor(reward, dtype=torch.float32) log_prob = torch.log(policy_network(state_tensor)[action_tensor]) policy_gradients.append(-log_prob * reward_tensor) return torch.stack(policy_gradients).sum()