/
GraphTreeHeap
/
NeuroFighter
Обзор
Документация
Войти
/
GraphTreeHeap
/
NeuroFighter
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/neural_ppo.py
353 строки
16 KB
FedMam
moving to the new project version
11 окт 2025, 04:16
11 окт 2025, 04:16
00409b0
Код
Авторство
О чём код?
import torch import torch.nn as nn import numpy as np import copy import random import tqdm from collections import deque from mechanics import * from robotics import * from typing import * from torch.utils.tensorboard import SummaryWriter class PPOAgent(nn.Module): def __init__(self, state_dim: int, action_dim: tuple, hidden_dim=256): super().__init__() self.encoder = nn.Sequential(*[ nn.Linear(state_dim, hidden_dim), nn.ReLU(), nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.Tanh(), nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.Tanh() ]) self.actor_heads = nn.ModuleList([nn.Linear(hidden_dim, head_dim) for head_dim in action_dim]) self.critic_head = nn.Linear(hidden_dim, 1) def forward(self, state): state = self.encoder(state) logits = [actor_head(state).squeeze() for actor_head in self.actor_heads] value = self.critic_head(state).squeeze() return logits, value, def act(self, state: torch.Tensor, temperature: float | None=None): with torch.no_grad(): logits, value = self.forward(state) dists = [torch.distributions.Categorical(logits=(l / temperature) if temperature is not None else l) for l in logits] actions = [dist.sample() for dist in dists] log_probs = [dists[i].log_prob(actions[i]) for i in range(len(dists))] entropy = sum([dist.entropy().mean() for dist in dists]) return actions, log_probs, value.item(), entropy def collect_trajectory_ppo(env: Environment, agent: PPOAgent, opponent: PPOAgent | Bot, agent_is_player: bool, device: str='cpu'): states, actions, rewards, dones, next_states, old_values, old_log_probs = [], [], [], [], [], [], [] agent.to(device='cpu') if not isinstance(opponent, Bot): opponent.to(device='cpu') while not env.done(): agent_state = torch.FloatTensor(env.get_state(player=agent_is_player)).unsqueeze(0) opponent_state = torch.FloatTensor(env.get_state(player=not agent_is_player)).unsqueeze(0) # Current agent acts ag_actions, log_probs, value, _ = agent.act(agent_state) env.take_action(agent_is_player, [action.item() for action in ag_actions]) # Opponent acts if isinstance(opponent, Bot): opponent.decide_outer(env, not agent_is_player) else: opp_actions, _, _, _ = opponent.act(opponent_state) env.take_action(not agent_is_player, [action.item() for action in opp_actions]) # Step environment player_reward, opponent_reward = env.step() reward = (player_reward if agent_is_player else opponent_reward) next_state_agent = torch.FloatTensor(env.get_state(player=agent_is_player)).unsqueeze(0) done = env.done() states.append(agent_state) actions.append(ag_actions) rewards.append(reward) dones.append(done) next_states.append(next_state_agent) old_values.append(value) old_log_probs.append(log_probs) agent.to(device=device) return { 'states': torch.cat(states, dim=0).to(device=device), 'actions': torch.tensor(actions, dtype=torch.long, device=device), 'rewards': torch.tensor(rewards, dtype=torch.float32, device=device), 'dones': torch.tensor(dones, dtype=torch.float32, device=device), 'next_states': torch.cat(next_states, dim=0).to(device=device), 'old_values': torch.tensor(old_values, dtype=torch.float32, device=device), 'old_log_probs': torch.tensor(old_log_probs, dtype=torch.float32, device=device) } class PPOAgentWrapper: def __init__(self, learning_rate: float=3e-4, weight_decay: float=0.0, gamma: float=0.99, gae_lambda: float=0.95, clip_epsilon: float=0.2, entropy_coef_start: float=0.01, entropy_coef_min: float=0.001, entropy_coef_decay_episodes: int=1000, batch_size: int=64, update_epochs: int=4, hidden_dim: int=256, logging: bool=True, device: str='cpu', seed: int | None=None): self.hidden_dim = hidden_dim self.agent = PPOAgent(STATE_DIM, ACTION_DIM, hidden_dim) self.device = device self.agent.to(device=device) self.gamma = gamma self.gae_lambda = gae_lambda self.clip_epsilon = clip_epsilon self.entropy_coef_max = entropy_coef_start self.entropy_coef_min = entropy_coef_min self.entropy_coef_decay_episodes = entropy_coef_decay_episodes self.entropy_coef = entropy_coef_start self.episodes_count = 0 self.batch_size = batch_size self.update_epochs = update_epochs self.optimizer = torch.optim.Adam(self.agent.parameters(), lr=learning_rate, weight_decay=weight_decay) self.rand = random.Random(seed) self.experiment_id = random.randint(0, 99999) # use random instead of rand for different IDs for same experiments self.results_dir = f'trained-{self.experiment_id}' self.logging = logging if self.logging: self.logger = SummaryWriter(f'{self.results_dir}/tensorboard') self.logger_global_step = 0 def load_model(self, filename): self.agent.load_state_dict(torch.load(filename, map_location=torch.device(self.device))) def compute_advantages(self, rewards: torch.Tensor, values: torch.Tensor, dones: torch.Tensor): advantages = torch.zeros_like(rewards) last_advantage = 0 next_values = torch.cat([values[1:], torch.tensor([0.0], device=self.device)]) # Backward computation of GAE for t in reversed(range(len(rewards))): delta = rewards[t] + self.gamma * next_values[t] * (1 - dones[t]) - values[t] advantages[t] = last_advantage = delta + self.gamma * self.gae_lambda * (1 - dones[t]) * last_advantage returns = advantages + values return advantages, returns def update(self, data: dict[str, torch.Tensor]): states, actions, rewards, dones, next_states, old_values, old_log_probs = data.values() advantages, returns = self.compute_advantages(rewards, old_values, dones) # Normalize advantages advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) # Repeat for multi-headed policy advantages = advantages.unsqueeze(1).repeat(1, len(ACTION_DIM)) # PPO epochs for epoch in range(self.update_epochs): # Truncated BPTT: ensure full sequences indices = torch.randperm(len(states)) # For logging loss_data, actor_loss_data, critic_loss_data = [], [], [] entropy_data, ratio_data, kl_data, grad_norm_data, ev_data = [], [], [], [], [] for batch_i in range(0, len(states), self.batch_size): idxs = indices[batch_i:batch_i+self.batch_size] # [B, L, DIM] seq_states = states[idxs] seq_actions = actions[idxs] seq_old_log_probs = old_log_probs[idxs] seq_old_values = old_values[idxs] seq_returns = returns[idxs] seq_advantages = advantages[idxs] # Evaluate current policy logits, new_values = self.agent(seq_states) dists = [torch.distributions.Categorical(logits=l) for l in logits] new_log_probs = torch.stack([dists[i].log_prob(seq_actions[:, i]) for i in range(len(ACTION_DIM))], dim=1) entropy = torch.stack([dist.entropy().mean() for dist in dists], dim=-1).mean() # Actor loss ratio = torch.exp(new_log_probs - seq_old_log_probs) surr1 = ratio * seq_advantages surr2 = torch.clamp(ratio, 1 - self.clip_epsilon, 1 + self.clip_epsilon) * seq_advantages actor_loss = -torch.min(surr1, surr2).mean() # Critic loss critic_loss = (new_values - seq_returns).pow(2).mean() # Total loss loss = actor_loss + 0.5 * critic_loss - self.entropy_coef * entropy # Update self.optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(self.agent.parameters(), 0.5) self.optimizer.step() # For logging loss_data.append(loss.item()) actor_loss_data.append(actor_loss.item()) critic_loss_data.append(critic_loss.item()) entropy_data.append(entropy.item()) ratio_data.append(ratio.mean().item()) kl_data.append(torch.mean(seq_old_log_probs - new_log_probs).item()) grad_norm_data.append(sum([param.grad.data.norm(2).item() ** 2 for param in self.agent.parameters() if param.grad is not None]) ** (1/2)) ev_data.append(1 - (seq_returns - seq_old_values).var().item() / (returns.var().item() + 1e-8)) # Logging if self.logging: self.logger.add_scalar('Loss/Total', np.mean(loss_data), self.logger_global_step) self.logger.add_scalar('Loss/Actor', np.mean(actor_loss_data), self.logger_global_step) self.logger.add_scalar('Loss/Critic', np.mean(critic_loss_data), self.logger_global_step) self.logger.add_scalar('Policy/Entropy', np.mean(entropy_data), self.logger_global_step) self.logger.add_scalar('Policy/Ratio', np.mean(ratio_data), self.logger_global_step) self.logger.add_scalar('Policy/KLDivergence', np.mean(kl_data), self.logger_global_step) self.logger.add_scalar('Reward/Mean', torch.mean(rewards), self.logger_global_step) self.logger.add_scalar('Advantage/Mean', advantages.mean(), self.logger_global_step) self.logger.add_scalar('Grad/Norm', np.mean(grad_norm_data), self.logger_global_step) self.logger.add_scalar('Value/Predicted', old_values.mean(), self.logger_global_step) self.logger.add_scalar('Value/Actual', returns.mean(), self.logger_global_step) self.logger.add_scalar("Value/ExplainedVariance", np.mean(ev_data), self.logger_global_step) ''' self.logger.add_scalar('MyChar/HSpeed', states[0, STATE_MY_CHAR_SPEED].item(), self.logger_global_step) self.logger.add_scalar('MyChar/JumpHeight', states[0, STATE_MY_CHAR_JUMP_HEIGHT].item(), self.logger_global_step) self.logger.add_scalar('MyChar/ShootDamage', states[0, STATE_MY_CHAR_SHOOT_DAMAGE].item(), self.logger_global_step) self.logger.add_scalar('MyChar/ShootCooldown', states[0, STATE_MY_CHAR_SHOOT_COOLDOWN].item(), self.logger_global_step) self.logger.add_scalar('MyChar/ShootSpeed', states[0, STATE_MY_CHAR_SHOOT_SPEED].item(), self.logger_global_step) self.logger.add_scalar('OppChar/HSpeed', states[0, STATE_OPP_CHAR_SPEED].item(), self.logger_global_step) self.logger.add_scalar('OppChar/JumpHeight', states[0, STATE_OPP_CHAR_JUMP_HEIGHT].item(), self.logger_global_step) self.logger.add_scalar('OppChar/ShootDamage', states[0, STATE_OPP_CHAR_SHOOT_DAMAGE].item(), self.logger_global_step) self.logger.add_scalar('OppChar/ShootCooldown', states[0, STATE_OPP_CHAR_SHOOT_COOLDOWN].item(), self.logger_global_step) self.logger.add_scalar('OppChar/ShootSpeed', states[0, STATE_OPP_CHAR_SHOOT_SPEED].item(), self.logger_global_step) ''' for i, action_dim in enumerate(ACTION_DIM): self.logger.add_histogram(f'Action/Dim_{i}', actions[:, i], self.logger_global_step) self.logger_global_step += 1 def train(self, opponents: Iterable[Any], iterations: int=1000, trajs_per_iter: int=10, save_state_dict_freq: int | None=50, verbose: bool=False): if verbose: print('Experiment ID:', self.experiment_id) self.agent.train() for iter in tqdm.tqdm(range(iterations), 'Training PPO agent') if verbose else range(iterations): total_data = None for traj in range(trajs_per_iter): env = Environment(player_desc=CHARACTERS[self.rand.randint(0, len(CHARACTERS)-1)], opponent_desc=CHARACTERS[self.rand.randint(0, len(CHARACTERS)-1)]) data = collect_trajectory_ppo(env=env, agent=self.agent, opponent=self.rand.choice(opponents), agent_is_player=self.rand.randint(0, 1) == 0, device=self.device) if total_data is None: total_data = data else: total_data = {k: torch.cat((total_data[k], data[k]), dim=0) for k in total_data.keys()} # Update current agent self.update(total_data) # Decay entropy coef self.episodes_count += 1 if self.episodes_count >= self.entropy_coef_decay_episodes: self.entropy_coef = self.entropy_coef_min else: self.entropy_coef = self.entropy_coef_min + 0.5 * (self.entropy_coef_max - self.entropy_coef_min) * (1 + math.cos(self.episodes_count / self.entropy_coef_decay_episodes * math.pi)) if save_state_dict_freq is not None and ((iter + 1) % save_state_dict_freq == 0 or iter == iterations - 1): torch.save(self.agent.state_dict(), f'{self.results_dir}/agent.pth') def train_self_play(self, epochs: int=100, iterations_per_epoch: int=10, trajs_per_iter: int=10, opponent_pool_size: int=5, verbose: bool=False): if verbose: print('Experiment ID:', self.experiment_id) self.agent.train() opponent_pool = deque(maxlen=opponent_pool_size) for epoch in tqdm.tqdm(range(epochs), 'Training self-play PPO agent') if verbose else range(epochs): old_version = PPOAgent(STATE_DIM, ACTION_DIM, self.hidden_dim) old_version.load_state_dict(copy.deepcopy(self.agent.state_dict())) old_version.to(device='cpu') old_version.eval() opponent_pool.append(old_version) self.train(opponent_pool, iterations=iterations_per_epoch, trajs_per_iter=trajs_per_iter, save_state_dict_freq=None, verbose=False) torch.save(self.agent.state_dict(), f'{self.results_dir}/agent.pth') def play_ppo(ppo_agent: PPOAgentWrapper, env: Environment): import graphics ppo_agent.agent.to(device='cpu') ppo_agent.agent.eval() graphics.init_graphics() tv = graphics.GameTV(env, player_controlled=True) while not env.done(): agent_state = torch.FloatTensor(env.get_state(player=False)).unsqueeze(0).unsqueeze(1) actions, _, _, _ = ppo_agent.agent.act(agent_state) env.take_action(False, actions) # DEBUG # print(agent_hidden) env.step() tv.step() graphics.deinit_graphics()