/
GraphTreeHeap
/
NeuroFighter
Обзор
Документация
Войти
/
GraphTreeHeap
/
NeuroFighter
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/neural_ppo_singlehead.py
460 строк
21 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: int, 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.ReLU(), nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.Tanh() ]) self.actor_head = nn.Linear(hidden_dim, action_dim) self.critic_head = nn.Linear(hidden_dim, 1) def forward(self, state): state = self.encoder(state) logits = self.actor_head(state).squeeze() value = self.critic_head(state).squeeze() return logits, value def act(self, state: torch.Tensor, temperature: float=None): with torch.no_grad(): logits, value = self.forward(state) dist = torch.distributions.Categorical(logits=(logits / temperature if temperature is not None else logits)) action = dist.sample() log_prob = dist.log_prob(action) entropy = dist.entropy().mean() return action.item(), log_prob, value.item(), entropy def collect_trajectory_ppo(env: Environment, agent: PPOAgent, opponent: PPOAgent | Bot, agent_is_player: bool, device: str='cpu', show_my_characteristics_for_agent: bool=True, show_opp_characteristics_for_agent: bool=True, show_my_characteristics_for_opponent: bool=True, show_opp_characteristics_for_opponent: bool=True, tv: Any=None): states, actions, rewards, dones, 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, my_characteristics=show_my_characteristics_for_agent, opp_characteristics=show_opp_characteristics_for_agent)).unsqueeze(0) opponent_state = torch.FloatTensor(env.get_state(player=not agent_is_player, my_characteristics=show_my_characteristics_for_opponent, opp_characteristics=show_opp_characteristics_for_opponent)).unsqueeze(0) # Current agent acts ag_action, log_prob, value, _ = agent.act(agent_state) env.take_action(agent_is_player, ag_action) # Opponent acts if isinstance(opponent, Bot): opponent.decide_outer(env, not agent_is_player) else: opp_action, _, _, _ = opponent.act(opponent_state) env.take_action(not agent_is_player, opp_action) # Step environment player_reward, opponent_reward = env.step() reward = (player_reward if agent_is_player else opponent_reward) if tv is not None: tv.step() states.append(agent_state) actions.append(ag_action) rewards.append(reward) dones.append(False) old_values.append(value) old_log_probs.append(log_prob) # Final state (done) final_state = torch.FloatTensor(env.get_state(player=agent_is_player, my_characteristics=show_my_characteristics_for_agent, opp_characteristics=show_opp_characteristics_for_agent)).unsqueeze(0) ag_action, log_prob, value, _ = agent.act(final_state) states.append(final_state) actions.append(ag_action) rewards.append(0) dones.append(True) old_values.append(value) old_log_probs.append(log_prob) 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), '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, reward_system: RewardSystem=REWARD_SYSTEM_DEFAULT, logging: bool=True, device: str='cpu', my_character_subset: list[CharacterDescription]=None, opp_character_subset: list[CharacterDescription]=None, experiment_id: str=None, seed: int | None=None): self.my_character_subset = my_character_subset or CHARACTERS self.opp_character_subset = opp_character_subset or CHARACTERS self.state_dim = STATE_DIM if len(self.my_character_subset) == 1: self.state_dim -= CHARACTERISTICS_DIM if len(self.opp_character_subset) == 1: self.state_dim -= CHARACTERISTICS_DIM self.hidden_dim = hidden_dim self.agent = PPOAgent(self.state_dim, ACTION_DIM_S, 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.learning_rate = learning_rate self.optimizer = torch.optim.Adam(self.agent.parameters(), lr=learning_rate, weight_decay=weight_decay) self.reward_system = reward_system self.rand = random.Random(seed) self.experiment_id = experiment_id if experiment_id is not None else str(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, device=self.device) 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, 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) # PPO epochs for epoch in range(self.update_epochs): 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) dist = torch.distributions.Categorical(logits=logits) new_log_probs = dist.log_prob(seq_actions) entropy = dist.entropy().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) ''' self.logger.add_histogram(f'Action/Chosen', actions, 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): # FIXME: this is a rather stupid code for character subsets other than CHARACTERS agent_is_player = self.rand.randint(0, 1) == 0 env = Environment(player_desc=self.rand.choice(self.my_character_subset if agent_is_player else self.opp_character_subset), opponent_desc=self.rand.choice(self.opp_character_subset if agent_is_player else self.my_character_subset), reward_system=self.reward_system) data = collect_trajectory_ppo(env=env, agent=self.agent, opponent=self.rand.choice(opponents), agent_is_player=agent_is_player, device=self.device, show_my_characteristics_for_agent=len(self.my_character_subset) > 1, show_opp_characteristics_for_agent=len(self.opp_character_subset) > 1, show_my_characteristics_for_opponent=len(self.my_character_subset) > 1, show_opp_characteristics_for_opponent=len(self.opp_character_subset) > 1) 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.entropy_coef_decay_episodes is not None: 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(self.state_dim, ACTION_DIM_S, 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 train_universal_bots(self, iterations: int=1000, trajs_per_iter: int=10, save_state_dict_freq=10, verbose: bool=False): if verbose: print('Experiment ID:', self.experiment_id) self.agent.train() for iter in tqdm.tqdm(range(iterations), 'Training PPO agent on UniversalBot™\'s') if verbose else range(iterations): bot = UniversalBot(strategy_seed=self.rand.randint(0, 0xffffffff), action_seed=self.rand.randint(0, 0xffffffff)) self.train([bot], iterations=1, trajs_per_iter=trajs_per_iter, save_state_dict_freq=None, verbose=False) if (iter + 1) % save_state_dict_freq == 0 or iter == iterations - 1: torch.save(self.agent.state_dict(), f'{self.results_dir}/agent.pth') def play_ppo(ppo_agent: PPOAgentWrapper, env: Environment, temperature: float | None=None, tv: Any | None=None): import graphics ppo_agent.agent.to(device='cpu') ppo_agent.agent.eval() graphics.init_graphics() if tv is None: tv = graphics.GameTV(env, player_controlled=True) while not env.done(): agent_state = torch.FloatTensor(env.get_state(player=False, my_characteristics=len(ppo_agent.my_character_subset) > 1, opp_characteristics=len(ppo_agent.opp_character_subset) > 1)).unsqueeze(0) action, _, _, _ = ppo_agent.agent.act(agent_state, temperature) env.take_action(False, action) # DEBUG # print(agent_hidden) env.step() tv.step() graphics.deinit_graphics() def play_ppo_2(agent: PPOAgentWrapper, opponent: PPOAgentWrapper | Bot, env: Environment, watch: bool=False, temperature: float | None=None): if watch: import graphics agent.agent.to(device='cpu') agent.agent.eval() if not isinstance(opponent, Bot): opponent.agent.to(device='cpu') opponent.agent.eval() if watch: graphics.init_graphics() tv = graphics.GameTV(env, player_controlled=False, opponent_controlled=False) while not env.done(): agent_state = torch.FloatTensor(env.get_state(player=True)).unsqueeze(0) opponent_state = torch.FloatTensor(env.get_state(player=False)).unsqueeze(0) ag_action, *_ = agent.agent.act(agent_state, temperature) env.take_action(True, ag_action) if isinstance(opponent, Bot): opponent.decide_outer(env, False) else: opp_action, *_ = opponent.agent.act(opponent_state, temperature) env.take_action(False, opp_action) env.step() if watch: tv.step() if watch: graphics.deinit_graphics()