/
GraphTreeHeap
/
NeuroFighter
Обзор
Документация
Войти
/
GraphTreeHeap
/
NeuroFighter
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/neural_sb.py
93 строки
4 KB
FedMam
more updates!
12 окт 2025, 23:45
12 окт 2025, 23:45
e25bf29
Код
Авторство
О чём код?
import numpy as np import gymnasium as gym import random import onnxruntime as ort import os from gymnasium import spaces from mechanics import * from scipy.special import softmax from robotics import * from typing import * STATE_SPACE_BOX = [ (0., 1.), # player flag *[(0., 1.) for _ in range(N_BUTTONS)], # buttons *[(0., 1.) for _ in range(5)], # my state 1 *[(0., 1.) for _ in range(7)], # my characteristics (-1., 1.), # my facing right flag *[(0., 1.) for _ in range(4)], # my state 2 (0., 1.), # enemy absolute pos (-1., 1.), # enemy relative pos *[(0., 1.) for _ in range(4)], # enemy state 1 *[(0., 1.) for _ in range(7)], # enemy characteristics (-1., 1.), # enemy facing right flag *[(0., 1.) for _ in range(4)], # enemy state 2 *[(-1., 1.) for _ in range(3 * SHOW_K_PROJECTILES)], # my projectiles *[(-1., 1.) for _ in range(3 * SHOW_K_PROJECTILES)], # enemy projectiles (0., 1.) # time left ] SB_GAME_OPPONENTS = [ *[bot(seed=i*12345) for i, bot in enumerate(ALL_BOTS)], *[ort.InferenceSession(f'trained/onnx/{onnx_file}') for onnx_file in os.listdir('trained/onnx')] ] class SBGymEnv(gym.Env): def __init__(self): self.action_space = spaces.Discrete(ACTION_DIM_S) self.observation_space = spaces.Box(low=np.array([l for l, _ in STATE_SPACE_BOX]), high=np.array([h for _, h in STATE_SPACE_BOX]), dtype=np.float32) self.opponent_hidden = None self.np_rand = np.random.default_rng() self.reset() def reset(self, seed: int | None=None): rand = random.Random(seed) self.env = Environment(player_desc=rand.choice(CHARACTERS), opponent_desc=rand.choice(CHARACTERS), reward_system=REWARD_SYSTEM_DEFAULT) self.agent_is_player = rand.randint(0, 1) == 1 self.opponent = self.choose_opponent(rand) self.np_rand = np.random.default_rng(rand.randint(0, 0xffffffff)) if isinstance(self.opponent, ort.InferenceSession) and len(self.opponent.get_inputs()) > 1: self.opponent_hidden = (np.zeros((1, 1, 256), dtype=np.float32), np.zeros((1, 1, 256), dtype=np.float32)) else: self.opponent_hidden = None return self.get_observation(), {} def choose_opponent(self, rand: random.Random): return rand.choice(SB_GAME_OPPONENTS) def get_observation(self): return np.array(self.env.get_state(player=self.agent_is_player), dtype=np.float32) def step(self, action): self.env.take_action(self.agent_is_player, action) if isinstance(self.opponent, Bot): self.opponent.decide_outer(self.env, i_am_player=not self.agent_is_player) elif isinstance(self.opponent, ort.InferenceSession): if len(self.opponent.get_inputs()) == 1: # standard PPO state = np.expand_dims(np.array(self.env.get_state(player=not self.agent_is_player), dtype=np.float32), axis=0) logits, _ = self.opponent.run(None, {'state': state}) else: # PPO+LSTM state = np.expand_dims(np.expand_dims(np.array(self.env.get_state(player=not self.agent_is_player), dtype=np.float32), axis=0), axis=1) logits, _, hidden_h, hidden_c = self.opponent.run(None, {'state': state, 'hidden_in': self.opponent_hidden[0], 'onnx::LSTM_2': self.opponent_hidden[1]}) self.opponent_hidden = (hidden_h, hidden_c) action = self.np_rand.multinomial(1, softmax(logits.squeeze().astype(np.float64), axis=-1)).argmax() self.env.take_action(not self.agent_is_player, action) player_reward, opponent_reward = self.env.step() return self.get_observation(), (player_reward if self.agent_is_player else opponent_reward), self.env.done(), False, {}