/
GraphTreeHeap
/
NeuroFighter
Обзор
Документация
Войти
/
GraphTreeHeap
/
NeuroFighter
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/util.py
171 строка
7 KB
FedMam
more updates!
12 окт 2025, 23:45
12 окт 2025, 23:45
e25bf29
Код
Авторство
О чём код?
from mechanics import * import robotics import neural_ppo_singlehead import neural_ppolstm_singlehead import random import numpy as np import torch import stable_baselines3 import sb3_contrib from typing import * READY_FRAMES = FPS * 1 WINS_FRAMES = FPS * 2 def hold_fight_universal(visual: bool, player_character: Character, opponent_character: Character, player_network: Any | None, # None for player opponent_network: Any | None, # None for player player_opponent_reversed: bool, round: int, logging: bool=False) -> int | tuple[int, dict]: if visual: import pygame import graphics if logging: logs = { 'action': ([], []), 'action_dist': ([], []), 'action_entropy': ([], []) } env = Environment(player_desc=player_character if not player_opponent_reversed else opponent_character, opponent_desc=opponent_character if not player_opponent_reversed else player_character) if visual: tv = graphics.GameTV(env, player_controlled=(player_network is None and not player_opponent_reversed) or (opponent_network is None and player_opponent_reversed), opponent_controlled=(opponent_network is None and not player_opponent_reversed) or (player_network is None and player_opponent_reversed), opponent_background=not player_opponent_reversed) # ready... for _ in range(READY_FRAMES): tv.step(ready_round=round) # the game players = (player_network if not player_opponent_reversed else opponent_network, opponent_network if not player_opponent_reversed else player_network) hiddens = [None, None] for player_i, player in enumerate(players): if isinstance(player, neural_ppolstm_singlehead.PPOAgentLSTM) or \ isinstance(player, neural_ppolstm_singlehead.PPOAgentLSTMResidual): hiddens[player_i] = (torch.zeros((player.lstm_layers, 1, player.lstm_hidden_dim)), torch.zeros((player.lstm_layers, 1, player.lstm_hidden_dim))) elif isinstance(player, sb3_contrib.RecurrentPPO): hiddens[player_i] = (np.zeros((1, 1, 128)), np.zeros((1, 1, 128))) while not env.done(): for player_i, player in enumerate(players): action_dist = None action_entropy = 0.0 if player is None: pass elif isinstance(player, robotics.Bot): player.decide_outer(env, i_am_player=player_i == 0) elif isinstance(player, neural_ppo_singlehead.PPOAgent): state = torch.FloatTensor(env.get_state(player=player_i == 0)).unsqueeze(0) # I've copied the act() method code here on purpose to have access # to all the data for logging purposes with torch.no_grad(): logits, value = player.forward(state) dist = torch.distributions.Categorical(logits=logits) action = dist.sample() log_prob = dist.log_prob(action) entropy = dist.entropy().mean() action_dist = torch.softmax(logits, dim=-1).squeeze().numpy(force=True) action_entropy = entropy env.take_action(player_i == 0, action.item()) elif isinstance(player, neural_ppolstm_singlehead.PPOAgentLSTM) or \ isinstance(player, neural_ppolstm_singlehead.PPOAgentLSTMResidual): state = torch.FloatTensor(env.get_state(player=player_i == 0)).unsqueeze(0).unsqueeze(1) with torch.no_grad(): logits, value, hiddens[player_i] = player.forward(state, hiddens[player_i]) dist = torch.distributions.Categorical(logits=logits) action = dist.sample() log_prob = dist.log_prob(action) entropy = dist.entropy().mean() action_dist = torch.softmax(logits, dim=-1).squeeze().numpy(force=True) action_entropy = entropy env.take_action(player_i == 0, action.item()) elif isinstance(player, stable_baselines3.PPO): obs = env.get_state(player=player_i == 0) action = player.predict(obs)[0].item() env.take_action(player=player_i == 0, action=action) elif isinstance(player, sb3_contrib.RecurrentPPO): obs = np.array(env.get_state(player=player_i == 0), dtype=np.float32) action, hiddens[player_i] = player.predict(obs, hiddens[player_i]) env.take_action(player=player_i == 0, action=action.item()) elif isinstance(player, stable_baselines3.DQN): obs = np.array(env.get_state(player=player_i == 0), dtype=np.float32) action = player.predict(obs)[0].item() env.take_action(player=player_i == 0, action=action) else: raise NotImplementedError(f'Error: game mechanics for class {player.__class__.__name__} are not implemented') # logging if logging: action = env.get_current_action(player=player_i == 0, multi_discrete=False) logs['action'][player_i].append(action) if action_dist is None: logs['action_dist'][player_i].append(np.array([(1 if i == action else 0) for i in range(ACTION_DIM_S)], dtype=np.float32)) else: logs['action_dist'][player_i].append(action_dist) logs['action_entropy'][player_i].append(action_entropy) env.step() if visual: tv.step() # winning conditions if env.player.hp <= 0 and env.opponent.hp <= 0: # double KO result = 0 elif env.opponent.hp <= 0: # win result = 1 elif env.player.hp <= 0: # loss result = -1 else: # time over result = 0 if player_opponent_reversed: result = -result if logging: for k in logs.keys(): logs[k] = (logs[k][1], logs[k][0]) # proclaim winner if visual: if graphics.ENABLE_MUSIC: pygame.mixer.music.stop() for _ in range(WINS_FRAMES): tv.step() if logging: return result, logs return result