diff --git a/CLAUDE.md b/CLAUDE.md index 3b7bfb10..4df40beb 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -9,7 +9,7 @@ En cas de doute ou d'alternative de design : **poser la question d'abord**, prop options, attendre la réponse. Les corrections de bugs internes (fichiers `src/`) qui ne changent ni contrat ni comportement voulu restent autorisées. -A neural network (SAC) trained to play a tank-arena game. +A neural network (PPO) trained to play a tank-arena game. Jolt physics, offscreen Vulkan rendering (the agent's vision), NN via LibTorch. ## Architecture (CMake modules) @@ -21,7 +21,9 @@ Dependency order: utils → model → view → core → agent / desktop - **arenai_view** — Vulkan rendering; `offscreen_renderer` = offscreen render for the agent's vision - **arenai_core** — `BaseTanksEnvironment`, enemy_handler, thread_pool (RL env loop) - **arenai_controller** — input handling -- **arenai_agent** — SAC: agents, networks, replay_buffer, reward_transforms + `main.cpp` (training executable) +- **arenai_agent** — agents (sac, ppo, ppo_liquid), networks, core (train env + spawn curriculum), + metrics + `main.cpp` (training executable). The reward is purely event-based + (hit/kill/received/death/win — no dense shaping). - **arenai_desktop** — the playable game executable. Its `src/` folders are hexagons of their own: `gui/` (RmlUi main menu — RmlUi types must never leak out of it, other code only includes `gui/menu.h`), `controller/`, `core/`. Menu assets live in `resources/menu/` + `resources/font/`. diff --git a/arenai_agent/CMakeLists.txt b/arenai_agent/CMakeLists.txt index 6ff658bf..8cb5dd3c 100644 --- a/arenai_agent/CMakeLists.txt +++ b/arenai_agent/CMakeLists.txt @@ -70,4 +70,6 @@ arenai_copy_runtime_dlls(arenai_agent_train TORCH) # Library add_library(arenai_agent SHARED ${ARENAI_AGENT_SRC}) target_link_libraries(arenai_agent PRIVATE ${ARENAI_AGENT_LINK_LIBS} "${stb_SOURCE_DIR}") +# nlohmann::json appears in the public factory.h contract +target_link_libraries(arenai_agent PUBLIC nlohmann_json::nlohmann_json) target_include_directories(arenai_agent PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}/include") diff --git a/arenai_agent/include/arenai_agent/factory.h b/arenai_agent/include/arenai_agent/factory.h index 75bdab64..9e599124 100644 --- a/arenai_agent/include/arenai_agent/factory.h +++ b/arenai_agent/include/arenai_agent/factory.h @@ -5,64 +5,52 @@ #ifndef ARENAI_AGENT_HOST_FACTORY_H #define ARENAI_AGENT_HOST_FACTORY_H -#include -#include #include #include +#include + #include "./agent.h" namespace arenai::agent { + enum AgentAlgorithm { PPO, PPO_LIQUID }; + class AgentFactory { public: - virtual ~AgentFactory() = default; - - explicit AgentFactory(const std::map &arguments); + // config: the content of a training run's config.json — the network + // hyper-parameters come from its "agent" section, the vision size and + // the control frequency from its "environment" section + explicit AgentFactory(const nlohmann::json &config); std::shared_ptr get_agent( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_actions); - - protected: - template - T get_value(const std::string &argument_name, T default_value) { - if (!arguments.contains(argument_name)) return default_value; - - const std::string value_as_string = arguments[argument_name]; - std::stringstream ss(value_as_string); - T value; - ss >> value; + AgentAlgorithm algorithm, const int &nb_sensors, const int &nb_continuous_actions, + const int &nb_discrete_actions, bool cuda); - if (ss.fail() || !ss.eof()) - throw std::runtime_error(std::format( - R"(Wrong value for "{}" : "{}", example : "{}")", argument_name, - value_as_string, default_value)); - - arguments.erase(arguments.find(argument_name)); - - return value; - } + int get_vision_height() const; + int get_vision_width() const; + float get_wanted_frequency() const; + private: + // every hyper-parameter is required: a config.json that misses one + // does not describe the run it claims to (at() throws on a missing key) template - T get_value( - const std::string &argument_name, const std::function &parse_fn, - T default_value) { - if (!arguments.contains(argument_name)) return default_value; - - const std::string value_as_string = arguments[argument_name]; - - arguments.erase(arguments.find(argument_name)); - - return parse_fn(value_as_string); + T get_value(const std::string &argument_name) { + return agent_arguments.at(argument_name).get(); } - virtual std::shared_ptr get_agent_impl( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_action) = 0; + std::shared_ptr create_ppo_agent( + const int &nb_sensors, const int &nb_continuous_actions, const int &nb_discrete_action, + bool cuda); - private: - std::map arguments; + std::shared_ptr create_liquid_ppo_agent( + const int &nb_sensors, const int &nb_continuous_actions, const int &nb_discrete_action, + bool cuda); + + nlohmann::json agent_arguments; + int vision_height; + int vision_width; + float wanted_frequency; }; }// namespace arenai::agent diff --git a/arenai_agent/include/arenai_agent/factory_set.h b/arenai_agent/include/arenai_agent/factory_set.h deleted file mode 100644 index f06d0cad..00000000 --- a/arenai_agent/include/arenai_agent/factory_set.h +++ /dev/null @@ -1,24 +0,0 @@ -// -// Created by samuel on 11/03/2026. -// - -#ifndef ARENAI_AGENT_HOST_FACTORY_SET_H -#define ARENAI_AGENT_HOST_FACTORY_SET_H - -#include "./factory.h" - -namespace arenai::agent { - - class ActorAgentFactory : public AgentFactory { - public: - explicit ActorAgentFactory(const std::map &arguments); - - protected: - std::shared_ptr get_agent_impl( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_action) override; - }; - -}// namespace arenai::agent - -#endif//ARENAI_AGENT_HOST_FACTORY_SET_H diff --git a/arenai_agent/python/.gitignore b/arenai_agent/python/.gitignore index c18dd8d8..6bbbd0dc 100644 --- a/arenai_agent/python/.gitignore +++ b/arenai_agent/python/.gitignore @@ -1 +1,2 @@ __pycache__/ +.ipynb_checkpoints/ diff --git a/arenai_agent/python/notebooks/convert_liquid_cell_fusion.ipynb b/arenai_agent/python/notebooks/convert_liquid_cell_fusion.ipynb new file mode 100644 index 00000000..a691cb2f --- /dev/null +++ b/arenai_agent/python/notebooks/convert_liquid_cell_fusion.ipynb @@ -0,0 +1,206 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "intro", + "metadata": {}, + "source": [ + "# Liquid checkpoint conversion — `LiquidRecurrent` / `LiquidCell` fusion\n", + "\n", + "The classes `LiquidRecurrent` and `LiquidCell` were merged into a single `LiquidCell`\n", + "(`arenai_agent/src/networks/recurrent/liquid_cell.h`). The intermediate `cell` submodule\n", + "disappeared, so every parameter saved under `liquid.cell.*` now lives under `liquid.*`:\n", + "\n", + "| old key | new key |\n", + "|---|---|\n", + "| `liquid.cell.a` | `liquid.a` |\n", + "| `liquid.cell.raw_tau` | `liquid.raw_tau` |\n", + "| `liquid.cell.f.*` | `liquid.f.*` |\n", + "| `liquid.to_output.*` | unchanged |\n", + "\n", + "This notebook reads an `actor.pt` saved **before** the fusion (via\n", + "`torch::serialize::OutputArchive`), remaps the keys and writes a checkpoint loadable by the\n", + "fused C++ code (`torch::serialize::InputArchive`). The same mapping applies to `critic.pt`.\n", + "\n", + "Notes:\n", + "- the parameter **order** is unchanged by the fusion, so `actor_optim.pt` / `critic_optim.pt`\n", + " need no conversion;\n", + "- requires `torch` (`pip install torch`, the CPU wheel is enough — it is not part of the\n", + " `arenai_agent/python` requirements)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "paths", + "metadata": {}, + "outputs": [], + "source": [ + "import re\n", + "from pathlib import Path\n", + "\n", + "import torch\n", + "from torch import nn\n", + "\n", + "OLD_CHECKPOINT = Path(\n", + " \"/home/samuel/CLionProjects/ArenAI/outputs/train_400_ppo_liquid_beta_silu_truncation-no-penalty/save_97/actor.pt\"\n", + ")\n", + "# written under a `converted/` sibling folder, same file name\n", + "NEW_CHECKPOINT = OLD_CHECKPOINT.parent / \"converted\" / OLD_CHECKPOINT.name" + ] + }, + { + "cell_type": "markdown", + "id": "load-md", + "metadata": {}, + "source": [ + "## Read the old state dict\n", + "\n", + "A module saved by LibTorch's `OutputArchive` is a TorchScript archive: `torch.jit.load`\n", + "gives back the full module hierarchy with its named parameters." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "load", + "metadata": {}, + "outputs": [], + "source": [ + "old_module = torch.jit.load(str(OLD_CHECKPOINT), map_location=\"cpu\")\n", + "\n", + "old_params = {name: p.detach().clone() for name, p in old_module.named_parameters()}\n", + "old_buffers = dict(old_module.named_buffers())\n", + "assert not old_buffers, f\"unexpected buffers, extend the notebook: {sorted(old_buffers)}\"\n", + "\n", + "for name, p in old_params.items():\n", + " print(f\"{name:45s} {tuple(p.shape)}\")" + ] + }, + { + "cell_type": "markdown", + "id": "convert-md", + "metadata": {}, + "source": [ + "## Remap and rebuild\n", + "\n", + "The C++ `Module::load` walks the archive by **module hierarchy**, and requires every\n", + "serialized submodule to exist — including parameterless ones (SiLU, Sigmoid, LayerNorm\n", + "positions in the `Sequential`s...). So the converted archive must replicate the full module\n", + "tree, not just the parameter paths: the old `liquid.cell` node is dropped and its children\n", + "are re-attached directly to `liquid`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "convert", + "metadata": {}, + "outputs": [], + "source": [ + "def convert(name: str) -> str:\n", + " return re.sub(r\"^liquid\\.cell\\.\", \"liquid.\", name)\n", + "\n", + "\n", + "class Node(nn.Module):\n", + " def forward(self) -> None:\n", + " return None\n", + "\n", + "\n", + "def ensure_path(root: Node, dotted: str) -> Node:\n", + " module = root\n", + " for part in dotted.split(\".\"):\n", + " if part not in module._modules:\n", + " module.add_module(part, Node())\n", + " module = module._modules[part]\n", + " return module\n", + "\n", + "\n", + "root = Node()\n", + "for name, _ in old_module.named_modules():\n", + " if name and name != \"liquid.cell\":\n", + " ensure_path(root, convert(name))\n", + "for name, tensor in old_params.items():\n", + " *path, leaf = convert(name).split(\".\")\n", + " ensure_path(root, \".\".join(path)).register_parameter(leaf, nn.Parameter(tensor))\n", + "\n", + "NEW_CHECKPOINT.parent.mkdir(parents=True, exist_ok=True)\n", + "torch.jit.save(torch.jit.script(root), str(NEW_CHECKPOINT))\n", + "print(f\"written: {NEW_CHECKPOINT}\")" + ] + }, + { + "cell_type": "markdown", + "id": "verify-md", + "metadata": {}, + "source": [ + "## Verify the round trip" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "verify", + "metadata": {}, + "outputs": [], + "source": [ + "reloaded = torch.jit.load(str(NEW_CHECKPOINT), map_location=\"cpu\")\n", + "\n", + "new_params = dict(reloaded.named_parameters())\n", + "expected = {convert(name): t for name, t in old_params.items()}\n", + "\n", + "assert set(new_params) == set(expected), set(new_params) ^ set(expected)\n", + "for name, t in expected.items():\n", + " assert torch.equal(new_params[name], t), name\n", + "\n", + "# structure the C++ InputArchive will walk\n", + "liquid = reloaded.liquid\n", + "for attr in (\"a\", \"raw_tau\", \"f\", \"to_output\"):\n", + " assert hasattr(liquid, attr), attr\n", + "assert not hasattr(liquid, \"cell\")\n", + "\n", + "print(f\"OK — {len(new_params)} parameters remapped\")\n", + "for name in sorted(new_params):\n", + " if name.startswith(\"liquid\"):\n", + " print(f\" {name:35s} {tuple(new_params[name].shape)}\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "98c45cbd-59e2-4ae7-8918-54e749510121", + "metadata": {}, + "outputs": [], + "source": [] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "0cf3755b-aebe-45f6-b8ee-c8caad225b5c", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.14.7" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/arenai_agent/src/agents/agent_cli.cpp b/arenai_agent/src/agents/agent_cli.cpp index 7e0d671b..b81c6c11 100644 --- a/arenai_agent/src/agents/agent_cli.cpp +++ b/arenai_agent/src/agents/agent_cli.cpp @@ -9,8 +9,8 @@ #include "../utils/cli_fields.h" #include "./ppo/ppo_factory.h" #include "./ppo/ppo_hyperparams.h" -#include "./sac/sac_factory.h" -#include "./sac/sac_hyperparams.h" +#include "./ppo_liquid/liquid_ppo_factory.h" +#include "./ppo_liquid/liquid_ppo_hyperparams.h" namespace arenai::agent { @@ -39,10 +39,10 @@ namespace arenai::agent { std::vector make_agent_clis() { std::vector algorithms; - algorithms.push_back( - make_agent_cli("sac", sac_cli_fields())); algorithms.push_back( make_agent_cli("ppo", ppo_cli_fields())); + algorithms.push_back(make_agent_cli( + "ppo_liquid", liquid_ppo_cli_fields())); return algorithms; } diff --git a/arenai_agent/src/agents/factory.cpp b/arenai_agent/src/agents/factory.cpp index 78d7d88c..878de7b9 100644 --- a/arenai_agent/src/agents/factory.cpp +++ b/arenai_agent/src/agents/factory.cpp @@ -2,33 +2,69 @@ // Created by samuel on 22/01/2026. // -#include -#include -#include -#include +#include +#include #include +#include "./ppo/ppo_agent.h" +#include "./ppo_liquid/liquid_ppo_agent.h" + using namespace arenai; using namespace arenai::agent; namespace arenai::agent { - AgentFactory::AgentFactory(const std::map &arguments) - : arguments(arguments) {} + AgentFactory::AgentFactory(const nlohmann::json &config) + : agent_arguments(config.at("agent")), + vision_height(config.at("environment").at("vision_height").get()), + vision_width(config.at("environment").at("vision_width").get()), + wanted_frequency(config.at("environment").at("wanted_frequency").get()) {} + + int AgentFactory::get_vision_height() const { return vision_height; } + int AgentFactory::get_vision_width() const { return vision_width; } + float AgentFactory::get_wanted_frequency() const { return wanted_frequency; } std::shared_ptr AgentFactory::get_agent( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_actions) { - const auto agent = get_agent_impl( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions); - - if (!arguments.empty()) { - std::cerr << "Invalid argument(s) : " << std::get<0>(*arguments.begin()) << std::endl; - throw std::runtime_error("Invalid argument(s)"); + const AgentAlgorithm algorithm, const int &nb_sensors, const int &nb_continuous_actions, + const int &nb_discrete_actions, const bool cuda) { + switch (algorithm) { + case PPO: + return create_ppo_agent( + nb_sensors, nb_continuous_actions, nb_discrete_actions, cuda); + case PPO_LIQUID: + return create_liquid_ppo_agent( + nb_sensors, nb_continuous_actions, nb_discrete_actions, cuda); + default: throw std::runtime_error("Unknown agent algorithm"); } + } + + std::shared_ptr AgentFactory::create_ppo_agent( + const int &nb_sensors, const int &nb_continuous_actions, const int &nb_discrete_action, + const bool cuda) { + return std::make_shared( + std::make_shared( + vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_action, + get_value("hidden_size_sensors"), + get_value>("actor_hidden_sizes"), + get_value>>("vision_channels"), + get_value>("group_norm_nums"), 0.f, std::vector{0.5f, 0.5f}), + cuda ? torch::kCUDA : torch::kCPU); + } + + std::shared_ptr AgentFactory::create_liquid_ppo_agent( + const int &nb_sensors, const int &nb_continuous_actions, const int &nb_discrete_action, + const bool cuda) { + const auto actor = std::make_shared( + vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_action, + get_value("hidden_size_sensors"), + get_value>>("vision_channels"), + get_value>("group_norm_nums"), get_value("neuron_number"), + get_value("unfolding_steps"), get_value("delta_t"), 0.f, + std::vector{0.5f, 0.5f}); - return agent; + return std::make_shared( + actor, std::make_shared(actor), cuda ? torch::kCUDA : torch::kCPU); } }// namespace arenai::agent diff --git a/arenai_agent/src/agents/factory_set.cpp b/arenai_agent/src/agents/factory_set.cpp deleted file mode 100644 index 17e5f142..00000000 --- a/arenai_agent/src/agents/factory_set.cpp +++ /dev/null @@ -1,39 +0,0 @@ -// -// Created by samuel on 11/03/2026. -// - -#include - -#include "../utils/cli_parser.h" -#include "./ppo/ppo_agent.h" -#include "./sac/sac_agent.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - ActorAgentFactory::ActorAgentFactory(const std::map &arguments) - : AgentFactory(arguments) {} - - std::shared_ptr ActorAgentFactory::get_agent_impl( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_action) { - return std::make_shared( - std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_action, - get_value("hidden_size_sensors", 128), - get_value("hidden_sizes", parse_cli_hidden_layer, {{1024, 512}}) - .layers, - get_value( - "vision_channels", parse_cli_vision_channels, - {{{3, 8}, {8, 16}, {16, 24}, {24, 32}, {32, 48}, {48, 64}}}) - .channels, - get_value( - "group_norm_nums", parse_cli_group_norms, {{{1, 2, 3, 4, 6, 8}}}) - .groups, - 0.f, 0.f), - get_value("cuda", false) ? torch::kCUDA : torch::kCPU); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo/ppo_agent.cpp b/arenai_agent/src/agents/ppo/ppo_agent.cpp index 439bbb61..23b3e8fe 100644 --- a/arenai_agent/src/agents/ppo/ppo_agent.cpp +++ b/arenai_agent/src/agents/ppo/ppo_agent.cpp @@ -4,8 +4,8 @@ #include "./ppo_agent.h" -#include "../../distributions/multinomial.h" -#include "../../distributions/truncated_normal.h" +#include "../../distributions/bernoulli.h" +#include "../../distributions/beta_law.h" #include "../../networks/constants.h" #include "../../networks_utils/torch_converter.h" #include "../../networks_utils/torch_loader.h" @@ -42,22 +42,22 @@ namespace arenai::agent { torch::NoGradGuard guard; const auto &[vision, sensors] = state; - const auto &[mu, sigma, discrete_proba] = actor->act(vision, sensors); + const auto &[mode, concentration, discrete_proba] = actor->act(vision, sensors); if (sample) { - action.continuous_action = truncated_normal_sample(mu, sigma); - action.discrete_action = multinomial_sample(discrete_proba); + action.continuous_action = beta_law_sample(mode, concentration); + action.discrete_action = bernoulli_sample(discrete_proba); } else { - action.continuous_action = truncated_normal_mean(mu, sigma); - action.discrete_action = multinomial_max_action(discrete_proba); + action.continuous_action = beta_law_mode_action(mode); + action.discrete_action = bernoulli_max_action(discrete_proba); } // old log-probabilities, kept for the PPO importance ratio continuous_log_prob = - truncated_normal_log_pdf(action.continuous_action, mu, sigma).sum(-1, true); + beta_law_log_proba(action.continuous_action, mode, concentration).sum(-1, true); - const auto clamped_proba = torch::clamp(discrete_proba, EPSILON, 1.0 - EPSILON); - discrete_log_prob = (action.discrete_action * torch::log(clamped_proba)).sum(-1, true); + discrete_log_prob = + bernoulli_log_proba(action.discrete_action, discrete_proba).sum(-1, true); } if (collector.has_value()) diff --git a/arenai_agent/src/agents/ppo/ppo_collector.cpp b/arenai_agent/src/agents/ppo/ppo_collector.cpp index b0e30543..8f87e0cc 100644 --- a/arenai_agent/src/agents/ppo/ppo_collector.cpp +++ b/arenai_agent/src/agents/ppo/ppo_collector.cpp @@ -21,14 +21,16 @@ namespace arenai::agent { last_discrete_log_prob = discrete_log_prob; } - void PpoStepCollector::on_transition(const torch::Tensor &rewards, const torch::Tensor &done) { + void PpoStepCollector::on_transition( + const torch::Tensor &rewards, const torch::Tensor &done, const torch::Tensor &truncated) { rollout_buffer->add( {.state = last_state, .action = last_action, .continuous_log_prob = last_continuous_log_prob, .discrete_log_prob = last_discrete_log_prob, .reward = rewards, - .done = done}); + .done = done, + .truncated = truncated}); } void PpoStepCollector::on_episode_end(const TorchState &final_state) { diff --git a/arenai_agent/src/agents/ppo/ppo_collector.h b/arenai_agent/src/agents/ppo/ppo_collector.h index f98cb8a8..89c7a1c6 100644 --- a/arenai_agent/src/agents/ppo/ppo_collector.h +++ b/arenai_agent/src/agents/ppo/ppo_collector.h @@ -22,7 +22,9 @@ namespace arenai::agent { const TorchState &state, const TorchAction &action, const torch::Tensor &continuous_log_prob, const torch::Tensor &discrete_log_prob); - void on_transition(const torch::Tensor &rewards, const torch::Tensor &done) override; + void on_transition( + const torch::Tensor &rewards, const torch::Tensor &done, + const torch::Tensor &truncated) override; void on_episode_end(const TorchState &final_state) override; diff --git a/arenai_agent/src/agents/ppo/ppo_factory.cpp b/arenai_agent/src/agents/ppo/ppo_factory.cpp index 9edec026..5d12ed27 100644 --- a/arenai_agent/src/agents/ppo/ppo_factory.cpp +++ b/arenai_agent/src/agents/ppo/ppo_factory.cpp @@ -13,11 +13,12 @@ namespace arenai::agent { const int vision_height, const int vision_width, const int nb_sensors, const int nb_continuous_actions, const int nb_discrete_actions, const torch::Device device, const PpoHyperParams ¶ms) - : config(cli_fields_to_map(ppo_cli_fields(), params)), + : config(cli_fields_to_json(ppo_cli_fields(), params)), actor(std::make_shared( vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, params.hidden_size_sensors, params.actor_hidden_sizes, params.vision_channels, - params.group_norm_nums, params.initial_sigma, params.initial_fire_proba)), + params.group_norm_nums, params.initial_sigma, + std::vector{params.initial_fire_proba, params.initial_zoom_proba})), rollout_buffer(std::make_shared()), collector(std::make_shared(rollout_buffer)), agent(std::make_shared(actor, device, collector)), @@ -27,7 +28,7 @@ namespace arenai::agent { params.hidden_size_sensors, params.critic_hidden_sizes, params.vision_channels, params.group_norm_nums, device, params.metric_window_size, params.gamma, params.gae_lambda, params.clip_epsilon, params.target_kl, params.grad_norm_max, - params.continuous_target_entropy, params.discrete_target_entropy_factor, + params.continuous_target_entropy, params.discrete_target_entropy_factors, params.epochs, params.rollout_size, params.minibatch_size)) {} std::shared_ptr PpoTorchAgentFactory::get_agent() { return agent; } @@ -38,6 +39,6 @@ namespace arenai::agent { std::shared_ptr PpoTorchAgentFactory::get_trainer() { return trainer; } - std::map PpoTorchAgentFactory::get_config() const { return config; } + nlohmann::json PpoTorchAgentFactory::get_config() const { return config; } }// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo/ppo_factory.h b/arenai_agent/src/agents/ppo/ppo_factory.h index 5bcdbb9e..c70361b1 100644 --- a/arenai_agent/src/agents/ppo/ppo_factory.h +++ b/arenai_agent/src/agents/ppo/ppo_factory.h @@ -24,10 +24,10 @@ namespace arenai::agent { std::shared_ptr get_collector() override; std::shared_ptr get_trainer() override; - std::map get_config() const override; + nlohmann::json get_config() const override; private: - std::map config; + nlohmann::json config; // triad built once, sharing actor + rollout_buffer std::shared_ptr actor; diff --git a/arenai_agent/src/agents/ppo/ppo_hyperparams.cpp b/arenai_agent/src/agents/ppo/ppo_hyperparams.cpp index 73349a59..9aad8072 100644 --- a/arenai_agent/src/agents/ppo/ppo_hyperparams.cpp +++ b/arenai_agent/src/agents/ppo/ppo_hyperparams.cpp @@ -17,6 +17,7 @@ namespace arenai::agent { {.name = "--group_norm_nums", .member = &PpoHyperParams::group_norm_nums}, {.name = "--initial_sigma", .member = &PpoHyperParams::initial_sigma}, {.name = "--initial_fire_proba", .member = &PpoHyperParams::initial_fire_proba}, + {.name = "--initial_zoom_proba", .member = &PpoHyperParams::initial_zoom_proba}, {.name = "--metric_window_size", .member = &PpoHyperParams::metric_window_size}, {.name = "--gamma", .member = &PpoHyperParams::gamma}, {.name = "--gae_lambda", .member = &PpoHyperParams::gae_lambda}, @@ -25,8 +26,8 @@ namespace arenai::agent { {.name = "--grad_norm_max", .member = &PpoHyperParams::grad_norm_max}, {.name = "--continuous_target_entropy", .member = &PpoHyperParams::continuous_target_entropy}, - {.name = "--discrete_target_entropy_factor", - .member = &PpoHyperParams::discrete_target_entropy_factor}, + {.name = "--discrete_target_entropy_factors", + .member = &PpoHyperParams::discrete_target_entropy_factors}, {.name = "--epochs", .member = &PpoHyperParams::epochs}, {.name = "--rollout_size", .member = &PpoHyperParams::rollout_size}, {.name = "--minibatch_size", .member = &PpoHyperParams::minibatch_size}, diff --git a/arenai_agent/src/agents/ppo/ppo_hyperparams.h b/arenai_agent/src/agents/ppo/ppo_hyperparams.h index bad6b3bc..3e24ccff 100644 --- a/arenai_agent/src/agents/ppo/ppo_hyperparams.h +++ b/arenai_agent/src/agents/ppo/ppo_hyperparams.h @@ -17,21 +17,24 @@ namespace arenai::agent { float actor_learning_rate = 1e-4f; float critic_learning_rate = 3e-4f; int hidden_size_sensors = 128; - std::vector actor_hidden_sizes = {1024, 512}; - std::vector critic_hidden_sizes = {1024, 512}; - std::vector> vision_channels = {{3, 8}, {8, 16}, {16, 24}, - {24, 32}, {32, 48}, {48, 64}}; - std::vector group_norm_nums = {1, 2, 3, 4, 6, 8}; + std::vector actor_hidden_sizes = {256, 128}; + std::vector critic_hidden_sizes = {256, 128}; + std::vector> vision_channels = {{3, 16}, {16, 24}, {24, 32}, {32, 48}, + {48, 64}, {64, 96}, {96, 128}}; + std::vector group_norm_nums = {2, 3, 4, 6, 8, 12, 16}; float initial_sigma = 0.5f; float initial_fire_proba = 0.4f; + float initial_zoom_proba = 0.25f; int metric_window_size = 256; float gamma = 0.997f; float gae_lambda = 0.99f; float clip_epsilon = 0.2f; float target_kl = 0.05f; float grad_norm_max = 0.5f; - float continuous_target_entropy = 0.4f; - float discrete_target_entropy_factor = 0.2f; + // per continuous action: direction x, direction y, canon x, canon y + std::vector continuous_target_entropy = {-0.2f, -0.2f, -0.88f, -0.88f}; + // factors of the Bernoulli maximum entropy, per discrete action: fire, zoom + std::vector discrete_target_entropy_factors = {0.4f, 0.4f}; int epochs = 2; int rollout_size = 30 * 30; int minibatch_size = 1024; diff --git a/arenai_agent/src/agents/ppo/ppo_rollout_buffer.cpp b/arenai_agent/src/agents/ppo/ppo_rollout_buffer.cpp index cf382337..1f7bbc05 100644 --- a/arenai_agent/src/agents/ppo/ppo_rollout_buffer.cpp +++ b/arenai_agent/src/agents/ppo/ppo_rollout_buffer.cpp @@ -38,7 +38,8 @@ namespace arenai::agent { .continuous_log_prob = step.continuous_log_prob.detach().cpu(), .discrete_log_prob = step.discrete_log_prob.detach().cpu(), .reward = step.reward.detach().cpu(), - .done = step.done.detach().cpu()}, + .done = step.done.detach().cpu(), + .truncated = step.truncated.detach().cpu()}, .valid = valid}); // the freshly added step is pending: its closing observation is not known yet @@ -88,6 +89,7 @@ namespace arenai::agent { stack([](const StoredStep &s) { return s.step.discrete_log_prob; }), .rewards = stack([](const StoredStep &s) { return s.step.reward; }), .dones = stack([](const StoredStep &s) { return s.step.done; }), + .truncateds = stack([](const StoredStep &s) { return s.step.truncated; }), .bootstrap_state = bootstrap_state, .valids = stack([](const StoredStep &s) { return s.valid; }).unsqueeze(-1)}; diff --git a/arenai_agent/src/agents/ppo/ppo_rollout_buffer.h b/arenai_agent/src/agents/ppo/ppo_rollout_buffer.h index 924a5941..973f8772 100644 --- a/arenai_agent/src/agents/ppo/ppo_rollout_buffer.h +++ b/arenai_agent/src/agents/ppo/ppo_rollout_buffer.h @@ -21,6 +21,8 @@ namespace arenai::agent { torch::Tensor discrete_log_prob; torch::Tensor reward; torch::Tensor done; + // done whose value target must bootstrap instead of cutting the return (famine) + torch::Tensor truncated; }; // On-policy rollout stacked on the time dimension: every tensor is [T, nb_tanks, ...] @@ -31,6 +33,8 @@ namespace arenai::agent { torch::Tensor discrete_log_probs; torch::Tensor rewards; torch::Tensor dones; + // dones whose value target must bootstrap instead of cutting the return + torch::Tensor truncateds; // [nb_tanks, ...] observation closing the last step, for the value bootstrap TorchState bootstrap_state; // [T, nb_tanks, 1] whether the (step, tank) pair is a live transition diff --git a/arenai_agent/src/agents/ppo/ppo_trainer.cpp b/arenai_agent/src/agents/ppo/ppo_trainer.cpp index 94bc7b8d..11dd68b1 100644 --- a/arenai_agent/src/agents/ppo/ppo_trainer.cpp +++ b/arenai_agent/src/agents/ppo/ppo_trainer.cpp @@ -7,8 +7,8 @@ #include #include -#include "../../distributions/multinomial.h" -#include "../../distributions/truncated_normal.h" +#include "../../distributions/bernoulli.h" +#include "../../distributions/beta_law.h" #include "../../metrics/mean_metric.h" #include "../../networks/constants.h" #include "../../networks_utils/print_module.h" @@ -54,17 +54,23 @@ namespace arenai::agent { const std::vector &group_norm_nums, const torch::Device device, const int metric_window_size, const float gamma, const float gae_lambda, const float clip_epsilon, const float target_kl, const float grad_norm_max, - const float continuous_target_entropy, const float discrete_target_entropy_factor, - const int epochs, const int rollout_size, const int minibatch_size) + const std::vector &continuous_target_entropy, + const std::vector &discrete_target_entropy_factors, const int epochs, + const int rollout_size, const int minibatch_size) : actor(actor), rollout_buffer(rollout_buffer), continuous_alpha(std::make_unique( CONTINUOUS_ALPHA_K_P, CONTINUOUS_ALPHA_K_I, CONTINUOUS_ALPHA_K_D, ALPHA_INITIAL, nb_continuous_actions)), discrete_alpha(std::make_unique( - DISCRETE_ALPHA_K_P, DISCRETE_ALPHA_K_I, DISCRETE_ALPHA_K_D, ALPHA_INITIAL, 1)), + DISCRETE_ALPHA_K_P, DISCRETE_ALPHA_K_I, DISCRETE_ALPHA_K_D, ALPHA_INITIAL, + nb_discrete_action)), continuous_target_entropy(continuous_target_entropy), - discrete_target_entropy( - discrete_target_entropy_factor * multinomial_maximum_entropy(nb_discrete_action)), + // per-action target: each discrete action is an independent Bernoulli + discrete_target_entropy([&discrete_target_entropy_factors] { + std::vector targets = discrete_target_entropy_factors; + for (auto &target: targets) target *= bernoulli_maximum_entropy(); + return targets; + }()), critic(std::make_shared( vision_height, vision_width, nb_sensors, hidden_size_sensors, critic_hidden_sizes, vision_channels, group_norm_nums)), @@ -85,6 +91,13 @@ namespace arenai::agent { gamma(gamma), gae_lambda(gae_lambda), clip_epsilon(clip_epsilon), target_kl(target_kl), grad_norm_max(grad_norm_max), epochs(epochs), rollout_size(rollout_size), minibatch_size(minibatch_size) { + TORCH_CHECK( + static_cast(continuous_target_entropy.size()) == nb_continuous_actions, + "continuous_target_entropy needs one target per continuous action"); + TORCH_CHECK( + static_cast(discrete_target_entropy_factors.size()) == nb_discrete_action, + "discrete_target_entropy_factors needs one factor per discrete action"); + to(device); set_train(false); @@ -152,14 +165,13 @@ namespace arenai::agent { const torch::Tensor &old_log_probs, const torch::Tensor &advantages) const { const auto device = actor->parameters().back().device(); - const auto [mu, sigma, discrete_proba] = actor->act(vision, proprioception); + const auto [mode, concentration, discrete_proba] = actor->act(vision, proprioception); const auto curr_continuous_log_probs = - truncated_normal_log_pdf(continuous_actions, mu, sigma).sum(-1, true); + beta_law_log_proba(continuous_actions, mode, concentration).sum(-1, true); - const auto clamped_proba = torch::clamp(discrete_proba, EPSILON, 1.0 - EPSILON); const auto curr_discrete_log_probs = - torch::sum(discrete_actions * torch::log(clamped_proba), -1, true); + bernoulli_log_proba(discrete_actions, discrete_proba).sum(-1, true); const auto log_ratio = torch::clamp( curr_continuous_log_probs + curr_discrete_log_probs - old_log_probs, -LOG_RATIO_MAX_ABS, @@ -167,8 +179,8 @@ namespace arenai::agent { const auto ratio = torch::exp(log_ratio); - const auto continuous_entropy = truncated_normal_entropy(mu, sigma); - const auto discrete_entropy = multinomial_entropy(discrete_proba); + const auto continuous_entropy = beta_law_entropy(mode, concentration); + const auto discrete_entropy = bernoulli_entropy(discrete_proba); const auto kl_per_row = (ratio - 1.f - log_ratio).flatten(); @@ -184,7 +196,7 @@ namespace arenai::agent { const auto entropy_bonus = torch::sum(continuous_alpha->alpha().detach() * continuous_entropy, -1) - + discrete_alpha->alpha().squeeze(1).detach() * discrete_entropy; + + torch::sum(discrete_alpha->alpha().detach() * discrete_entropy, -1); if (!kl_exceeded) { const auto clipped_ratio = torch::clamp(ratio, 1.f - clip_epsilon, 1.f + clip_epsilon); @@ -276,9 +288,13 @@ namespace arenai::agent { const auto rewards = rollout.rewards.to(torch::kFloat); const auto dones = rollout.dones.to(torch::kFloat); + const auto truncateds = rollout.truncateds.to(torch::kFloat); const auto valids = rollout.valids.to(torch::kFloat); - const auto deltas = rewards + gamma * next_values * (1.f - dones) - values; + // a truncated step bootstraps on its own value (the last alive observation): + // the post-mortem next state is not a state the counterfactual life would reach + const auto deltas = + rewards + gamma * (next_values * (1.f - dones) + values * truncateds) - values; auto advantages = torch::zeros_like(deltas); auto gae = torch::zeros({nb_tanks, 1}, deltas.options()); diff --git a/arenai_agent/src/agents/ppo/ppo_trainer.h b/arenai_agent/src/agents/ppo/ppo_trainer.h index 93021335..957330cb 100644 --- a/arenai_agent/src/agents/ppo/ppo_trainer.h +++ b/arenai_agent/src/agents/ppo/ppo_trainer.h @@ -30,8 +30,9 @@ namespace arenai::agent { const std::vector> &vision_channels, const std::vector &group_norm_nums, torch::Device device, int metric_window_size, float gamma, float gae_lambda, float clip_epsilon, float target_kl, float grad_norm_max, - float continuous_target_entropy, float discrete_target_entropy_factor, int epochs, - int rollout_size, int minibatch_size); + const std::vector &continuous_target_entropy, + const std::vector &discrete_target_entropy_factors, int epochs, int rollout_size, + int minibatch_size); void step() override; @@ -48,8 +49,8 @@ namespace arenai::agent { std::unique_ptr continuous_alpha; std::unique_ptr discrete_alpha; - float continuous_target_entropy; - float discrete_target_entropy; + std::vector continuous_target_entropy; + std::vector discrete_target_entropy; std::shared_ptr critic; diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.cpp new file mode 100644 index 00000000..9befd04c --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.cpp @@ -0,0 +1,25 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_hidden_state.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + LiquidHiddenState::LiquidHiddenState(std::shared_ptr actor) + : actor(std::move(actor)) {} + + torch::Tensor LiquidHiddenState::get(const long batch_size) { + if (!x.defined() || x.size(0) != batch_size) + x = actor->initial_state(static_cast(batch_size)); + return x; + } + + void LiquidHiddenState::set(const torch::Tensor &next_x) { x = next_x.detach(); } + + void LiquidHiddenState::reset() { x = torch::Tensor(); } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.h b/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.h new file mode 100644 index 00000000..51a8e904 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_hidden_state.h @@ -0,0 +1,38 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_HIDDEN_STATE_H +#define ARENAI_LIQUID_HIDDEN_STATE_H + +#include + +#include + +#include "../../networks/recurrent/liquid_actor.h" + +namespace arenai::agent { + + // Actor liquid state carried across act() calls, shared between the agent + // (reads it before each step, advances it after) and the collector (resets + // it when the episode ends). + class LiquidHiddenState { + public: + explicit LiquidHiddenState(std::shared_ptr actor); + + // the state opening the next step; re-drawn on first use, on a batch + // size change and after reset() + torch::Tensor get(long batch_size); + void set(const torch::Tensor &next_x); + + // fresh episode: the next get() re-draws every tank's initial state + void reset(); + + private: + std::shared_ptr actor; + torch::Tensor x; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_HIDDEN_STATE_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.cpp new file mode 100644 index 00000000..41d1f049 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.cpp @@ -0,0 +1,79 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_agent.h" + +#include "../../distributions/bernoulli.h" +#include "../../distributions/beta_law.h" +#include "../../networks/constants.h" +#include "../../networks_utils/torch_converter.h" +#include "../../networks_utils/torch_loader.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + /* + * Torch liquid PPO agent + */ + + TorchLiquidPpoAgent::TorchLiquidPpoAgent( + const std::shared_ptr &actor, + const std::shared_ptr &hidden_state, const torch::Device device, + std::optional> collector) + : actor(actor), hidden_state(hidden_state), collector(std::move(collector)) { + actor->to(device); + } + + std::vector TorchLiquidPpoAgent::act( + const std::vector &states, const int vision_height, const int vision_width) { + const auto [continuous_action, discrete_action] = + act(states_to_tensor(states, vision_height, vision_width), false); + return tensor_to_actions(continuous_action, discrete_action); + } + + TorchAction TorchLiquidPpoAgent::act(const TorchState &state, const bool sample) { + TorchAction action; + torch::Tensor continuous_log_prob; + torch::Tensor discrete_log_prob; + torch::Tensor x_t; + + { + torch::NoGradGuard guard; + + const auto &[vision, sensors] = state; + + x_t = hidden_state->get(vision.size(0)); + const auto &[mode, concentration, discrete_proba, next_x] = + actor->act(vision, sensors, x_t); + hidden_state->set(next_x); + + if (sample) { + action.continuous_action = beta_law_sample(mode, concentration); + action.discrete_action = bernoulli_sample(discrete_proba); + } else { + action.continuous_action = beta_law_mode_action(mode); + action.discrete_action = bernoulli_max_action(discrete_proba); + } + + // old log-probabilities, kept for the PPO importance ratio + continuous_log_prob = + beta_law_log_proba(action.continuous_action, mode, concentration).sum(-1, true); + + discrete_log_prob = + bernoulli_log_proba(action.discrete_action, discrete_proba).sum(-1, true); + } + + if (collector.has_value()) + collector.value()->on_act(state, action, continuous_log_prob, discrete_log_prob, x_t); + + return action; + } + + void TorchLiquidPpoAgent::load(const std::filesystem::path &agent_folder) { + load_torch(agent_folder, actor, "actor.pt"); + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.h new file mode 100644 index 00000000..bace4608 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_agent.h @@ -0,0 +1,39 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_AGENT_H +#define ARENAI_LIQUID_PPO_AGENT_H + +#include + +#include "../../networks/recurrent/liquid_actor.h" +#include "../torch_agent.h" +#include "./liquid_hidden_state.h" +#include "./liquid_ppo_collector.h" + +namespace arenai::agent { + + class TorchLiquidPpoAgent final : public virtual AbstractAgent, + public virtual AbstractTorchAgent { + public: + TorchLiquidPpoAgent( + const std::shared_ptr &actor, + const std::shared_ptr &hidden_state, torch::Device device, + std::optional> collector = std::nullopt); + + TorchAction act(const TorchState &state, bool sample) override; + + std::vector + act(const std::vector &states, int vision_height, int vision_width) override; + void load(const std::filesystem::path &agent_folder) override; + + private: + std::shared_ptr actor; + std::shared_ptr hidden_state; + std::optional> collector; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_AGENT_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.cpp new file mode 100644 index 00000000..d0931a89 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.cpp @@ -0,0 +1,52 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_collector.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + LiquidPpoStepCollector::LiquidPpoStepCollector( + std::shared_ptr rollout_buffer, + std::shared_ptr hidden_state) + : rollout_buffer(std::move(rollout_buffer)), hidden_state(std::move(hidden_state)) {} + + void LiquidPpoStepCollector::on_act( + const TorchState &state, const TorchAction &action, + const torch::Tensor &continuous_log_prob, const torch::Tensor &discrete_log_prob, + const torch::Tensor &actor_hidden) { + last_state = state; + last_action = action; + last_continuous_log_prob = continuous_log_prob; + last_discrete_log_prob = discrete_log_prob; + last_actor_hidden = actor_hidden; + } + + void LiquidPpoStepCollector::on_transition( + const torch::Tensor &rewards, const torch::Tensor &done, const torch::Tensor &truncated) { + rollout_buffer->add( + {.state = last_state, + .action = last_action, + .continuous_log_prob = last_continuous_log_prob, + .discrete_log_prob = last_discrete_log_prob, + .reward = rewards, + .done = done, + .truncated = truncated, + .actor_hidden = last_actor_hidden, + .episode_start = next_episode_start}); + + next_episode_start = false; + } + + void LiquidPpoStepCollector::on_episode_end(const TorchState &final_state) { + rollout_buffer->finish_episode(final_state); + + // the environment resets every tank: re-draw the liquid states + hidden_state->reset(); + next_episode_start = true; + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.h new file mode 100644 index 00000000..65c42e5b --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_collector.h @@ -0,0 +1,51 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_COLLECTOR_H +#define ARENAI_LIQUID_PPO_COLLECTOR_H + +#include + +#include "../step_collector.h" +#include "../torch_types.h" +#include "./liquid_hidden_state.h" +#include "./liquid_ppo_rollout_buffer.h" + +namespace arenai::agent { + + class LiquidPpoStepCollector final : public AbstractStepCollector { + public: + LiquidPpoStepCollector( + std::shared_ptr rollout_buffer, + std::shared_ptr hidden_state); + + // concrete act-time channel, called by TorchLiquidPpoAgent::act + void on_act( + const TorchState &state, const TorchAction &action, + const torch::Tensor &continuous_log_prob, const torch::Tensor &discrete_log_prob, + const torch::Tensor &actor_hidden); + + void on_transition( + const torch::Tensor &rewards, const torch::Tensor &done, + const torch::Tensor &truncated) override; + + void on_episode_end(const TorchState &final_state) override; + + private: + std::shared_ptr rollout_buffer; + std::shared_ptr hidden_state; + + TorchState last_state; + TorchAction last_action; + torch::Tensor last_continuous_log_prob; + torch::Tensor last_discrete_log_prob; + torch::Tensor last_actor_hidden; + + // raised at construction and by on_episode_end, consumed by the next on_transition + bool next_episode_start = true; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_COLLECTOR_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.cpp new file mode 100644 index 00000000..e0ffcf8d --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.cpp @@ -0,0 +1,46 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_factory.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + LiquidPpoTorchAgentFactory::LiquidPpoTorchAgentFactory( + const int vision_height, const int vision_width, const int nb_sensors, + const int nb_continuous_actions, const int nb_discrete_actions, const torch::Device device, + const LiquidPpoHyperParams ¶ms) + : config(cli_fields_to_json(liquid_ppo_cli_fields(), params)), + actor(std::make_shared( + vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, + params.hidden_size_sensors, params.vision_channels, params.group_norm_nums, + params.neuron_number, params.unfolding_steps, params.delta_t, params.initial_sigma, + std::vector{params.initial_fire_proba, params.initial_zoom_proba})), + hidden_state(std::make_shared(actor)), + rollout_buffer(std::make_shared()), + collector(std::make_shared(rollout_buffer, hidden_state)), + agent(std::make_shared(actor, hidden_state, device, collector)), + trainer(std::make_shared( + actor, rollout_buffer, vision_height, vision_width, nb_sensors, nb_continuous_actions, + nb_discrete_actions, params.actor_learning_rate, params.critic_learning_rate, + params.hidden_size_sensors, params.vision_channels, params.group_norm_nums, + params.neuron_number, params.unfolding_steps, params.delta_t, device, + params.metric_window_size, params.gamma, params.gae_lambda, params.clip_epsilon, + params.target_kl, params.grad_norm_max, params.continuous_target_entropy, + params.discrete_target_entropy_factors, params.epochs, params.rollout_size, + params.minibatch_size, params.chunk_size)) {} + + std::shared_ptr LiquidPpoTorchAgentFactory::get_agent() { return agent; } + + std::shared_ptr LiquidPpoTorchAgentFactory::get_collector() { + return collector; + } + + std::shared_ptr LiquidPpoTorchAgentFactory::get_trainer() { return trainer; } + + nlohmann::json LiquidPpoTorchAgentFactory::get_config() const { return config; } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.h new file mode 100644 index 00000000..aa60cc8f --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_factory.h @@ -0,0 +1,45 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_FACTORY_H +#define ARENAI_LIQUID_PPO_FACTORY_H + +#include "../torch_factory.h" +#include "./liquid_hidden_state.h" +#include "./liquid_ppo_agent.h" +#include "./liquid_ppo_collector.h" +#include "./liquid_ppo_hyperparams.h" +#include "./liquid_ppo_rollout_buffer.h" +#include "./liquid_ppo_trainer.h" + +namespace arenai::agent { + + class LiquidPpoTorchAgentFactory final : public AbstractTorchAgentFactory { + public: + LiquidPpoTorchAgentFactory( + int vision_height, int vision_width, int nb_sensors, int nb_continuous_actions, + int nb_discrete_actions, torch::Device device, const LiquidPpoHyperParams ¶ms); + + std::shared_ptr get_agent() override; + std::shared_ptr get_collector() override; + std::shared_ptr get_trainer() override; + + nlohmann::json get_config() const override; + + private: + nlohmann::json config; + + // triad built once, sharing actor + hidden state + rollout_buffer + std::shared_ptr actor; + std::shared_ptr hidden_state; + + std::shared_ptr rollout_buffer; + std::shared_ptr collector; + std::shared_ptr agent; + std::shared_ptr trainer; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_FACTORY_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.cpp new file mode 100644 index 00000000..52a94886 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.cpp @@ -0,0 +1,40 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_hyperparams.h" + +namespace arenai::agent { + + std::vector> liquid_ppo_cli_fields() { + return { + {.name = "--actor_learning_rate", .member = &LiquidPpoHyperParams::actor_learning_rate}, + {.name = "--critic_learning_rate", + .member = &LiquidPpoHyperParams::critic_learning_rate}, + {.name = "--hidden_size_sensors", .member = &LiquidPpoHyperParams::hidden_size_sensors}, + {.name = "--vision_channels", .member = &LiquidPpoHyperParams::vision_channels}, + {.name = "--group_norm_nums", .member = &LiquidPpoHyperParams::group_norm_nums}, + {.name = "--neuron_number", .member = &LiquidPpoHyperParams::neuron_number}, + {.name = "--unfolding_steps", .member = &LiquidPpoHyperParams::unfolding_steps}, + {.name = "--delta_t", .member = &LiquidPpoHyperParams::delta_t}, + {.name = "--chunk_size", .member = &LiquidPpoHyperParams::chunk_size}, + {.name = "--initial_sigma", .member = &LiquidPpoHyperParams::initial_sigma}, + {.name = "--initial_fire_proba", .member = &LiquidPpoHyperParams::initial_fire_proba}, + {.name = "--initial_zoom_proba", .member = &LiquidPpoHyperParams::initial_zoom_proba}, + {.name = "--metric_window_size", .member = &LiquidPpoHyperParams::metric_window_size}, + {.name = "--gamma", .member = &LiquidPpoHyperParams::gamma}, + {.name = "--gae_lambda", .member = &LiquidPpoHyperParams::gae_lambda}, + {.name = "--clip_epsilon", .member = &LiquidPpoHyperParams::clip_epsilon}, + {.name = "--target_kl", .member = &LiquidPpoHyperParams::target_kl}, + {.name = "--grad_norm_max", .member = &LiquidPpoHyperParams::grad_norm_max}, + {.name = "--continuous_target_entropy", + .member = &LiquidPpoHyperParams::continuous_target_entropy}, + {.name = "--discrete_target_entropy_factors", + .member = &LiquidPpoHyperParams::discrete_target_entropy_factors}, + {.name = "--epochs", .member = &LiquidPpoHyperParams::epochs}, + {.name = "--rollout_size", .member = &LiquidPpoHyperParams::rollout_size}, + {.name = "--minibatch_size", .member = &LiquidPpoHyperParams::minibatch_size}, + }; + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.h new file mode 100644 index 00000000..b4f9e148 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_hyperparams.h @@ -0,0 +1,49 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_HYPERPARAMS_H +#define ARENAI_LIQUID_PPO_HYPERPARAMS_H + +#include +#include + +#include "../../utils/cli_fields.h" + +namespace arenai::agent { + + // Member initializers are the CLI defaults (single source of truth). + struct LiquidPpoHyperParams { + float actor_learning_rate = 1e-4f; + float critic_learning_rate = 3e-4f; + int hidden_size_sensors = 128; + std::vector> vision_channels = {{3, 16}, {16, 24}, {24, 32}, {32, 48}, + {48, 64}, {64, 96}, {96, 128}}; + std::vector group_norm_nums = {2, 3, 4, 6, 8, 12, 16}; + int neuron_number = 128; + int unfolding_steps = 6; + float delta_t = 1.f / 30.f; + int chunk_size = 30; + float initial_sigma = 0.5f; + float initial_fire_proba = 0.4f; + float initial_zoom_proba = 0.25f; + int metric_window_size = 256; + float gamma = 0.997f; + float gae_lambda = 0.99f; + float clip_epsilon = 0.2f; + float target_kl = 0.05f; + float grad_norm_max = 0.5f; + // per continuous action: direction x, direction y, canon x, canon y + std::vector continuous_target_entropy = {-0.2f, -0.2f, -0.88f, -0.88f}; + // factors of the Bernoulli maximum entropy, per discrete action: fire, zoom + std::vector discrete_target_entropy_factors = {0.4f, 0.4f}; + int epochs = 2; + int rollout_size = 30 * 30; + int minibatch_size = 1024; + }; + + std::vector> liquid_ppo_cli_fields(); + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_HYPERPARAMS_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.cpp new file mode 100644 index 00000000..300d3083 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.cpp @@ -0,0 +1,110 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_rollout_buffer.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + namespace { + TorchState to_cpu(const TorchState &state) { + return { + .vision = state.vision.detach().cpu(), + .proprioception = state.proprioception.detach().cpu()}; + } + }// namespace + + void LiquidPpoRolloutBuffer::add(const LiquidPpoInputStep &step) { + const auto nb_tanks = step.state.vision.size(0); + + if (!already_terminated_.defined()) + already_terminated_ = torch::zeros( + {nb_tanks}, torch::TensorOptions().dtype(torch::kBool).device(torch::kCPU)); + + // tanks already terminated before this step have no valid transition to store + const auto valid = already_terminated_.logical_not(); + already_terminated_.logical_or_( + step.done.detach().cpu().to(torch::kBool).reshape({nb_tanks})); + + steps_.push_back( + {.step = + {.state = to_cpu(step.state), + .action = + {.continuous_action = step.action.continuous_action.detach().cpu(), + .discrete_action = step.action.discrete_action.detach().cpu()}, + .continuous_log_prob = step.continuous_log_prob.detach().cpu(), + .discrete_log_prob = step.discrete_log_prob.detach().cpu(), + .reward = step.reward.detach().cpu(), + .done = step.done.detach().cpu(), + .truncated = step.truncated.detach().cpu(), + .actor_hidden = step.actor_hidden.detach().cpu(), + .episode_start = step.episode_start}, + .valid = valid}); + + // the freshly added step is pending: its closing observation is not known yet + final_state_.reset(); + } + + void LiquidPpoRolloutBuffer::finish_episode(const TorchState &final_state) { + if (!steps_.empty()) final_state_ = to_cpu(final_state); + + if (already_terminated_.defined()) already_terminated_.fill_(false); + } + + size_t LiquidPpoRolloutBuffer::nb_complete_steps() const { + if (steps_.empty()) return 0; + return final_state_.has_value() ? steps_.size() : steps_.size() - 1; + } + + LiquidPpoRollout LiquidPpoRolloutBuffer::get_rollout() { + const auto nb_steps = nb_complete_steps(); + TORCH_CHECK(nb_steps > 0, "LiquidPpoRolloutBuffer: no complete step to train on"); + + // the observation closing the last consumed step: the episode's final state, + // or the pending step's own observation + const auto bootstrap_state = + final_state_.has_value() ? *final_state_ : steps_[nb_steps].step.state; + + const auto stack = [&](const auto &get_field) { + std::vector tensors; + tensors.reserve(nb_steps); + for (size_t i = 0; i < nb_steps; i++) tensors.push_back(get_field(steps_[i])); + return torch::stack(tensors, 0); + }; + + auto episode_starts = torch::zeros({static_cast(nb_steps)}, torch::kBool); + for (size_t i = 0; i < nb_steps; i++) + if (steps_[i].step.episode_start) episode_starts[static_cast(i)] = true; + + LiquidPpoRollout rollout{ + .states = + {.vision = stack([](const StoredStep &s) { return s.step.state.vision; }), + .proprioception = + stack([](const StoredStep &s) { return s.step.state.proprioception; })}, + .actions = + {.continuous_action = + stack([](const StoredStep &s) { return s.step.action.continuous_action; }), + .discrete_action = + stack([](const StoredStep &s) { return s.step.action.discrete_action; })}, + .continuous_log_probs = + stack([](const StoredStep &s) { return s.step.continuous_log_prob; }), + .discrete_log_probs = + stack([](const StoredStep &s) { return s.step.discrete_log_prob; }), + .rewards = stack([](const StoredStep &s) { return s.step.reward; }), + .dones = stack([](const StoredStep &s) { return s.step.done; }), + .truncateds = stack([](const StoredStep &s) { return s.step.truncated; }), + .bootstrap_state = bootstrap_state, + .valids = stack([](const StoredStep &s) { return s.valid; }).unsqueeze(-1), + .actor_hiddens = stack([](const StoredStep &s) { return s.step.actor_hidden; }), + .episode_starts = episode_starts}; + + steps_.erase(steps_.begin(), steps_.begin() + static_cast(nb_steps)); + final_state_.reset(); + + return rollout; + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.h new file mode 100644 index 00000000..9f3bb8ca --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_rollout_buffer.h @@ -0,0 +1,81 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_ROLLOUT_BUFFER_H +#define ARENAI_LIQUID_PPO_ROLLOUT_BUFFER_H + +#include +#include + +#include + +#include "../torch_types.h" + +namespace arenai::agent { + + struct LiquidPpoInputStep { + TorchState state; + TorchAction action; + torch::Tensor continuous_log_prob; + torch::Tensor discrete_log_prob; + torch::Tensor reward; + torch::Tensor done; + // done whose value target must bootstrap instead of cutting the return (famine) + torch::Tensor truncated; + // [nb_tanks, neuron_number] actor liquid state opening the step + torch::Tensor actor_hidden; + // the step opens a fresh episode (the liquid states were re-drawn) + bool episode_start; + }; + + // On-policy rollout stacked on the time dimension: every tensor is [T, nb_tanks, ...] + struct LiquidPpoRollout { + TorchState states; + TorchAction actions; + torch::Tensor continuous_log_probs; + torch::Tensor discrete_log_probs; + torch::Tensor rewards; + torch::Tensor dones; + // dones whose value target must bootstrap instead of cutting the return + torch::Tensor truncateds; + // [nb_tanks, ...] observation closing the last step, for the value bootstrap + TorchState bootstrap_state; + // [T, nb_tanks, 1] whether the (step, tank) pair is a live transition + torch::Tensor valids; + // [T, nb_tanks, neuron_number] actor liquid state opening each step + torch::Tensor actor_hiddens; + // [T] bool, the step opens a fresh episode + torch::Tensor episode_starts; + }; + + // Sequential on-policy buffer. A step is "complete" once the observation that + // follows it is known - the next add()'s state, or finish_episode()'s final state. + class LiquidPpoRolloutBuffer { + public: + void add(const LiquidPpoInputStep &step); + void finish_episode(const TorchState &final_state); + + size_t nb_complete_steps() const; + + // stacks and removes every complete step; the pending one (if any) stays + LiquidPpoRollout get_rollout(); + + private: + struct StoredStep { + LiquidPpoInputStep step; + torch::Tensor valid; + }; + + std::vector steps_; + + // observation closing the last stored step, set by finish_episode() + std::optional final_state_; + + // [nb_tanks] tanks already done in the current episode + torch::Tensor already_terminated_; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_ROLLOUT_BUFFER_H diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.cpp b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.cpp new file mode 100644 index 00000000..d3b13920 --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.cpp @@ -0,0 +1,445 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_ppo_trainer.h" + +#include +#include + +#include "../../distributions/bernoulli.h" +#include "../../distributions/beta_law.h" +#include "../../metrics/mean_metric.h" +#include "../../networks/constants.h" +#include "../../networks_utils/print_module.h" +#include "../../networks_utils/torch_saver.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + namespace { + // merges the [T, nb_tanks] leading dimensions into a single row dimension + torch::Tensor flatten_steps(const torch::Tensor &tensor) { + auto sizes = tensor.sizes().vec(); + sizes.erase(sizes.begin()); + sizes[0] = tensor.size(0) * tensor.size(1); + return tensor.reshape(sizes); + } + + constexpr float LOG_RATIO_MAX_ABS = 3.f; + + constexpr float KL_TRIM_FRACTION = 0.01f; + + constexpr float CONTINUOUS_ALPHA_K_P = 2e-1f; + constexpr float CONTINUOUS_ALPHA_K_I = 5e-3f; + constexpr float CONTINUOUS_ALPHA_K_D = 1.f; + + constexpr float DISCRETE_ALPHA_K_P = 2e-1f; + constexpr float DISCRETE_ALPHA_K_I = 1e-2f; + constexpr float DISCRETE_ALPHA_K_D = 1.f; + + constexpr float ALPHA_INITIAL = 1e-3f; + + }// namespace + + LiquidPpoTrainer::LiquidPpoTrainer( + const std::shared_ptr &actor, + const std::shared_ptr &rollout_buffer, const int vision_height, + const int vision_width, const int nb_sensors, const int nb_continuous_actions, + const int nb_discrete_action, const float actor_learning_rate, + const float critic_learning_rate, const int hidden_size_sensors, + const std::vector> &vision_channels, + const std::vector &group_norm_nums, const int neuron_number, const int unfolding_steps, + const float delta_t, const torch::Device device, const int metric_window_size, + const float gamma, const float gae_lambda, const float clip_epsilon, const float target_kl, + const float grad_norm_max, const std::vector &continuous_target_entropy, + const std::vector &discrete_target_entropy_factors, const int epochs, + const int rollout_size, const int minibatch_size, const int chunk_size) + : actor(actor), rollout_buffer(rollout_buffer), + continuous_alpha(std::make_unique( + CONTINUOUS_ALPHA_K_P, CONTINUOUS_ALPHA_K_I, CONTINUOUS_ALPHA_K_D, ALPHA_INITIAL, + nb_continuous_actions)), + discrete_alpha(std::make_unique( + DISCRETE_ALPHA_K_P, DISCRETE_ALPHA_K_I, DISCRETE_ALPHA_K_D, ALPHA_INITIAL, + nb_discrete_action)), + continuous_target_entropy(continuous_target_entropy), + // per-action target: each discrete action is an independent Bernoulli + discrete_target_entropy([&discrete_target_entropy_factors] { + std::vector targets = discrete_target_entropy_factors; + for (auto &target: targets) target *= bernoulli_maximum_entropy(); + return targets; + }()), + critic(std::make_shared( + vision_height, vision_width, nb_sensors, hidden_size_sensors, vision_channels, + group_norm_nums, neuron_number, unfolding_steps, delta_t)), + actor_optim( + std::make_unique(this->actor->parameters(), actor_learning_rate)), + critic_optim( + std::make_unique(critic->parameters(), critic_learning_rate)), + actor_mean_loss_metric(std::make_shared("π", metric_window_size)), + critic_mean_loss_metric(std::make_shared("v", metric_window_size)), + explained_variance_metric(std::make_shared("ev", metric_window_size)), + continuous_entropy_metric(std::make_shared("Hc", metric_window_size)), + discrete_entropy_metric(std::make_shared("Hd", metric_window_size)), + continuous_alpha_metric(std::make_shared("α_c", metric_window_size, 2, true)), + discrete_alpha_metric(std::make_shared("α_d", metric_window_size, 2, true)), + clip_fraction_metric(std::make_shared("clip", metric_window_size)), + kl_metric(std::make_shared("kl", metric_window_size, 2, true)), + skip_fraction_metric(std::make_shared("skip", metric_window_size)), + gamma(gamma), gae_lambda(gae_lambda), clip_epsilon(clip_epsilon), target_kl(target_kl), + grad_norm_max(grad_norm_max), epochs(epochs), rollout_size(rollout_size), + minibatch_size(minibatch_size), chunk_size(chunk_size) { + TORCH_CHECK( + static_cast(continuous_target_entropy.size()) == nb_continuous_actions, + "continuous_target_entropy needs one target per continuous action"); + TORCH_CHECK( + static_cast(discrete_target_entropy_factors.size()) == nb_discrete_action, + "discrete_target_entropy_factors needs one factor per discrete action"); + + to(device); + + set_train(false); + } + + void LiquidPpoTrainer::step() { + if (rollout_buffer->nb_complete_steps() >= static_cast(rollout_size)) train(); + } + + void LiquidPpoTrainer::train() { + const auto device = actor->parameters().back().device(); + + const auto rollout = rollout_buffer->get_rollout(); + + set_train(false); + const auto [advantages, returns, critic_hiddens] = compute_gae(rollout, device); + + set_train(true); + + const auto nb_steps = rollout.rewards.size(0); + const auto nb_tanks = rollout.rewards.size(1); + + // contiguous chunk_size-long windows per tank; the (at most chunk_size - 1) + // trailing steps that cannot form a full chunk are dropped + const auto nb_chunks_per_tank = nb_steps / chunk_size; + if (nb_chunks_per_tank == 0) return; + const auto nb_kept_steps = nb_chunks_per_tank * chunk_size; + + // the rollout stays on CPU as flat [T * nb_tanks, ...] views; only the + // minibatches hit the device + const auto flat = [&](const torch::Tensor &tensor) { + return flatten_steps(tensor.slice(0, 0, nb_kept_steps)); + }; + + const auto flat_vision = flat(rollout.states.vision); + const auto flat_proprioception = flat(rollout.states.proprioception); + const auto flat_continuous_actions = flat(rollout.actions.continuous_action); + const auto flat_discrete_actions = flat(rollout.actions.discrete_action); + const auto flat_old_log_probs = + flat(rollout.continuous_log_probs) + flat(rollout.discrete_log_probs); + const auto flat_advantages = flat(advantages); + const auto flat_returns = flat(returns); + const auto flat_valids = flat(rollout.valids); + const auto flat_actor_hiddens = flat(rollout.actor_hiddens); + const auto flat_critic_hiddens = flat(critic_hiddens); + + // chunk unit u = chunk * nb_tanks + tank; the ones without any live + // transition bring nothing to train on + const auto chunk_valid_counts = rollout.valids.slice(0, 0, nb_kept_steps) + .reshape({nb_chunks_per_tank, chunk_size, nb_tanks}) + .sum(1) + .flatten(); + const auto valid_chunk_idx = torch::nonzero(chunk_valid_counts > 0).squeeze(-1); + const auto nb_valid_chunks = valid_chunk_idx.size(0); + if (nb_valid_chunks == 0) return; + + const auto chunks_per_minibatch = std::max(1, minibatch_size / chunk_size); + const auto step_offsets = torch::arange(chunk_size); + + for (int e = 0; e < epochs; e++) { + const auto perm = valid_chunk_idx.index_select(0, torch::randperm(nb_valid_chunks)); + + for (int64_t start = 0; start < nb_valid_chunks; start += chunks_per_minibatch) { + const auto idx = perm.slice( + 0, start, std::min(start + chunks_per_minibatch, nb_valid_chunks)); + + const auto tank = idx.remainder(nb_tanks); + const auto start_step = idx.div(nb_tanks, "floor") * chunk_size; + + // flat rows of the minibatch's chunks and of their first steps + const auto rows = ((start_step.unsqueeze(1) + step_offsets.unsqueeze(0)) * nb_tanks + + tank.unsqueeze(1)) + .flatten(); + const auto start_rows = start_step * nb_tanks + tank; + + // [nb_chunks, chunk_size, ...] + const auto select_chunks = [&](const torch::Tensor &flat_tensor) { + auto sizes = flat_tensor.sizes().vec(); + sizes[0] = idx.size(0); + sizes.insert(sizes.begin() + 1, chunk_size); + return flat_tensor.index_select(0, rows).reshape(sizes).to(device); + }; + // [nb_chunks, neuron_number] + const auto select_hiddens = [&](const torch::Tensor &flat_hiddens) { + return flat_hiddens.index_select(0, start_rows).to(device); + }; + + const auto mb_vision = select_chunks(flat_vision); + const auto mb_proprioception = select_chunks(flat_proprioception); + const auto mb_valids = select_chunks(flat_valids); + + train_actor( + mb_vision, mb_proprioception, select_hiddens(flat_actor_hiddens), + select_chunks(flat_continuous_actions), select_chunks(flat_discrete_actions), + select_chunks(flat_old_log_probs), select_chunks(flat_advantages), mb_valids); + + train_critic( + mb_vision, mb_proprioception, select_hiddens(flat_critic_hiddens), + select_chunks(flat_returns), mb_valids); + } + } + + set_train(false); + } + + bool LiquidPpoTrainer::train_actor( + const torch::Tensor &vision, const torch::Tensor &proprioception, + const torch::Tensor &hidden, const torch::Tensor &continuous_actions, + const torch::Tensor &discrete_actions, const torch::Tensor &old_log_probs, + const torch::Tensor &advantages, const torch::Tensor &valids) const { + const auto device = actor->parameters().back().device(); + + const auto out = actor->act_sequence(vision, proprioception, hidden); + + // recurrence done: fold time into rows and keep the live transitions only + const auto valid_idx = torch::nonzero(valids.flatten(0, 1).squeeze(-1)).squeeze(-1); + const auto rows = [&](const torch::Tensor &tensor) { + return tensor.flatten(0, 1).index_select(0, valid_idx); + }; + + const auto mode = rows(out.mode); + const auto concentration = rows(out.concentration); + const auto discrete_proba = rows(out.discrete); + + const auto curr_continuous_log_probs = + beta_law_log_proba(rows(continuous_actions), mode, concentration).sum(-1, true); + + const auto curr_discrete_log_probs = + bernoulli_log_proba(rows(discrete_actions), discrete_proba).sum(-1, true); + + const auto log_ratio = torch::clamp( + curr_continuous_log_probs + curr_discrete_log_probs - rows(old_log_probs), + -LOG_RATIO_MAX_ABS, LOG_RATIO_MAX_ABS); + + const auto ratio = torch::exp(log_ratio); + + const auto continuous_entropy = beta_law_entropy(mode, concentration); + const auto discrete_entropy = bernoulli_entropy(discrete_proba); + + const auto kl_per_row = (ratio - 1.f - log_ratio).flatten(); + + const auto nb_kept = std::max( + 1, static_cast( + static_cast(kl_per_row.size(0)) * (1.f - KL_TRIM_FRACTION))); + + // per-row KL is non-negative: sorting ascending puts the outliers past nb_kept + const auto approx_kl = + std::get<0>(torch::sort(kl_per_row)).slice(0, 0, nb_kept).mean().item(); + + const bool kl_exceeded = target_kl > 0.f && approx_kl > 1.5f * target_kl; + + const auto entropy_bonus = + torch::sum(continuous_alpha->alpha().detach() * continuous_entropy, -1) + + torch::sum(discrete_alpha->alpha().detach() * discrete_entropy, -1); + + if (!kl_exceeded) { + const auto clipped_ratio = torch::clamp(ratio, 1.f - clip_epsilon, 1.f + clip_epsilon); + const auto valid_advantages = rows(advantages); + const auto surrogate = + torch::min(ratio * valid_advantages, clipped_ratio * valid_advantages); + + const auto actor_loss = -torch::mean(surrogate + entropy_bonus); + + actor_optim->zero_grad(); + actor_loss.backward(); + torch::nn::utils::clip_grad_norm_(actor->parameters(), grad_norm_max); + actor_optim->step(); + + // actor metrics + actor_mean_loss_metric->add(actor_loss.cpu().item()); + } + + // adjust alphas + continuous_alpha->update( + continuous_entropy, torch::tensor(continuous_target_entropy, device)); + discrete_alpha->update(discrete_entropy, torch::tensor(discrete_target_entropy, device)); + + // metrics + continuous_entropy_metric->add(continuous_entropy.mean().item()); + discrete_entropy_metric->add(discrete_entropy.mean().item()); + + continuous_alpha_metric->add(continuous_alpha->alpha().mean().item()); + discrete_alpha_metric->add(discrete_alpha->alpha().mean().item()); + + kl_metric->add(approx_kl); + + clip_fraction_metric->add( + ((ratio - 1.f).abs() > clip_epsilon).to(torch::kFloat).mean().item()); + + skip_fraction_metric->add(kl_exceeded ? 1.f : 0.f); + + return kl_exceeded; + } + + void LiquidPpoTrainer::train_critic( + const torch::Tensor &vision, const torch::Tensor &proprioception, + const torch::Tensor &hidden, const torch::Tensor &returns, + const torch::Tensor &valids) const { + const auto [value, next_x] = critic->value_sequence(vision, proprioception, hidden); + + // recurrence done: fold time into rows and keep the live transitions only + const auto valid_idx = torch::nonzero(valids.flatten(0, 1).squeeze(-1)).squeeze(-1); + const auto values = value.flatten(0, 1).index_select(0, valid_idx); + const auto valid_returns = returns.flatten(0, 1).index_select(0, valid_idx); + + const auto critic_loss = torch::mse_loss(values, valid_returns, at::Reduction::Mean); + + critic_optim->zero_grad(); + critic_loss.backward(); + torch::nn::utils::clip_grad_norm_(critic->parameters(), grad_norm_max); + critic_optim->step(); + + critic_mean_loss_metric->add(critic_loss.cpu().item()); + + const auto residual_var = (valid_returns - values.detach()).var(false); + const auto returns_var = valid_returns.var(false); + explained_variance_metric->add( + (1.f - residual_var / returns_var.clamp_min(EPSILON)).item()); + } + + LiquidGaeResult + LiquidPpoTrainer::compute_gae(const LiquidPpoRollout &rollout, const torch::Device device) { + torch::NoGradGuard no_grad; + + const auto nb_steps = rollout.rewards.size(0); + const auto nb_tanks = rollout.rewards.size(1); + + const auto episode_starts = rollout.episode_starts.accessor(); + + if (!critic_carry.defined() || critic_carry.size(0) != nb_tanks) + critic_carry = critic->initial_state(static_cast(nb_tanks)); + + // sequential critic pass: the values and the liquid state opening each step + std::vector value_list; + std::vector hidden_list; + value_list.reserve(nb_steps); + hidden_list.reserve(nb_steps); + + for (int64_t t = 0; t < nb_steps; t++) { + if (episode_starts[t]) critic_carry = critic->initial_state(static_cast(nb_tanks)); + + hidden_list.push_back(critic_carry.cpu()); + + const auto out = critic->value( + rollout.states.vision[t].to(device), rollout.states.proprioception[t].to(device), + critic_carry); + + value_list.push_back(out.value.cpu()); + critic_carry = out.next_x; + } + + const auto values = torch::stack(value_list, 0); + const auto critic_hiddens = torch::stack(hidden_list, 0); + + // the carry is NOT advanced past the bootstrap: its observation opens the + // next rollout's first step + const auto bootstrap_value = + critic + ->value( + rollout.bootstrap_state.vision.to(device), + rollout.bootstrap_state.proprioception.to(device), critic_carry) + .value.cpu() + .unsqueeze(0); + + // next values are the values shifted by one step, closed by the bootstrap state + const auto next_values = torch::cat({values.slice(0, 1), bootstrap_value}, 0); + + const auto rewards = rollout.rewards.to(torch::kFloat); + const auto dones = rollout.dones.to(torch::kFloat); + const auto truncateds = rollout.truncateds.to(torch::kFloat); + const auto valids = rollout.valids.to(torch::kFloat); + + // a truncated step bootstraps on its own value (the last alive observation): + // the post-mortem next state is not a state the counterfactual life would reach + const auto deltas = + rewards + gamma * (next_values * (1.f - dones) + values * truncateds) - values; + + auto advantages = torch::zeros_like(deltas); + auto gae = torch::zeros({nb_tanks, 1}, deltas.options()); + for (int64_t t = nb_steps - 1; t >= 0; t--) { + gae = deltas[t] + gamma * gae_lambda * (1.f - dones[t]) * gae; + advantages[t] = gae; + } + + const auto returns = advantages + values; + + const auto nb_valid = valids.sum().clamp_min(1.f); + const auto advantage_mean = torch::sum(advantages * valids) / nb_valid; + const auto advantage_std = + torch::sqrt(torch::sum(torch::square(advantages - advantage_mean) * valids) / nb_valid); + advantages = (advantages - advantage_mean) / (advantage_std + EPSILON); + + return {.advantages = advantages, .returns = returns, .critic_hiddens = critic_hiddens}; + } + + std::vector> LiquidPpoTrainer::get_metrics() { + return {actor_mean_loss_metric, critic_mean_loss_metric, explained_variance_metric, + continuous_entropy_metric, continuous_alpha_metric, discrete_entropy_metric, + discrete_alpha_metric, clip_fraction_metric, kl_metric, + skip_fraction_metric}; + } + + void LiquidPpoTrainer::save(const std::filesystem::path &output_folder) { + // Models + save_torch(output_folder, actor, "actor.pt"); + save_torch(output_folder, critic, "critic.pt"); + + // Optimizers + save_torch(output_folder, actor_optim, "actor_optim.pt"); + save_torch(output_folder, critic_optim, "critic_optim.pt"); + + // string repr + std::ostringstream actor_repr_oss; + dump_module_tree(actor, actor_repr_oss, 0, "actor"); + std::ofstream actor_repr_file(output_folder / "actor_repr.txt"); + actor_repr_file << actor_repr_oss.str(); + actor_repr_file.close(); + + std::ostringstream critic_repr_oss; + dump_module_tree(critic, critic_repr_oss, 0, "critic"); + std::ofstream critic_repr_file(output_folder / "critic_repr.txt"); + critic_repr_file << critic_repr_oss.str(); + critic_repr_file.close(); + } + + void LiquidPpoTrainer::set_train(const bool train) const { + actor->train(train); + critic->train(train); + + continuous_alpha->train(train); + discrete_alpha->train(train); + } + + void LiquidPpoTrainer::to(const torch::Device device) const { + actor->to(device); + critic->to(device); + + continuous_alpha->to(device); + discrete_alpha->to(device); + } + + int LiquidPpoTrainer::count_parameters() { + return count_parameters_impl(actor->parameters()) + + count_parameters_impl(critic->parameters()); + } +}// namespace arenai::agent diff --git a/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.h b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.h new file mode 100644 index 00000000..2f44baba --- /dev/null +++ b/arenai_agent/src/agents/ppo_liquid/liquid_ppo_trainer.h @@ -0,0 +1,123 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_PPO_TRAINER_H +#define ARENAI_LIQUID_PPO_TRAINER_H + +#include "../../networks/entropy.h" +#include "../../networks/recurrent/liquid_actor.h" +#include "../../networks/recurrent/liquid_critic.h" +#include "../trainer.h" +#include "./liquid_ppo_rollout_buffer.h" + +namespace arenai::agent { + + struct LiquidGaeResult { + torch::Tensor advantages; + torch::Tensor returns; + // [T, nb_tanks, neuron_number] critic liquid state opening each step + torch::Tensor critic_hiddens; + }; + + // Recurrent PPO: the updates run on contiguous [chunk_size]-long sequences + // (truncated BPTT), each initialized with the liquid state recorded when the + // chunk's first step was collected. + class LiquidPpoTrainer final : public AbstractTrainer { + public: + LiquidPpoTrainer( + const std::shared_ptr &actor, + const std::shared_ptr &rollout_buffer, int vision_height, + int vision_width, int nb_sensors, int nb_continuous_actions, int nb_discrete_action, + float actor_learning_rate, float critic_learning_rate, int hidden_size_sensors, + const std::vector> &vision_channels, + const std::vector &group_norm_nums, int neuron_number, int unfolding_steps, + float delta_t, torch::Device device, int metric_window_size, float gamma, + float gae_lambda, float clip_epsilon, float target_kl, float grad_norm_max, + const std::vector &continuous_target_entropy, + const std::vector &discrete_target_entropy_factors, int epochs, int rollout_size, + int minibatch_size, int chunk_size); + + void step() override; + + std::vector> get_metrics() override; + + void save(const std::filesystem::path &output_folder) override; + + int count_parameters() override; + + private: + std::shared_ptr actor; + std::shared_ptr rollout_buffer; + + std::unique_ptr continuous_alpha; + std::unique_ptr discrete_alpha; + + // per-dimension targets, aligned with the per-dimension alphas + std::vector continuous_target_entropy; + std::vector discrete_target_entropy; + + std::shared_ptr critic; + + std::unique_ptr actor_optim; + std::unique_ptr critic_optim; + + std::shared_ptr actor_mean_loss_metric; + std::shared_ptr critic_mean_loss_metric; + + // share of the return variance the critic explains, per critic minibatch + std::shared_ptr explained_variance_metric; + + // both regulated by their constant entropy bonus + std::shared_ptr continuous_entropy_metric; + std::shared_ptr discrete_entropy_metric; + + std::shared_ptr continuous_alpha_metric; + std::shared_ptr discrete_alpha_metric; + + // both recorded on every attempted minibatch, skipped ones included + std::shared_ptr clip_fraction_metric; + std::shared_ptr kl_metric; + + // fraction of minibatches the KL threshold skipped + std::shared_ptr skip_fraction_metric; + + float gamma; + float gae_lambda; + float clip_epsilon; + float target_kl; + + float grad_norm_max; + + int epochs; + int rollout_size; + int minibatch_size; + int chunk_size; + + // critic liquid state opening the next rollout's first step, carried + // across train() calls; re-drawn at every episode start + torch::Tensor critic_carry; + + void train(); + + bool train_actor( + const torch::Tensor &vision, const torch::Tensor &proprioception, + const torch::Tensor &hidden, const torch::Tensor &continuous_actions, + const torch::Tensor &discrete_actions, const torch::Tensor &old_log_probs, + const torch::Tensor &advantages, const torch::Tensor &valids) const; + + void train_critic( + const torch::Tensor &vision, const torch::Tensor &proprioception, + const torch::Tensor &hidden, const torch::Tensor &returns, + const torch::Tensor &valids) const; + + // sequential critic pass over the rollout, advancing critic_carry + LiquidGaeResult compute_gae(const LiquidPpoRollout &rollout, torch::Device device); + + void set_train(bool train) const; + void to(torch::Device device) const; + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_PPO_TRAINER_H diff --git a/arenai_agent/src/agents/sac/sac_agent.cpp b/arenai_agent/src/agents/sac/sac_agent.cpp deleted file mode 100644 index b2d3963c..00000000 --- a/arenai_agent/src/agents/sac/sac_agent.cpp +++ /dev/null @@ -1,64 +0,0 @@ -// -// Created by samuel on 21/01/2026. -// - -#include "./sac_agent.h" - -#include "../../distributions/multinomial.h" -#include "../../distributions/truncated_normal.h" -#include "../../networks_utils/torch_converter.h" -#include "../../networks_utils/torch_loader.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - /* - * Torch SAC agent - */ - - TorchSacAgent::TorchSacAgent( - const std::shared_ptr &actor, const torch::Device device, - std::optional> collector) - : actor(actor), collector(std::move(collector)) { - actor->to(device); - } - - std::vector TorchSacAgent::act( - const std::vector &states, const int vision_height, const int vision_width) { - const auto [continuous_action, discrete_action] = - act(states_to_tensor(states, vision_height, vision_width), false); - return tensor_to_actions(continuous_action, discrete_action); - } - - TorchAction TorchSacAgent::act(const TorchState &state, const bool sample) { - TorchAction action; - - { - torch::NoGradGuard guard; - - actor->train(false); - - const auto &[vision, sensors] = state; - const auto &[mu, sigma, discrete_proba] = actor->act(vision, sensors); - - if (sample) { - action.continuous_action = truncated_normal_sample(mu, sigma); - action.discrete_action = multinomial_sample(discrete_proba); - } else { - action.continuous_action = mu; - action.discrete_action = multinomial_max_action(discrete_proba); - } - } - - if (collector.has_value()) collector.value()->on_act(state, action); - - return action; - } - - void TorchSacAgent::load(const std::filesystem::path &agent_folder) { - load_torch(agent_folder, actor, "actor.pt"); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_agent.h b/arenai_agent/src/agents/sac/sac_agent.h deleted file mode 100644 index 25d4481b..00000000 --- a/arenai_agent/src/agents/sac/sac_agent.h +++ /dev/null @@ -1,35 +0,0 @@ -// -// Created by samuel on 21/01/2026. -// - -#ifndef ARENAI_AGENT_HOST_SAC_H -#define ARENAI_AGENT_HOST_SAC_H - -#include - -#include "../../networks/actor.h" -#include "../torch_agent.h" -#include "./sac_collector.h" - -namespace arenai::agent { - - class TorchSacAgent final : public virtual AbstractAgent, public virtual AbstractTorchAgent { - public: - TorchSacAgent( - const std::shared_ptr &actor, torch::Device device, - std::optional> collector = std::nullopt); - - TorchAction act(const TorchState &state, bool sample) override; - - std::vector - act(const std::vector &states, int vision_height, int vision_width) override; - void load(const std::filesystem::path &agent_folder) override; - - private: - std::shared_ptr actor; - std::optional> collector; - }; - -}// namespace arenai::agent - -#endif//ARENAI_AGENT_HOST_SAC_H diff --git a/arenai_agent/src/agents/sac/sac_collector.cpp b/arenai_agent/src/agents/sac/sac_collector.cpp deleted file mode 100644 index e8ef7c12..00000000 --- a/arenai_agent/src/agents/sac/sac_collector.cpp +++ /dev/null @@ -1,29 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#include "./sac_collector.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - SacStepCollector::SacStepCollector(std::shared_ptr replay_buffer) - : replay_buffer(std::move(replay_buffer)) {} - - void SacStepCollector::on_act(const TorchState &state, const TorchAction &action) { - last_state = state; - last_action = action; - } - - void SacStepCollector::on_transition(const torch::Tensor &rewards, const torch::Tensor &done) { - replay_buffer->add( - {.state = last_state, .action = last_action, .reward = rewards, .done = done}); - } - - void SacStepCollector::on_episode_end(const TorchState &final_state) { - replay_buffer->finish_episode(final_state); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_collector.h b/arenai_agent/src/agents/sac/sac_collector.h deleted file mode 100644 index c93bd6e1..00000000 --- a/arenai_agent/src/agents/sac/sac_collector.h +++ /dev/null @@ -1,36 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#ifndef ARENAI_SAC_COLLECTOR_H -#define ARENAI_SAC_COLLECTOR_H - -#include - -#include "../step_collector.h" -#include "../torch_types.h" -#include "./sac_replay_buffer.h" - -namespace arenai::agent { - - class SacStepCollector final : public AbstractStepCollector { - public: - explicit SacStepCollector(std::shared_ptr replay_buffer); - - // concrete act-time channel, called by TrainableSacAgent::act - void on_act(const TorchState &state, const TorchAction &action); - - void on_transition(const torch::Tensor &rewards, const torch::Tensor &done) override; - - void on_episode_end(const TorchState &final_state) override; - - private: - std::shared_ptr replay_buffer; - - TorchState last_state; - TorchAction last_action; - }; - -}// namespace arenai::agent - -#endif//ARENAI_SAC_COLLECTOR_H diff --git a/arenai_agent/src/agents/sac/sac_factory.cpp b/arenai_agent/src/agents/sac/sac_factory.cpp deleted file mode 100644 index 5f5049c1..00000000 --- a/arenai_agent/src/agents/sac/sac_factory.cpp +++ /dev/null @@ -1,43 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#include "./sac_factory.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - SacTorchAgentFactory::SacTorchAgentFactory( - const int vision_height, const int vision_width, const int nb_sensors, - const int nb_continuous_actions, const int nb_discrete_actions, const torch::Device device, - const SacHyperParams ¶ms) - : config(cli_fields_to_map(sac_cli_fields(), params)), - actor(std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, - params.hidden_size_sensors, params.actor_hidden_sizes, params.vision_channels, - params.group_norm_nums, params.initial_sigma, params.initial_fire_proba)), - replay_buffer(std::make_shared(params.replay_buffer_size)), - collector(std::make_shared(replay_buffer)), - agent(std::make_shared(actor, device, collector)), - trainer(std::make_shared( - actor, replay_buffer, vision_height, vision_width, nb_sensors, nb_continuous_actions, - nb_discrete_actions, params.actor_learning_rate, params.critic_learning_rate, - params.alpha_learning_rate, params.hidden_size_sensors, params.hidden_size_actions, - params.critic_hidden_sizes, params.vision_channels, params.group_norm_nums, device, - params.metric_window_size, params.tau, params.gamma, params.train_every, - params.epochs, params.batch_size, params.continuous_target_entropy, - params.discrete_target_entropy_factor)) {} - - std::shared_ptr SacTorchAgentFactory::get_agent() { return agent; } - - std::shared_ptr SacTorchAgentFactory::get_collector() { - return collector; - } - - std::shared_ptr SacTorchAgentFactory::get_trainer() { return trainer; } - - std::map SacTorchAgentFactory::get_config() const { return config; } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_factory.h b/arenai_agent/src/agents/sac/sac_factory.h deleted file mode 100644 index d45ff3bf..00000000 --- a/arenai_agent/src/agents/sac/sac_factory.h +++ /dev/null @@ -1,43 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#ifndef ARENAI_SAC_FACTORY_H -#define ARENAI_SAC_FACTORY_H - -#include "../torch_factory.h" -#include "./sac_agent.h" -#include "./sac_collector.h" -#include "./sac_hyperparams.h" -#include "./sac_replay_buffer.h" -#include "./sac_trainer.h" - -namespace arenai::agent { - - class SacTorchAgentFactory final : public AbstractTorchAgentFactory { - public: - SacTorchAgentFactory( - int vision_height, int vision_width, int nb_sensors, int nb_continuous_actions, - int nb_discrete_actions, torch::Device device, const SacHyperParams ¶ms); - - std::shared_ptr get_agent() override; - std::shared_ptr get_collector() override; - std::shared_ptr get_trainer() override; - - std::map get_config() const override; - - private: - std::map config; - - // triad built once, sharing actor + replay_buffer - std::shared_ptr actor; - - std::shared_ptr replay_buffer; - std::shared_ptr collector; - std::shared_ptr agent; - std::shared_ptr trainer; - }; - -}// namespace arenai::agent - -#endif//ARENAI_SAC_FACTORY_H diff --git a/arenai_agent/src/agents/sac/sac_hyperparams.cpp b/arenai_agent/src/agents/sac/sac_hyperparams.cpp deleted file mode 100644 index 8ca658ae..00000000 --- a/arenai_agent/src/agents/sac/sac_hyperparams.cpp +++ /dev/null @@ -1,36 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#include "./sac_hyperparams.h" - -namespace arenai::agent { - - std::vector> sac_cli_fields() { - return { - {.name = "--actor_learning_rate", .member = &SacHyperParams::actor_learning_rate}, - {.name = "--critic_learning_rate", .member = &SacHyperParams::critic_learning_rate}, - {.name = "--alpha_learning_rate", .member = &SacHyperParams::alpha_learning_rate}, - {.name = "--hidden_size_sensors", .member = &SacHyperParams::hidden_size_sensors}, - {.name = "--hidden_size_actions", .member = &SacHyperParams::hidden_size_actions}, - {.name = "--actor_hidden_sizes", .member = &SacHyperParams::actor_hidden_sizes}, - {.name = "--critic_hidden_sizes", .member = &SacHyperParams::critic_hidden_sizes}, - {.name = "--vision_channels", .member = &SacHyperParams::vision_channels}, - {.name = "--group_norm_nums", .member = &SacHyperParams::group_norm_nums}, - {.name = "--initial_sigma", .member = &SacHyperParams::initial_sigma}, - {.name = "--initial_fire_proba", .member = &SacHyperParams::initial_fire_proba}, - {.name = "--continuous_target_entropy", - .member = &SacHyperParams::continuous_target_entropy}, - {.name = "--discrete_target_entropy_factor", - .member = &SacHyperParams::discrete_target_entropy_factor}, - {.name = "--metric_window_size", .member = &SacHyperParams::metric_window_size}, - {.name = "--tau", .member = &SacHyperParams::tau}, - {.name = "--gamma", .member = &SacHyperParams::gamma}, - {.name = "--replay_buffer_size", .member = &SacHyperParams::replay_buffer_size}, - {.name = "--train_every", .member = &SacHyperParams::train_every}, - {.name = "--epochs", .member = &SacHyperParams::epochs}, - {.name = "--batch_size", .member = &SacHyperParams::batch_size}, - }; - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_hyperparams.h b/arenai_agent/src/agents/sac/sac_hyperparams.h deleted file mode 100644 index acf6fe30..00000000 --- a/arenai_agent/src/agents/sac/sac_hyperparams.h +++ /dev/null @@ -1,44 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#ifndef ARENAI_SAC_HYPERPARAMS_H -#define ARENAI_SAC_HYPERPARAMS_H - -#include -#include - -#include "../../utils/cli_fields.h" - -namespace arenai::agent { - - // Member initializers are the CLI defaults (single source of truth). - struct SacHyperParams { - float actor_learning_rate = 1e-4f; - float critic_learning_rate = 3e-4f; - float alpha_learning_rate = 3e-5f; - int hidden_size_sensors = 128; - int hidden_size_actions = 32; - std::vector actor_hidden_sizes = {1024, 512}; - std::vector critic_hidden_sizes = {1024, 512}; - std::vector> vision_channels = {{3, 8}, {8, 16}, {16, 24}, - {24, 32}, {32, 48}, {48, 64}}; - std::vector group_norm_nums = {1, 2, 3, 4, 6, 8}; - float initial_sigma = 0.5f; - float initial_fire_proba = 0.5f; - float continuous_target_entropy = -1.f; - float discrete_target_entropy_factor = 0.3f; - int metric_window_size = 256; - float tau = 0.005f; - float gamma = 0.997f; - int replay_buffer_size = 300000; - int train_every = 256; - int epochs = 128; - int batch_size = 256; - }; - - std::vector> sac_cli_fields(); - -}// namespace arenai::agent - -#endif//ARENAI_SAC_HYPERPARAMS_H diff --git a/arenai_agent/src/agents/sac/sac_replay_buffer.cpp b/arenai_agent/src/agents/sac/sac_replay_buffer.cpp deleted file mode 100644 index 0a4b0942..00000000 --- a/arenai_agent/src/agents/sac/sac_replay_buffer.cpp +++ /dev/null @@ -1,134 +0,0 @@ -// -// Created by samuel on 03/10/2025. -// - -#include "./sac_replay_buffer.h" - -#include - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - SacReplayBuffer::SacReplayBuffer(const int memory_size) - : initialized_(false), memory_size_(memory_size), nb_steps_(0), write_idx_(0), size_(0), - nb_tanks_(0) {} - - void SacReplayBuffer::initialize(const SacInputStep &first_step) { - constexpr auto cpu = torch::kCPU; - - nb_tanks_ = first_step.state.vision.size(0); - nb_steps_ = memory_size_ / static_cast(nb_tanks_); - TORCH_CHECK( - nb_steps_ > 0, "SacReplayBuffer: memory_size (", memory_size_, - ") must be >= nb_tanks (", nb_tanks_, ")"); - - const auto mem = static_cast(nb_steps_); - - const auto make_storage = [&](const torch::Tensor &ref) { - auto sizes = ref.sizes().vec(); - sizes.insert(sizes.begin(), mem); - return torch::empty(sizes, ref.options().device(cpu).requires_grad(false)); - }; - - store_vision_ = make_storage(first_step.state.vision); - store_proprioception_ = make_storage(first_step.state.proprioception); - store_cont_action_ = make_storage(first_step.action.continuous_action); - store_disc_action_ = make_storage(first_step.action.discrete_action); - store_reward_ = make_storage(first_step.reward); - store_done_ = make_storage(first_step.done); - - const auto bool_cpu = torch::TensorOptions().dtype(torch::kBool).device(cpu); - store_sampleable_ = torch::zeros({mem, nb_tanks_}, bool_cpu); - already_terminated_ = torch::zeros({nb_tanks_}, bool_cpu); - - initialized_ = true; - } - - void SacReplayBuffer::add(const SacInputStep &step) { - if (!initialized_) initialize(step); - - const auto idx = static_cast(write_idx_); - - const auto done_bool = step.done.to(torch::kBool); - - store_vision_[idx].copy_(step.state.vision); - store_proprioception_[idx].copy_(step.state.proprioception); - store_cont_action_[idx].copy_(step.action.continuous_action.detach()); - store_disc_action_[idx].copy_(step.action.discrete_action.detach()); - store_reward_[idx].copy_(step.reward); - store_done_[idx].copy_(done_bool); - - // tanks already terminated before this step have no valid transition to store - store_sampleable_[idx].copy_(already_terminated_.logical_not()); - already_terminated_.logical_or_(done_bool.reshape({nb_tanks_})); - - advance_write_idx(); - } - - void SacReplayBuffer::finish_episode(const TorchState &final_step) { - if (!initialized_) return; - - // the final observation is only read as next_state of the episode's last transitions: - // stored in the ring but never sampled as a starting state - const auto idx = static_cast(write_idx_); - - store_vision_[idx].copy_(final_step.vision); - store_proprioception_[idx].copy_(final_step.proprioception); - store_cont_action_[idx].zero_(); - store_disc_action_[idx].zero_(); - store_reward_[idx].zero_(); - store_done_[idx].zero_(); - - store_sampleable_[idx].fill_(false); - already_terminated_.fill_(false); - - advance_write_idx(); - } - - void SacReplayBuffer::advance_write_idx() { - write_idx_ = (write_idx_ + 1) % nb_steps_; - if (size_ < nb_steps_) size_++; - } - - SacTrainStep SacReplayBuffer::sample(int batch_size, const torch::Device device) const { - // valid starting pairs: sampleable, and their following slot is already written - const auto valid = store_sampleable_.clone(); - const auto last_written = static_cast((write_idx_ + nb_steps_ - 1) % nb_steps_); - valid[last_written] = false; - - const auto valid_pairs = torch::nonzero(valid); - const auto nb_transitions = valid_pairs.size(0); - TORCH_CHECK(nb_transitions > 0, "SacReplayBuffer: no transition to sample"); - - batch_size = std::max(1, std::min(batch_size, static_cast(nb_transitions))); - - const auto pick = torch::randint( - nb_transitions, {batch_size}, torch::TensorOptions().dtype(torch::kInt64)); - const auto chosen = valid_pairs.index_select(0, pick); - const auto step_idx = chosen.select(1, 0); - const auto tank_idx = chosen.select(1, 1); - const auto next_idx = (step_idx + 1).remainder(static_cast(nb_steps_)); - - const auto take = [&](const torch::Tensor &store, const torch::Tensor &steps) { - return store.index({steps, tank_idx}).to(device); - }; - - return { - .state = - {.vision = take(store_vision_, step_idx), - .proprioception = take(store_proprioception_, step_idx)}, - .action = - {.continuous_action = take(store_cont_action_, step_idx), - .discrete_action = take(store_disc_action_, step_idx)}, - .reward = take(store_reward_, step_idx), - .done = take(store_done_, step_idx), - .next_state = { - .vision = take(store_vision_, next_idx), - .proprioception = take(store_proprioception_, next_idx)}}; - } - - size_t SacReplayBuffer::size() const { return size_; } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_replay_buffer.h b/arenai_agent/src/agents/sac/sac_replay_buffer.h deleted file mode 100644 index 9f86e96c..00000000 --- a/arenai_agent/src/agents/sac/sac_replay_buffer.h +++ /dev/null @@ -1,71 +0,0 @@ -// -// Created by samuel on 03/10/2025. -// - -#ifndef ARENAI_AGENT_HOST_REPLAY_BUFFER_H -#define ARENAI_AGENT_HOST_REPLAY_BUFFER_H - -#include - -#include "../torch_types.h" - -namespace arenai::agent { - - struct SacInputStep { - TorchState state; - TorchAction action; - torch::Tensor reward; - torch::Tensor done; - }; - - struct SacTrainStep { - TorchState state; - TorchAction action; - torch::Tensor reward; - torch::Tensor done; - TorchState next_state; - }; - - class SacReplayBuffer { - public: - virtual ~SacReplayBuffer() = default; - - explicit SacReplayBuffer(int memory_size); - - SacTrainStep sample(int batch_size, torch::Device device) const; - - void add(const SacInputStep &step); - void finish_episode(const TorchState &final_step); - - size_t size() const; - - private: - void initialize(const SacInputStep &first_step); - void advance_write_idx(); - - bool initialized_; - - // total transition budget: the ring holds memory_size_ / nb_tanks_ steps - size_t memory_size_; - size_t nb_steps_; - size_t write_idx_; - size_t size_; - - int64_t nb_tanks_; - - torch::Tensor store_vision_; - torch::Tensor store_proprioception_; - torch::Tensor store_cont_action_; - torch::Tensor store_disc_action_; - torch::Tensor store_reward_; - torch::Tensor store_done_; - - // [mem, nb_tanks] whether the (step, tank) pair can start a sampled transition - torch::Tensor store_sampleable_; - // [nb_tanks] tanks already done in the current episode - torch::Tensor already_terminated_; - }; - -}// namespace arenai::agent - -#endif// ARENAI_AGENT_HOST_REPLAY_BUFFER_H diff --git a/arenai_agent/src/agents/sac/sac_trainer.cpp b/arenai_agent/src/agents/sac/sac_trainer.cpp deleted file mode 100644 index 4a4d8d98..00000000 --- a/arenai_agent/src/agents/sac/sac_trainer.cpp +++ /dev/null @@ -1,328 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#include "./sac_trainer.h" - -#include - -#include "../../distributions/multinomial.h" -#include "../../distributions/truncated_normal.h" -#include "../../metrics/last_metric.h" -#include "../../metrics/mean_metric.h" -#include "../../metrics/std_metric.h" -#include "../../networks/constants.h" -#include "../../networks_utils/print_module.h" -#include "../../networks_utils/target_update.h" -#include "../../networks_utils/torch_loader.h" -#include "../../networks_utils/torch_saver.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - SacTrainer::SacTrainer( - std::shared_ptr actor, std::shared_ptr replay_buffer, - const int vision_height, const int vision_width, const int nb_sensors, - const int nb_continuous_actions, const int nb_discrete_actions, - const float actor_learning_rate, const float critic_learning_rate, - const float alpha_learning_rate, const int hidden_size_sensors, - const int hidden_size_actions, const std::vector &critic_hidden_sizes, - const std::vector> &vision_channels, - const std::vector &group_norm_nums, const torch::Device device, - const int metric_window_size, const float tau, const float gamma, const int train_every, - const int epochs, const int batch_size, const float continuous_target_entropy, - const float discrete_target_entropy_factor) - : actor(std::move(actor)), replay_buffer(std::move(replay_buffer)), - critic_1(std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, - hidden_size_sensors, hidden_size_actions, critic_hidden_sizes, vision_channels, - group_norm_nums)), - critic_2(std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, - hidden_size_sensors, hidden_size_actions, critic_hidden_sizes, vision_channels, - group_norm_nums)), - target_critic_1(std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, - hidden_size_sensors, hidden_size_actions, critic_hidden_sizes, vision_channels, - group_norm_nums)), - target_critic_2(std::make_shared( - vision_height, vision_width, nb_sensors, nb_continuous_actions, nb_discrete_actions, - hidden_size_sensors, hidden_size_actions, critic_hidden_sizes, vision_channels, - group_norm_nums)), - alpha_continuous(std::make_shared(5e-2f, nb_continuous_actions)), - alpha_discrete(std::make_shared(5e-2f, 1)), - continuous_target_entropy( - std::make_unique(continuous_target_entropy)), - discrete_target_entropy(std::make_unique( - discrete_target_entropy_factor * multinomial_maximum_entropy(nb_discrete_actions))), - actor_optim(std::make_unique( - this->actor->parameters(), torch::optim::AdamOptions(actor_learning_rate))), - critic_1_optim(std::make_unique( - critic_1->parameters(), torch::optim::AdamOptions(critic_learning_rate))), - critic_2_optim(std::make_unique( - critic_2->parameters(), torch::optim::AdamOptions(critic_learning_rate))), - alpha_continuous_optim(std::make_unique( - alpha_continuous->parameters(), torch::optim::AdamOptions(alpha_learning_rate))), - alpha_discrete_optim(std::make_unique( - alpha_discrete->parameters(), torch::optim::AdamOptions(alpha_learning_rate))), - actor_mean_loss_metric(std::make_shared("π_μ", metric_window_size)), - actor_std_loss_metric(std::make_shared("π_σ", metric_window_size)), - critic_1_mean_loss_metric(std::make_shared("q1_μ", metric_window_size)), - critic_1_std_loss_metric(std::make_shared("q1_σ", metric_window_size)), - critic_2_mean_loss_metric(std::make_shared("q2_μ", metric_window_size)), - critic_2_std_loss_metric(std::make_shared("q2_σ", metric_window_size)), - continuous_entropy_metric(std::make_shared("Hc", metric_window_size)), - discrete_entropy_metric(std::make_shared("Hd", metric_window_size)), - alpha_continuous_metric(std::make_shared("α_c", metric_window_size, 2, true)), - alpha_discrete_metric(std::make_shared("α_d", metric_window_size, 2, true)), - continuous_target_entropy_metric(std::make_shared("Hc_t")), - discrete_target_entropy_metric(std::make_shared("Hd_t")), tau(tau), - gamma(gamma), train_every(train_every), train_counter(0), epochs(epochs), - batch_size(batch_size) { - - hard_update(target_critic_1, critic_1); - hard_update(target_critic_2, critic_2); - - to(device); - - set_train(false); - } - - void SacTrainer::step() { - if (train_counter == train_every - 1) train(); - train_counter = (train_counter + 1) % train_every; - } - - void SacTrainer::train() const { - - set_train(true); - - for (int e = 0; e < epochs; e++) { - const auto [state, action, reward, done, next_state] = - replay_buffer->sample(batch_size, actor->parameters().back().device()); - - torch::Tensor target_q_values; - { - torch::NoGradGuard no_grad; - - const auto [next_mu, next_sigma, next_discrete_proba] = - actor->act(next_state.vision, next_state.proprioception); - - const auto next_continuous_action = truncated_normal_sample(next_mu, next_sigma); - const auto next_continuous_entropy = truncated_normal_entropy(next_mu, next_sigma); - - const auto next_discrete_entropy = multinomial_entropy(next_discrete_proba); - - const auto next_target_q_values_1 = target_critic_1->value_per_discrete_action( - next_state.vision, next_state.proprioception, next_continuous_action); - const auto next_target_q_values_2 = target_critic_2->value_per_discrete_action( - next_state.vision, next_state.proprioception, next_continuous_action); - - const auto next_min_q_value = torch::sum( - next_discrete_proba - * torch::min(next_target_q_values_1, next_target_q_values_2), - -1, true); - - const auto target_v_value = - next_min_q_value - + torch::sum(alpha_continuous->alpha() * next_continuous_entropy, -1, true) - + torch::sum(alpha_discrete->alpha() * next_discrete_entropy, -1, true); - - target_q_values = reward + (1.f - done.to(torch::kFloat)) * gamma * target_v_value; - } - - // critic 1 - const auto q_value_1 = critic_1->value_ohe( - state.vision, state.proprioception, action.continuous_action, - action.discrete_action); - const auto critic_1_loss = - torch::mse_loss(q_value_1, target_q_values, at::Reduction::Mean); - - critic_1_optim->zero_grad(); - critic_1_loss.backward(); - torch::nn::utils::clip_grad_norm_(critic_1->parameters(), GRAD_NORM_MAX); - critic_1_optim->step(); - - // critic 2 - const auto q_value_2 = critic_2->value_ohe( - state.vision, state.proprioception, action.continuous_action, - action.discrete_action); - const auto critic_2_loss = - torch::mse_loss(q_value_2, target_q_values, at::Reduction::Mean); - - critic_2_optim->zero_grad(); - critic_2_loss.backward(); - torch::nn::utils::clip_grad_norm_(critic_2->parameters(), GRAD_NORM_MAX); - critic_2_optim->step(); - - // target value soft update - soft_update(target_critic_1, critic_1, tau); - soft_update(target_critic_2, critic_2, tau); - - // policy - const auto [curr_mu, curr_sigma, curr_discrete_proba] = - actor->act(state.vision, state.proprioception); - - const auto curr_continuous_action = truncated_normal_sample(curr_mu, curr_sigma); - const auto curr_continuous_entropy = truncated_normal_entropy(curr_mu, curr_sigma); - - const auto curr_discrete_entropy = multinomial_entropy(curr_discrete_proba); - - const auto curr_q_values_1 = critic_1->value_per_discrete_action( - state.vision, state.proprioception, curr_continuous_action); - const auto curr_q_values_2 = critic_2->value_per_discrete_action( - state.vision, state.proprioception, curr_continuous_action); - const auto q_value = torch::sum( - curr_discrete_proba * torch::min(curr_q_values_1, curr_q_values_2), -1, true); - - const auto actor_loss = -torch::mean( - torch::sum(alpha_continuous->alpha().detach() * curr_continuous_entropy, -1, true) - + torch::sum(alpha_discrete->alpha().detach() * curr_discrete_entropy, -1, true) - + q_value); - - actor_optim->zero_grad(); - actor_loss.backward(); - torch::nn::utils::clip_grad_norm_(actor->parameters(), GRAD_NORM_MAX); - actor_optim->step(); - - // continuous entropy - const auto alpha_continuous_loss = - torch::sum( - alpha_continuous->log_alpha() - * torch::detach( - curr_continuous_entropy - continuous_target_entropy->target_entropy()), - -1) - .mean(); - - alpha_continuous_optim->zero_grad(); - alpha_continuous_loss.backward(); - alpha_continuous_optim->step(); - - // discrete entropy - const auto alpha_discrete_loss = - torch::sum( - alpha_discrete->log_alpha() - * torch::detach( - curr_discrete_entropy - discrete_target_entropy->target_entropy()), - -1) - .mean(); - - alpha_discrete_optim->zero_grad(); - alpha_discrete_loss.backward(); - alpha_discrete_optim->step(); - - // metrics - actor_mean_loss_metric->add(actor_loss.cpu().item()); - actor_std_loss_metric->add(actor_loss.cpu().item()); - - continuous_entropy_metric->add(curr_continuous_entropy.mean().item()); - discrete_entropy_metric->add(curr_discrete_entropy.mean().item()); - - continuous_target_entropy_metric->add( - continuous_target_entropy->target_entropy().mean().item()); - discrete_target_entropy_metric->add( - discrete_target_entropy->target_entropy().mean().item()); - - critic_1_mean_loss_metric->add(critic_1_loss.cpu().item()); - critic_1_std_loss_metric->add(critic_1_loss.cpu().item()); - critic_2_mean_loss_metric->add(critic_2_loss.cpu().item()); - critic_2_std_loss_metric->add(critic_2_loss.cpu().item()); - - alpha_continuous_metric->add(alpha_continuous->alpha().mean().item()); - alpha_discrete_metric->add(alpha_discrete->alpha().mean().item()); - } - - continuous_target_entropy->step(train_every); - discrete_target_entropy->step(train_every); - - set_train(false); - } - - std::vector> SacTrainer::get_metrics() { - return { - actor_mean_loss_metric, actor_std_loss_metric, critic_1_mean_loss_metric, - critic_1_std_loss_metric, critic_2_mean_loss_metric, critic_2_std_loss_metric, - continuous_target_entropy_metric, continuous_entropy_metric, alpha_continuous_metric, - discrete_target_entropy_metric, discrete_entropy_metric, alpha_discrete_metric}; - } - - void SacTrainer::save(const std::filesystem::path &output_folder) { - // Models - save_torch(output_folder, actor, "actor.pt"); - - save_torch(output_folder, critic_1, "critic_1.pt"); - save_torch(output_folder, critic_2, "critic_2.pt"); - - save_torch(output_folder, target_critic_1, "target_critic_1.pt"); - save_torch(output_folder, target_critic_2, "target_critic_2.pt"); - - save_torch(output_folder, alpha_continuous, "alpha_continuous.pt"); - save_torch(output_folder, alpha_discrete, "alpha_discrete.pt"); - - // Optimizers - save_torch(output_folder, actor_optim, "actor_optim.pt"); - - save_torch(output_folder, critic_1_optim, "critic_1_optim.pt"); - save_torch(output_folder, critic_2_optim, "critic_2_optim.pt"); - - save_torch(output_folder, alpha_continuous_optim, "alpha_continuous_optim.pt"); - save_torch(output_folder, alpha_discrete_optim, "alpha_discrete_optim.pt"); - - // string repr - std::ostringstream actor_repr_oss; - dump_module_tree(actor, actor_repr_oss, 0, "actor"); - std::ofstream actor_repr_file(output_folder / "actor_repr.txt"); - actor_repr_file << actor_repr_oss.str(); - actor_repr_file.close(); - - std::ostringstream critic_repr_oss; - dump_module_tree(critic_1, critic_repr_oss, 0, "critic"); - std::ofstream critic_repr_file(output_folder / "critic_repr.txt"); - critic_repr_file << critic_repr_oss.str(); - critic_repr_file.close(); - } - - void SacTrainer::set_train(const bool train) const { - actor->train(train); - - critic_1->train(train); - critic_2->train(train); - - alpha_continuous->train(train); - alpha_discrete->train(train); - - continuous_target_entropy->train(train); - discrete_target_entropy->train(train); - - // force eval for target critics - target_critic_1->train(false); - target_critic_2->train(false); - } - - void SacTrainer::to(const torch::Device device) const { - actor->to(device); - - critic_1->to(device); - critic_2->to(device); - - target_critic_1->to(device); - target_critic_2->to(device); - - alpha_continuous->to(device); - alpha_discrete->to(device); - - continuous_target_entropy->to(device); - discrete_target_entropy->to(device); - } - - int SacTrainer::count_parameters() { - return count_parameters_impl(actor->parameters()) - + count_parameters_impl(critic_1->parameters()) - + count_parameters_impl(critic_2->parameters()) - + count_parameters_impl(alpha_continuous->parameters()) - + count_parameters_impl(alpha_discrete->parameters()); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/agents/sac/sac_trainer.h b/arenai_agent/src/agents/sac/sac_trainer.h deleted file mode 100644 index ef4953c9..00000000 --- a/arenai_agent/src/agents/sac/sac_trainer.h +++ /dev/null @@ -1,97 +0,0 @@ -// -// Created by claude on 22/07/2026. -// - -#ifndef ARENAI_SAC_TRAINER_H -#define ARENAI_SAC_TRAINER_H - -#include "../../networks/actor.h" -#include "../../networks/entropy.h" -#include "../../networks/q_function.h" -#include "../trainer.h" -#include "./sac_replay_buffer.h" - -namespace arenai::agent { - - class SacTrainer final : public AbstractTrainer { - public: - SacTrainer( - std::shared_ptr actor, std::shared_ptr replay_buffer, - int vision_height, int vision_width, int nb_sensors, int nb_continuous_actions, - int nb_discrete_actions, float actor_learning_rate, float critic_learning_rate, - float alpha_learning_rate, int hidden_size_sensors, int hidden_size_actions, - const std::vector &critic_hidden_sizes, - const std::vector> &vision_channels, - const std::vector &group_norm_nums, torch::Device device, int metric_window_size, - float tau, float gamma, int train_every, int epochs, int batch_size, - float continuous_target_entropy, float discrete_target_entropy_factor); - - void step() override; - - std::vector> get_metrics() override; - - void save(const std::filesystem::path &output_folder) override; - - int count_parameters() override; - - private: - static constexpr double GRAD_NORM_MAX = 1.0; - - std::shared_ptr actor; - std::shared_ptr replay_buffer; - - std::shared_ptr critic_1; - std::shared_ptr critic_2; - - std::shared_ptr target_critic_1; - std::shared_ptr target_critic_2; - - std::shared_ptr alpha_continuous; - std::shared_ptr alpha_discrete; - - std::shared_ptr continuous_target_entropy; - std::shared_ptr discrete_target_entropy; - - std::shared_ptr actor_optim; - std::shared_ptr critic_1_optim; - std::shared_ptr critic_2_optim; - - std::shared_ptr alpha_continuous_optim; - std::shared_ptr alpha_discrete_optim; - - std::shared_ptr actor_mean_loss_metric; - std::shared_ptr actor_std_loss_metric; - - std::shared_ptr critic_1_mean_loss_metric; - std::shared_ptr critic_1_std_loss_metric; - - std::shared_ptr critic_2_mean_loss_metric; - std::shared_ptr critic_2_std_loss_metric; - - std::shared_ptr continuous_entropy_metric; - std::shared_ptr discrete_entropy_metric; - - std::shared_ptr alpha_continuous_metric; - std::shared_ptr alpha_discrete_metric; - - std::shared_ptr continuous_target_entropy_metric; - std::shared_ptr discrete_target_entropy_metric; - - float tau; - float gamma; - - int train_every; - int train_counter; - - int epochs; - int batch_size; - - void train() const; - - void set_train(bool train) const; - void to(torch::Device device) const; - }; - -}// namespace arenai::agent - -#endif//ARENAI_SAC_TRAINER_H diff --git a/arenai_agent/src/agents/step_collector.h b/arenai_agent/src/agents/step_collector.h index 9ab8975e..74a1a2d5 100644 --- a/arenai_agent/src/agents/step_collector.h +++ b/arenai_agent/src/agents/step_collector.h @@ -16,7 +16,9 @@ namespace arenai::agent { public: virtual ~AbstractStepCollector() = default; - virtual void on_transition(const torch::Tensor &rewards, const torch::Tensor &done) = 0; + virtual void on_transition( + const torch::Tensor &rewards, const torch::Tensor &done, + const torch::Tensor &truncated) = 0; virtual void on_episode_end(const TorchState &final_state) = 0; }; diff --git a/arenai_agent/src/agents/torch_factory.h b/arenai_agent/src/agents/torch_factory.h index 68ba4daf..6f67c900 100644 --- a/arenai_agent/src/agents/torch_factory.h +++ b/arenai_agent/src/agents/torch_factory.h @@ -5,9 +5,9 @@ #ifndef ARENAI_TORCH_FACTORY_H #define ARENAI_TORCH_FACTORY_H -#include #include -#include + +#include #include "./step_collector.h" #include "./torch_agent.h" @@ -27,7 +27,7 @@ namespace arenai::agent { // the algorithm's resolved hyper-parameters, keyed by CLI option name: // dumped next to the metrics so a run stays identifiable afterwards - virtual std::map get_config() const = 0; + virtual nlohmann::json get_config() const = 0; }; }// namespace arenai::agent diff --git a/arenai_agent/src/agents/torch_types.h b/arenai_agent/src/agents/torch_types.h index 7b6dab83..d18b902e 100644 --- a/arenai_agent/src/agents/torch_types.h +++ b/arenai_agent/src/agents/torch_types.h @@ -22,6 +22,7 @@ namespace arenai::agent { TorchState states; torch::Tensor rewards; torch::Tensor is_done; + torch::Tensor is_truncated; }; }// namespace arenai::agent diff --git a/arenai_agent/src/core/spawn_curriculum.cpp b/arenai_agent/src/core/spawn_curriculum.cpp new file mode 100644 index 00000000..3a28d12a --- /dev/null +++ b/arenai_agent/src/core/spawn_curriculum.cpp @@ -0,0 +1,57 @@ +// +// Created by samuel on 05/09/2026. +// + +#include "./spawn_curriculum.h" + +#include + +namespace arenai::agent { + + SpawnCurriculum::SpawnCurriculum( + const float delta_progress, const float ratio_low, const float ratio_high, + const int probe_window, const float boundary_proba, const std::uint64_t seed) + : upper(0.f), delta_progress(delta_progress), ratio_low(ratio_low), ratio_high(ratio_high), + probe_window(probe_window), boundary_proba(boundary_proba), last_was_probe(false), + nb_probe_episodes(0), sum_fires(0), sum_hits(0), rng(seed) {} + + float SpawnCurriculum::sample_progress() { + std::uniform_real_distribution unif(0.f, 1.f); + + last_was_probe = unif(rng) < boundary_proba; + if (last_was_probe) return upper; + + return upper * unif(rng); + } + + void SpawnCurriculum::on_episode_end(const int nb_fires, const int nb_hits) { + if (!last_was_probe) return; + + nb_probe_episodes++; + sum_fires += nb_fires; + sum_hits += nb_hits; + + if (nb_probe_episodes >= probe_window) attempt_update(); + } + + void SpawnCurriculum::attempt_update() { + // a window without a single fire counts as failure: an agent that stopped + // firing must not see its task keep hardening + const float ratio = + sum_fires > 0 ? static_cast(sum_hits) / static_cast(sum_fires) : 0.f; + + if (ratio > ratio_high) upper += delta_progress; + else if (ratio < ratio_low) upper -= delta_progress; + + upper = std::clamp(upper, 0.f, 1.f); + + nb_probe_episodes = 0; + sum_fires = 0; + sum_hits = 0; + } + + float SpawnCurriculum::upper_bound() const { return upper; } + + bool SpawnCurriculum::is_probe() const { return last_was_probe; } + +}// namespace arenai::agent diff --git a/arenai_agent/src/core/spawn_curriculum.h b/arenai_agent/src/core/spawn_curriculum.h new file mode 100644 index 00000000..ea0db398 --- /dev/null +++ b/arenai_agent/src/core/spawn_curriculum.h @@ -0,0 +1,60 @@ +// +// Created by samuel on 05/09/2026. +// + +#ifndef ARENAI_AGENT_HOST_SPAWN_CURRICULUM_H +#define ARENAI_AGENT_HOST_SPAWN_CURRICULUM_H + +#include +#include + +namespace arenai::agent { + + // ADR-style adaptive curriculum on the spawn-zone size. The difficulty is a + // progress fraction in [0, 1] mapped by the caller onto [initial, final] spawn + // sizes. Each episode either probes the current upper bound (probability + // boundary_proba) or draws uniformly below it — the mixture keeps easy episodes + // in the training data so long-range practice never erases short-range aim. + // Only probe episodes feed the controller: every probe_window of them, the + // aggregated hit/fire ratio moves the bound up past ratio_high, down below + // ratio_low, and not at all in between (hysteresis). + class SpawnCurriculum { + public: + SpawnCurriculum( + float delta_progress, float ratio_low, float ratio_high, int probe_window, + float boundary_proba, std::uint64_t seed); + + // draws the difficulty of the next episode, in [0, 1] + float sample_progress(); + + // per-episode totals across every tank; ignored unless the last sampled + // episode was a probe + void on_episode_end(int nb_fires, int nb_hits); + + float upper_bound() const; + + // whether the last sampled episode probes the upper bound + bool is_probe() const; + + private: + void attempt_update(); + + float upper; + + float delta_progress; + float ratio_low; + float ratio_high; + int probe_window; + float boundary_proba; + + bool last_was_probe; + int nb_probe_episodes; + long sum_fires; + long sum_hits; + + std::mt19937 rng; + }; + +}// namespace arenai::agent + +#endif// ARENAI_AGENT_HOST_SPAWN_CURRICULUM_H diff --git a/arenai_agent/src/core/train_environment.cpp b/arenai_agent/src/core/train_environment.cpp index bd574fe1..af4bbf1b 100644 --- a/arenai_agent/src/core/train_environment.cpp +++ b/arenai_agent/src/core/train_environment.cpp @@ -26,13 +26,10 @@ namespace arenai::agent { const int vision_num_threads) : BaseTanksEnvironment( std::make_shared(android_assets_path), graphics_backend, - nb_tanks, wanted_frequency, vision_height, vision_width, vision_num_threads, false), - wanted_frequency(wanted_frequency), - max_frames_without_hit(static_cast(30.f / wanted_frequency)), - remaining_frames(nb_tanks, max_frames_without_hit), - nb_frames_added_when_hit(static_cast(3.f / wanted_frequency)), - nb_frames_added_when_kill(static_cast(15.f / wanted_frequency)), nb_tanks(nb_tanks), - nb_steps(0), done(nb_tanks, false), already_done(nb_tanks, false), + nb_tanks, wanted_frequency, vision_height, vision_width, vision_num_threads, false, + true), + wanted_frequency(wanted_frequency), nb_tanks(nb_tanks), nb_steps(0), + done(nb_tanks, false), already_done(nb_tanks, false), max_episode_steps(max_episode_steps), nb_hits_per_tanks(nb_tanks, 0), nb_kills_per_tanks(nb_tanks, 0), reward_metric(std::make_shared( "r", 4 * nb_tanks * max_episode_steps, 1, true)), @@ -45,10 +42,11 @@ namespace arenai::agent { miss_distance_metric(std::make_shared("miss", 1024 * nb_tanks, 1)), episode_step_mean_nb_metric(std::make_shared("s", 32, 1)), fire_metric(std::make_shared("fire", 256, 2)), - hit_metric(std::make_shared("hit", 256, 2, true)), - kill_metric(std::make_shared("kill", 16, 1)), nb_kills_episode(0) {} + hit_metric(std::make_shared("hit", 16, 2, true)), + kill_metric(std::make_shared("kill", 16, 1)), nb_kills_episode(0), + nb_fires_episode(0), nb_hits_episode(0) {} - std::vector> + std::vector> TrainTankEnvironment::step(const float time_delta, const std::vector &actions) { // tanks flagged done on a previous step already emitted their terminal transition: @@ -97,7 +95,14 @@ namespace arenai::agent { return is_suicide_result; }); - // fire / hit frequencies (per second, per tank that acted this step) + const auto is_timeout = apply_on_enemies>([&](const auto &factories) { + std::vector is_timeout_result; + is_timeout_result.reserve(nb_tanks); + for (const auto &factory: factories) is_timeout_result.push_back(factory->is_timeout()); + return is_timeout_result; + }); + + // fire frequency (per second, per tank that acted this step) int nb_acting = 0, nb_fires = 0, nb_hits = 0; for (int i = 0; i < nb_tanks; i++) { if (already_done[i]) continue; @@ -106,39 +111,22 @@ namespace arenai::agent { nb_hits += has_hit[i] ? 1 : 0; } - if (nb_acting > 0) { + if (nb_acting > 0) fire_metric->add( static_cast(nb_fires) / (static_cast(nb_acting) * wanted_frequency)); - hit_metric->add( - static_cast(nb_hits) / (static_cast(nb_acting) * wanted_frequency)); - } - - // step over tanks (remaining steps, hits and kills counters + detect and apply timeout + detect death) - for (int i = 0; i < step_result.size(); i++) { - remaining_frames[i]--; - - if (has_hit[i]) { - remaining_frames[i] += nb_frames_added_when_hit; - nb_hits_per_tanks[i] += 1; - } - if (has_kill[i]) { - remaining_frames[i] += nb_frames_added_when_kill; - nb_kills_per_tanks[i] += 1; - } - const auto &[state, reward, is_done] = step_result[i]; + nb_fires_episode += nb_fires; + nb_hits_episode += nb_hits; - // detect death (kill or suicide) - if (is_done) { - if (!already_done[i] && !is_suicide[i]) nb_kills_episode++; - done[i] = true; - } + // step over tanks (hits and kills counters) + for (int i = 0; i < step_result.size(); i++) { - // starving out (no hit for too long) is a real death: penalized and terminal - if (!done[i] && remaining_frames[i] <= 0) { - constexpr float timeout_penalty = 1.f; + if (has_hit[i]) { nb_hits_per_tanks[i] += 1; } + if (has_kill[i]) { nb_kills_per_tanks[i] += 1; } - step_result[i] = {state, reward - timeout_penalty, true}; + // detect death (kill, suicide or timeout) + if (const auto &[state, reward, is_done, is_truncated] = step_result[i]; is_done) { + if (!already_done[i] && !is_suicide[i] && !is_timeout[i]) nb_kills_episode++; done[i] = true; } } @@ -151,10 +139,11 @@ namespace arenai::agent { if (tanks_not_done_indexes.size() == 1) { const auto winner_index = tanks_not_done_indexes[0]; - const auto &[state, reward, is_done] = step_result[winner_index]; + const auto &[state, reward, is_done, is_truncated] = step_result[winner_index]; - constexpr float win_reward = 2.f; - step_result[winner_index] = {state, reward + win_reward, true}; + const float win_reward = nb_kills_per_tanks[winner_index] > 0 ? 2.f : 0.f; + // winning is a genuine termination, never a truncation + step_result[winner_index] = {state, reward + win_reward, true, false}; done[winner_index] = true; } @@ -190,10 +179,16 @@ namespace arenai::agent { void TrainTankEnvironment::on_reset_physics( const std::unique_ptr &engine) { - remaining_frames = std::vector(nb_tanks, max_frames_without_hit); - // close the previous episode's counter (skip the very first reset) - if (nb_steps > 0) kill_metric->add(static_cast(nb_kills_episode)); + // close the previous episode's counters (skip the very first reset) + if (nb_steps > 0) { + kill_metric->add(static_cast(nb_kills_episode)); + + // hit accuracy is undefined on an episode without a single fire + if (nb_fires_episode > 0) + hit_metric->add( + static_cast(nb_hits_episode) / static_cast(nb_fires_episode)); + } nb_kills_episode = 0; nb_steps = 0; @@ -203,8 +198,15 @@ namespace arenai::agent { nb_hits_per_tanks = std::vector(nb_tanks, 0); nb_kills_per_tanks = std::vector(nb_tanks, 0); + + nb_fires_episode = 0; + nb_hits_episode = 0; } + int TrainTankEnvironment::episode_nb_fires() const { return nb_fires_episode; } + + int TrainTankEnvironment::episode_nb_hits() const { return nb_hits_episode; } + bool TrainTankEnvironment::are_all_done() { return std::accumulate( done.begin(), done.end(), true, diff --git a/arenai_agent/src/core/train_environment.h b/arenai_agent/src/core/train_environment.h index 070c30fe..02478f6f 100644 --- a/arenai_agent/src/core/train_environment.h +++ b/arenai_agent/src/core/train_environment.h @@ -18,13 +18,17 @@ namespace arenai::agent { const std::filesystem::path &android_assets_path, float wanted_frequency, int max_episode_steps, int vision_height, int vision_width, int vision_num_threads); - std::vector> + std::vector> step(float time_delta, const std::vector &actions) override; std::vector> get_metrics() const; bool is_episode_terminated(); + // totals across every tank since the last reset, for the spawn curriculum + int episode_nb_fires() const; + int episode_nb_hits() const; + static void reset_singleton(); protected: @@ -38,10 +42,6 @@ namespace arenai::agent { private: float wanted_frequency; - int max_frames_without_hit; - std::vector remaining_frames; - int nb_frames_added_when_hit; - int nb_frames_added_when_kill; int nb_tanks; int nb_steps; @@ -71,6 +71,8 @@ namespace arenai::agent { std::shared_ptr kill_metric; int nb_kills_episode; + int nb_fires_episode; + int nb_hits_episode; bool are_all_done(); }; diff --git a/arenai_agent/src/distributions/bernoulli.cpp b/arenai_agent/src/distributions/bernoulli.cpp new file mode 100644 index 00000000..500d888d --- /dev/null +++ b/arenai_agent/src/distributions/bernoulli.cpp @@ -0,0 +1,51 @@ +// +// Created by samuel on 18/09/2026. +// + +#include "./bernoulli.h" + +#include + +#include "../networks/constants.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace { + + // float32: 1 - EPSILON (1e-8) rounds back to 1 and log(1 - p) would still + // hit -inf when the sigmoid saturates, so the upper clamp needs its own margin + constexpr float PROBA_MAX = 1.f - 1e-6f; + + torch::Tensor clamp_proba(const torch::Tensor &probabilities) { + return torch::clamp(probabilities, agent::EPSILON, PROBA_MAX); + } + +}// namespace + +namespace arenai::agent { + + torch::Tensor bernoulli_sample(const torch::Tensor &probabilities) { + return torch::bernoulli(clamp_proba(probabilities)); + } + + torch::Tensor bernoulli_max_action(const torch::Tensor &probabilities) { + return (probabilities > 0.5f).to(probabilities.dtype()); + } + + torch::Tensor + bernoulli_log_proba(const torch::Tensor &actions, const torch::Tensor &probabilities) { + const auto clamped_proba = clamp_proba(probabilities); + return actions * torch::log(clamped_proba) + + (1.f - actions) * torch::log(1.f - clamped_proba); + } + + torch::Tensor bernoulli_entropy(const torch::Tensor &probabilities) { + const auto clamped_proba = clamp_proba(probabilities); + return -clamped_proba * torch::log(clamped_proba) + - (1.f - clamped_proba) * torch::log(1.f - clamped_proba); + } + + float bernoulli_maximum_entropy() { return std::log(2.f); } + +}// namespace arenai::agent diff --git a/arenai_agent/src/distributions/bernoulli.h b/arenai_agent/src/distributions/bernoulli.h new file mode 100644 index 00000000..7d746240 --- /dev/null +++ b/arenai_agent/src/distributions/bernoulli.h @@ -0,0 +1,24 @@ +// +// Created by samuel on 18/09/2026. +// + +#ifndef ARENAI_AGENT_HOST_BERNOULLI_H +#define ARENAI_AGENT_HOST_BERNOULLI_H + +#include + +namespace arenai::agent { + + torch::Tensor bernoulli_sample(const torch::Tensor &probabilities); + torch::Tensor bernoulli_max_action(const torch::Tensor &probabilities); + + torch::Tensor + bernoulli_log_proba(const torch::Tensor &actions, const torch::Tensor &probabilities); + + torch::Tensor bernoulli_entropy(const torch::Tensor &probabilities); + + float bernoulli_maximum_entropy(); + +}// namespace arenai::agent + +#endif//ARENAI_AGENT_HOST_BERNOULLI_H diff --git a/arenai_agent/src/distributions/beta_law.cpp b/arenai_agent/src/distributions/beta_law.cpp index 96126655..d960d714 100644 --- a/arenai_agent/src/distributions/beta_law.cpp +++ b/arenai_agent/src/distributions/beta_law.cpp @@ -13,13 +13,24 @@ namespace arenai::agent { // the log-proba and entropy carry the log(2) change of scale from [0, 1] to [-1, 1] + torch::Tensor to_alpha(const torch::Tensor &mode, const torch::Tensor &concentration) { + return mode * (concentration - 2.0) + 1.0; + } + + torch::Tensor to_beta(const torch::Tensor &mode, const torch::Tensor &concentration) { + return (1.0 - mode) * (concentration - 2.0) + 1.0; + } + static torch::Tensor clamp_pos(const torch::Tensor &t) { return torch::clamp_min(t, EPSILON); } static torch::Tensor log_beta_function(const torch::Tensor &alpha, const torch::Tensor &beta) { return torch::lgamma(alpha) + torch::lgamma(beta) - torch::lgamma(alpha + beta); } - torch::Tensor beta_law_sample(const torch::Tensor &alpha, const torch::Tensor &beta) { + torch::Tensor beta_law_sample(const torch::Tensor &mode, const torch::Tensor &concentration) { + const auto alpha = to_alpha(mode, concentration); + const auto beta = to_beta(mode, concentration); + // Beta(α, β) = X / (X + Y) with X ~ Gamma(α, 1) and Y ~ Gamma(β, 1), // differentiable w.r.t. α and β through the implicit gradients of _standard_gamma const auto x = at::_standard_gamma(clamp_pos(alpha)); @@ -30,7 +41,10 @@ namespace arenai::agent { } torch::Tensor beta_law_log_proba( - const torch::Tensor &x, const torch::Tensor &alpha, const torch::Tensor &beta) { + const torch::Tensor &x, const torch::Tensor &mode, const torch::Tensor &concentration) { + const auto alpha = to_alpha(mode, concentration); + const auto beta = to_beta(mode, concentration); + const auto clamped_alpha = clamp_pos(alpha); const auto clamped_beta = clamp_pos(beta); @@ -41,7 +55,10 @@ namespace arenai::agent { - log_beta_function(clamped_alpha, clamped_beta) - std::log(2.0); } - torch::Tensor beta_law_entropy(const torch::Tensor &alpha, const torch::Tensor &beta) { + torch::Tensor beta_law_entropy(const torch::Tensor &mode, const torch::Tensor &concentration) { + const auto alpha = to_alpha(mode, concentration); + const auto beta = to_beta(mode, concentration); + const auto clamped_alpha = clamp_pos(alpha); const auto clamped_beta = clamp_pos(beta); @@ -52,15 +69,21 @@ namespace arenai::agent { + std::log(2.0); } - torch::Tensor beta_law_mean_action(const torch::Tensor &alpha, const torch::Tensor &beta) { + torch::Tensor + beta_law_mean_action(const torch::Tensor &mode, const torch::Tensor &concentration) { + const auto alpha = to_alpha(mode, concentration); + const auto beta = to_beta(mode, concentration); + const auto clamped_alpha = clamp_pos(alpha); const auto clamped_beta = clamp_pos(beta); return 2.f * clamped_alpha / (clamped_alpha + clamped_beta) - 1.f; } + torch::Tensor beta_law_mode_action(const torch::Tensor &mode) { return mode * 2.f - 1.f; } + float beta_law_target_entropy(const int &nb_actions) { - return beta_law_entropy(torch::tensor(1.f), torch::tensor(1.f)).item() + return beta_law_entropy(torch::tensor(0.5f), torch::tensor(2.f)).item() * static_cast(nb_actions); } diff --git a/arenai_agent/src/distributions/beta_law.h b/arenai_agent/src/distributions/beta_law.h index a4ea4716..3e92feff 100644 --- a/arenai_agent/src/distributions/beta_law.h +++ b/arenai_agent/src/distributions/beta_law.h @@ -11,12 +11,18 @@ namespace arenai::agent { // Beta distribution rescaled to the [-1, 1] action support - torch::Tensor beta_law_sample(const torch::Tensor &alpha, const torch::Tensor &beta); + torch::Tensor to_alpha(const torch::Tensor &mode, const torch::Tensor &concentration); + torch::Tensor to_beta(const torch::Tensor &mode, const torch::Tensor &concentration); + + torch::Tensor beta_law_sample(const torch::Tensor &mode, const torch::Tensor &concentration); torch::Tensor beta_law_log_proba( - const torch::Tensor &x, const torch::Tensor &alpha, const torch::Tensor &beta); - torch::Tensor beta_law_entropy(const torch::Tensor &alpha, const torch::Tensor &beta); + const torch::Tensor &x, const torch::Tensor &mode, const torch::Tensor &concentration); + torch::Tensor beta_law_entropy(const torch::Tensor &mode, const torch::Tensor &concentration); + + torch::Tensor + beta_law_mean_action(const torch::Tensor &mode, const torch::Tensor &concentration); - torch::Tensor beta_law_mean_action(const torch::Tensor &alpha, const torch::Tensor &beta); + torch::Tensor beta_law_mode_action(const torch::Tensor &mode); float beta_law_target_entropy(const int &nb_actions); diff --git a/arenai_agent/src/distributions/multinomial.cpp b/arenai_agent/src/distributions/multinomial.cpp deleted file mode 100644 index d7b80789..00000000 --- a/arenai_agent/src/distributions/multinomial.cpp +++ /dev/null @@ -1,43 +0,0 @@ -// -// Created by samuel on 22/02/2026. -// - -#include "./multinomial.h" - -#include "../networks/constants.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - torch::Tensor multinomial_sample(const torch::Tensor &probabilities) { - const auto clamped_proba = torch::clamp(probabilities, EPSILON, 1.0 - EPSILON); - const auto idx = torch::multinomial(clamped_proba, 1, false); - const auto one_hot = torch::zeros_like(clamped_proba).scatter_(1, idx, 1.0); - return one_hot; - } - - torch::Tensor multinomial_max_action(const torch::Tensor &probabilities) { - const auto clamped_proba = torch::clamp(probabilities, EPSILON, 1.0 - EPSILON); - const auto idx = torch::argmax(clamped_proba, 1, true); - const auto one_hot = torch::zeros_like(clamped_proba).scatter_(1, idx, 1.0); - return one_hot; - } - - torch::Tensor multinomial_entropy(const torch::Tensor &probabilities) { - const auto clamped_proba = torch::clamp(probabilities, EPSILON, 1.0 - EPSILON); - return -torch::sum(clamped_proba * torch::log(clamped_proba), -1, true); - } - - float multinomial_maximum_entropy(const int &nb_actions) { - return multinomial_entropy(torch::ones({nb_actions}) / static_cast(nb_actions)) - .item(); - } - - float multinomial_target_entropy(const float &shoot_probability) { - return multinomial_entropy(torch::tensor({shoot_probability, 1.f - shoot_probability})) - .item(); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/distributions/multinomial.h b/arenai_agent/src/distributions/multinomial.h deleted file mode 100644 index d4b94877..00000000 --- a/arenai_agent/src/distributions/multinomial.h +++ /dev/null @@ -1,23 +0,0 @@ -// -// Created by samuel on 22/02/2026. -// - -#ifndef ARENAI_AGENT_HOST_MULTINOMIAL_H -#define ARENAI_AGENT_HOST_MULTINOMIAL_H - -#include - -namespace arenai::agent { - - torch::Tensor multinomial_sample(const torch::Tensor &probabilities); - torch::Tensor multinomial_max_action(const torch::Tensor &probabilities); - - torch::Tensor multinomial_entropy(const torch::Tensor &probabilities); - - float multinomial_maximum_entropy(const int &nb_actions); - - float multinomial_target_entropy(const float &shoot_probability); - -}// namespace arenai::agent - -#endif//ARENAI_AGENT_HOST_MULTINOMIAL_H diff --git a/arenai_agent/src/main.cpp b/arenai_agent/src/main.cpp index 065de690..903c245a 100644 --- a/arenai_agent/src/main.cpp +++ b/arenai_agent/src/main.cpp @@ -34,10 +34,15 @@ int main(const int argc, char **argv) { parser.add_argument("--nb_tanks").scan<'i', int>().default_value(32); parser.add_argument("--vision_height").scan<'i', int>().default_value(128); parser.add_argument("--vision_width").scan<'i', int>().default_value(256); - parser.add_argument("--initial_spawn_width").scan<'g', float>().default_value(500.f); - parser.add_argument("--initial_spawn_height").scan<'g', float>().default_value(500.f); - parser.add_argument("--final_spawn_width").scan<'g', float>().default_value(500.f); - parser.add_argument("--final_spawn_height").scan<'g', float>().default_value(500.f); + parser.add_argument("--initial_spawn_width").scan<'g', float>().default_value(250.f); + parser.add_argument("--initial_spawn_height").scan<'g', float>().default_value(250.f); + parser.add_argument("--final_spawn_width").scan<'g', float>().default_value(1000.f); + parser.add_argument("--final_spawn_height").scan<'g', float>().default_value(1000.f); + parser.add_argument("--curriculum_delta").scan<'g', float>().default_value(50.f); + parser.add_argument("--curriculum_ratio_low").scan<'g', float>().default_value(0.03f); + parser.add_argument("--curriculum_ratio_high").scan<'g', float>().default_value(0.06f); + parser.add_argument("--curriculum_probe_window").scan<'i', int>().default_value(24); + parser.add_argument("--curriculum_boundary_proba").scan<'g', float>().default_value(0.2f); parser.add_argument("--vision_num_threads") .scan<'i', int>() .default_value(static_cast(std::thread::hardware_concurrency())); @@ -73,6 +78,11 @@ int main(const int argc, char **argv) { .initial_spawn_height = parser.get("--initial_spawn_height"), .final_spawn_width = parser.get("--final_spawn_width"), .final_spawn_height = parser.get("--final_spawn_height"), + .curriculum_delta = parser.get("--curriculum_delta"), + .curriculum_ratio_low = parser.get("--curriculum_ratio_low"), + .curriculum_ratio_high = parser.get("--curriculum_ratio_high"), + .curriculum_probe_window = parser.get("--curriculum_probe_window"), + .curriculum_boundary_proba = parser.get("--curriculum_boundary_proba"), .num_threads = parser.get("--vision_num_threads")}, {.output_folder = std::filesystem::path(parser.get("--output_folder")), .resources_folder = std::filesystem::path(parser.get("--resources_folder")), diff --git a/arenai_agent/src/networks/actor.cpp b/arenai_agent/src/networks/actor.cpp index 32e83bc9..53a7997a 100644 --- a/arenai_agent/src/networks/actor.cpp +++ b/arenai_agent/src/networks/actor.cpp @@ -19,7 +19,7 @@ namespace arenai::agent { const int &hidden_size_sensors, const std::vector &hidden_sizes, const std::vector> &vision_channels, const std::vector &group_norm_nums, const float &initial_sigma, - const float &initial_fire_proba) + const std::vector &initial_discrete_probas) : vision_encoder(register_module( "vision_encoder", std::make_shared( vision_height, vision_width, vision_channels, group_norm_nums))), @@ -31,18 +31,19 @@ namespace arenai::agent { torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_sensors})), torch::nn::SiLU()))), head(register_module("head", torch::nn::Sequential())), - mu(register_module( - "mu", torch::nn::Sequential( - torch::nn::Linear(hidden_sizes.back(), nb_continuous_actions), - torch::nn::Tanh()))), - sigma(register_module( - "sigma", torch::nn::Sequential( - torch::nn::Linear(hidden_sizes.back(), nb_continuous_actions), - std::make_shared(SIGMA_MIN, SIGMA_MAX)))), + mode(register_module( + "mode", torch::nn::Sequential( + torch::nn::Linear(hidden_sizes.back(), nb_continuous_actions), + torch::nn::Sigmoid()))), + concentration(register_module( + "concentration", + torch::nn::Sequential( + torch::nn::Linear(hidden_sizes.back(), nb_continuous_actions), + std::make_shared(CONCENTRATION_MIN, CONCENTRATION_MAX)))), discrete(register_module( "discrete", torch::nn::Sequential( torch::nn::Linear(hidden_sizes.back(), nb_discrete_actions), - torch::nn::Softmax(-1)))) { + torch::nn::Sigmoid()))) { head->push_back(torch::nn::Linear( torch::nn::LinearOptions( @@ -64,11 +65,13 @@ namespace arenai::agent { sensors_encoder->apply(init_hidden_weights); head->apply(init_hidden_weights); - mu->apply(init_mu_output_weights); - sigma->apply([initial_sigma](Module &m) { init_sigma_output_weights(m, initial_sigma); }); + // zero bias + sigmoid puts the initial mode at the action-range center + mode->apply(init_mu_output_weights); + concentration->apply( + [initial_sigma](Module &m) { init_concentration_output_weights(m, initial_sigma); }); - discrete->apply([initial_fire_proba](Module &m) { - init_discrete_output_weights(m, initial_fire_proba); + discrete->apply([&initial_discrete_probas](Module &m) { + init_discrete_output_weights(m, initial_discrete_probas); }); } @@ -77,8 +80,8 @@ namespace arenai::agent { auto sensors_encoded = sensors_encoder->forward(sensors); auto encoded = head->forward(torch::cat({vision_encoded, sensors_encoded}, 1)); return { - .mu = mu->forward(encoded), - .sigma = sigma->forward(encoded), + .mode = mode->forward(encoded), + .concentration = concentration->forward(encoded), .discrete = discrete->forward(encoded)}; } diff --git a/arenai_agent/src/networks/actor.h b/arenai_agent/src/networks/actor.h index 549eddcf..5257c221 100644 --- a/arenai_agent/src/networks/actor.h +++ b/arenai_agent/src/networks/actor.h @@ -6,6 +6,7 @@ #define ARENAI_AGENT_HOST_ACTOR_H #include +#include #include @@ -14,8 +15,8 @@ namespace arenai::agent { struct ActorRawOutput { - torch::Tensor mu; - torch::Tensor sigma; + torch::Tensor mode; + torch::Tensor concentration; torch::Tensor discrete; }; @@ -27,7 +28,7 @@ namespace arenai::agent { const int &hidden_size_sensors, const std::vector &hidden_sizes, const std::vector> &vision_channels, const std::vector &group_norm_nums, const float &initial_sigma, - const float &initial_fire_proba); + const std::vector &initial_discrete_probas); ActorRawOutput act(const torch::Tensor &vision, const torch::Tensor &sensors); private: @@ -36,8 +37,8 @@ namespace arenai::agent { torch::nn::Sequential head; - torch::nn::Sequential mu; - torch::nn::Sequential sigma; + torch::nn::Sequential mode; + torch::nn::Sequential concentration; torch::nn::Sequential discrete; }; diff --git a/arenai_agent/src/networks/constants.h b/arenai_agent/src/networks/constants.h index 2af7d3aa..5549665e 100644 --- a/arenai_agent/src/networks/constants.h +++ b/arenai_agent/src/networks/constants.h @@ -10,6 +10,11 @@ namespace arenai::agent { constexpr float SIGMA_MIN = 1e-4f; constexpr float SIGMA_MAX = 1.f; + + // Beta mode/concentration policy: κ ≥ 2 keeps the density unimodal (α, β ≥ 1); + // the floor sits just above the uniform, the ceiling caps the sharpness (σ ≈ 0.02) + constexpr float CONCENTRATION_MIN = 2.01f; + constexpr float CONCENTRATION_MAX = 2000.f; }// namespace arenai::agent #endif//ARENAI_CONSTANTS_H diff --git a/arenai_agent/src/networks/entropy.cpp b/arenai_agent/src/networks/entropy.cpp index ac5be82a..386c925d 100644 --- a/arenai_agent/src/networks/entropy.cpp +++ b/arenai_agent/src/networks/entropy.cpp @@ -13,56 +13,6 @@ using namespace arenai::agent; namespace arenai::agent { - /* - * Alpha parameters [0; +inf[ - */ - - AlphaParameters::AlphaParameters(const float initial_alpha, const int nb_alphas) - : log_alpha_tensor(register_parameter( - "log_alpha", - torch::tensor(std::vector(nb_alphas, std::log(initial_alpha))).unsqueeze(0))) {} - - torch::Tensor AlphaParameters::log_alpha() { return log_alpha_tensor; } - - torch::Tensor AlphaParameters::alpha() { return log_alpha().exp(); } - - /* - * Clamped alpha parameters - */ - - ClampedAlphaParameters::ClampedAlphaParameters( - const float initial_alpha, const float min_alpha, const float max_alpha, - const int nb_alphas) - : AlphaParameters(std::clamp(initial_alpha, min_alpha, max_alpha), nb_alphas), - min_log_alpha(std::log(min_alpha)), max_log_alpha(std::log(max_alpha)) {} - - torch::Tensor ClampedAlphaParameters::log_alpha() { - auto curr_log_alpha = AlphaParameters::log_alpha(); - - { - const torch::NoGradGuard no_grad; - curr_log_alpha.data().clamp_(min_log_alpha, max_log_alpha); - } - - return curr_log_alpha; - } - - /* - * Abstract target entropy - */ - - // a schedule-less target ignores the progression - void AbstractTargetEntropy::step(int64_t) {} - - /* - * Constant target entropy - */ - - ConstantTargetEntropy::ConstantTargetEntropy(const float initial_target) - : initial_target(register_buffer("initial_target", torch::tensor(initial_target))) {} - - torch::Tensor ConstantTargetEntropy::target_entropy() const { return initial_target; } - /* * PID Lagrangian */ diff --git a/arenai_agent/src/networks/entropy.h b/arenai_agent/src/networks/entropy.h index 92120759..2eafdcf5 100644 --- a/arenai_agent/src/networks/entropy.h +++ b/arenai_agent/src/networks/entropy.h @@ -9,58 +9,6 @@ namespace arenai::agent { - /* - * Base class - */ - - class AlphaParameters : public torch::nn::Module { - public: - explicit AlphaParameters(float initial_alpha, int nb_alphas); - - virtual torch::Tensor log_alpha(); - torch::Tensor alpha(); - - private: - torch::Tensor log_alpha_tensor; - }; - - class ClampedAlphaParameters final : public AlphaParameters { - public: - explicit ClampedAlphaParameters( - float initial_alpha, float min_alpha, float max_alpha, int nb_alphas); - - torch::Tensor log_alpha() override; - - private: - float min_log_alpha; - float max_log_alpha; - }; - - class AbstractTargetEntropy : public torch::nn::Module { - public: - // the current target: pure, safe to call several times inside the same rollout - virtual torch::Tensor target_entropy() const = 0; - - // advances the schedule by nb_env_steps environment steps. Called once per rollout, - // so a schedule is expressed in the same unit as the training progress reported in - // metrics.csv — independent of minibatch_size, epochs, nb_tanks and tank mortality. - virtual void step(int64_t nb_env_steps); - }; - - /* - * Constants - */ - - class ConstantTargetEntropy : public AbstractTargetEntropy { - public: - explicit ConstantTargetEntropy(float initial_target); - - torch::Tensor target_entropy() const override; - - private: - torch::Tensor initial_target; - }; - /* * Lagrangian */ diff --git a/arenai_agent/src/networks/misc.cpp b/arenai_agent/src/networks/misc.cpp index 929aec17..db374772 100644 --- a/arenai_agent/src/networks/misc.cpp +++ b/arenai_agent/src/networks/misc.cpp @@ -15,7 +15,7 @@ namespace arenai::agent { torch::Tensor Exp::forward(const torch::Tensor &x) { return torch::exp(x); } - void Exp::pretty_print(std::ostream &stream) { stream << name() << "()"; } + void Exp::pretty_print(std::ostream &stream) const { stream << name() << "()"; } /* * Clamp @@ -28,7 +28,7 @@ namespace arenai::agent { return torch::clamp(x, lower_bound, upper_bound); } - void Clamp::pretty_print(std::ostream &stream) { + void Clamp::pretty_print(std::ostream &stream) const { stream << name() << "(min=" << lower_bound << ", max=" << upper_bound << ")"; } @@ -45,9 +45,29 @@ namespace arenai::agent { return torch::exp(log_sigma); } - void SigmaOutput::pretty_print(std::ostream &stream) { + void SigmaOutput::pretty_print(std::ostream &stream) const { stream << name() << "(min=" << std::exp(min_log_sigma) << ", max=" << std::exp(max_log_sigma) << ")"; } + /* + * Concentration of Beta distribution output layer + */ + + ConcentrationOutput::ConcentrationOutput( + const float min_concentration, const float max_concentration) + : min_log_excess(std::log(min_concentration - 2.f)), + max_log_excess(std::log(max_concentration - 2.f)) {} + + torch::Tensor ConcentrationOutput::forward(const torch::Tensor &input) { + const auto log_excess = + min_log_excess + (max_log_excess - min_log_excess) * torch::sigmoid(input); + return 2.f + torch::exp(log_excess); + } + + void ConcentrationOutput::pretty_print(std::ostream &stream) const { + stream << name() << "(min=" << 2.f + std::exp(min_log_excess) + << ", max=" << 2.f + std::exp(max_log_excess) << ")"; + } + }// namespace arenai::agent diff --git a/arenai_agent/src/networks/misc.h b/arenai_agent/src/networks/misc.h index 70daf773..af10c519 100644 --- a/arenai_agent/src/networks/misc.h +++ b/arenai_agent/src/networks/misc.h @@ -13,7 +13,7 @@ namespace arenai::agent { public: virtual torch::Tensor forward(const torch::Tensor &input) = 0; - virtual void pretty_print(std::ostream &stream) = 0; + void pretty_print(std::ostream &stream) const override = 0; }; class Clamp : public AbstractFunctionModule { @@ -22,7 +22,7 @@ namespace arenai::agent { torch::Tensor forward(const torch::Tensor &x) override; - void pretty_print(std::ostream &stream) override; + void pretty_print(std::ostream &stream) const override; private: float lower_bound; @@ -33,7 +33,7 @@ namespace arenai::agent { public: torch::Tensor forward(const torch::Tensor &x) override; - void pretty_print(std::ostream &stream) override; + void pretty_print(std::ostream &stream) const override; }; class SigmaOutput : public AbstractFunctionModule { @@ -42,13 +42,27 @@ namespace arenai::agent { torch::Tensor forward(const torch::Tensor &input) override; - void pretty_print(std::ostream &stream) override; + void pretty_print(std::ostream &stream) const override; private: float min_log_sigma; float max_log_sigma; }; + // Beta concentration κ, log-scale on the excess κ - 2 bounded between min and max + class ConcentrationOutput : public AbstractFunctionModule { + public: + ConcentrationOutput(float min_concentration, float max_concentration); + + torch::Tensor forward(const torch::Tensor &input) override; + + void pretty_print(std::ostream &stream) const override; + + private: + float min_log_excess; + float max_log_excess; + }; + }// namespace arenai::agent #endif//ARENAI_AGENT_HOST_MISC_H diff --git a/arenai_agent/src/networks/q_function.cpp b/arenai_agent/src/networks/q_function.cpp deleted file mode 100644 index f5816fda..00000000 --- a/arenai_agent/src/networks/q_function.cpp +++ /dev/null @@ -1,117 +0,0 @@ -// -// Created by samuel on 22/01/2026. -// - -#include "./q_function.h" - -#include "../networks_utils/init.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - QFunction::QFunction( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_actions, - const int &hidden_size_sensors, const int &hidden_size_actions, - const std::vector &hidden_sizes, - const std::vector> &vision_channels, - const std::vector &group_norm_nums) - : nb_discrete_actions(nb_discrete_actions), - vision_encoder(register_module( - "vision_encoder", - std::make_shared( - vision_height, vision_width, vision_channels, group_norm_nums))), - sensors_encoder(register_module( - "sensors_encoder", - torch::nn::Sequential( - torch::nn::Linear( - torch::nn::LinearOptions(nb_sensors, hidden_size_sensors).bias(false)), - torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_sensors})), - torch::nn::SiLU()))), - continuous_action_encoder(register_module( - "continuous_action_encoder", - torch::nn::Sequential( - torch::nn::Linear(nb_continuous_actions, hidden_size_actions), - torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_actions})), - torch::nn::SiLU()))), - discrete_action_encoder(register_module( - "discrete_action_encoder", - torch::nn::Sequential( - torch::nn::Linear(nb_discrete_actions, hidden_size_actions), - torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_actions})), - torch::nn::SiLU()))), - head(register_module("head", torch::nn::Sequential())), - to_value(register_module("to_value", torch::nn::Linear(hidden_sizes.back(), 1))) { - - head->push_back(torch::nn::Linear( - torch::nn::LinearOptions( - 2 * hidden_size_actions + hidden_size_sensors + vision_encoder->get_output_size(), - hidden_sizes.front()) - .bias(false))); - head->push_back(torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_sizes.front()}))); - head->push_back(torch::nn::SiLU()); - - for (int i = 1; i < hidden_sizes.size(); i++) { - const auto curr_size = hidden_sizes[i - 1]; - const auto next_size = hidden_sizes[i]; - head->push_back( - torch::nn::Linear(torch::nn::LinearOptions(curr_size, next_size).bias(false))); - head->push_back(torch::nn::LayerNorm(torch::nn::LayerNormOptions({next_size}))); - head->push_back(torch::nn::SiLU()); - } - - vision_encoder->apply(init_hidden_weights); - sensors_encoder->apply(init_hidden_weights); - continuous_action_encoder->apply(init_hidden_weights); - discrete_action_encoder->apply(init_hidden_weights); - head->apply(init_hidden_weights); - - to_value->apply(init_value_output_weights); - } - - torch::Tensor QFunction::value_ohe( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions, const torch::Tensor &discrete_action_ohe) { - const auto common_encoded = encode_common(vision, sensors, continuous_actions); - const auto discrete_action_encoded = discrete_action_encoder->forward(discrete_action_ohe); - - const auto encoded_hidden = torch::cat({common_encoded, discrete_action_encoded}, 1); - - return to_value->forward(head->forward(encoded_hidden)); - } - - torch::Tensor QFunction::value_per_discrete_action( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions) { - const auto batch_size = vision.size(0); - - const auto common_encoded = encode_common(vision, sensors, continuous_actions); - const auto one_hots = torch::eye(nb_discrete_actions, common_encoded.options()); - - std::vector q_values; - q_values.reserve(nb_discrete_actions); - - for (int a = 0; a < nb_discrete_actions; a++) { - const auto discrete_encoded = - discrete_action_encoder->forward(one_hots[a].unsqueeze(0).expand({batch_size, -1})); - - q_values.push_back(to_value->forward( - head->forward(torch::cat({common_encoded, discrete_encoded}, 1)))); - } - - return torch::cat(q_values, 1); - } - - torch::Tensor QFunction::encode_common( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions) { - const auto vision_encoded = vision_encoder->forward(vision); - const auto sensors_encoded = sensors_encoder->forward(sensors); - const auto action_encoded = continuous_action_encoder->forward(continuous_actions); - - return torch::cat({vision_encoded, sensors_encoded, action_encoded}, 1); - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/networks/q_function.h b/arenai_agent/src/networks/q_function.h deleted file mode 100644 index fc7ac36a..00000000 --- a/arenai_agent/src/networks/q_function.h +++ /dev/null @@ -1,49 +0,0 @@ -// -// Created by samuel on 22/01/2026. -// - -#ifndef ARENAI_AGENT_HOST_Q_FUNCTION_H -#define ARENAI_AGENT_HOST_Q_FUNCTION_H - -#include - -#include - -#include "./vision.h" - -namespace arenai::agent { - - class QFunction final : public torch::nn::Module { - public: - QFunction( - const int &vision_height, const int &vision_width, const int &nb_sensors, - const int &nb_continuous_actions, const int &nb_discrete_actions, - const int &hidden_size_sensors, const int &hidden_size_actions, - const std::vector &hidden_sizes, - const std::vector> &vision_channels, - const std::vector &group_norm_nums); - torch::Tensor value_ohe( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions, const torch::Tensor &discrete_action_ohe); - - torch::Tensor value_per_discrete_action( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions); - - private: - int nb_discrete_actions; - std::shared_ptr vision_encoder; - torch::nn::Sequential sensors_encoder; - torch::nn::Sequential continuous_action_encoder; - torch::nn::Sequential discrete_action_encoder; - torch::nn::Sequential head; - torch::nn::Linear to_value; - - torch::Tensor encode_common( - const torch::Tensor &vision, const torch::Tensor &sensors, - const torch::Tensor &continuous_actions); - }; - -}// namespace arenai::agent - -#endif//ARENAI_AGENT_HOST_Q_FUNCTION_H diff --git a/arenai_agent/src/networks/recurrent/liquid_actor.cpp b/arenai_agent/src/networks/recurrent/liquid_actor.cpp new file mode 100644 index 00000000..4893b21c --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_actor.cpp @@ -0,0 +1,113 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_actor.h" + +#include "../../networks_utils/init.h" +#include "../constants.h" +#include "../misc.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + LiquidActor::LiquidActor( + const int &vision_height, const int &vision_width, const int &nb_sensors, + const int &nb_continuous_actions, const int &nb_discrete_actions, + const int &hidden_size_sensors, const std::vector> &vision_channels, + const std::vector &group_norm_nums, const int &neuron_number, + const int &unfolding_steps, const float &delta_t, const float &initial_sigma, + const std::vector &initial_discrete_probas) + : vision_encoder(register_module( + "vision_encoder", std::make_shared( + vision_height, vision_width, vision_channels, group_norm_nums))), + sensors_encoder(register_module( + "sensors_encoder", + torch::nn::Sequential( + torch::nn::Linear( + torch::nn::LinearOptions(nb_sensors, hidden_size_sensors).bias(false)), + torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_sensors})), + torch::nn::SiLU()))), + liquid(register_module( + "liquid", std::make_shared( + neuron_number, hidden_size_sensors + vision_encoder->get_output_size(), + neuron_number, unfolding_steps, + [](const torch::Tensor &t) { return torch::silu(t); }, delta_t))), + mode(register_module( + "mode", + torch::nn::Sequential( + torch::nn::Linear(neuron_number, nb_continuous_actions), torch::nn::Sigmoid()))), + concentration(register_module( + "concentration", + torch::nn::Sequential( + torch::nn::Linear(neuron_number, nb_continuous_actions), + std::make_shared(CONCENTRATION_MIN, CONCENTRATION_MAX)))), + discrete(register_module( + "discrete", + torch::nn::Sequential( + torch::nn::Linear(neuron_number, nb_discrete_actions), torch::nn::Sigmoid()))) { + + vision_encoder->apply(init_hidden_weights); + sensors_encoder->apply(init_hidden_weights); + + // zero bias + sigmoid puts the initial mode at the action-range center + mode->apply(init_mu_output_weights); + concentration->apply( + [initial_sigma](Module &m) { init_concentration_output_weights(m, initial_sigma); }); + + discrete->apply([&initial_discrete_probas](Module &m) { + init_discrete_output_weights(m, initial_discrete_probas); + }); + } + + torch::Tensor LiquidActor::encode(const torch::Tensor &vision, const torch::Tensor &sensors) { + return torch::cat({vision_encoder->forward(vision), sensors_encoder->forward(sensors)}, 1); + } + + LiquidActorOutput LiquidActor::act( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t) { + const auto [output, next_x] = liquid->forward_step(x_t, encode(vision, sensors)); + return { + .mode = mode->forward(output), + .concentration = concentration->forward(output), + .discrete = discrete->forward(output), + .next_x = next_x}; + } + + LiquidActorOutput LiquidActor::act_sequence( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t) { + const auto batch_size = vision.size(0); + const auto nb_steps = vision.size(1); + + // the encoders are step-wise: fold time into the row dimension + const auto encoded = + encode(vision.flatten(0, 1), sensors.flatten(0, 1)).reshape({batch_size, nb_steps, -1}); + + auto x = x_t; + + std::vector outputs; + outputs.reserve(nb_steps); + + for (auto t = 0; t < nb_steps; t++) { + auto [output, next_x] = liquid->forward_step( + x, encoded.index({at::indexing::Slice(), t, at::indexing::Slice()})); + + x = next_x; + outputs.push_back(output); + } + + const auto output = torch::stack(outputs, 1); + return { + .mode = mode->forward(output), + .concentration = concentration->forward(output), + .discrete = discrete->forward(output), + .next_x = x}; + } + + torch::Tensor LiquidActor::initial_state(const int batch_size) const { + return liquid->get_first_x(batch_size); + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/networks/recurrent/liquid_actor.h b/arenai_agent/src/networks/recurrent/liquid_actor.h new file mode 100644 index 00000000..1a5216f6 --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_actor.h @@ -0,0 +1,68 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_ACTOR_H +#define ARENAI_LIQUID_ACTOR_H + +#include + +#include + +#include "../vision.h" +#include "./liquid_cell.h" + +namespace arenai::agent { + + struct LiquidActorOutput { + // Beta distribution on [-1, 1]: mode in [0, 1] (on the underlying [0, 1] + // support) and concentration κ = α + β + torch::Tensor mode; + torch::Tensor concentration; + // independent Bernoulli probabilities, one per discrete action + torch::Tensor discrete; + // liquid state after the last processed step + torch::Tensor next_x; + }; + + // CNN on the frames, MLP on the proprioception, both concatenated into the + // recurrent input of each step, then run through the liquid network + class LiquidActor final : public torch::nn::Module { + public: + explicit LiquidActor( + const int &vision_height, const int &vision_width, const int &nb_sensors, + const int &nb_continuous_actions, const int &nb_discrete_actions, + const int &hidden_size_sensors, + const std::vector> &vision_channels, + const std::vector &group_norm_nums, const int &neuron_number, + const int &unfolding_steps, const float &delta_t, const float &initial_sigma, + const std::vector &initial_discrete_probas); + + // one env step: vision [B, C, H, W], sensors [B, S], x_t [B, neuron_number] + LiquidActorOutput + act(const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t); + + // BPTT chunk: vision [B, T, C, H, W], sensors [B, T, S], x_t [B, neuron_number]; + // outputs are [B, T, ...] + LiquidActorOutput act_sequence( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t); + + torch::Tensor initial_state(int batch_size) const; + + private: + std::shared_ptr vision_encoder; + torch::nn::Sequential sensors_encoder; + + std::shared_ptr liquid; + + torch::nn::Sequential mode; + torch::nn::Sequential concentration; + torch::nn::Sequential discrete; + + // [rows, features] recurrent input: encoded vision and sensors, concatenated + torch::Tensor encode(const torch::Tensor &vision, const torch::Tensor &sensors); + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_ACTOR_H diff --git a/arenai_agent/src/networks/recurrent/liquid_cell.cpp b/arenai_agent/src/networks/recurrent/liquid_cell.cpp new file mode 100644 index 00000000..1d265dd1 --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_cell.cpp @@ -0,0 +1,105 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_cell.h" + +#include "../../networks_utils/init.h" + +using namespace arenai; +using namespace arenai::agent; + +/* + * Cell model + */ + +CellModel::CellModel( + const int neuron_number, const int input_size, + const std::function &activation_function) + : weights(register_module( + "weights", + torch::nn::Linear(torch::nn::LinearOptions(input_size, neuron_number).bias(false)))), + recurrent_weights(register_module( + "recurrent_weights", + torch::nn::Linear(torch::nn::LinearOptions(neuron_number, neuron_number).bias(false)))), + biases(register_parameter("biases", torch::zeros({1, neuron_number}))), + activation_function(activation_function) { + + weights->apply(init_liquid_weights); + recurrent_weights->apply(init_liquid_weights); +} + +torch::Tensor CellModel::forward(const torch::Tensor &x_t, const torch::Tensor &input_t) { + return activation_function(recurrent_weights(x_t) + weights(input_t) + biases); +} + +/* + * Liquid cell + */ + +LiquidCell::LiquidCell( + const int neuron_number, const int input_size, const int output_size, const int unfolding_steps, + const std::function &activation_function, + const float delta_t) + : a(register_parameter("a", torch::ones({1, neuron_number}))), + raw_tau(register_parameter("raw_tau", torch::zeros({1, neuron_number}))), + f(register_module( + "f", std::make_shared(neuron_number, input_size, activation_function))), + unfolding_steps(unfolding_steps), delta_t(delta_t), neuron_number(neuron_number), + to_output(register_module( + "to_output", + torch::nn::Sequential( + torch::nn::Linear(torch::nn::LinearOptions(neuron_number, output_size).bias(false)), + torch::nn::LayerNorm( + torch::nn::LayerNormOptions({output_size}).elementwise_affine(true)), + torch::nn::SiLU()))) {} + +torch::Tensor LiquidCell::forward(const torch::Tensor &inputs) { + TORCH_CHECK( + inputs.sizes().size() == 3, + "Processed input needs to have 3 dimensions (Batch, Time, Features)"); + + const auto batch_size = inputs.size(0); + const auto nb_steps = inputs.size(1); + + auto x_t = get_first_x(static_cast(batch_size)); + + std::vector results; + results.reserve(nb_steps); + + for (auto t = 0; t < nb_steps; t++) { + auto [output, x_t_next] = + forward_step(x_t, inputs.index({at::indexing::Slice(), t, at::indexing::Slice()})); + + x_t = x_t_next; + results.push_back(output); + } + + // (batch, time, output_features) + return torch::stack(results, 1); +} + +std::tuple +LiquidCell::forward_step(const torch::Tensor &x_t, const torch::Tensor &input_t) { + const auto x_t_next = next_x(x_t, input_t); + return {to_output->forward(x_t_next), x_t_next}; +} + +torch::Tensor LiquidCell::next_x(const torch::Tensor &x_t, const torch::Tensor &input_t) const { + auto x_t_next = x_t; + const auto curr_delta_t = delta_t / static_cast(unfolding_steps); + + for (int i = 0; i < unfolding_steps; ++i) { + const auto f_output = f->forward(x_t_next, input_t); + x_t_next = (x_t_next + curr_delta_t * f_output * a) + / (1.0 + curr_delta_t * (1.0 / tau() + f_output)); + } + + return x_t_next; +} + +torch::Tensor LiquidCell::tau() const { return torch::nn::functional::softplus(raw_tau) + 1e-3; } + +torch::Tensor LiquidCell::get_first_x(int batch_size) { + return 1e-1 * torch::randn({batch_size, neuron_number}, parameters().back().device()); +} diff --git a/arenai_agent/src/networks/recurrent/liquid_cell.h b/arenai_agent/src/networks/recurrent/liquid_cell.h new file mode 100644 index 00000000..ca1f6cb5 --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_cell.h @@ -0,0 +1,59 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_CELL_H +#define ARENAI_LIQUID_CELL_H + +#include + +namespace arenai::agent { + class CellModel : public torch::nn::Module { + public: + CellModel( + int neuron_number, int input_size, + const std::function &activation_function); + + torch::Tensor forward(const torch::Tensor &x_t, const torch::Tensor &input_t); + + private: + torch::nn::Linear weights; + torch::nn::Linear recurrent_weights; + torch::Tensor biases; + + std::function activation_function; + }; + + class LiquidCell : public torch::nn::Module { + public: + LiquidCell( + int neuron_number, int input_size, int output_size, int unfolding_steps, + const std::function &activation_function, + float delta_t); + + torch::Tensor forward(const torch::Tensor &inputs); + + // one recurrent step, the state is handled by the caller: (output, x_t_next) + std::tuple + forward_step(const torch::Tensor &x_t, const torch::Tensor &input_t); + + torch::Tensor get_first_x(int batch_size); + + private: + torch::Tensor a; + torch::Tensor raw_tau; + + std::shared_ptr f; + + int unfolding_steps; + float delta_t; + int neuron_number; + + torch::nn::Sequential to_output; + + torch::Tensor tau() const; + torch::Tensor next_x(const torch::Tensor &x_t, const torch::Tensor &input_t) const; + }; +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_CELL_H diff --git a/arenai_agent/src/networks/recurrent/liquid_critic.cpp b/arenai_agent/src/networks/recurrent/liquid_critic.cpp new file mode 100644 index 00000000..2a70fbb0 --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_critic.cpp @@ -0,0 +1,81 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./liquid_critic.h" + +#include "../../networks_utils/init.h" + +using namespace arenai; +using namespace arenai::agent; + +namespace arenai::agent { + + LiquidCritic::LiquidCritic( + const int &vision_height, const int &vision_width, const int &nb_sensors, + const int &hidden_size_sensors, const std::vector> &vision_channels, + const std::vector &group_norm_nums, const int &neuron_number, + const int &unfolding_steps, const float &delta_t) + : vision_encoder(register_module( + "vision_encoder", std::make_shared( + vision_height, vision_width, vision_channels, group_norm_nums))), + sensors_encoder(register_module( + "sensors_encoder", + torch::nn::Sequential( + torch::nn::Linear( + torch::nn::LinearOptions(nb_sensors, hidden_size_sensors).bias(false)), + torch::nn::LayerNorm(torch::nn::LayerNormOptions({hidden_size_sensors})), + torch::nn::SiLU()))), + liquid(register_module( + "liquid", std::make_shared( + neuron_number, hidden_size_sensors + vision_encoder->get_output_size(), + neuron_number, unfolding_steps, + [](const torch::Tensor &t) { return torch::silu(t); }, delta_t))), + to_value(register_module("to_value", torch::nn::Linear(neuron_number, 1))) { + + vision_encoder->apply(init_hidden_weights); + sensors_encoder->apply(init_hidden_weights); + + to_value->apply(init_value_output_weights); + } + + torch::Tensor LiquidCritic::encode(const torch::Tensor &vision, const torch::Tensor &sensors) { + return torch::cat({vision_encoder->forward(vision), sensors_encoder->forward(sensors)}, 1); + } + + LiquidCriticOutput LiquidCritic::value( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t) { + const auto [output, next_x] = liquid->forward_step(x_t, encode(vision, sensors)); + return {.value = to_value->forward(output), .next_x = next_x}; + } + + LiquidCriticOutput LiquidCritic::value_sequence( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t) { + const auto batch_size = vision.size(0); + const auto nb_steps = vision.size(1); + + // the encoders are step-wise: fold time into the row dimension + const auto encoded = + encode(vision.flatten(0, 1), sensors.flatten(0, 1)).reshape({batch_size, nb_steps, -1}); + + auto x = x_t; + + std::vector outputs; + outputs.reserve(nb_steps); + + for (auto t = 0; t < nb_steps; t++) { + auto [output, next_x] = liquid->forward_step( + x, encoded.index({at::indexing::Slice(), t, at::indexing::Slice()})); + + x = next_x; + outputs.push_back(output); + } + + return {.value = to_value->forward(torch::stack(outputs, 1)), .next_x = x}; + } + + torch::Tensor LiquidCritic::initial_state(const int batch_size) const { + return liquid->get_first_x(batch_size); + } + +}// namespace arenai::agent diff --git a/arenai_agent/src/networks/recurrent/liquid_critic.h b/arenai_agent/src/networks/recurrent/liquid_critic.h new file mode 100644 index 00000000..d5286525 --- /dev/null +++ b/arenai_agent/src/networks/recurrent/liquid_critic.h @@ -0,0 +1,58 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_LIQUID_CRITIC_H +#define ARENAI_LIQUID_CRITIC_H + +#include + +#include + +#include "../vision.h" +#include "./liquid_cell.h" + +namespace arenai::agent { + + struct LiquidCriticOutput { + torch::Tensor value; + // liquid state after the last processed step + torch::Tensor next_x; + }; + + // same encoders as the liquid actor, ending on a state-value head + class LiquidCritic final : public torch::nn::Module { + public: + explicit LiquidCritic( + const int &vision_height, const int &vision_width, const int &nb_sensors, + const int &hidden_size_sensors, + const std::vector> &vision_channels, + const std::vector &group_norm_nums, const int &neuron_number, + const int &unfolding_steps, const float &delta_t); + + // one env step: vision [B, C, H, W], sensors [B, S], x_t [B, neuron_number] + LiquidCriticOutput + value(const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t); + + // BPTT chunk: vision [B, T, C, H, W], sensors [B, T, S], x_t [B, neuron_number]; + // value is [B, T, 1] + LiquidCriticOutput value_sequence( + const torch::Tensor &vision, const torch::Tensor &sensors, const torch::Tensor &x_t); + + torch::Tensor initial_state(int batch_size) const; + + private: + std::shared_ptr vision_encoder; + torch::nn::Sequential sensors_encoder; + + std::shared_ptr liquid; + + torch::nn::Linear to_value; + + // [rows, features] recurrent input: encoded vision and sensors, concatenated + torch::Tensor encode(const torch::Tensor &vision, const torch::Tensor &sensors); + }; + +}// namespace arenai::agent + +#endif//ARENAI_LIQUID_CRITIC_H diff --git a/arenai_agent/src/networks/vision.cpp b/arenai_agent/src/networks/vision.cpp index a51039fa..5aee5a4a 100644 --- a/arenai_agent/src/networks/vision.cpp +++ b/arenai_agent/src/networks/vision.cpp @@ -46,4 +46,59 @@ namespace arenai::agent { int ConvolutionNetwork::get_output_size() const { return output_size; } + /* + * Residual + */ + + ResidualBlock::ResidualBlock(const int channels) + : conv_block(register_module( + "conv_block", + torch::nn::Sequential( + torch::nn::SiLU(), + torch::nn::Conv2d( + torch::nn::Conv2dOptions(channels, channels, 3).stride(1).padding(1)), + torch::nn::SiLU(), + torch::nn::Conv2d( + torch::nn::Conv2dOptions(channels, channels, 3).stride(1).padding(1))))) {} + + torch::Tensor ResidualBlock::forward(const torch::Tensor &input) { + return conv_block->forward(input) + input; + } + + /* + * Impala CNN + */ + + ImpalaConvolutionNetwork::ImpalaConvolutionNetwork( + const int vision_height, const int vision_width, + const std::vector> &channels) + : cnn(register_module("cnn", torch::nn::Sequential())) { + int w = vision_width, h = vision_height; + + for (auto [c_i, c_o]: channels) { + constexpr int padding = 1, stride = 2, kernel = 3; + + w = (w - kernel + 2 * padding) / stride + 1; + h = (h - kernel + 2 * padding) / stride + 1; + + cnn->push_back(torch::nn::Conv2d( + torch::nn::Conv2dOptions(c_i, c_o, kernel).stride(1).padding(padding))); + cnn->push_back(torch::nn::MaxPool2d( + torch::nn::MaxPool2dOptions(kernel).stride(stride).padding(padding))); + cnn->push_back(std::make_shared(c_o)); + } + + cnn->push_back(torch::nn::SiLU()); + cnn->push_back(torch::nn::Flatten(torch::nn::FlattenOptions().start_dim(1).end_dim(-1))); + + output_size = w * h * std::get<1>(channels.back()); + } + + torch::Tensor ImpalaConvolutionNetwork::forward(const torch::Tensor &input) { + TORCH_CHECK(input.dtype() == torch::kUInt8, "Input must be UInt8"); + + return cnn->forward(input.to(torch::kFloat).mul_(2.0f / 255.0f).add_(-1.0f)); + } + + int ImpalaConvolutionNetwork::get_output_size() const { return output_size; } }// namespace arenai::agent diff --git a/arenai_agent/src/networks/vision.h b/arenai_agent/src/networks/vision.h index d74920af..93ad3811 100644 --- a/arenai_agent/src/networks/vision.h +++ b/arenai_agent/src/networks/vision.h @@ -24,6 +24,30 @@ namespace arenai::agent { int output_size; }; + class ResidualBlock final : public torch::nn::Module { + public: + explicit ResidualBlock(int channels); + + torch::Tensor forward(const torch::Tensor &input); + + private: + torch::nn::Sequential conv_block{nullptr}; + }; + + class ImpalaConvolutionNetwork final : public torch::nn::Module { + public: + ImpalaConvolutionNetwork( + int vision_height, int vision_width, const std::vector> &channels); + + torch::Tensor forward(const torch::Tensor &input); + + int get_output_size() const; + + private: + torch::nn::Sequential cnn{nullptr}; + int output_size; + }; + }// namespace arenai::agent #endif// ARENAI_AGENT_HOST_VISION_H diff --git a/arenai_agent/src/networks_utils/init.cpp b/arenai_agent/src/networks_utils/init.cpp index eb04e53f..71aa8c60 100644 --- a/arenai_agent/src/networks_utils/init.cpp +++ b/arenai_agent/src/networks_utils/init.cpp @@ -4,6 +4,9 @@ #include "./init.h" +#include +#include + #include "../networks/constants.h" using namespace arenai; @@ -38,6 +41,12 @@ namespace arenai::agent { } } + void init_liquid_weights(torch::nn::Module &module) { + if (const auto *lin = module.as()) { + torch::nn::init::normal_(lin->weight, 0.f, 1e-2f); + } + } + void init_sigma_output_weights(torch::nn::Module &module, const float wanted_sigma) { const float min_log_sigma = std::log(SIGMA_MIN); const float max_log_sigma = std::log(SIGMA_MAX); @@ -53,18 +62,38 @@ namespace arenai::agent { } } - void - init_discrete_output_weights(torch::nn::Module &module, const float initial_fire_probability) { + void init_concentration_output_weights(torch::nn::Module &module, const float wanted_sigma) { + // Beta on [-1, 1]: var = 4 μ(1-μ) / (κ+1), so at μ = 0.5 a wanted action + // std σ maps to κ = 1/σ² - 1 + const auto wanted_concentration = std::clamp( + 1.f / (wanted_sigma * wanted_sigma) - 1.f, CONCENTRATION_MIN, CONCENTRATION_MAX); + + const float min_log_excess = std::log(CONCENTRATION_MIN - 2.f); + const float max_log_excess = std::log(CONCENTRATION_MAX - 2.f); + + const auto initial_sigmoid = (std::log(wanted_concentration - 2.f) - min_log_excess) + / (max_log_excess - min_log_excess); + const auto initial_logit = std::log(initial_sigmoid / (1.f - initial_sigmoid)); + + if (auto *lin = module.as()) { + torch::nn::init::orthogonal_(lin->weight, 0.01f); + if (lin->options.bias()) torch::nn::init::constant_(lin->bias, initial_logit); + } + } + + void init_discrete_output_weights( + torch::nn::Module &module, const std::vector &initial_probabilities) { if (auto *lin = module.as()) { torch::nn::init::orthogonal_(lin->weight, 0.01f); if (lin->options.bias()) { - torch::nn::init::zeros_(lin->bias); + // sigmoid head: the bias is the logit of the wanted probability + std::vector logits; + logits.reserve(initial_probabilities.size()); + for (const auto probability: initial_probabilities) + logits.push_back(std::log(probability / (1.f - probability))); - lin->bias.data().index_fill_( - 0, torch::tensor({0}), std::log(initial_fire_probability)); - lin->bias.data().index_fill_( - 0, torch::tensor({1}), std::log(1.f - initial_fire_probability)); + lin->bias.data().copy_(torch::tensor(logits)); } } } diff --git a/arenai_agent/src/networks_utils/init.h b/arenai_agent/src/networks_utils/init.h index 5ef95f31..cac0a4a5 100644 --- a/arenai_agent/src/networks_utils/init.h +++ b/arenai_agent/src/networks_utils/init.h @@ -5,15 +5,22 @@ #ifndef ARENAI_AGENT_HOST_INIT_H #define ARENAI_AGENT_HOST_INIT_H +#include + #include namespace arenai::agent { void init_hidden_weights(torch::nn::Module &module); + void init_liquid_weights(torch::nn::Module &module); + void init_mu_output_weights(torch::nn::Module &module); void init_sigma_output_weights(torch::nn::Module &module, float wanted_sigma); - void init_discrete_output_weights(torch::nn::Module &module, float initial_fire_probability); + void init_concentration_output_weights(torch::nn::Module &module, float wanted_sigma); + // sigmoid head: one initial engage probability per discrete action + void init_discrete_output_weights( + torch::nn::Module &module, const std::vector &initial_probabilities); void init_value_output_weights(torch::nn::Module &module); diff --git a/arenai_agent/src/networks_utils/torch_converter.cpp b/arenai_agent/src/networks_utils/torch_converter.cpp index afb05379..b2e149db 100644 --- a/arenai_agent/src/networks_utils/torch_converter.cpp +++ b/arenai_agent/src/networks_utils/torch_converter.cpp @@ -28,12 +28,15 @@ namespace arenai::agent { for (int i = 0; i < batch_size; i++) { const controller::joystick joystick_direction{.x = cont_acc[i][0], .y = cont_acc[i][1]}; const controller::joystick joystick_canon{.x = cont_acc[i][2], .y = cont_acc[i][3]}; - const controller::button fire_button(disc_acc[i][0] > disc_acc[i][1]); + // binary Bernoulli actions: one neuron per discrete action + const controller::button fire_button(disc_acc[i][0] > 0.5f); + const controller::button zoom_button(disc_acc[i][1] > 0.5f); actions.push_back( {.left_joystick = joystick_direction, .right_joystick = joystick_canon, - .fire_button = fire_button}); + .fire_button = fire_button, + .zoom_button = zoom_button}); } return actions; @@ -81,23 +84,28 @@ namespace arenai::agent { } TorchStep steps_to_tensor( - const std::vector> &steps, + const std::vector> + &steps, const int vision_height, const int vision_width) { std::vector states; std::vector rewards; std::vector are_done; + std::vector are_truncated; - for (const auto &[state, reward, is_done]: steps) { + for (const auto &[state, reward, is_done, is_truncated]: steps) { states.push_back(state); rewards.push_back(torch::tensor({reward}, torch::TensorOptions().dtype(torch::kFloat))); are_done.push_back( torch::tensor({is_done}, torch::TensorOptions().dtype(torch::kBool))); + are_truncated.push_back( + torch::tensor({is_truncated}, torch::TensorOptions().dtype(torch::kBool))); } return { .states = states_to_tensor(states, vision_height, vision_width), .rewards = torch::stack(rewards), - .is_done = torch::stack(are_done)}; + .is_done = torch::stack(are_done), + .is_truncated = torch::stack(are_truncated)}; } }// namespace arenai::agent diff --git a/arenai_agent/src/networks_utils/torch_converter.h b/arenai_agent/src/networks_utils/torch_converter.h index 3063e1dd..53c93b03 100644 --- a/arenai_agent/src/networks_utils/torch_converter.h +++ b/arenai_agent/src/networks_utils/torch_converter.h @@ -23,7 +23,8 @@ namespace arenai::agent { TorchState state_to_tensor(const core::State &state, int vision_height, int vision_width); TorchStep steps_to_tensor( - const std::vector> &steps, + const std::vector> + &steps, int vision_height, int vision_width); }// namespace arenai::agent diff --git a/arenai_agent/src/train.cpp b/arenai_agent/src/train.cpp index 19eb90df..84bf2115 100644 --- a/arenai_agent/src/train.cpp +++ b/arenai_agent/src/train.cpp @@ -4,9 +4,10 @@ #include "./train.h" +#include +#include #include #include -#include #include #include @@ -15,7 +16,9 @@ #include +#include "./core/spawn_curriculum.h" #include "./core/train_environment.h" +#include "./metrics/last_metric.h" #include "./metrics/metric_saver.h" #include "./networks_utils/torch_converter.h" #include "./networks_utils/torch_saver.h" @@ -30,7 +33,7 @@ namespace arenai::agent { // identifiable once its command line is forgotten void save_run_config( const EnvironmentOptions &environment_options, const TrainOptions &train_options, - const std::map &agent_config) { + const nlohmann::json &agent_config) { const nlohmann::json config = { {"train", @@ -49,6 +52,11 @@ namespace arenai::agent { {"initial_spawn_height", environment_options.initial_spawn_height}, {"final_spawn_width", environment_options.final_spawn_width}, {"final_spawn_height", environment_options.final_spawn_height}, + {"curriculum_delta", environment_options.curriculum_delta}, + {"curriculum_ratio_low", environment_options.curriculum_ratio_low}, + {"curriculum_ratio_high", environment_options.curriculum_ratio_high}, + {"curriculum_probe_window", environment_options.curriculum_probe_window}, + {"curriculum_boundary_proba", environment_options.curriculum_boundary_proba}, {"vision_num_threads", environment_options.num_threads}}}, {"agent", agent_config}}; @@ -84,15 +92,23 @@ namespace arenai::agent { train_options.max_episode_steps, environment_options.vision_height, environment_options.vision_width, environment_options.num_threads); - const float spawn_width_increase = - (environment_options.final_spawn_width - environment_options.initial_spawn_width) - / static_cast(train_options.nb_episodes); - const float spawn_height_increase = - (environment_options.final_spawn_height - environment_options.initial_spawn_height) - / static_cast(train_options.nb_episodes); + const float initial_spawn_side = std::sqrt( + environment_options.initial_spawn_width * environment_options.initial_spawn_height); + const float final_spawn_side = std::sqrt( + environment_options.final_spawn_width * environment_options.final_spawn_height); - float spawn_width = environment_options.initial_spawn_width; - float spawn_height = environment_options.initial_spawn_height; + // delta is given in meters on the spawn side; a non-growing range disables + // the curriculum (progress stays at 0, i.e. the initial spawn size) + const float curriculum_delta_progress = + final_spawn_side > initial_spawn_side + ? environment_options.curriculum_delta / (final_spawn_side - initial_spawn_side) + : 0.f; + + constexpr std::uint64_t curriculum_seed = 1337; + SpawnCurriculum curriculum( + curriculum_delta_progress, environment_options.curriculum_ratio_low, + environment_options.curriculum_ratio_high, environment_options.curriculum_probe_window, + environment_options.curriculum_boundary_proba, curriculum_seed); const auto agent = agent_factory->get_agent(); const auto collector = agent_factory->get_collector(); @@ -106,12 +122,18 @@ namespace arenai::agent { // metrics - const auto sac_metrics = trainer->get_metrics(); + const auto trainer_metrics = trainer->get_metrics(); const auto env_metrics = env->get_metrics(); + // curriculum observability: the side played this episode and the current bound + const auto spawn_side_metric = std::make_shared("side", 0); + const auto spawn_bound_metric = std::make_shared("D", 0); + std::vector> metrics; metrics.insert(metrics.end(), env_metrics.begin(), env_metrics.end()); - metrics.insert(metrics.end(), sac_metrics.begin(), sac_metrics.end()); + metrics.push_back(spawn_side_metric); + metrics.push_back(spawn_bound_metric); + metrics.insert(metrics.end(), trainer_metrics.begin(), trainer_metrics.end()); MetricCsvSaver metric_csv_saver( train_options.output_folder, metrics, @@ -139,8 +161,20 @@ namespace arenai::agent { int print_counter = 0; for (int episode_index = 0; episode_index < train_options.nb_episodes; episode_index++) { + const float progress = curriculum.sample_progress(); + + const float spawn_width = std::lerp( + environment_options.initial_spawn_width, environment_options.final_spawn_width, + progress); + const float spawn_height = std::lerp( + environment_options.initial_spawn_height, environment_options.final_spawn_height, + progress); const float spawn_side = std::sqrt(spawn_width * spawn_height); + spawn_side_metric->add(spawn_side); + spawn_bound_metric->add( + std::lerp(initial_spawn_side, final_spawn_side, curriculum.upper_bound())); + // set variable for episode bool is_done = false; @@ -165,11 +199,12 @@ namespace arenai::agent { // step environment const auto steps = env->step(environment_options.wanted_frequency, actions_for_env); - const auto [torch_next_states, torch_rewards, torch_are_done] = steps_to_tensor( - steps, environment_options.vision_height, environment_options.vision_width); + const auto [torch_next_states, torch_rewards, torch_are_done, torch_are_truncated] = + steps_to_tensor( + steps, environment_options.vision_height, environment_options.vision_width); // complete the pending transition - maybe train - collector->on_transition(torch_rewards, torch_are_done); + collector->on_transition(torch_rewards, torch_are_done, torch_are_truncated); trainer->step(); // step ending stuff @@ -200,8 +235,7 @@ namespace arenai::agent { env->stop_drawing(); - spawn_width += spawn_width_increase; - spawn_height += spawn_height_increase; + curriculum.on_episode_end(env->episode_nb_fires(), env->episode_nb_hits()); p_bar.tick(); } diff --git a/arenai_agent/src/train.h b/arenai_agent/src/train.h index 639f9f7d..18536bc7 100644 --- a/arenai_agent/src/train.h +++ b/arenai_agent/src/train.h @@ -30,6 +30,11 @@ namespace arenai::agent { float initial_spawn_height; float final_spawn_width; float final_spawn_height; + float curriculum_delta; + float curriculum_ratio_low; + float curriculum_ratio_high; + int curriculum_probe_window; + float curriculum_boundary_proba; int num_threads; }; diff --git a/arenai_agent/src/utils/cli_fields.h b/arenai_agent/src/utils/cli_fields.h index 51bfa6b9..7cc0b419 100644 --- a/arenai_agent/src/utils/cli_fields.h +++ b/arenai_agent/src/utils/cli_fields.h @@ -5,17 +5,14 @@ #ifndef ARENAI_AGENT_HOST_CLI_FIELDS_H #define ARENAI_AGENT_HOST_CLI_FIELDS_H -#include -#include -#include +#include #include #include #include #include #include - -#include "./cli_parser.h" +#include namespace arenai::agent { @@ -26,10 +23,18 @@ namespace arenai::agent { struct CliField { std::string name; std::variant< - int S::*, float S::*, std::vector S::*, std::vector> S::*> + int S::*, float S::*, std::vector S::*, std::vector S::*, + std::vector> S::*> member; }; + // argparse's get unpacks container types element-wise, so a + // vector produced by an action must travel boxed in a scalar type + template + struct CliJsonValue { + T value; + }; + /* * Add field */ @@ -44,22 +49,38 @@ namespace arenai::agent { parser.add_argument(name).scan<'g', float>().default_value(default_value); } - inline void add_cli_field( - argparse::ArgumentParser &parser, const std::string &name, - const std::vector &default_value) { + template + void add_cli_json_field( + argparse::ArgumentParser &parser, const std::string &name, const T &default_value) { parser.add_argument(name) - .default_value({default_value}) + .default_value(CliJsonValue{default_value}) .action([name](const std::string &value) { - return hidden_layers{parse_int_vector(value, name, "[256, 128, ..., 64]")}; + try { + return CliJsonValue{nlohmann::json::parse(value).get()}; + } catch (const nlohmann::json::exception &) { + throw std::invalid_argument( + "invalid " + name + " value, usage : " + nlohmann::json(T{}).dump() + + " (JSON), actual value = \"" + value + "\""); + } }); } + inline void add_cli_field( + argparse::ArgumentParser &parser, const std::string &name, + const std::vector &default_value) { + add_cli_json_field(parser, name, default_value); + } + + inline void add_cli_field( + argparse::ArgumentParser &parser, const std::string &name, + const std::vector &default_value) { + add_cli_json_field(parser, name, default_value); + } + inline void add_cli_field( argparse::ArgumentParser &parser, const std::string &name, const std::vector> &default_value) { - parser.add_argument(name) - .default_value({default_value}) - .action(parse_cli_vision_channels); + add_cli_json_field(parser, name, default_value); } /* @@ -78,44 +99,19 @@ namespace arenai::agent { inline void read_cli_field( const argparse::ArgumentParser &parser, const std::string &name, std::vector &output) { - output = parser.get(name).layers; + output = parser.get>>(name).value; } inline void read_cli_field( const argparse::ArgumentParser &parser, const std::string &name, - std::vector> &output) { - output = parser.get(name).channels; - } - - /* - * Format fields - */ - - inline std::string format_cli_value(const int value) { return std::to_string(value); } - - inline std::string format_cli_value(const float value) { - std::ostringstream stream; - stream << std::setprecision(6) << value; - return stream.str(); - } - - inline std::string format_cli_value(const std::vector &value) { - std::ostringstream stream; - stream << "["; - for (int i = 0; i < value.size(); i++) stream << (i ? ", " : "") << value[i]; - stream << "]"; - return stream.str(); + std::vector &output) { + output = parser.get>>(name).value; } - inline std::string format_cli_value(const std::vector> &value) { - std::ostringstream stream; - stream << "["; - for (int i = 0; i < value.size(); i++) { - const auto &[in_channels, out_channels] = value[i]; - stream << (i ? ", " : "") << "(" << in_channels << ", " << out_channels << ")"; - } - stream << "]"; - return stream.str(); + inline void read_cli_field( + const argparse::ArgumentParser &parser, const std::string &name, + std::vector> &output) { + output = parser.get>>>(name).value; } /* @@ -134,16 +130,14 @@ namespace arenai::agent { } // the resolved hyper-parameters keyed by option name, dashes stripped: what the - // run was actually launched with + // run was actually launched with, as native JSON values template - std::map - cli_fields_to_map(const std::vector> &fields, const S ¶ms) { - std::map config; + nlohmann::json cli_fields_to_json(const std::vector> &fields, const S ¶ms) { + nlohmann::json config; for (const auto &field: fields) std::visit( [&](const auto member) { - config[field.name.substr(field.name.find_first_not_of('-'))] = - format_cli_value(params.*member); + config[field.name.substr(field.name.find_first_not_of('-'))] = params.*member; }, field.member); return config; diff --git a/arenai_agent/src/utils/cli_parser.cpp b/arenai_agent/src/utils/cli_parser.cpp deleted file mode 100644 index 0abc5a69..00000000 --- a/arenai_agent/src/utils/cli_parser.cpp +++ /dev/null @@ -1,72 +0,0 @@ -// -// Created by samuel on 11/03/2026. -// - -#include "./cli_parser.h" - -#include - -using namespace arenai; -using namespace arenai::agent; - -namespace arenai::agent { - - vision_channels parse_cli_vision_channels(const std::string &value) { - const std::regex regex_match( - R"(^ *\[(?: *\( *\d+ *, *\d+ *\) *,)* *\( *\d+ *, *\d+ *\) *] *$)"); - const std::regex regex_layer(R"(\( *\d+ *, *\d+ *\))"); - const std::regex regex_channel(R"(\d+)"); - - if (!std::regex_match(value.begin(), value.end(), regex_match)) - throw std::invalid_argument( - "invalid --vision_channels format, usage : [(10, 20), (20, 40), ...], actual value " - "= \"" - + value + "\""); - - vision_channels vision_channels; - - std::sregex_iterator it_layer(value.begin(), value.end(), regex_layer); - for (const std::sregex_iterator end; it_layer != end; ++it_layer) { - const auto layer_str = it_layer->str(); - - const std::sregex_iterator it_channel( - layer_str.begin(), layer_str.end(), regex_channel); - - const int c_i = std::stoi(it_channel->str()); - const int c_o = std::stoi(std::next(it_channel)->str()); - - vision_channels.channels.emplace_back(c_i, c_o); - } - - return vision_channels; - } - - std::vector parse_int_vector( - const std::string &value, const std::string &cli_arg_name, - const std::string &cli_arg_value_suggestion) { - const std::regex regex_match(R"(^ *\[(?: *\d+ *,)* *\d+ *] *$)"); - const std::regex regex_groups(R"(\d+)"); - - if (!std::regex_match(value.begin(), value.end(), regex_match)) - throw std::invalid_argument( - "invalid " + cli_arg_name + " format, usage : " + cli_arg_value_suggestion - + ", actual value = \"" + value + "\""); - - std::vector int_vector; - - std::sregex_iterator it_layer(value.begin(), value.end(), regex_groups); - for (const std::sregex_iterator end; it_layer != end; ++it_layer) - int_vector.emplace_back(std::stoi(it_layer->str())); - - return int_vector; - } - - group_norm_nums parse_cli_group_norms(const std::string &value) { - return {parse_int_vector(value, "group_norm_nums", "[4, 8, 16, ...]")}; - } - - hidden_layers parse_cli_hidden_layer(const std::string &value) { - return {parse_int_vector(value, "hidden_layers", "[256, 128, ..., 64]")}; - } - -}// namespace arenai::agent diff --git a/arenai_agent/src/utils/cli_parser.h b/arenai_agent/src/utils/cli_parser.h deleted file mode 100644 index a1ef4eac..00000000 --- a/arenai_agent/src/utils/cli_parser.h +++ /dev/null @@ -1,36 +0,0 @@ -// -// Created by samuel on 11/03/2026. -// - -#ifndef ARENAI_AGENT_HOST_CLI_PARSER_H -#define ARENAI_AGENT_HOST_CLI_PARSER_H - -#include -#include - -namespace arenai::agent { - - struct vision_channels { - std::vector> channels; - }; - - struct group_norm_nums { - std::vector groups; - }; - - struct hidden_layers { - std::vector layers; - }; - - vision_channels parse_cli_vision_channels(const std::string &value); - - group_norm_nums parse_cli_group_norms(const std::string &value); - hidden_layers parse_cli_hidden_layer(const std::string &value); - - std::vector parse_int_vector( - const std::string &value, const std::string &cli_arg_name, - const std::string &cli_arg_value_suggestion); - -}// namespace arenai::agent - -#endif//ARENAI_AGENT_HOST_CLI_PARSER_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo.h b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo.h new file mode 100644 index 00000000..43331bad --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo.h @@ -0,0 +1,44 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_TESTS_LIQUID_PPO_H +#define ARENAI_TESTS_LIQUID_PPO_H + +#include +#include + +#include +#include + +struct LiquidPpoTestConfig { + int vision_height; + int vision_width; + int nb_sensors; + int nb_continuous_actions; + int nb_discrete_actions; +}; + +class LiquidPpoAgentTest : public testing::Test { +protected: + void SetUp() override; + void TearDown() override; + + std::unique_ptr + make_factory(const LiquidPpoTestConfig &cfg) const; + + static arenai::agent::TorchState make_state(const LiquidPpoTestConfig &cfg, int batch); + + std::filesystem::path tmp_dir; + torch::Device device{torch::kCPU}; +}; + +typedef LiquidPpoTestConfig LiquidPpoActShapeParam; + +class LiquidPpoActShapeParamTest : public LiquidPpoAgentTest, + public testing::WithParamInterface {}; + +class LiquidPpoSaveLoadParamTest : public LiquidPpoAgentTest, + public testing::WithParamInterface {}; + +#endif//ARENAI_TESTS_LIQUID_PPO_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_rollout_buffer.h b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_rollout_buffer.h new file mode 100644 index 00000000..634d08ce --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_rollout_buffer.h @@ -0,0 +1,28 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_TESTS_LIQUID_PPO_ROLLOUT_BUFFER_H +#define ARENAI_TESTS_LIQUID_PPO_ROLLOUT_BUFFER_H + +#include +#include + +class LiquidPpoRolloutBufferTest : public testing::Test { +protected: + static constexpr int NB_TANKS = 2; + static constexpr int VISION_SIZE = 4; + static constexpr int NB_SENSORS = 3; + static constexpr int NB_CONTINUOUS_ACTIONS = 2; + static constexpr int NB_DISCRETE_ACTIONS = 2; + static constexpr int NEURON_NUMBER = 4; + + static arenai::agent::TorchState make_state(); + + static arenai::agent::LiquidPpoInputStep make_step( + const arenai::agent::TorchState &state, const torch::Tensor &done, bool episode_start); + + static arenai::agent::LiquidPpoInputStep make_step(const arenai::agent::TorchState &state); +}; + +#endif//ARENAI_TESTS_LIQUID_PPO_ROLLOUT_BUFFER_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_training.h b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_training.h new file mode 100644 index 00000000..6b58a0f5 --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_liquid_ppo_training.h @@ -0,0 +1,41 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_TESTS_LIQUID_PPO_TRAINING_H +#define ARENAI_TESTS_LIQUID_PPO_TRAINING_H + +#include + +#include +#include + +struct LiquidPpoTrainingTestConfig { + int vision_height; + int vision_width; + int nb_sensors; + int nb_continuous_actions; + int nb_discrete_actions; +}; + +class LiquidPpoTrainingTest : public testing::Test { +protected: + // small enough for the trainer to trigger during the test loop + static constexpr int ROLLOUT_SIZE = 4; + // smaller than the number of valid rows so the loop exercises several minibatches + static constexpr int MINIBATCH_SIZE = 4; + static constexpr int CHUNK_SIZE = 2; + static constexpr int NEURON_NUMBER = 16; + static constexpr int UNFOLDING_STEPS = 2; + static constexpr float DELTA_T = 1.f / 30.f; + + torch::Device device{torch::kCPU}; + + std::unique_ptr + make_factory(const LiquidPpoTrainingTestConfig &cfg) const; + + static arenai::agent::TorchState + make_state(const LiquidPpoTrainingTestConfig &cfg, int nb_tanks); +}; + +#endif//ARENAI_TESTS_LIQUID_PPO_TRAINING_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac.h b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac.h deleted file mode 100644 index f26ba22b..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac.h +++ /dev/null @@ -1,44 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_TESTS_SAC_H -#define ARENAI_TESTS_SAC_H - -#include -#include - -#include -#include - -struct SacTestConfig { - int vision_height; - int vision_width; - int nb_sensors; - int nb_continuous_actions; - int nb_discrete_actions; -}; - -class SacAgentTest : public testing::Test { -protected: - void SetUp() override; - void TearDown() override; - - std::unique_ptr - make_factory(const SacTestConfig &cfg) const; - - static arenai::agent::TorchState make_state(const SacTestConfig &cfg, int batch); - - std::filesystem::path tmp_dir; - torch::Device device{torch::kCPU}; -}; - -typedef SacTestConfig ActShapeParam; - -class SacActShapeParamTest : public SacAgentTest, - public testing::WithParamInterface {}; - -class SacSaveLoadParamTest : public SacAgentTest, - public testing::WithParamInterface {}; - -#endif//ARENAI_TESTS_SAC_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac_training.h b/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac_training.h deleted file mode 100644 index 4f36fd38..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_agents/tests_sac_training.h +++ /dev/null @@ -1,31 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#ifndef ARENAI_TESTS_SAC_TRAINING_H -#define ARENAI_TESTS_SAC_TRAINING_H - -#include - -#include -#include - -struct SacTrainingTestConfig { - int vision_height; - int vision_width; - int nb_sensors; - int nb_continuous_actions; - int nb_discrete_actions; -}; - -class SacTrainingTest : public testing::Test { -protected: - torch::Device device{torch::kCPU}; - - std::unique_ptr - make_factory(const SacTrainingTestConfig &cfg) const; - - static arenai::agent::TorchState make_state(const SacTrainingTestConfig &cfg); -}; - -#endif//ARENAI_TESTS_SAC_TRAINING_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli.h b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli.h new file mode 100644 index 00000000..0aa6581f --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli.h @@ -0,0 +1,16 @@ +// +// Created by samuel on 18/09/2026. +// + +#ifndef ARENAI_TESTS_BERNOULLI_H +#define ARENAI_TESTS_BERNOULLI_H + +#include + +typedef int NbActions; + +class BernoulliTest : public testing::Test {}; +class BernoulliShapeParamTest + : public testing::TestWithParam> {}; + +#endif//ARENAI_TESTS_BERNOULLI_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli_edge.h b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli_edge.h new file mode 100644 index 00000000..34ff5bd3 --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_bernoulli_edge.h @@ -0,0 +1,12 @@ +// +// Created by samuel on 18/09/2026. +// + +#ifndef ARENAI_TESTS_BERNOULLI_EDGE_H +#define ARENAI_TESTS_BERNOULLI_EDGE_H + +#include + +class BernoulliEdgeTest : public testing::Test {}; + +#endif//ARENAI_TESTS_BERNOULLI_EDGE_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial.h b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial.h deleted file mode 100644 index 22147aac..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial.h +++ /dev/null @@ -1,19 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_TESTS_MULTINOMIAL_H -#define ARENAI_TESTS_MULTINOMIAL_H - -#include - -typedef int NbActions; -typedef float ShootProbability; - -class MultinomialTest : public testing::Test {}; -class MultinomialShapeParamTest - : public testing::TestWithParam> {}; -class MultinomialMaxEntropyParamTest : public testing::TestWithParam {}; -class MultinomialTargetEntropyParamTest : public testing::TestWithParam {}; - -#endif//ARENAI_TESTS_MULTINOMIAL_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial_edge.h b/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial_edge.h deleted file mode 100644 index 8a4f135f..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_distributions/tests_multinomial_edge.h +++ /dev/null @@ -1,12 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#ifndef ARENAI_TESTS_MULTINOMIAL_EDGE_H -#define ARENAI_TESTS_MULTINOMIAL_EDGE_H - -#include - -class MultinomialEdgeTest : public testing::Test {}; - -#endif//ARENAI_TESTS_MULTINOMIAL_EDGE_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_entropy.h b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_entropy.h index 479c5651..21532f1b 100644 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_entropy.h +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_entropy.h @@ -7,12 +7,6 @@ #include -class AlphaParameterTest : public testing::Test {}; - -class ConstantTargetEntropyTest : public testing::Test {}; - -class CosineAnnealingTargetEntropyTest : public testing::Test {}; - class PidLagrangianAlphaParameterTest : public testing::Test {}; #endif//ARENAI_TESTS_ENTROPY_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_liquid_cell.h b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_liquid_cell.h new file mode 100644 index 00000000..554c1192 --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_liquid_cell.h @@ -0,0 +1,24 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_TESTS_LIQUID_CELL_H +#define ARENAI_TESTS_LIQUID_CELL_H + +#include + +typedef int NeuronNumber; +typedef int InputSize; +typedef int OutputSize; +typedef int UnfoldingSteps; +typedef int BatchSize; +typedef int TimeSteps; + +class CellModelTestParam + : public testing::TestWithParam> {}; + +class LiquidCellTestParam + : public testing::TestWithParam< + std::tuple> {}; + +#endif//ARENAI_TESTS_LIQUID_CELL_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function.h b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function.h deleted file mode 100644 index 44560098..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function.h +++ /dev/null @@ -1,25 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_TESTS_Q_FUNCTION_H -#define ARENAI_TESTS_Q_FUNCTION_H - -#include - -typedef std::vector HiddenLayers; -typedef int ContinuousActionsNb; -typedef int DiscreteActionsNb; - -typedef int SensorsNb; -typedef int SensorsHiddenSize; - -typedef int ActionsHiddenSize; - -typedef int BatchSize; - -class QFunctionTestParam : public testing::TestWithParam> {}; - -#endif//ARENAI_TESTS_Q_FUNCTION_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function_consistency.h b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function_consistency.h deleted file mode 100644 index bae3b46e..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_q_function_consistency.h +++ /dev/null @@ -1,14 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#ifndef ARENAI_TESTS_Q_FUNCTION_CONSISTENCY_H -#define ARENAI_TESTS_Q_FUNCTION_CONSISTENCY_H - -#include - -class QFunctionConsistencyTest : public testing::Test {}; -class QFunctionGradientTest : public testing::Test {}; -class ActorGradientTest : public testing::Test {}; - -#endif//ARENAI_TESTS_Q_FUNCTION_CONSISTENCY_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_vision_impala.h b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_vision_impala.h new file mode 100644 index 00000000..4f630429 --- /dev/null +++ b/arenai_agent/tests/include/arenai_agent_tests/tests_networks/tests_vision_impala.h @@ -0,0 +1,24 @@ +// +// Created by claude on 20/09/2026. +// + +#ifndef ARENAI_TESTS_VISION_IMPALA_H +#define ARENAI_TESTS_VISION_IMPALA_H + +#include + +typedef int ImpalaVisionWidth; +typedef int ImpalaVisionHeight; +typedef int ImpalaVisionChannel; + +typedef std::vector ImpalaOutputConvChannels; + +typedef int ImpalaBatchSize; + +class ImpalaVisionTestParam : public testing::TestWithParam> {}; + +class ImpalaVisionEdgeTest : public testing::Test {}; + +#endif//ARENAI_TESTS_VISION_IMPALA_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/replay_buffer_test_param.h b/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/replay_buffer_test_param.h deleted file mode 100644 index 470fbd30..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/replay_buffer_test_param.h +++ /dev/null @@ -1,16 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_REPLAY_BUFFER_TEST_PARAM_H -#define ARENAI_REPLAY_BUFFER_TEST_PARAM_H - -#include - -typedef uint32_t MemorySize; - -template -class ReplayBufferTestParam : public testing::TestWithParam> { -}; - -#endif//ARENAI_REPLAY_BUFFER_TEST_PARAM_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_add.h b/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_add.h deleted file mode 100644 index 1926197d..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_add.h +++ /dev/null @@ -1,15 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_TEST_REPLAY_BUFFER_ADD_H -#define ARENAI_TEST_REPLAY_BUFFER_ADD_H - -#include "./replay_buffer_test_param.h" - -typedef uint32_t StepsNbToAdd; - -class ReplayBufferAddNormalTestParam : public ReplayBufferTestParam {}; -class ReplayBufferAddOverflowTestParam : public ReplayBufferTestParam {}; - -#endif//ARENAI_TEST_REPLAY_BUFFER_ADD_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_edge.h b/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_edge.h deleted file mode 100644 index 5461c55d..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_edge.h +++ /dev/null @@ -1,12 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#ifndef ARENAI_TESTS_REPLAY_BUFFER_EDGE_H -#define ARENAI_TESTS_REPLAY_BUFFER_EDGE_H - -#include - -class ReplayBufferEdgeTest : public testing::Test {}; - -#endif//ARENAI_TESTS_REPLAY_BUFFER_EDGE_H diff --git a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_sample.h b/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_sample.h deleted file mode 100644 index 3be37865..00000000 --- a/arenai_agent/tests/include/arenai_agent_tests/tests_replay_buffer/tests_replay_buffer_sample.h +++ /dev/null @@ -1,21 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_TESTS_REPLAY_BUFFER_SAMPLE_H -#define ARENAI_TESTS_REPLAY_BUFFER_SAMPLE_H - -#include "./replay_buffer_test_param.h" - -typedef uint32_t BatchSize; -typedef uint32_t StepsNbToAdd; - -class ReplayBufferSampleNormalTestParam : public ReplayBufferTestParam {}; -class ReplayBufferSampleOverflowTestParam : public ReplayBufferTestParam { -}; -class ReplayBufferSampleUnderflowTestParam : public ReplayBufferTestParam { -}; -class ReplayBufferSampleDoubleOverflowTestParam - : public ReplayBufferTestParam {}; - -#endif//ARENAI_TESTS_REPLAY_BUFFER_SAMPLE_H diff --git a/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_0.json b/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_0.json index 3a7a886a..5ca4ab2c 100644 --- a/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_0.json +++ b/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_0.json @@ -1 +1 @@ -[93,93,93,94,95,95,95,161,153,147,148,62,71,124,132,137,93,93,93,93,94,94,95,95,170,150,62,68,71,67,132,150,93,93,93,93,93,94,94,95,95,57,68,70,69,60,116,139,92,92,92,92,94,94,94,94,94,91,41,69,63,49,105,120,92,92,92,94,94,94,94,94,94,91,91,93,55,93,97,99,92,92,94,94,94,94,94,93,93,93,91,93,93,95,86,100,92,94,94,94,94,94,94,93,93,93,92,94,94,94,88,88,93,93,93,94,94,90,93,94,46,94,94,94,94,93,94,88,90,90,90,90,90,90,38,48,47,63,94,94,94,93,93,94,90,90,90,90,90,151,152,152,152,152,152,94,94,93,93,90,86,90,90,90,154,110,110,110,110,108,154,137,94,93,90,90,86,90,90,25,154,98,98,98,98,98,154,159,85,85,90,90,86,86,90,22,154,98,98,98,98,98,154,128,84,85,85,90,86,86,86,18,154,121,76,76,76,76,154,105,81,84,85,85,74,74,74,73,14,76,76,76,76,76,73,126,81,81,81,85,74,74,74,73,73,73,73,73,76,76,76,76,76,81,81,81,118,118,118,120,120,120,120,181,176,171,170,14,14,155,163,165,118,118,118,118,120,120,120,120,188,172,14,14,14,14,163,176,118,118,118,118,118,120,120,120,120,13,13,13,14,14,147,166,117,117,117,118,120,120,120,120,119,116,12,13,13,14,136,151,117,117,117,120,120,120,120,120,119,115,115,118,13,129,133,137,117,117,120,120,120,120,120,119,119,119,115,118,118,131,118,137,117,119,120,120,120,120,119,119,119,119,117,120,120,119,122,120,119,119,119,119,120,115,119,119,140,120,120,120,120,119,119,120,115,115,115,115,115,115,117,147,144,194,119,119,120,119,119,119,115,115,115,115,115,126,127,127,127,127,127,119,119,119,119,115,110,115,115,115,128,93,93,93,93,92,129,115,119,118,114,115,110,115,115,51,128,84,83,83,83,83,129,23,108,108,114,114,110,110,115,46,128,84,83,83,83,83,128,120,108,108,108,114,110,110,110,36,128,102,98,98,98,98,128,100,104,108,108,108,96,96,96,95,27,98,98,98,98,98,71,119,105,105,105,108,96,96,96,95,95,95,95,95,98,98,98,99,99,105,105,105,205,205,205,207,207,207,207,195,189,185,184,198,217,173,181,183,205,205,205,205,207,207,207,207,198,186,197,212,217,210,181,191,205,205,205,206,206,207,207,207,207,188,212,216,215,193,165,183,205,205,205,205,208,208,208,208,206,201,154,214,202,169,159,171,205,205,205,208,208,208,208,208,206,202,201,204,184,153,157,160,205,205,208,208,208,208,208,206,206,206,201,204,204,156,143,160,205,208,208,208,208,208,208,206,206,206,204,207,207,206,149,145,208,208,208,208,208,202,207,208,120,208,207,207,207,206,206,145,202,202,203,203,203,203,108,124,122,147,208,207,207,206,206,206,203,203,203,203,203,170,171,171,171,171,171,208,207,206,206,201,197,203,203,203,172,140,140,140,140,138,172,160,208,206,201,201,197,203,203,88,172,131,131,131,131,131,172,217,194,194,201,201,197,197,203,83,172,131,131,131,131,131,172,126,194,194,194,201,198,197,197,71,172,148,182,182,182,183,172,111,190,195,195,195,180,180,180,178,60,183,183,183,183,183,91,125,190,190,190,195,181,181,180,179,178,178,178,178,183,183,183,183,183,191,191,190] +[40,39,39,39,39,39,39,66,63,60,60,61,62,61,61,61,40,40,39,39,39,39,39,39,67,60,58,61,60,62,60,60,40,40,40,40,39,39,39,39,39,39,64,57,59,58,62,61,40,40,40,40,40,39,39,39,39,39,39,65,58,65,59,68,40,40,40,40,40,40,40,39,39,39,39,39,63,66,62,70,41,40,40,40,40,40,40,40,39,39,39,39,39,39,60,66,41,41,41,40,40,40,40,40,40,40,39,39,39,39,39,60,42,42,41,41,40,40,40,40,40,40,40,39,39,39,39,39,138,132,42,42,41,41,40,40,40,40,40,74,40,39,39,39,130,63,63,63,62,62,62,62,62,62,63,63,63,63,39,39,62,62,62,62,62,62,62,62,62,63,63,63,63,62,62,62,63,63,63,63,63,63,63,65,58,58,58,64,64,63,66,66,63,63,63,63,63,63,64,65,63,58,58,58,64,63,63,63,63,63,63,63,63,62,64,65,63,58,58,58,64,63,63,63,63,63,63,63,63,62,64,64,63,63,58,58,64,63,63,63,63,63,63,63,63,62,64,64,36,63,63,58,64,63,63,63,53,53,53,53,53,53,53,77,72,69,67,68,69,68,68,68,53,53,53,53,53,53,53,53,78,71,65,68,67,69,67,67,53,53,53,53,53,53,53,53,53,52,76,69,69,65,70,70,54,53,53,53,53,53,53,53,53,53,52,73,65,75,66,77,54,54,53,53,53,53,53,53,53,53,53,52,72,75,70,80,54,54,54,54,53,53,53,53,53,53,53,53,53,52,68,78,55,54,54,54,54,53,53,53,53,53,53,53,53,53,52,69,56,55,55,54,54,54,54,53,53,53,53,53,53,53,53,52,166,159,56,55,55,54,54,54,54,53,53,232,53,53,53,53,158,198,198,198,196,197,197,197,197,197,198,199,199,199,53,53,196,196,197,197,197,197,197,197,197,198,198,198,198,196,196,196,198,198,198,198,198,198,198,207,185,185,185,201,201,199,207,207,198,198,198,198,198,198,203,207,198,185,185,185,201,199,199,199,198,198,198,198,198,197,203,207,198,185,185,185,201,199,199,199,198,198,198,198,198,197,203,203,198,198,185,185,202,199,199,199,198,198,198,198,198,197,203,203,118,198,198,185,202,199,199,199,127,127,127,127,127,127,126,92,87,84,83,86,87,86,86,86,127,127,127,127,127,127,127,126,87,85,83,84,83,87,85,85,127,127,127,127,127,127,127,127,126,126,88,83,85,81,87,86,127,127,127,127,127,127,127,127,127,126,126,90,81,86,82,92,127,127,127,127,127,127,127,127,127,127,127,126,87,90,83,91,127,127,127,127,127,127,127,127,127,127,127,127,126,126,83,90,128,127,127,127,127,127,127,127,127,127,127,127,127,126,126,84,128,128,127,127,127,127,127,127,127,127,127,127,127,127,127,126,183,176,128,128,127,127,127,127,127,127,127,165,127,127,127,127,170,148,148,148,147,148,148,148,148,148,148,149,149,149,127,127,147,147,147,147,147,147,148,148,148,148,148,148,148,147,147,147,148,148,148,148,148,148,148,153,141,141,141,150,150,149,153,153,148,148,148,148,148,148,151,153,148,141,141,141,150,149,149,149,148,148,148,148,148,147,151,153,148,141,141,141,150,149,149,149,148,148,148,148,148,148,151,151,148,148,141,141,150,149,149,149,148,148,148,148,148,148,151,151,108,148,148,141,150,149,149,149] diff --git a/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_1.json b/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_1.json index 8dddac6b..f0a1f5f8 100644 --- a/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_1.json +++ b/arenai_agent/tests/resources/golden_images/golden_agent_input_step30_tank_1.json @@ -1 +1 @@ -[56,56,56,54,53,53,53,53,52,52,52,50,50,49,49,49,55,54,54,53,53,53,53,52,50,50,50,50,49,49,49,48,54,54,53,53,53,53,52,50,50,50,50,49,49,48,48,48,54,55,53,53,52,94,100,77,165,50,49,48,50,50,50,50,55,55,53,119,117,117,118,155,146,194,50,50,50,50,50,49,55,100,149,118,117,117,117,131,124,114,50,50,49,49,49,49,53,81,97,118,82,82,82,82,107,103,176,49,49,49,49,48,53,67,99,149,82,82,82,82,80,160,149,49,49,52,50,50,52,57,54,87,149,82,82,82,82,134,127,208,52,52,52,50,52,52,52,21,19,148,82,82,82,82,82,110,52,52,51,51,52,52,52,18,17,24,148,82,82,82,82,86,137,51,51,51,52,52,48,48,48,74,108,148,51,51,51,51,51,51,51,51,52,48,48,48,48,48,51,51,51,51,51,51,51,51,51,51,51,50,50,50,50,53,49,49,48,48,48,48,48,48,48,48,50,50,50,50,53,53,49,48,48,48,48,48,48,48,48,48,50,50,50,53,53,53,53,48,48,48,48,48,48,48,48,48,12,11,10,10,9,9,8,8,7,7,7,7,6,6,6,6,11,10,10,9,9,8,8,7,7,7,7,6,6,6,6,6,10,10,9,9,8,8,7,7,7,6,6,6,6,6,5,5,10,9,9,8,8,119,127,99,122,6,6,6,6,5,5,5,9,9,8,149,147,147,147,115,108,142,6,5,5,5,5,5,9,127,185,149,147,147,147,98,93,86,5,5,5,5,5,5,8,132,123,149,105,105,105,105,81,79,44,5,5,5,4,4,8,111,160,185,105,105,105,105,102,41,38,5,5,4,4,4,8,95,90,142,185,105,105,105,105,35,33,52,4,4,4,4,7,7,7,143,133,184,105,105,105,105,105,25,4,4,4,4,7,7,6,121,116,170,184,105,105,105,105,20,30,4,4,4,7,7,6,6,5,7,9,184,4,4,4,4,4,4,4,3,7,6,6,5,5,5,4,4,4,4,4,4,4,3,3,3,6,6,5,5,5,4,4,4,4,4,4,3,3,3,3,3,6,5,5,5,4,4,4,4,4,4,4,3,3,3,3,3,5,5,5,4,4,4,4,4,4,4,3,3,3,3,3,3,103,104,103,101,101,100,100,100,99,99,99,97,97,96,96,96,101,101,101,100,100,100,100,99,97,97,97,97,96,96,96,96,101,101,100,100,100,100,100,97,97,97,97,96,96,95,95,95,101,103,100,100,100,6,5,5,160,97,96,95,97,97,97,97,103,103,100,6,6,6,6,154,149,177,97,97,97,97,97,97,102,6,6,6,6,6,6,140,136,130,97,97,97,97,97,97,100,201,6,6,6,5,5,5,126,124,21,97,97,97,97,96,100,181,229,6,6,6,6,6,5,20,19,97,97,100,97,97,99,165,161,211,7,6,6,6,6,18,18,23,100,100,100,97,99,99,99,28,27,7,6,6,6,6,6,78,100,100,100,100,99,99,99,26,25,31,7,6,6,6,6,68,89,99,99,99,99,99,95,95,95,90,113,7,99,99,99,99,99,99,99,99,99,95,95,95,95,95,99,99,99,99,99,99,99,99,99,99,98,98,97,97,97,101,96,96,96,96,96,96,96,96,96,96,98,97,97,97,101,101,96,96,96,96,96,96,96,96,96,96,97,97,97,101,101,101,101,96,96,96,96,96,96,96,96,96] +[136,140,137,146,136,132,124,110,118,133,151,208,155,228,144,195,136,136,156,160,152,137,128,118,134,149,145,155,234,212,198,171,142,132,125,133,119,112,116,123,142,143,143,145,158,156,180,166,101,141,136,125,123,123,123,134,142,141,140,155,152,144,147,161,83,89,92,104,135,121,122,124,127,135,140,146,59,77,77,148,75,80,82,83,85,90,89,126,127,56,58,58,61,67,68,68,75,98,27,77,71,70,70,56,55,58,58,61,67,67,67,55,69,99,98,77,78,63,55,59,58,61,61,61,67,55,56,54,72,73,73,63,61,54,59,57,60,60,60,55,55,55,53,53,61,68,68,68,59,59,59,59,59,63,63,55,53,53,53,51,69,69,56,63,63,59,63,63,63,63,55,53,51,51,59,59,63,63,62,62,63,63,23,30,30,36,51,51,59,59,59,59,66,66,66,66,66,142,141,139,140,140,140,58,58,58,58,58,60,60,60,60,140,140,155,140,141,137,156,156,58,58,58,58,66,66,66,66,156,155,102,119,102,102,156,157,54,67,67,67,66,66,66,103,155,155,102,102,102,102,156,156,204,67,67,67,154,157,155,166,158,152,146,133,140,152,169,219,172,233,165,206,153,153,170,177,170,161,157,144,156,167,164,173,240,219,207,185,158,149,146,156,148,142,144,146,161,162,162,164,175,172,191,182,87,158,153,149,151,151,146,156,161,160,159,173,170,163,165,177,51,61,61,81,158,150,150,151,150,154,157,164,12,11,12,165,30,40,43,52,56,63,66,154,154,13,12,11,10,10,10,10,20,26,147,22,22,23,23,13,11,11,10,10,10,10,10,9,16,25,25,20,20,13,12,10,10,10,9,9,9,9,8,8,13,13,13,12,12,10,10,9,9,9,9,8,8,8,7,7,12,11,11,11,10,9,9,9,8,8,8,8,7,7,7,6,10,10,9,9,8,8,8,8,8,7,7,7,6,6,6,6,7,7,8,8,8,7,26,34,34,40,6,6,6,5,5,5,6,6,6,6,6,176,175,173,174,175,175,5,5,5,5,5,5,5,5,5,174,174,192,175,175,170,193,193,5,5,5,5,4,4,5,5,193,192,129,150,129,129,194,194,4,4,4,4,4,4,4,167,193,193,129,129,129,129,193,193,149,4,4,4,168,172,169,177,171,163,160,147,154,167,181,226,186,237,179,214,169,169,182,186,180,173,171,159,170,181,178,187,246,223,212,195,173,165,165,172,160,152,157,160,176,176,177,178,188,185,200,194,139,173,169,165,163,163,160,170,176,175,174,187,184,177,179,190,126,131,134,143,172,164,164,164,164,169,173,178,106,128,128,181,122,125,127,126,127,131,129,169,168,103,106,106,110,117,117,117,124,94,79,125,119,117,117,103,102,106,106,109,117,117,117,102,117,95,94,127,127,111,102,107,106,109,109,109,116,102,103,101,122,123,123,111,109,101,107,105,109,109,109,103,103,103,101,101,109,117,117,117,107,107,108,108,108,113,113,103,101,101,101,98,119,119,103,112,112,107,112,112,112,112,103,101,98,98,108,108,113,113,112,112,112,112,31,36,36,41,98,98,107,108,108,108,117,117,117,117,117,7,7,7,7,7,7,107,107,107,107,107,109,109,109,109,6,6,7,6,6,6,7,7,107,107,107,107,117,117,117,117,7,7,6,6,6,6,7,7,102,118,118,118,117,117,117,236,7,7,6,6,6,6,7,7,183,118,118,118] diff --git a/arenai_agent/tests/src/tests_agents/probe_act_sample.cpp b/arenai_agent/tests/src/tests_agents/probe_act_sample.cpp deleted file mode 100644 index 4d172dac..00000000 --- a/arenai_agent/tests/src/tests_agents/probe_act_sample.cpp +++ /dev/null @@ -1,123 +0,0 @@ -// TEMPORARY diagnostic probe - to be deleted -#include - -#include -#include -#include -#include - -using namespace arenai; -using namespace arenai::agent; - -namespace { - - TorchState probe_state(const int batch, const int h, const int w, const int nb_sensors) { - torch::manual_seed(42); - return { - .vision = torch::randint(0, 255, {batch, 3, h, w}, torch::kUInt8), - .proprioception = torch::randn({batch, nb_sensors})}; - } - -}// namespace - -TEST(ProbeActSample, SacStatsBiasedMu) { - torch::manual_seed(1234); - constexpr int h = 8, w = 8, nb_sensors = 4, nb_cont = 2, nb_disc = 2; - - const auto actor = std::make_shared( - h, w, nb_sensors, nb_cont, nb_disc, 8, std::vector{16}, - std::vector>{{3, 4}}, std::vector{2}, 0.3f, 0.1f); - - // push mu away from 0 : mu ~ tanh(bias) - { - torch::NoGradGuard guard; - for (auto &p: actor->named_parameters()) - if (p.key().find("mu") != std::string::npos - && p.key().find("bias") != std::string::npos) - p.value().copy_(torch::tensor({std::atanh(0.7f), std::atanh(-0.4f)})); - } - - const auto agent = std::make_shared(actor, torch::Device(torch::kCPU)); - const auto state = probe_state(1, h, w, nb_sensors); - - torch::NoGradGuard guard; - const auto [mu, sigma, disc] = actor->act(state.vision, state.proprioception); - std::cout << "mu: " << mu << "\nsigma: " << sigma << std::endl; - - constexpr int N = 4000; - std::vector conts; - for (int i = 0; i < N; i++) conts.push_back(agent->act(state, true).continuous_action); - const auto cont_all = torch::cat(conts, 0); - std::cout << "act(true) cont mean: " << cont_all.mean(0) - << "\nact(true) cont std: " << cont_all.std(0) << std::endl; -} - -TEST(ProbeActSample, PpoLogProbs) { - torch::manual_seed(99); - constexpr int h = 8, w = 8, nb_sensors = 4, nb_cont = 2, nb_disc = 2; - - const auto actor = std::make_shared( - h, w, nb_sensors, nb_cont, nb_disc, 8, std::vector{16}, - std::vector>{{3, 4}}, std::vector{2}, 0.3f, 0.2f); - const auto rollout_buffer = std::make_shared(); - const auto collector = std::make_shared(rollout_buffer); - const auto agent = - std::make_shared(actor, torch::Device(torch::kCPU), collector); - - const auto state = probe_state(3, h, w, nb_sensors); - - const auto [continuous_action, discrete_action] = agent->act(state, true); - collector->on_transition(torch::randn({3, 1}), torch::zeros({3, 1})); - collector->on_episode_end(state); - - torch::NoGradGuard guard; - const auto [mu, sigma, disc] = actor->act(state.vision, state.proprioception); - const auto expected_cont = truncated_normal_log_pdf(continuous_action, mu, sigma).sum(-1, true); - const auto expected_disc = - (discrete_action * torch::log(torch::clamp(disc, 1e-8, 1.0 - 1e-8))).sum(-1, true); - - const auto rollout = rollout_buffer->get_rollout(); - std::cout << "stored cont lp: " << rollout.continuous_log_probs - << "\nexpected cont lp: " << expected_cont - << "\nstored disc lp: " << rollout.discrete_log_probs - << "\nexpected disc lp: " << expected_disc << std::endl; -} - -TEST(ProbeActSample, SacStats) { - torch::manual_seed(1234); - constexpr int h = 8, w = 8, nb_sensors = 4, nb_cont = 3, nb_disc = 2; - - const auto actor = std::make_shared( - h, w, nb_sensors, nb_cont, nb_disc, 8, std::vector{16}, - std::vector>{{3, 4}}, std::vector{2}, 0.4f, 0.1f); - const auto agent = std::make_shared(actor, torch::Device(torch::kCPU)); - - const auto state = probe_state(1, h, w, nb_sensors); - - // raw actor output - torch::NoGradGuard guard; - const auto [mu, sigma, disc] = actor->act(state.vision, state.proprioception); - std::cout << "mu: " << mu << "\nsigma: " << sigma << "\ndisc_proba: " << disc << std::endl; - - // deterministic - const auto [continuous_action, discrete_action] = agent->act(state, false); - std::cout << "act(false) cont: " << continuous_action - << "\nact(false) disc: " << discrete_action << std::endl; - std::cout << "cont == mu ? " << torch::allclose(continuous_action, mu) << std::endl; - - // stochastic stats - constexpr int N = 2000; - std::vector conts, discs; - for (int i = 0; i < N; i++) { - const auto [curr_continuous_action, curr_discrete_action] = agent->act(state, true); - conts.push_back(curr_continuous_action); - discs.push_back(curr_discrete_action); - } - const auto cont_all = torch::cat(conts, 0); - const auto disc_all = torch::cat(discs, 0); - std::cout << "act(true) cont mean: " << cont_all.mean(0) - << "\nact(true) cont std: " << cont_all.std(0) - << "\nact(true) cont min: " << std::get<0>(cont_all.min(0)) - << "\nact(true) cont max: " << std::get<0>(cont_all.max(0)) - << "\nact(true) disc freq: " << disc_all.mean(0) << std::endl; -} diff --git a/arenai_agent/tests/src/tests_agents/tests_sac.cpp b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo.cpp similarity index 59% rename from arenai_agent/tests/src/tests_agents/tests_sac.cpp rename to arenai_agent/tests/src/tests_agents/tests_liquid_ppo.cpp index d812a2d9..fc093f4e 100644 --- a/arenai_agent/tests/src/tests_agents/tests_sac.cpp +++ b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo.cpp @@ -1,8 +1,8 @@ // -// Created by samuel on 30/06/2026. +// Created by samuel on 06/09/2026. // -#include +#include using namespace arenai; using namespace arenai::agent; @@ -11,38 +11,41 @@ using namespace arenai::agent; // Fixture helpers // ======================================================================== -void SacAgentTest::SetUp() { - tmp_dir = std::filesystem::temp_directory_path() / "arenai_test_sac"; +void LiquidPpoAgentTest::SetUp() { + tmp_dir = std::filesystem::temp_directory_path() / "arenai_test_liquid_ppo"; std::filesystem::create_directories(tmp_dir); } -void SacAgentTest::TearDown() { std::filesystem::remove_all(tmp_dir); } +void LiquidPpoAgentTest::TearDown() { std::filesystem::remove_all(tmp_dir); } -std::unique_ptr SacAgentTest::make_factory(const SacTestConfig &cfg) const { - const SacHyperParams params{ +std::unique_ptr +LiquidPpoAgentTest::make_factory(const LiquidPpoTestConfig &cfg) const { + const LiquidPpoHyperParams params{ .actor_learning_rate = 1e-3f, .critic_learning_rate = 1e-3f, - .alpha_learning_rate = 1e-3f, .hidden_size_sensors = 16, - .hidden_size_actions = 16, - .actor_hidden_sizes = {32}, - .critic_hidden_sizes = {32}, .vision_channels = {{3, 8}}, .group_norm_nums = {4}, + .neuron_number = 16, + .unfolding_steps = 2, + .chunk_size = 2, .metric_window_size = 10, - .tau = 0.005f, .gamma = 0.99f, - .replay_buffer_size = 10, - .train_every = 1, + .gae_lambda = 0.95f, + .clip_epsilon = 0.2f, + .grad_norm_max = 1.f, + .continuous_target_entropy = std::vector(cfg.nb_continuous_actions, -0.88f), + .discrete_target_entropy_factors = std::vector(cfg.nb_discrete_actions, 0.2f), .epochs = 1, - .batch_size = 1}; + .rollout_size = 8, + .minibatch_size = 10}; - return std::make_unique( + return std::make_unique( cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, cfg.nb_discrete_actions, device, params); } -TorchState SacAgentTest::make_state(const SacTestConfig &cfg, const int batch) { +TorchState LiquidPpoAgentTest::make_state(const LiquidPpoTestConfig &cfg, const int batch) { return { .vision = torch::randint(0, 255, {batch, 3, cfg.vision_height, cfg.vision_width}, torch::kUInt8), @@ -53,8 +56,8 @@ TorchState SacAgentTest::make_state(const SacTestConfig &cfg, const int batch) { // Fixed tests // ======================================================================== -TEST_F(SacAgentTest, ParameterCountPositive) { - constexpr SacTestConfig cfg{ +TEST_F(LiquidPpoAgentTest, ParameterCountPositive) { + constexpr LiquidPpoTestConfig cfg{ .vision_height = 8, .vision_width = 8, .nb_sensors = 10, @@ -65,8 +68,8 @@ TEST_F(SacAgentTest, ParameterCountPositive) { ASSERT_GT(factory->get_trainer()->count_parameters(), 0); } -TEST_F(SacAgentTest, MetricsNotEmpty) { - constexpr SacTestConfig cfg{ +TEST_F(LiquidPpoAgentTest, MetricsNotEmpty) { + constexpr LiquidPpoTestConfig cfg{ .vision_height = 8, .vision_width = 8, .nb_sensors = 10, @@ -76,14 +79,14 @@ TEST_F(SacAgentTest, MetricsNotEmpty) { const auto metrics = factory->get_trainer()->get_metrics(); - ASSERT_EQ(metrics.size(), 12); + ASSERT_EQ(metrics.size(), 10); } // ======================================================================== // Parameterized: act shape tests // ======================================================================== -TEST_P(SacActShapeParamTest, ActOutputShapes) { +TEST_P(LiquidPpoActShapeParamTest, ActOutputShapes) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -98,7 +101,7 @@ TEST_P(SacActShapeParamTest, ActOutputShapes) { ASSERT_EQ(discrete_action.size(1), cfg.nb_discrete_actions); } -TEST_P(SacActShapeParamTest, ActContinuousFinite) { +TEST_P(LiquidPpoActShapeParamTest, ActContinuousFinite) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -109,7 +112,7 @@ TEST_P(SacActShapeParamTest, ActContinuousFinite) { ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); } -TEST_P(SacActShapeParamTest, ActDiscreteIsOneHot) { +TEST_P(LiquidPpoActShapeParamTest, ActDiscreteIsBinary) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -117,25 +120,42 @@ TEST_P(SacActShapeParamTest, ActDiscreteIsOneHot) { const auto [continuous_action, discrete_action] = factory->get_agent()->act(make_state(cfg, batch), true); - const auto row_sums = torch::sum(discrete_action, -1); - ASSERT_TRUE(torch::allclose(row_sums, torch::ones({batch}))); - + // independent Bernoulli actions: each entry is 0 or 1, no one-hot constraint const auto is_binary = torch::logical_or(torch::eq(discrete_action, 0.0f), torch::eq(discrete_action, 1.0f)); ASSERT_TRUE(torch::all(is_binary).item()); } +TEST_P(LiquidPpoActShapeParamTest, ActKeepsShapesOverConsecutiveSteps) { + const auto cfg = GetParam(); + const auto factory = make_factory(cfg); + + // the liquid state is carried across the calls: every step must stay consistent + constexpr int batch = 4; + for (int t = 0; t < 3; t++) { + const auto [continuous_action, discrete_action] = + factory->get_agent()->act(make_state(cfg, batch), true); + + ASSERT_EQ(continuous_action.size(0), batch); + ASSERT_EQ(continuous_action.size(1), cfg.nb_continuous_actions); + ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); + + ASSERT_EQ(discrete_action.size(0), batch); + ASSERT_EQ(discrete_action.size(1), cfg.nb_discrete_actions); + } +} + INSTANTIATE_TEST_SUITE_P( - SacAgent, SacActShapeParamTest, + LiquidPpoAgent, LiquidPpoActShapeParamTest, testing::Values( - SacTestConfig{8, 8, 10, 4, 2}, SacTestConfig{8, 8, 5, 2, 3}, - SacTestConfig{16, 16, 20, 6, 4}, SacTestConfig{8, 12, 10, 4, 2})); + LiquidPpoTestConfig{8, 8, 10, 4, 2}, LiquidPpoTestConfig{8, 8, 5, 2, 2}, + LiquidPpoTestConfig{16, 16, 20, 6, 2}, LiquidPpoTestConfig{8, 12, 10, 4, 2})); // ======================================================================== // Parameterized: save / load tests // ======================================================================== -TEST_P(SacSaveLoadParamTest, SaveCreatesExpectedFiles) { +TEST_P(LiquidPpoSaveLoadParamTest, SaveCreatesExpectedFiles) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -145,27 +165,15 @@ TEST_P(SacSaveLoadParamTest, SaveCreatesExpectedFiles) { factory->get_trainer()->save(save_dir); const std::vector expected_files = { - "actor.pt", - "critic_1.pt", - "critic_2.pt", - "target_critic_1.pt", - "target_critic_2.pt", - "alpha_continuous.pt", - "alpha_discrete.pt", - "actor_optim.pt", - "critic_1_optim.pt", - "critic_2_optim.pt", - "alpha_continuous_optim.pt", - "alpha_discrete_optim.pt", - "actor_repr.txt", - "critic_repr.txt", + "actor.pt", "critic.pt", "actor_optim.pt", + "critic_optim.pt", "actor_repr.txt", "critic_repr.txt", }; for (const auto &f: expected_files) ASSERT_TRUE(std::filesystem::exists(save_dir / f)) << "Missing file: " << f; } -TEST_P(SacSaveLoadParamTest, SavedFilesNonEmpty) { +TEST_P(LiquidPpoSaveLoadParamTest, SavedFilesNonEmpty) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -181,7 +189,7 @@ TEST_P(SacSaveLoadParamTest, SavedFilesNonEmpty) { } INSTANTIATE_TEST_SUITE_P( - SacAgent, SacSaveLoadParamTest, + LiquidPpoAgent, LiquidPpoSaveLoadParamTest, testing::Values( - SacTestConfig{8, 8, 10, 4, 2}, SacTestConfig{8, 8, 5, 2, 3}, - SacTestConfig{16, 16, 20, 6, 4})); + LiquidPpoTestConfig{8, 8, 10, 4, 2}, LiquidPpoTestConfig{8, 8, 5, 2, 2}, + LiquidPpoTestConfig{16, 16, 20, 6, 2})); diff --git a/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_rollout_buffer.cpp b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_rollout_buffer.cpp new file mode 100644 index 00000000..54364a26 --- /dev/null +++ b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_rollout_buffer.cpp @@ -0,0 +1,264 @@ +// +// Created by samuel on 06/09/2026. +// + +#include + +using namespace arenai; +using namespace arenai::agent; + +// ======================================================================== +// Fixture helpers +// ======================================================================== + +TorchState LiquidPpoRolloutBufferTest::make_state() { + return { + .vision = torch::randn({NB_TANKS, 3, VISION_SIZE, VISION_SIZE}), + .proprioception = torch::randn({NB_TANKS, NB_SENSORS})}; +} + +LiquidPpoInputStep LiquidPpoRolloutBufferTest::make_step( + const TorchState &state, const torch::Tensor &done, const bool episode_start) { + return { + .state = state, + .action = + {.continuous_action = torch::randn({NB_TANKS, NB_CONTINUOUS_ACTIONS}), + .discrete_action = torch::eye(NB_DISCRETE_ACTIONS) + .index_select( + 0, torch::randint( + NB_DISCRETE_ACTIONS, {NB_TANKS}, + torch::TensorOptions().dtype(torch::kInt64)))}, + .continuous_log_prob = torch::randn({NB_TANKS, 1}), + .discrete_log_prob = torch::randn({NB_TANKS, 1}), + .reward = torch::randn({NB_TANKS, 1}), + .done = done, + .truncated = torch::zeros({NB_TANKS, 1}), + .actor_hidden = torch::randn({NB_TANKS, NEURON_NUMBER}), + .episode_start = episode_start}; +} + +LiquidPpoInputStep LiquidPpoRolloutBufferTest::make_step(const TorchState &state) { + return make_step(state, torch::zeros({NB_TANKS, 1}), false); +} + +// ======================================================================== +// Completion counting +// ======================================================================== + +TEST_F(LiquidPpoRolloutBufferTest, EmptyBufferHasNoCompleteStep) { + LiquidPpoRolloutBuffer buffer; + + ASSERT_EQ(buffer.nb_complete_steps(), 0); + ASSERT_THROW(buffer.get_rollout(), c10::Error); +} + +TEST_F(LiquidPpoRolloutBufferTest, LastAddedStepStaysPending) { + LiquidPpoRolloutBuffer buffer; + + buffer.add(make_step(make_state())); + ASSERT_EQ(buffer.nb_complete_steps(), 0); + + buffer.add(make_step(make_state())); + ASSERT_EQ(buffer.nb_complete_steps(), 1); +} + +TEST_F(LiquidPpoRolloutBufferTest, FinishEpisodeCompletesPendingStep) { + LiquidPpoRolloutBuffer buffer; + + buffer.add(make_step(make_state())); + buffer.add(make_step(make_state())); + buffer.finish_episode(make_state()); + + ASSERT_EQ(buffer.nb_complete_steps(), 2); +} + +TEST_F(LiquidPpoRolloutBufferTest, GetRolloutKeepsPendingStep) { + LiquidPpoRolloutBuffer buffer; + + buffer.add(make_step(make_state())); + buffer.add(make_step(make_state())); + buffer.add(make_step(make_state())); + + const auto rollout = buffer.get_rollout(); + ASSERT_EQ(rollout.rewards.size(0), 2); + + // the pending step stays and is closed by the next observation + ASSERT_EQ(buffer.nb_complete_steps(), 0); + buffer.add(make_step(make_state())); + ASSERT_EQ(buffer.nb_complete_steps(), 1); +} + +// ======================================================================== +// Rollout content +// ======================================================================== + +TEST_F(LiquidPpoRolloutBufferTest, RolloutShapes) { + LiquidPpoRolloutBuffer buffer; + + constexpr int nb_steps = 3; + for (int t = 0; t < nb_steps; t++) buffer.add(make_step(make_state())); + buffer.finish_episode(make_state()); + + const auto rollout = buffer.get_rollout(); + + const auto expected_vision = + std::vector{nb_steps, NB_TANKS, 3, VISION_SIZE, VISION_SIZE}; + ASSERT_EQ(rollout.states.vision.sizes().vec(), expected_vision); + + // the bootstrap state has no time dimension + ASSERT_EQ( + rollout.bootstrap_state.vision.sizes().vec(), + (std::vector{NB_TANKS, 3, VISION_SIZE, VISION_SIZE})); + + const auto expected_scalar = std::vector{nb_steps, NB_TANKS, 1}; + ASSERT_EQ(rollout.rewards.sizes().vec(), expected_scalar); + ASSERT_EQ(rollout.continuous_log_probs.sizes().vec(), expected_scalar); + ASSERT_EQ(rollout.discrete_log_probs.sizes().vec(), expected_scalar); + ASSERT_EQ(rollout.valids.sizes().vec(), expected_scalar); + + ASSERT_EQ( + rollout.actions.continuous_action.sizes().vec(), + (std::vector{nb_steps, NB_TANKS, NB_CONTINUOUS_ACTIONS})); + + // liquid extras: the per-step actor states and the episode-start flags + ASSERT_EQ( + rollout.actor_hiddens.sizes().vec(), + (std::vector{nb_steps, NB_TANKS, NEURON_NUMBER})); + ASSERT_EQ(rollout.episode_starts.sizes().vec(), (std::vector{nb_steps})); + ASSERT_EQ(rollout.episode_starts.dtype(), torch::kBool); +} + +TEST_F(LiquidPpoRolloutBufferTest, ActorHiddenRoundTrip) { + LiquidPpoRolloutBuffer buffer; + + const auto step_0 = make_step(make_state()); + const auto step_1 = make_step(make_state()); + + buffer.add(step_0); + buffer.add(step_1); + buffer.finish_episode(make_state()); + + const auto rollout = buffer.get_rollout(); + + ASSERT_TRUE(torch::allclose(rollout.actor_hiddens[0], step_0.actor_hidden)); + ASSERT_TRUE(torch::allclose(rollout.actor_hiddens[1], step_1.actor_hidden)); +} + +TEST_F(LiquidPpoRolloutBufferTest, EpisodeStartFlagsRoundTrip) { + LiquidPpoRolloutBuffer buffer; + + const auto zeros = torch::zeros({NB_TANKS, 1}); + buffer.add(make_step(make_state(), zeros, true)); + buffer.add(make_step(make_state(), zeros, false)); + buffer.add(make_step(make_state(), zeros, true)); + buffer.finish_episode(make_state()); + + const auto episode_starts = buffer.get_rollout().episode_starts; + + ASSERT_TRUE(episode_starts[0].item()); + ASSERT_FALSE(episode_starts[1].item()); + ASSERT_TRUE(episode_starts[2].item()); +} + +TEST_F(LiquidPpoRolloutBufferTest, BootstrapStateIsEpisodeFinalObservation) { + LiquidPpoRolloutBuffer buffer; + + const auto state_0 = make_state(); + const auto state_1 = make_state(); + const auto final_state = make_state(); + + buffer.add(make_step(state_0)); + buffer.add(make_step(state_1)); + buffer.finish_episode(final_state); + + const auto rollout = buffer.get_rollout(); + + ASSERT_EQ(rollout.rewards.size(0), 2); + ASSERT_TRUE(torch::allclose(rollout.bootstrap_state.vision, final_state.vision)); +} + +TEST_F(LiquidPpoRolloutBufferTest, BootstrapStateIsPendingObservationMidEpisode) { + LiquidPpoRolloutBuffer buffer; + + const auto state_0 = make_state(); + const auto state_1 = make_state(); + + buffer.add(make_step(state_0)); + buffer.add(make_step(state_1)); + + const auto rollout = buffer.get_rollout(); + + // only the first step is complete; the pending step's own observation closes it + ASSERT_EQ(rollout.rewards.size(0), 1); + ASSERT_TRUE(torch::allclose(rollout.bootstrap_state.vision, state_1.vision)); +} + +// ======================================================================== +// Validity mask +// ======================================================================== + +TEST_F(LiquidPpoRolloutBufferTest, TerminatedTankInvalidatesFollowingSteps) { + LiquidPpoRolloutBuffer buffer; + + // tank 0 dies at the first step + const auto done = torch::cat({torch::ones({1, 1}), torch::zeros({1, 1})}, 0); + + buffer.add(make_step(make_state(), done, false)); + buffer.add(make_step(make_state())); + buffer.add(make_step(make_state())); + buffer.finish_episode(make_state()); + + const auto valids = buffer.get_rollout().valids.squeeze(-1); + + // the dying transition itself is valid, the following ones are not for tank 0 + ASSERT_TRUE(valids[0][0].item()); + ASSERT_FALSE(valids[1][0].item()); + ASSERT_FALSE(valids[2][0].item()); + + // tank 1 stays valid the whole rollout + ASSERT_TRUE(valids[0][1].item()); + ASSERT_TRUE(valids[1][1].item()); + ASSERT_TRUE(valids[2][1].item()); +} + +TEST_F(LiquidPpoRolloutBufferTest, TruncatedStepIsStackedAndStaysValid) { + LiquidPpoRolloutBuffer buffer; + + // tank 0 starves at the first step: done and truncated together + const auto done = torch::cat({torch::ones({1, 1}), torch::zeros({1, 1})}, 0); + + auto step_0 = make_step(make_state(), done, false); + step_0.truncated = done.clone(); + + buffer.add(step_0); + buffer.add(make_step(make_state())); + buffer.finish_episode(make_state()); + + const auto rollout = buffer.get_rollout(); + const auto truncateds = rollout.truncateds.squeeze(-1); + const auto valids = rollout.valids.squeeze(-1); + + ASSERT_TRUE(truncateds[0][0].item()); + ASSERT_FALSE(truncateds[0][1].item()); + ASSERT_FALSE(truncateds[1][0].item()); + + // the truncated (dying) transition itself is valid, the following one is not + ASSERT_TRUE(valids[0][0].item()); + ASSERT_FALSE(valids[1][0].item()); +} + +TEST_F(LiquidPpoRolloutBufferTest, FinishEpisodeResetsTermination) { + LiquidPpoRolloutBuffer buffer; + + buffer.add(make_step(make_state(), torch::ones({NB_TANKS, 1}), false)); + buffer.finish_episode(make_state()); + + // new episode: every tank is alive again + buffer.add(make_step(make_state())); + buffer.finish_episode(make_state()); + + const auto valids = buffer.get_rollout().valids.squeeze(-1); + + ASSERT_TRUE(torch::all(valids[0]).item()); + ASSERT_TRUE(torch::all(valids[1]).item()); +} diff --git a/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_training.cpp b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_training.cpp new file mode 100644 index 00000000..e31d59e9 --- /dev/null +++ b/arenai_agent/tests/src/tests_agents/tests_liquid_ppo_training.cpp @@ -0,0 +1,186 @@ +// +// Created by samuel on 06/09/2026. +// + +#include + +using namespace arenai; +using namespace arenai::agent; + +std::unique_ptr +LiquidPpoTrainingTest::make_factory(const LiquidPpoTrainingTestConfig &cfg) const { + const LiquidPpoHyperParams params{ + .actor_learning_rate = 1e-3f, + .critic_learning_rate = 1e-3f, + .hidden_size_sensors = 8, + .vision_channels = {{3, 4}}, + .group_norm_nums = {2}, + .neuron_number = NEURON_NUMBER, + .unfolding_steps = UNFOLDING_STEPS, + .delta_t = DELTA_T, + .chunk_size = CHUNK_SIZE, + .metric_window_size = 10, + .gamma = 0.99f, + .gae_lambda = 0.95f, + .clip_epsilon = 0.2f, + .continuous_target_entropy = std::vector(cfg.nb_continuous_actions, -0.88f), + .discrete_target_entropy_factors = std::vector(cfg.nb_discrete_actions, 0.2f), + .epochs = 2, + .rollout_size = ROLLOUT_SIZE, + .minibatch_size = MINIBATCH_SIZE}; + + return std::make_unique( + cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, + cfg.nb_discrete_actions, device, params); +} + +TorchState +LiquidPpoTrainingTest::make_state(const LiquidPpoTrainingTestConfig &cfg, const int nb_tanks) { + return { + .vision = torch::randint( + 0, 255, {nb_tanks, 3, cfg.vision_height, cfg.vision_width}, torch::kUInt8), + .proprioception = torch::randn({nb_tanks, cfg.nb_sensors})}; +} + +TEST_F(LiquidPpoTrainingTest, ActProducesValidOutput) { + constexpr LiquidPpoTrainingTestConfig cfg{ + .vision_height = 8, + .vision_width = 8, + .nb_sensors = 3, + .nb_continuous_actions = 2, + .nb_discrete_actions = 2}; + const auto factory = make_factory(cfg); + + const auto [continuous_action, discrete_action] = + factory->get_agent()->act(make_state(cfg, 1), true); + + ASSERT_EQ(continuous_action.size(0), 1); + ASSERT_EQ(continuous_action.size(1), 2); + ASSERT_EQ(discrete_action.size(0), 1); + ASSERT_EQ(discrete_action.size(1), 2); + + ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); + ASSERT_TRUE(torch::all(torch::isfinite(discrete_action)).item()); +} + +TEST_F(LiquidPpoTrainingTest, CountParametersPositive) { + constexpr LiquidPpoTrainingTestConfig cfg{ + .vision_height = 8, + .vision_width = 8, + .nb_sensors = 3, + .nb_continuous_actions = 2, + .nb_discrete_actions = 2}; + const auto factory = make_factory(cfg); + + ASSERT_GT(factory->get_trainer()->count_parameters(), 0) + << "Agent should have a positive number of parameters"; +} + +TEST_F(LiquidPpoTrainingTest, HiddenStateAdvancesAndResetsAcrossEpisodes) { + // build the triad by hand to keep a handle on the rollout buffer + constexpr LiquidPpoTrainingTestConfig cfg{ + .vision_height = 8, + .vision_width = 8, + .nb_sensors = 3, + .nb_continuous_actions = 2, + .nb_discrete_actions = 3}; + + constexpr int nb_tanks = 2; + const std::vector> vision_channels{{3, 4}}; + const std::vector group_norm_nums{2}; + + const auto actor = std::make_shared( + cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, + cfg.nb_discrete_actions, 8, vision_channels, group_norm_nums, NEURON_NUMBER, + UNFOLDING_STEPS, DELTA_T, 0.1f, std::vector(cfg.nb_discrete_actions, 0.2f)); + const auto hidden_state = std::make_shared(actor); + const auto rollout_buffer = std::make_shared(); + const auto collector = std::make_shared(rollout_buffer, hidden_state); + const auto agent = + std::make_shared(actor, hidden_state, device, collector); + + // first episode: two steps + agent->act(make_state(cfg, nb_tanks), true); + collector->on_transition( + torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); + agent->act(make_state(cfg, nb_tanks), true); + collector->on_transition( + torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); + collector->on_episode_end(make_state(cfg, nb_tanks)); + + // second episode: one step + agent->act(make_state(cfg, nb_tanks), true); + collector->on_transition( + torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); + collector->on_episode_end(make_state(cfg, nb_tanks)); + + const auto rollout = rollout_buffer->get_rollout(); + + // within the episode, the liquid state advanced between the steps + ASSERT_FALSE(torch::allclose(rollout.actor_hiddens[1], rollout.actor_hiddens[0])); + + // episode boundaries: first step of each episode is flagged + ASSERT_TRUE(rollout.episode_starts[0].item()); + ASSERT_FALSE(rollout.episode_starts[1].item()); + ASSERT_TRUE(rollout.episode_starts[2].item()); +} + +TEST_F(LiquidPpoTrainingTest, TrainingUpdatesActorParameters) { + // build the triad by hand to keep a handle on the actor's parameters + constexpr LiquidPpoTrainingTestConfig cfg{ + .vision_height = 8, + .vision_width = 8, + .nb_sensors = 3, + .nb_continuous_actions = 2, + .nb_discrete_actions = 3}; + + const std::vector> vision_channels{{3, 4}}; + const std::vector group_norm_nums{2}; + + const auto actor = std::make_shared( + cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, + cfg.nb_discrete_actions, 8, vision_channels, group_norm_nums, NEURON_NUMBER, + UNFOLDING_STEPS, DELTA_T, 0.1f, std::vector(cfg.nb_discrete_actions, 0.2f)); + const auto hidden_state = std::make_shared(actor); + const auto rollout_buffer = std::make_shared(); + const auto collector = std::make_shared(rollout_buffer, hidden_state); + const auto agent = + std::make_shared(actor, hidden_state, device, collector); + // target_kl = 0 : early stop disabled so every minibatch applies its update + const auto trainer = std::make_shared( + actor, rollout_buffer, cfg.vision_height, cfg.vision_width, cfg.nb_sensors, + cfg.nb_continuous_actions, cfg.nb_discrete_actions, 1e-3f, 1e-3f, 8, vision_channels, + group_norm_nums, NEURON_NUMBER, UNFOLDING_STEPS, DELTA_T, device, 10, 0.99f, 0.95f, 0.2f, + 0.f, 1.f, std::vector(cfg.nb_continuous_actions, 0.25f), + std::vector(cfg.nb_discrete_actions, 0.98f), 2, ROLLOUT_SIZE, MINIBATCH_SIZE, CHUNK_SIZE); + + std::vector initial_parameters; + for (const auto ¶meter: actor->parameters()) + initial_parameters.push_back(parameter.detach().clone()); + + // env loop: act -> transition -> maybe train, one more step than the rollout + // horizon so that the batch is complete when the trainer checks + for (int t = 0; t < ROLLOUT_SIZE + 2; t++) { + constexpr int nb_tanks = 2; + + agent->act(make_state(cfg, nb_tanks), true); + collector->on_transition( + torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); + trainer->step(); + } + + // the rollout has been consumed by the training + ASSERT_LT(rollout_buffer->nb_complete_steps(), static_cast(ROLLOUT_SIZE)); + + const auto parameters = actor->parameters(); + ASSERT_EQ(parameters.size(), initial_parameters.size()); + + bool any_changed = false; + for (size_t i = 0; i < parameters.size(); i++) + if (!torch::allclose(parameters[i], initial_parameters[i])) { + any_changed = true; + break; + } + + ASSERT_TRUE(any_changed) << "Training should update the actor's parameters"; +} diff --git a/arenai_agent/tests/src/tests_agents/tests_ppo.cpp b/arenai_agent/tests/src/tests_agents/tests_ppo.cpp index 154bb1b1..741c00d4 100644 --- a/arenai_agent/tests/src/tests_agents/tests_ppo.cpp +++ b/arenai_agent/tests/src/tests_agents/tests_ppo.cpp @@ -32,6 +32,8 @@ std::unique_ptr PpoAgentTest::make_factory(const PpoTestCo .gae_lambda = 0.95f, .clip_epsilon = 0.2f, .grad_norm_max = 1.f, + .continuous_target_entropy = std::vector(cfg.nb_continuous_actions, -0.88f), + .discrete_target_entropy_factors = std::vector(cfg.nb_discrete_actions, 0.5f), .epochs = 1, .rollout_size = 8, .minibatch_size = 10}; @@ -108,7 +110,7 @@ TEST_P(PpoActShapeParamTest, ActContinuousFinite) { ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); } -TEST_P(PpoActShapeParamTest, ActDiscreteIsOneHot) { +TEST_P(PpoActShapeParamTest, ActDiscreteIsBinary) { const auto cfg = GetParam(); const auto factory = make_factory(cfg); @@ -116,9 +118,7 @@ TEST_P(PpoActShapeParamTest, ActDiscreteIsOneHot) { const auto [continuous_action, discrete_action] = factory->get_agent()->act(make_state(cfg, batch), true); - const auto row_sums = torch::sum(discrete_action, -1); - ASSERT_TRUE(torch::allclose(row_sums, torch::ones({batch}))); - + // independent Bernoulli actions: each entry is 0 or 1, no one-hot constraint const auto is_binary = torch::logical_or(torch::eq(discrete_action, 0.0f), torch::eq(discrete_action, 1.0f)); ASSERT_TRUE(torch::all(is_binary).item()); @@ -127,8 +127,8 @@ TEST_P(PpoActShapeParamTest, ActDiscreteIsOneHot) { INSTANTIATE_TEST_SUITE_P( PpoAgent, PpoActShapeParamTest, testing::Values( - PpoTestConfig{8, 8, 10, 4, 2}, PpoTestConfig{8, 8, 5, 2, 3}, - PpoTestConfig{16, 16, 20, 6, 4}, PpoTestConfig{8, 12, 10, 4, 2})); + PpoTestConfig{8, 8, 10, 4, 2}, PpoTestConfig{8, 8, 5, 2, 2}, + PpoTestConfig{16, 16, 20, 6, 2}, PpoTestConfig{8, 12, 10, 4, 2})); // ======================================================================== // Parameterized: save / load tests @@ -170,5 +170,5 @@ TEST_P(PpoSaveLoadParamTest, SavedFilesNonEmpty) { INSTANTIATE_TEST_SUITE_P( PpoAgent, PpoSaveLoadParamTest, testing::Values( - PpoTestConfig{8, 8, 10, 4, 2}, PpoTestConfig{8, 8, 5, 2, 3}, - PpoTestConfig{16, 16, 20, 6, 4})); + PpoTestConfig{8, 8, 10, 4, 2}, PpoTestConfig{8, 8, 5, 2, 2}, + PpoTestConfig{16, 16, 20, 6, 2})); diff --git a/arenai_agent/tests/src/tests_agents/tests_ppo_rollout_buffer.cpp b/arenai_agent/tests/src/tests_agents/tests_ppo_rollout_buffer.cpp index e4029a07..dacb4434 100644 --- a/arenai_agent/tests/src/tests_agents/tests_ppo_rollout_buffer.cpp +++ b/arenai_agent/tests/src/tests_agents/tests_ppo_rollout_buffer.cpp @@ -30,7 +30,8 @@ PpoInputStep PpoRolloutBufferTest::make_step(const TorchState &state, const torc .continuous_log_prob = torch::randn({NB_TANKS, 1}), .discrete_log_prob = torch::randn({NB_TANKS, 1}), .reward = torch::randn({NB_TANKS, 1}), - .done = done}; + .done = done, + .truncated = torch::zeros({NB_TANKS, 1})}; } PpoInputStep PpoRolloutBufferTest::make_step(const TorchState &state) { diff --git a/arenai_agent/tests/src/tests_agents/tests_ppo_training.cpp b/arenai_agent/tests/src/tests_agents/tests_ppo_training.cpp index 59103955..557f2a2c 100644 --- a/arenai_agent/tests/src/tests_agents/tests_ppo_training.cpp +++ b/arenai_agent/tests/src/tests_agents/tests_ppo_training.cpp @@ -21,6 +21,8 @@ PpoTrainingTest::make_factory(const PpoTrainingTestConfig &cfg) const { .gamma = 0.99f, .gae_lambda = 0.95f, .clip_epsilon = 0.2f, + .continuous_target_entropy = std::vector(cfg.nb_continuous_actions, -0.88f), + .discrete_target_entropy_factors = std::vector(cfg.nb_discrete_actions, 0.5f), .epochs = 2, .rollout_size = ROLLOUT_SIZE, .minibatch_size = MINIBATCH_SIZE}; @@ -43,7 +45,7 @@ TEST_F(PpoTrainingTest, ActProducesValidOutput) { .vision_width = 8, .nb_sensors = 3, .nb_continuous_actions = 2, - .nb_discrete_actions = 3}; + .nb_discrete_actions = 2}; const auto factory = make_factory(cfg); const auto [continuous_action, discrete_action] = @@ -52,7 +54,7 @@ TEST_F(PpoTrainingTest, ActProducesValidOutput) { ASSERT_EQ(continuous_action.size(0), 1); ASSERT_EQ(continuous_action.size(1), 2); ASSERT_EQ(discrete_action.size(0), 1); - ASSERT_EQ(discrete_action.size(1), 3); + ASSERT_EQ(discrete_action.size(1), 2); ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); ASSERT_TRUE(torch::all(torch::isfinite(discrete_action)).item()); @@ -64,7 +66,7 @@ TEST_F(PpoTrainingTest, CountParametersPositive) { .vision_width = 8, .nb_sensors = 3, .nb_continuous_actions = 2, - .nb_discrete_actions = 3}; + .nb_discrete_actions = 2}; const auto factory = make_factory(cfg); ASSERT_GT(factory->get_trainer()->count_parameters(), 0) @@ -85,7 +87,8 @@ TEST_F(PpoTrainingTest, TrainingUpdatesActorParameters) { const auto actor = std::make_shared( cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, - cfg.nb_discrete_actions, 8, std::vector{16}, vision_channels, group_norm_nums, 0.1f, 0.2f); + cfg.nb_discrete_actions, 8, std::vector{16}, vision_channels, group_norm_nums, 0.1f, + std::vector(cfg.nb_discrete_actions, 0.2f)); const auto rollout_buffer = std::make_shared(); const auto collector = std::make_shared(rollout_buffer); const auto agent = std::make_shared(actor, device, collector); @@ -93,8 +96,9 @@ TEST_F(PpoTrainingTest, TrainingUpdatesActorParameters) { const auto trainer = std::make_shared( actor, rollout_buffer, cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, cfg.nb_discrete_actions, 1e-3f, 1e-3f, 8, std::vector{16}, - vision_channels, group_norm_nums, device, 10, 0.99f, 0.95f, 0.2f, 0.f, 1.f, 0.25f, 0.98f, 2, - ROLLOUT_SIZE, MINIBATCH_SIZE); + vision_channels, group_norm_nums, device, 10, 0.99f, 0.95f, 0.2f, 0.f, 1.f, + std::vector(cfg.nb_continuous_actions, 0.25f), std::vector(cfg.nb_discrete_actions, 0.98f), + 2, ROLLOUT_SIZE, MINIBATCH_SIZE); std::vector initial_parameters; for (const auto ¶meter: actor->parameters()) @@ -106,7 +110,8 @@ TEST_F(PpoTrainingTest, TrainingUpdatesActorParameters) { constexpr int nb_tanks = 2; agent->act(make_state(cfg, nb_tanks), true); - collector->on_transition(torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); + collector->on_transition( + torch::randn({nb_tanks, 1}), torch::zeros({nb_tanks, 1}), torch::zeros({nb_tanks, 1})); trainer->step(); } diff --git a/arenai_agent/tests/src/tests_agents/tests_sac_training.cpp b/arenai_agent/tests/src/tests_agents/tests_sac_training.cpp deleted file mode 100644 index a7f259a0..00000000 --- a/arenai_agent/tests/src/tests_agents/tests_sac_training.cpp +++ /dev/null @@ -1,74 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#include - -using namespace arenai; -using namespace arenai::agent; - -std::unique_ptr -SacTrainingTest::make_factory(const SacTrainingTestConfig &cfg) const { - const SacHyperParams params{ - .actor_learning_rate = 1e-3f, - .critic_learning_rate = 1e-3f, - .alpha_learning_rate = 1e-3f, - .hidden_size_sensors = 8, - .hidden_size_actions = 8, - .actor_hidden_sizes = {16}, - .critic_hidden_sizes = {16}, - .vision_channels = {{3, 4}}, - .group_norm_nums = {2}, - .metric_window_size = 10, - .tau = 0.005f, - .gamma = 0.99f, - .replay_buffer_size = 10, - .train_every = 1, - .epochs = 1, - .batch_size = 1}; - - return std::make_unique( - cfg.vision_height, cfg.vision_width, cfg.nb_sensors, cfg.nb_continuous_actions, - cfg.nb_discrete_actions, device, params); -} - -TorchState SacTrainingTest::make_state(const SacTrainingTestConfig &cfg) { - return { - .vision = - torch::randint(0, 255, {1, 3, cfg.vision_height, cfg.vision_width}, torch::kUInt8), - .proprioception = torch::randn({1, cfg.nb_sensors})}; -} - -TEST_F(SacTrainingTest, ActProducesValidOutput) { - constexpr SacTrainingTestConfig cfg{ - .vision_height = 8, - .vision_width = 8, - .nb_sensors = 3, - .nb_continuous_actions = 2, - .nb_discrete_actions = 3}; - const auto factory = make_factory(cfg); - - const auto [continuous_action, discrete_action] = - factory->get_agent()->act(make_state(cfg), true); - - ASSERT_EQ(continuous_action.size(0), 1); - ASSERT_EQ(continuous_action.size(1), 2); - ASSERT_EQ(discrete_action.size(0), 1); - ASSERT_EQ(discrete_action.size(1), 3); - - ASSERT_TRUE(torch::all(torch::isfinite(continuous_action)).item()); - ASSERT_TRUE(torch::all(torch::isfinite(discrete_action)).item()); -} - -TEST_F(SacTrainingTest, CountParametersPositive) { - constexpr SacTrainingTestConfig cfg{ - .vision_height = 8, - .vision_width = 8, - .nb_sensors = 3, - .nb_continuous_actions = 2, - .nb_discrete_actions = 3}; - const auto factory = make_factory(cfg); - - ASSERT_GT(factory->get_trainer()->count_parameters(), 0) - << "Agent should have a positive number of parameters"; -} diff --git a/arenai_agent/tests/src/tests_curriculum/tests_spawn_curriculum.cpp b/arenai_agent/tests/src/tests_curriculum/tests_spawn_curriculum.cpp new file mode 100644 index 00000000..b8e17b86 --- /dev/null +++ b/arenai_agent/tests/src/tests_curriculum/tests_spawn_curriculum.cpp @@ -0,0 +1,122 @@ +// +// Created by samuel on 05/09/2026. +// + +#include +#include + +using namespace arenai; +using namespace arenai::agent; + +namespace { + constexpr float DELTA = 0.05f; + constexpr float RATIO_LOW = 0.02f; + constexpr float RATIO_HIGH = 0.04f; + constexpr int PROBE_WINDOW = 4; + constexpr std::uint64_t SEED = 42; + + // feed exactly one controller window of probe episodes with the given ratio + void run_probe_window(SpawnCurriculum &curriculum, const int nb_fires, const int nb_hits) { + int done = 0; + while (done < PROBE_WINDOW) { + curriculum.sample_progress(); + const bool is_probe = curriculum.is_probe(); + + curriculum.on_episode_end(nb_fires, nb_hits); + if (is_probe) done++; + } + } +}// namespace + +TEST(SpawnCurriculumTest, StartsAtZero) { + const SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 0.f); +} + +TEST(SpawnCurriculumTest, SampleStaysWithinBound) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + run_probe_window(curriculum, 100, 10);// grow once so the bound is not 0 + + for (int i = 0; i < 1000; i++) { + const float upper = curriculum.upper_bound(); + const float progress = curriculum.sample_progress(); + ASSERT_GE(progress, 0.f); + ASSERT_LE(progress, upper); + } +} + +TEST(SpawnCurriculumTest, GrowsAboveHighRatio) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + run_probe_window(curriculum, 100, 10);// ratio 0.1 > high + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), DELTA); +} + +TEST(SpawnCurriculumTest, ShrinksBelowLowRatio) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + run_probe_window(curriculum, 100, 10); + run_probe_window(curriculum, 100, 10); + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 2.f * DELTA); + + run_probe_window(curriculum, 100, 1);// ratio 0.01 < low + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), DELTA); +} + +TEST(SpawnCurriculumTest, HoldsInsideHysteresisBand) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + run_probe_window(curriculum, 100, 10); + ASSERT_FLOAT_EQ(curriculum.upper_bound(), DELTA); + + run_probe_window(curriculum, 100, 3);// ratio 0.03, between low and high + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), DELTA); +} + +TEST(SpawnCurriculumTest, NoFireCountsAsFailure) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + run_probe_window(curriculum, 100, 10); + ASSERT_FLOAT_EQ(curriculum.upper_bound(), DELTA); + + run_probe_window(curriculum, 0, 0); + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 0.f); +} + +TEST(SpawnCurriculumTest, ClampsToUnitInterval) { + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + // more than enough successful windows to reach the top + for (int i = 0; i < 30; i++) run_probe_window(curriculum, 100, 10); + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 1.f); + + // and enough failures to reach the bottom + for (int i = 0; i < 30; i++) run_probe_window(curriculum, 100, 0); + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 0.f); +} + +TEST(SpawnCurriculumTest, NonProbeEpisodesDoNotFeedController) { + // boundary_proba 0: no episode is a probe, the controller must never move + // whatever the results + SpawnCurriculum curriculum(DELTA, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.f, SEED); + for (int i = 0; i < 100; i++) { + curriculum.sample_progress(); + curriculum.on_episode_end(100, 10); + } + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 0.f); +} + +TEST(SpawnCurriculumTest, ZeroDeltaKeepsBoundFixed) { + SpawnCurriculum curriculum(0.f, RATIO_LOW, RATIO_HIGH, PROBE_WINDOW, 0.2f, SEED); + + for (int i = 0; i < 10; i++) run_probe_window(curriculum, 100, 10); + + ASSERT_FLOAT_EQ(curriculum.upper_bound(), 0.f); +} diff --git a/arenai_agent/tests/src/tests_distributions/tests_bernoulli.cpp b/arenai_agent/tests/src/tests_distributions/tests_bernoulli.cpp new file mode 100644 index 00000000..f2cd43d6 --- /dev/null +++ b/arenai_agent/tests/src/tests_distributions/tests_bernoulli.cpp @@ -0,0 +1,108 @@ +// +// Created by samuel on 18/09/2026. +// + +#include + +#include + +using namespace arenai; +using namespace arenai::agent; + +// ======================================================================== +// Fixed tests +// ======================================================================== + +TEST_F(BernoulliTest, EntropyMaxAtHalf) { + const auto proba = torch::full({1, 3}, 0.5f); + + const auto entropy = bernoulli_entropy(proba); + + ASSERT_TRUE(torch::allclose(entropy, torch::full({1, 3}, std::log(2.f)), 1e-4f)) + << "Each action entropy should be log(2) at p=0.5"; +} + +TEST_F(BernoulliTest, EntropyMinAtDegenerate) { + const auto proba = torch::tensor({{0.f, 1.f}}); + + const auto entropy = bernoulli_entropy(proba); + + ASSERT_TRUE(torch::all(torch::lt(entropy, 1e-3f)).item()) + << "Entropy should be near 0 for degenerate probabilities"; +} + +TEST_F(BernoulliTest, MaximumEntropyEqualsLog2) { + ASSERT_NEAR(bernoulli_maximum_entropy(), std::log(2.f), 1e-6f); +} + +TEST_F(BernoulliTest, MaxActionThreshold) { + const auto proba = torch::tensor({{0.49f, 0.51f}, {0.9f, 0.1f}}); + + const auto action = bernoulli_max_action(proba); + + ASSERT_TRUE(torch::allclose(action, torch::tensor({{0.f, 1.f}, {1.f, 0.f}}))) + << "Actions should engage strictly above the 0.5 threshold"; +} + +TEST_F(BernoulliTest, LogProbaMatchesTakenActions) { + const auto proba = torch::tensor({{0.8f, 0.3f}}); + const auto action = torch::tensor({{1.f, 0.f}}); + + const auto log_proba = bernoulli_log_proba(action, proba); + + ASSERT_NEAR(log_proba[0][0].item(), std::log(0.8f), 1e-4f); + ASSERT_NEAR(log_proba[0][1].item(), std::log(0.7f), 1e-4f); +} + +// ======================================================================== +// Parameterized: sample shape and binary property +// ======================================================================== + +TEST_P(BernoulliShapeParamTest, SampleIsBinary) { + const auto [batch_size, nb_actions] = GetParam(); + + const auto proba = torch::sigmoid(torch::randn({batch_size, nb_actions})); + + const auto sample = bernoulli_sample(proba); + + ASSERT_EQ(sample.size(0), batch_size); + ASSERT_EQ(sample.size(1), nb_actions); + + // each element is 0 or 1, independently of the others + const auto is_binary = torch::logical_or(torch::eq(sample, 0.0f), torch::eq(sample, 1.0f)); + ASSERT_TRUE(torch::all(is_binary).item()) << "Sample should contain only 0s and 1s"; +} + +TEST_P(BernoulliShapeParamTest, EntropyShapeAndBounds) { + const auto [batch_size, nb_actions] = GetParam(); + + const auto proba = torch::sigmoid(torch::randn({batch_size, nb_actions})); + + const auto entropy = bernoulli_entropy(proba); + + ASSERT_EQ(entropy.size(0), batch_size); + ASSERT_EQ(entropy.size(1), nb_actions); + ASSERT_TRUE(torch::all(torch::ge(entropy, 0.0f)).item()) + << "Entropy should be non-negative"; + ASSERT_TRUE(torch::all(torch::le(entropy, bernoulli_maximum_entropy() + 1e-4f)).item()) + << "Each action entropy should be <= log(2)"; +} + +TEST_P(BernoulliShapeParamTest, LogProbaShapeAndFinite) { + const auto [batch_size, nb_actions] = GetParam(); + + const auto proba = torch::sigmoid(torch::randn({batch_size, nb_actions})); + const auto action = bernoulli_sample(proba); + + const auto log_proba = bernoulli_log_proba(action, proba); + + ASSERT_EQ(log_proba.size(0), batch_size); + ASSERT_EQ(log_proba.size(1), nb_actions); + ASSERT_TRUE(torch::all(torch::isfinite(log_proba)).item()); + ASSERT_TRUE(torch::all(torch::le(log_proba, 0.0f)).item()) + << "Log-probabilities should be <= 0"; +} + +INSTANTIATE_TEST_SUITE_P( + BernoulliShape, BernoulliShapeParamTest, + testing::Combine(testing::Values(1, 2, 8, 16), testing::Values(1, 2, 3, 5))); diff --git a/arenai_agent/tests/src/tests_distributions/tests_bernoulli_edge.cpp b/arenai_agent/tests/src/tests_distributions/tests_bernoulli_edge.cpp new file mode 100644 index 00000000..80cbd709 --- /dev/null +++ b/arenai_agent/tests/src/tests_distributions/tests_bernoulli_edge.cpp @@ -0,0 +1,65 @@ +// +// Created by samuel on 18/09/2026. +// + +#include + +#include + +using namespace arenai; +using namespace arenai::agent; + +TEST_F(BernoulliEdgeTest, EntropyWithBoundaryProbabilities) { + const auto proba = torch::tensor({{0.f, 1.f, 1e-10f, 1.f - 1e-10f}}); + + const auto entropy = bernoulli_entropy(proba); + + ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()) + << "Entropy should be finite with boundary probabilities"; + ASSERT_TRUE(torch::all(torch::ge(entropy, 0.0f)).item()) + << "Entropy should be non-negative"; +} + +TEST_F(BernoulliEdgeTest, LogProbaWithBoundaryProbabilities) { + const auto proba = torch::tensor({{0.f, 1.f}}); + const auto action = torch::tensor({{1.f, 0.f}}); + + const auto log_proba = bernoulli_log_proba(action, proba); + + ASSERT_TRUE(torch::all(torch::isfinite(log_proba)).item()) + << "Log-probability should be finite (clamped) even for impossible actions"; +} + +TEST_F(BernoulliEdgeTest, EntropyGradientFlowsThroughProbabilities) { + const auto logits = torch::randn({4, 3}, torch::TensorOptions().requires_grad(true)); + const auto proba = torch::sigmoid(logits); + + const auto entropy = bernoulli_entropy(proba); + const auto loss = entropy.sum(); + + loss.backward(); + + ASSERT_TRUE(logits.grad().defined()) << "Gradient should flow back through entropy"; + ASSERT_TRUE(torch::all(torch::isfinite(logits.grad())).item()) + << "Gradient should be finite"; +} + +TEST_F(BernoulliEdgeTest, SampleWithDeterministicProbabilities) { + const auto proba = torch::cat({torch::zeros({4, 1}), torch::ones({4, 1})}, -1); + + const auto sample = bernoulli_sample(proba); + + ASSERT_TRUE(torch::allclose(sample.slice(-1, 0, 1), torch::zeros({4, 1}))) + << "p=0 should never engage"; + ASSERT_TRUE(torch::allclose(sample.slice(-1, 1, 2), torch::ones({4, 1}))) + << "p=1 should always engage"; +} + +TEST_F(BernoulliEdgeTest, MaxActionAtExactHalf) { + const auto proba = torch::full({1, 2}, 0.5f); + + const auto action = bernoulli_max_action(proba); + + ASSERT_TRUE(torch::allclose(action, torch::zeros({1, 2}))) + << "p=0.5 exactly should not engage (strict > threshold)"; +} diff --git a/arenai_agent/tests/src/tests_distributions/tests_beta_law.cpp b/arenai_agent/tests/src/tests_distributions/tests_beta_law.cpp index 29aa0de8..3732e5ac 100644 --- a/arenai_agent/tests/src/tests_distributions/tests_beta_law.cpp +++ b/arenai_agent/tests/src/tests_distributions/tests_beta_law.cpp @@ -14,13 +14,22 @@ using namespace arenai::agent; // ======================================================================== TEST_F(BetaLawTest, UniformEntropyIsMaximal) { - // alpha=1, beta=1 → uniform → maximal entropy - const auto entropy_uniform = beta_law_entropy(torch::tensor({1.0f}), torch::tensor({1.0f})); - const auto entropy_peaked = beta_law_entropy(torch::tensor({5.0f}), torch::tensor({5.0f})); + // concentration=2 → alpha=beta=1 → uniform → maximal entropy + const auto entropy_uniform = beta_law_entropy(torch::tensor({0.5f}), torch::tensor({2.0f})); + const auto entropy_peaked = beta_law_entropy(torch::tensor({0.5f}), torch::tensor({10.0f})); ASSERT_GT(entropy_uniform.item(), entropy_peaked.item()); } +TEST_F(BetaLawTest, EntropyDecreasesWithConcentration) { + const auto mode = torch::tensor({0.3f}); + + const auto entropy_low = beta_law_entropy(mode, torch::tensor({5.0f})); + const auto entropy_high = beta_law_entropy(mode, torch::tensor({50.0f})); + + ASSERT_GT(entropy_low.item(), entropy_high.item()); +} + TEST_F(BetaLawTest, TargetEntropyProportionalToActions) { const auto t1 = beta_law_target_entropy(1); const auto t3 = beta_law_target_entropy(3); @@ -29,15 +38,35 @@ TEST_F(BetaLawTest, TargetEntropyProportionalToActions) { } TEST_F(BetaLawTest, LogProbaConsistentWithSample) { - const auto alpha = torch::ones({100}) * 2.0f; - const auto beta = torch::ones({100}) * 3.0f; + const auto mode = torch::ones({100}) * 0.4f; + const auto concentration = torch::ones({100}) * 5.0f; - const auto samples = beta_law_sample(alpha, beta); - const auto log_p = beta_law_log_proba(samples, alpha, beta); + const auto samples = beta_law_sample(mode, concentration); + const auto log_p = beta_law_log_proba(samples, mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(log_p)).item()); } +TEST_F(BetaLawTest, MeanActionMatchesModeWhenCentered) { + // mode=0.5 → alpha=beta → mean action 0 on [-1, 1] + const auto mode = torch::ones({10}) * 0.5f; + const auto concentration = torch::ones({10}) * 8.0f; + + const auto mean = beta_law_mean_action(mode, concentration); + + ASSERT_TRUE(torch::allclose(mean, torch::zeros({10}), 1e-5, 1e-5)); +} + +TEST_F(BetaLawTest, ModeActionSpansActionRange) { + // mode in [0, 1] maps linearly to [-1, 1], independent of concentration + const auto mode = torch::tensor({0.0f, 0.25f, 0.5f, 0.75f, 1.0f}); + + const auto action = beta_law_mode_action(mode); + + ASSERT_TRUE( + torch::allclose(action, torch::tensor({-1.0f, -0.5f, 0.0f, 0.5f, 1.0f}), 1e-5, 1e-5)); +} + // ======================================================================== // Parameterized: shape variations // ======================================================================== @@ -45,12 +74,12 @@ TEST_F(BetaLawTest, LogProbaConsistentWithSample) { TEST_P(BetaLawParamTest, SampleBounds) { const auto &shape = GetParam(); - const auto alpha = torch::rand(shape) * 4.0f + 0.5f; - const auto beta = torch::rand(shape) * 4.0f + 0.5f; + const auto mode = torch::rand(shape); + const auto concentration = torch::rand(shape) * 8.0f + 2.0f; - const auto samples = beta_law_sample(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); - ASSERT_EQ(samples.sizes(), alpha.sizes()); + ASSERT_EQ(samples.sizes(), mode.sizes()); ASSERT_TRUE(torch::all(torch::logical_and(torch::ge(samples, -1.0f), torch::le(samples, 1.0f))) .item()); } @@ -58,11 +87,11 @@ TEST_P(BetaLawParamTest, SampleBounds) { TEST_P(BetaLawParamTest, LogProbaShape) { const auto &shape = GetParam(); - const auto alpha = torch::rand(shape) * 4.0f + 0.5f; - const auto beta = torch::rand(shape) * 4.0f + 0.5f; - const auto samples = beta_law_sample(alpha, beta); + const auto mode = torch::rand(shape); + const auto concentration = torch::rand(shape) * 8.0f + 2.0f; + const auto samples = beta_law_sample(mode, concentration); - const auto log_p = beta_law_log_proba(samples, alpha, beta); + const auto log_p = beta_law_log_proba(samples, mode, concentration); ASSERT_EQ(log_p.sizes(), samples.sizes()); ASSERT_TRUE(torch::all(torch::isfinite(log_p)).item()); @@ -71,35 +100,36 @@ TEST_P(BetaLawParamTest, LogProbaShape) { TEST_P(BetaLawParamTest, EntropyShape) { const auto &shape = GetParam(); - const auto alpha = torch::rand(shape) * 4.0f + 0.5f; - const auto beta = torch::rand(shape) * 4.0f + 0.5f; + const auto mode = torch::rand(shape); + const auto concentration = torch::rand(shape) * 8.0f + 2.0f; - const auto entropy = beta_law_entropy(alpha, beta); + const auto entropy = beta_law_entropy(mode, concentration); - ASSERT_EQ(entropy.sizes(), alpha.sizes()); + ASSERT_EQ(entropy.sizes(), mode.sizes()); ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()); } -TEST_P(BetaLawParamTest, SampleNoNaNWithSmallParams) { +TEST_P(BetaLawParamTest, SampleNoNaNBelowUniformConcentration) { const auto &shape = GetParam(); - const auto alpha = torch::ones(shape) * 0.1f; - const auto beta = torch::ones(shape) * 0.1f; + // concentration below the κ=2 floor → alpha/beta below the clamp, still finite + const auto mode = torch::ones(shape) * 0.5f; + const auto concentration = torch::ones(shape) * 0.1f; - const auto samples = beta_law_sample(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()); ASSERT_TRUE(torch::all(torch::logical_and(torch::ge(samples, -1.0f), torch::le(samples, 1.0f))) .item()); } -TEST_P(BetaLawParamTest, SampleNoNaNWithLargeParams) { +TEST_P(BetaLawParamTest, SampleNoNaNWithLargeConcentration) { const auto &shape = GetParam(); - const auto alpha = torch::ones(shape) * 50.0f; - const auto beta = torch::ones(shape) * 50.0f; + const auto mode = torch::ones(shape) * 0.5f; + const auto concentration = torch::ones(shape) * 2000.0f; - const auto samples = beta_law_sample(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()); } diff --git a/arenai_agent/tests/src/tests_distributions/tests_beta_law_edge.cpp b/arenai_agent/tests/src/tests_distributions/tests_beta_law_edge.cpp index 457658c0..03f77aaf 100644 --- a/arenai_agent/tests/src/tests_distributions/tests_beta_law_edge.cpp +++ b/arenai_agent/tests/src/tests_distributions/tests_beta_law_edge.cpp @@ -13,44 +13,44 @@ using namespace arenai::agent; // Asymmetric parameter edge cases // ======================================================================== -TEST_F(BetaLawEdgeTest, VeryAsymmetricParamsAlphaLarge) { - const auto alpha = torch::ones({50}) * 100.0f; - const auto beta = torch::ones({50}) * 0.01f; +TEST_F(BetaLawEdgeTest, ModeNearUpperBound) { + const auto mode = torch::ones({50}) * 0.999f; + const auto concentration = torch::ones({50}) * 100.0f; - const auto samples = beta_law_sample(alpha, beta); - const auto log_p = beta_law_log_proba(samples, alpha, beta); - const auto entropy = beta_law_entropy(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); + const auto log_p = beta_law_log_proba(samples, mode, concentration); + const auto entropy = beta_law_entropy(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()) - << "Samples should be finite with alpha>>beta"; + << "Samples should be finite with mode near 1"; ASSERT_TRUE(torch::all(torch::isfinite(log_p)).item()) - << "Log-proba should be finite with alpha>>beta"; + << "Log-proba should be finite with mode near 1"; ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()) - << "Entropy should be finite with alpha>>beta"; + << "Entropy should be finite with mode near 1"; } -TEST_F(BetaLawEdgeTest, VeryAsymmetricParamsBetaLarge) { - const auto alpha = torch::ones({50}) * 0.01f; - const auto beta = torch::ones({50}) * 100.0f; +TEST_F(BetaLawEdgeTest, ModeNearLowerBound) { + const auto mode = torch::ones({50}) * 0.001f; + const auto concentration = torch::ones({50}) * 100.0f; - const auto samples = beta_law_sample(alpha, beta); - const auto log_p = beta_law_log_proba(samples, alpha, beta); - const auto entropy = beta_law_entropy(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); + const auto log_p = beta_law_log_proba(samples, mode, concentration); + const auto entropy = beta_law_entropy(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()) - << "Samples should be finite with beta>>alpha"; + << "Samples should be finite with mode near 0"; ASSERT_TRUE(torch::all(torch::isfinite(log_p)).item()) - << "Log-proba should be finite with beta>>alpha"; + << "Log-proba should be finite with mode near 0"; ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()) - << "Entropy should be finite with beta>>alpha"; + << "Entropy should be finite with mode near 0"; } TEST_F(BetaLawEdgeTest, ZeroParamsHandledGracefully) { - const auto alpha = torch::zeros({10}); - const auto beta = torch::zeros({10}); + const auto mode = torch::zeros({10}); + const auto concentration = torch::zeros({10}); - const auto samples = beta_law_sample(alpha, beta); - const auto entropy = beta_law_entropy(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); + const auto entropy = beta_law_entropy(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()) << "Samples should be finite with zero params (clamped to EPSILON)"; @@ -59,10 +59,10 @@ TEST_F(BetaLawEdgeTest, ZeroParamsHandledGracefully) { } TEST_F(BetaLawEdgeTest, NegativeParamsHandledGracefully) { - const auto alpha = torch::ones({10}) * -1.0f; - const auto beta = torch::ones({10}) * -1.0f; + const auto mode = torch::ones({10}) * -1.0f; + const auto concentration = torch::ones({10}) * -1.0f; - const auto samples = beta_law_sample(alpha, beta); + const auto samples = beta_law_sample(mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(samples)).item()) << "Samples should be finite with negative params (clamped to EPSILON)"; @@ -71,13 +71,13 @@ TEST_F(BetaLawEdgeTest, NegativeParamsHandledGracefully) { } TEST_F(BetaLawEdgeTest, LogProbAtBoundaryValues) { - const auto alpha = torch::ones({10}) * 2.0f; - const auto beta = torch::ones({10}) * 2.0f; + const auto mode = torch::ones({10}) * 0.5f; + const auto concentration = torch::ones({10}) * 6.0f; const auto x_near_minus1 = torch::ones({10}) * -0.999f; const auto x_near_plus1 = torch::ones({10}) * 0.999f; - const auto log_p_lo = beta_law_log_proba(x_near_minus1, alpha, beta); - const auto log_p_hi = beta_law_log_proba(x_near_plus1, alpha, beta); + const auto log_p_lo = beta_law_log_proba(x_near_minus1, mode, concentration); + const auto log_p_hi = beta_law_log_proba(x_near_plus1, mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(log_p_lo)).item()) << "Log-proba near -1 boundary should be finite"; @@ -86,13 +86,13 @@ TEST_F(BetaLawEdgeTest, LogProbAtBoundaryValues) { } TEST_F(BetaLawEdgeTest, LogProbAtExactBoundaryValues) { - const auto alpha = torch::ones({10}) * 2.0f; - const auto beta = torch::ones({10}) * 2.0f; + const auto mode = torch::ones({10}) * 0.5f; + const auto concentration = torch::ones({10}) * 6.0f; const auto x_minus1 = torch::ones({10}) * -1.0f; const auto x_plus1 = torch::ones({10}) * 1.0f; - const auto log_p_lo = beta_law_log_proba(x_minus1, alpha, beta); - const auto log_p_hi = beta_law_log_proba(x_plus1, alpha, beta); + const auto log_p_lo = beta_law_log_proba(x_minus1, mode, concentration); + const auto log_p_hi = beta_law_log_proba(x_plus1, mode, concentration); ASSERT_TRUE(torch::all(torch::isfinite(log_p_lo)).item()) << "Log-proba at exact -1 boundary should be finite (clamped)"; @@ -104,60 +104,61 @@ TEST_F(BetaLawEdgeTest, LogProbAtExactBoundaryValues) { // Gradient flow tests // ======================================================================== -TEST_F(BetaLawGradientTest, LogProbaGradientFlowsThroughAlpha) { - const auto alpha = torch::full({5}, 2.0f, torch::TensorOptions().requires_grad(true)); - const auto beta = torch::ones({5}) * 3.0f; +TEST_F(BetaLawGradientTest, LogProbaGradientFlowsThroughMode) { + const auto mode = torch::full({5}, 0.4f, torch::TensorOptions().requires_grad(true)); + const auto concentration = torch::ones({5}) * 6.0f; const auto x = torch::tensor({0.0f, 0.2f, -0.3f, 0.5f, -0.1f}); - const auto log_p = beta_law_log_proba(x, alpha, beta); + const auto log_p = beta_law_log_proba(x, mode, concentration); const auto loss = log_p.sum(); loss.backward(); - ASSERT_TRUE(alpha.grad().defined()) << "Gradient should flow back to alpha"; - ASSERT_TRUE(torch::all(torch::isfinite(alpha.grad())).item()) - << "Gradient w.r.t. alpha should be finite"; + ASSERT_TRUE(mode.grad().defined()) << "Gradient should flow back to mode"; + ASSERT_TRUE(torch::all(torch::isfinite(mode.grad())).item()) + << "Gradient w.r.t. mode should be finite"; } -TEST_F(BetaLawGradientTest, LogProbaGradientFlowsThroughBeta) { - const auto alpha = torch::ones({5}) * 2.0f; - const auto beta = torch::full({5}, 3.0f, torch::TensorOptions().requires_grad(true)); +TEST_F(BetaLawGradientTest, LogProbaGradientFlowsThroughConcentration) { + const auto mode = torch::ones({5}) * 0.4f; + const auto concentration = torch::full({5}, 6.0f, torch::TensorOptions().requires_grad(true)); const auto x = torch::tensor({0.0f, 0.2f, -0.3f, 0.5f, -0.1f}); - const auto log_p = beta_law_log_proba(x, alpha, beta); + const auto log_p = beta_law_log_proba(x, mode, concentration); const auto loss = log_p.sum(); loss.backward(); - ASSERT_TRUE(beta.grad().defined()) << "Gradient should flow back to beta"; - ASSERT_TRUE(torch::all(torch::isfinite(beta.grad())).item()) - << "Gradient w.r.t. beta should be finite"; + ASSERT_TRUE(concentration.grad().defined()) << "Gradient should flow back to concentration"; + ASSERT_TRUE(torch::all(torch::isfinite(concentration.grad())).item()) + << "Gradient w.r.t. concentration should be finite"; } -TEST_F(BetaLawGradientTest, EntropyGradientFlowsThroughAlpha) { - const auto alpha = torch::full({5}, 2.0f, torch::TensorOptions().requires_grad(true)); - const auto beta = torch::ones({5}) * 3.0f; +TEST_F(BetaLawGradientTest, EntropyGradientFlowsThroughMode) { + const auto mode = torch::full({5}, 0.4f, torch::TensorOptions().requires_grad(true)); + const auto concentration = torch::ones({5}) * 6.0f; - const auto entropy = beta_law_entropy(alpha, beta); + const auto entropy = beta_law_entropy(mode, concentration); const auto loss = entropy.sum(); loss.backward(); - ASSERT_TRUE(alpha.grad().defined()) << "Entropy gradient should flow back to alpha"; - ASSERT_TRUE(torch::all(torch::isfinite(alpha.grad())).item()) - << "Entropy gradient w.r.t. alpha should be finite"; + ASSERT_TRUE(mode.grad().defined()) << "Entropy gradient should flow back to mode"; + ASSERT_TRUE(torch::all(torch::isfinite(mode.grad())).item()) + << "Entropy gradient w.r.t. mode should be finite"; } -TEST_F(BetaLawGradientTest, EntropyGradientFlowsThroughBeta) { - const auto alpha = torch::ones({5}) * 2.0f; - const auto beta = torch::full({5}, 3.0f, torch::TensorOptions().requires_grad(true)); +TEST_F(BetaLawGradientTest, EntropyGradientFlowsThroughConcentration) { + const auto mode = torch::ones({5}) * 0.4f; + const auto concentration = torch::full({5}, 6.0f, torch::TensorOptions().requires_grad(true)); - const auto entropy = beta_law_entropy(alpha, beta); + const auto entropy = beta_law_entropy(mode, concentration); const auto loss = entropy.sum(); loss.backward(); - ASSERT_TRUE(beta.grad().defined()) << "Entropy gradient should flow back to beta"; - ASSERT_TRUE(torch::all(torch::isfinite(beta.grad())).item()) - << "Entropy gradient w.r.t. beta should be finite"; + ASSERT_TRUE(concentration.grad().defined()) + << "Entropy gradient should flow back to concentration"; + ASSERT_TRUE(torch::all(torch::isfinite(concentration.grad())).item()) + << "Entropy gradient w.r.t. concentration should be finite"; } diff --git a/arenai_agent/tests/src/tests_distributions/tests_multinomial.cpp b/arenai_agent/tests/src/tests_distributions/tests_multinomial.cpp deleted file mode 100644 index cc015342..00000000 --- a/arenai_agent/tests/src/tests_distributions/tests_multinomial.cpp +++ /dev/null @@ -1,148 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#include - -#include - -using namespace arenai; -using namespace arenai::agent; - -// ======================================================================== -// Fixed tests -// ======================================================================== - -TEST_F(MultinomialTest, EntropyMaxAtUniform) { - constexpr int n = 5; - const auto uniform = torch::ones({1, n}) / static_cast(n); - - const auto entropy = multinomial_entropy(uniform); - - ASSERT_NEAR(entropy.item(), std::log(static_cast(n)), 1e-4f); -} - -TEST_F(MultinomialTest, EntropyMinAtDegenerate) { - const auto proba = torch::zeros({1, 4}); - proba[0][0] = 1.0f; - - const auto entropy = multinomial_entropy(proba); - - ASSERT_NEAR(entropy.item(), 0.0f, 1e-3f); -} - -TEST_F(MultinomialTest, TargetEntropySymmetric) { - const auto target = multinomial_target_entropy(0.5f); - const auto max_2 = multinomial_maximum_entropy(2); - - ASSERT_NEAR(target, max_2, 1e-5f); -} - -TEST_F(MultinomialTest, MaxEntropyEqualsLogN) { - for (const int n: {2, 3, 5, 10}) { - const auto max_ent = multinomial_maximum_entropy(n); - ASSERT_NEAR(max_ent, std::log(static_cast(n)), 1e-4f) - << "Maximum entropy for n=" << n << " should be log(n)"; - } -} - -// ======================================================================== -// Parameterized: sample shape and one-hot property -// ======================================================================== - -TEST_P(MultinomialShapeParamTest, SampleIsOneHot) { - const auto [batch_size, nb_actions] = GetParam(); - - const auto proba = torch::softmax(torch::randn({batch_size, nb_actions}), -1); - - const auto sample = multinomial_sample(proba); - - ASSERT_EQ(sample.size(0), batch_size); - ASSERT_EQ(sample.size(1), nb_actions); - - // each row sums to 1 - const auto row_sums = torch::sum(sample, -1); - ASSERT_TRUE(torch::allclose(row_sums, torch::ones({batch_size}))) - << "Each sample row should sum to 1"; - - // each element is 0 or 1 - const auto is_binary = torch::logical_or(torch::eq(sample, 0.0f), torch::eq(sample, 1.0f)); - ASSERT_TRUE(torch::all(is_binary).item()) << "Sample should contain only 0s and 1s"; -} - -TEST_P(MultinomialShapeParamTest, EntropyShape) { - const auto [batch_size, nb_actions] = GetParam(); - - const auto proba = torch::softmax(torch::randn({batch_size, nb_actions}), -1); - - const auto entropy = multinomial_entropy(proba); - - ASSERT_EQ(entropy.size(0), batch_size); - ASSERT_EQ(entropy.size(1), 1); - ASSERT_TRUE(torch::all(torch::ge(entropy, 0.0f)).item()) - << "Entropy should be non-negative"; - ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()); -} - -TEST_P(MultinomialShapeParamTest, EntropyBoundedByLogN) { - const auto [batch_size, nb_actions] = GetParam(); - - const auto proba = torch::softmax(torch::randn({batch_size, nb_actions}), -1); - - const auto entropy = multinomial_entropy(proba); - const auto max_entropy = std::log(static_cast(nb_actions)); - - ASSERT_TRUE(torch::all(torch::le(entropy, max_entropy + 1e-4f)).item()) - << "Entropy should be <= log(nb_actions)"; -} - -INSTANTIATE_TEST_SUITE_P( - MultinomialShape, MultinomialShapeParamTest, - testing::Combine(testing::Values(1, 2, 8, 16), testing::Values(2, 3, 4, 5, 10))); - -// ======================================================================== -// Parameterized: maximum entropy monotone -// ======================================================================== - -TEST_P(MultinomialMaxEntropyParamTest, MonotoneIncreasing) { - const auto nb_actions = GetParam(); - - if (nb_actions < 2) return; - - const auto ent_n = multinomial_maximum_entropy(nb_actions); - const auto ent_n_minus_1 = multinomial_maximum_entropy(nb_actions - 1); - - ASSERT_GT(ent_n, ent_n_minus_1); -} - -INSTANTIATE_TEST_SUITE_P( - MultinomialMaxEntropy, MultinomialMaxEntropyParamTest, testing::Values(2, 3, 4, 5, 10, 20)); - -// ======================================================================== -// Parameterized: target entropy -// ======================================================================== - -TEST_P(MultinomialTargetEntropyParamTest, TargetEntropyBounded) { - const auto shoot_probability = GetParam(); - - const auto target = multinomial_target_entropy(shoot_probability); - - ASSERT_GE(target, 0.0f); - ASSERT_LE(target, multinomial_maximum_entropy(2) + 1e-5f); - ASSERT_TRUE(std::isfinite(target)); -} - -TEST_P(MultinomialTargetEntropyParamTest, HigherProbabilityLowerEntropy) { - const auto shoot_probability = GetParam(); - - if (shoot_probability >= 0.5f) return; - - const auto target_low = multinomial_target_entropy(shoot_probability); - const auto target_half = multinomial_target_entropy(0.5f); - - ASSERT_LT(target_low, target_half); -} - -INSTANTIATE_TEST_SUITE_P( - MultinomialTargetEntropy, MultinomialTargetEntropyParamTest, - testing::Values(0.1f, 0.2f, 0.3f, 0.4f, 0.5f, 0.8f, 0.9f)); diff --git a/arenai_agent/tests/src/tests_distributions/tests_multinomial_edge.cpp b/arenai_agent/tests/src/tests_distributions/tests_multinomial_edge.cpp deleted file mode 100644 index e13c1b03..00000000 --- a/arenai_agent/tests/src/tests_distributions/tests_multinomial_edge.cpp +++ /dev/null @@ -1,64 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#include - -#include - -using namespace arenai; -using namespace arenai::agent; - -TEST_F(MultinomialEdgeTest, EntropyWithNearZeroProbabilities) { - const auto proba = torch::zeros({1, 5}); - proba[0][0] = 1e-10f; - proba[0][1] = 1e-10f; - proba[0][2] = 1e-10f; - proba[0][3] = 1e-10f; - proba[0][4] = 1.f - 4e-10f; - - const auto entropy = multinomial_entropy(proba); - - ASSERT_TRUE(torch::all(torch::isfinite(entropy)).item()) - << "Entropy should be finite with near-zero probabilities"; - ASSERT_GE(entropy.item(), 0.0f) << "Entropy should be non-negative"; -} - -TEST_F(MultinomialEdgeTest, EntropyGradientFlowsThroughProbabilities) { - const auto logits = torch::randn({4, 3}, torch::TensorOptions().requires_grad(true)); - const auto proba = torch::softmax(logits, -1); - - const auto entropy = multinomial_entropy(proba); - const auto loss = entropy.sum(); - - loss.backward(); - - ASSERT_TRUE(logits.grad().defined()) << "Gradient should flow back through entropy"; - ASSERT_TRUE(torch::all(torch::isfinite(logits.grad())).item()) - << "Gradient should be finite"; -} - -TEST_F(MultinomialEdgeTest, SampleWithSingleAction) { - const auto proba = torch::ones({4, 1}); - - const auto sample = multinomial_sample(proba); - - ASSERT_EQ(sample.size(0), 4); - ASSERT_EQ(sample.size(1), 1); - ASSERT_TRUE(torch::allclose(sample, torch::ones({4, 1}))) - << "Single-action sample should always be 1"; -} - -TEST_F(MultinomialEdgeTest, MaximumEntropyWithSingleAction) { - const auto max_ent = multinomial_maximum_entropy(1); - - ASSERT_NEAR(max_ent, 0.0f, 1e-4f) << "Maximum entropy with 1 action should be 0 (log(1)=0)"; -} - -TEST_F(MultinomialEdgeTest, TargetEntropyBoundaryProbabilities) { - const auto target_0 = multinomial_target_entropy(0.0f); - const auto target_1 = multinomial_target_entropy(1.0f); - - ASSERT_TRUE(std::isfinite(target_0)) << "Target entropy with p=0 should be finite (clamped)"; - ASSERT_TRUE(std::isfinite(target_1)) << "Target entropy with p=1 should be finite (clamped)"; -} diff --git a/arenai_agent/tests/src/tests_distributions/tests_truncated_normal_edge.cpp b/arenai_agent/tests/src/tests_distributions/tests_truncated_normal_edge.cpp index c348b690..1118aa50 100644 --- a/arenai_agent/tests/src/tests_distributions/tests_truncated_normal_edge.cpp +++ b/arenai_agent/tests/src/tests_distributions/tests_truncated_normal_edge.cpp @@ -67,7 +67,7 @@ TEST_F(TruncatedNormalEdgeTest, LogPdfAndPdfConsistent) { } // ======================================================================== -// Gradient flow tests — critical for SAC training +// Gradient flow tests — critical for gradient-based training // ======================================================================== TEST_F(TruncatedNormalGradientTest, LogPdfGradientFlowsThroughMu) { diff --git a/arenai_agent/tests/src/tests_e2e/tests_e2e_agent_input.cpp b/arenai_agent/tests/src/tests_e2e/tests_e2e_agent_input.cpp index 4549d0be..a55be57c 100644 --- a/arenai_agent/tests/src/tests_e2e/tests_e2e_agent_input.cpp +++ b/arenai_agent/tests/src/tests_e2e/tests_e2e_agent_input.cpp @@ -51,7 +51,7 @@ namespace { const auto steps = env.step(FREQUENCY, actions); states.clear(); - for (const auto &[state, reward, done]: steps) states.push_back(state); + for (const auto &[state, reward, done, truncated]: steps) states.push_back(state); actions = agent.act(states, VISION_HEIGHT, VISION_WIDTH); } diff --git a/arenai_agent/tests/src/tests_networks/tests_actor.cpp b/arenai_agent/tests/src/tests_networks/tests_actor.cpp index 08c0b673..897ced7c 100644 --- a/arenai_agent/tests/src/tests_networks/tests_actor.cpp +++ b/arenai_agent/tests/src/tests_networks/tests_actor.cpp @@ -3,6 +3,7 @@ // #include +#include #include @@ -20,31 +21,36 @@ TEST_P(ActorTestParam, TestActorAct) { Actor actor( height, width, sensors_nb, cont_actions_nb, discrete_actions_nb, sensors_hidden_size, - layers, {{input_channels, 4}, {4, 8}}, {2, 4}, 0.1f, 0.2f); + layers, {{input_channels, 4}, {4, 8}}, {2, 4}, 0.1f, + std::vector(discrete_actions_nb, 0.2f)); const auto image = torch::randint( 255, {batch_size, input_channels, height, width}, torch::TensorOptions().dtype(torch::kUInt8)); const auto sensors = torch::randn({batch_size, sensors_nb}); - const auto [mu, sigma, discrete] = actor.act(image, sensors); + const auto [mode, concentration, discrete] = actor.act(image, sensors); - ASSERT_EQ(mu.ndimension(), 2); - ASSERT_EQ(mu.size(0), batch_size); - ASSERT_EQ(mu.size(1), cont_actions_nb); + ASSERT_EQ(mode.ndimension(), 2); + ASSERT_EQ(mode.size(0), batch_size); + ASSERT_EQ(mode.size(1), cont_actions_nb); ASSERT_TRUE( - torch::all(torch::logical_and(torch::ge(mu, -1.0), torch::le(mu, 1.0))).item()); + torch::all(torch::logical_and(torch::ge(mode, 0.0), torch::le(mode, 1.0))).item()); - ASSERT_EQ(sigma.ndimension(), 2); - ASSERT_EQ(sigma.size(0), batch_size); - ASSERT_EQ(sigma.size(1), cont_actions_nb); - ASSERT_TRUE( - torch::all(torch::logical_and(torch::gt(sigma, 0.0), torch::le(sigma, 1.0))).item()); + ASSERT_EQ(concentration.ndimension(), 2); + ASSERT_EQ(concentration.size(0), batch_size); + ASSERT_EQ(concentration.size(1), cont_actions_nb); + ASSERT_TRUE(torch::all(torch::logical_and( + torch::ge(concentration, CONCENTRATION_MIN), + torch::le(concentration, CONCENTRATION_MAX))) + .item()); ASSERT_EQ(discrete.ndimension(), 2); ASSERT_EQ(discrete.size(0), batch_size); ASSERT_EQ(discrete.size(1), discrete_actions_nb); - ASSERT_TRUE(torch::all(torch::abs(torch::sum(discrete, -1) - 1.0) < 1e-6).item()); + // independent Bernoulli probabilities: each in (0, 1), no sum constraint + ASSERT_TRUE(torch::all(torch::logical_and(torch::gt(discrete, 0.0), torch::lt(discrete, 1.0))) + .item()); } INSTANTIATE_TEST_SUITE_P( diff --git a/arenai_agent/tests/src/tests_networks/tests_entropy.cpp b/arenai_agent/tests/src/tests_networks/tests_entropy.cpp index f6bb8977..1a4cc5c7 100644 --- a/arenai_agent/tests/src/tests_networks/tests_entropy.cpp +++ b/arenai_agent/tests/src/tests_networks/tests_entropy.cpp @@ -9,48 +9,6 @@ using namespace arenai; using namespace arenai::agent; -TEST_F(AlphaParameterTest, AlphaAlwaysPositive) { - for (const float init: {0.01f, 0.1f, 1.0f, 10.0f}) { - AlphaParameters param(init, 1); - ASSERT_GT(param.alpha().item(), 0.0f) - << "alpha should be positive for initial_alpha=" << init; - } -} - -TEST_F(AlphaParameterTest, InitialValueMatchesInput) { - constexpr float initial = 0.2f; - AlphaParameters param(initial, 1); - - ASSERT_NEAR(param.alpha().item(), initial, 1e-6f); -} - -TEST_F(AlphaParameterTest, LogAlphaRequiresGrad) { - AlphaParameters param(1.0f, 1); - - ASSERT_TRUE(param.log_alpha().requires_grad()); -} - -TEST_F(AlphaParameterTest, LogAlphaConsistentWithAlpha) { - AlphaParameters param(0.5f, 1); - - const auto log_a = param.log_alpha().item(); - const auto a = param.alpha().item(); - - ASSERT_NEAR(std::exp(log_a), a, 1e-6f); -} - -/* - * Target entropy - */ - -TEST_F(ConstantTargetEntropyTest, StepIsANoOp) { - ConstantTargetEntropy target(0.117f); - - target.step(1000000); - - ASSERT_NEAR(target.target_entropy().item(), 0.117f, 1e-6f); -} - /* * PID Lagrangian */ diff --git a/arenai_agent/tests/src/tests_networks/tests_liquid_cell.cpp b/arenai_agent/tests/src/tests_networks/tests_liquid_cell.cpp new file mode 100644 index 00000000..13132d11 --- /dev/null +++ b/arenai_agent/tests/src/tests_networks/tests_liquid_cell.cpp @@ -0,0 +1,148 @@ +// +// Created by samuel on 06/09/2026. +// + +#include + +#include + +using namespace arenai; +using namespace arenai::agent; + +namespace { + constexpr float DELTA_T = 1.f; + + torch::Tensor tanh_activation(const torch::Tensor &t) { return torch::tanh(t); } + + void assert_all_parameters_receive_gradient(const torch::nn::Module &module) { + for (const auto &p: module.named_parameters()) { + ASSERT_TRUE(p.value().grad().defined()) + << "Parameter \"" << p.key() << "\" has no gradient"; + ASSERT_TRUE(torch::all(torch::isfinite(p.value().grad())).item()) + << "Parameter \"" << p.key() << "\" has non-finite gradients"; + ASSERT_GT(p.value().grad().abs().sum().item(), 0.f) + << "Parameter \"" << p.key() << "\" has an all-zero gradient"; + } + } +}// namespace + +/* + * Cell model + */ + +TEST_P(CellModelTestParam, OutputShape) { + const auto [neuron_number, input_size, batch_size] = GetParam(); + + CellModel cell_model(neuron_number, input_size, tanh_activation); + + const auto x_t = torch::randn({batch_size, neuron_number}); + const auto input_t = torch::randn({batch_size, input_size}); + + const auto output = cell_model.forward(x_t, input_t); + + ASSERT_EQ(output.ndimension(), 2); + ASSERT_EQ(output.size(0), batch_size); + ASSERT_EQ(output.size(1), neuron_number); +} + +TEST_P(CellModelTestParam, GradientFlows) { + const auto [neuron_number, input_size, batch_size] = GetParam(); + + CellModel cell_model(neuron_number, input_size, tanh_activation); + + const auto x_t = torch::randn({batch_size, neuron_number}, torch::requires_grad()); + const auto input_t = torch::randn({batch_size, input_size}, torch::requires_grad()); + + cell_model.forward(x_t, input_t).sum().backward(); + + assert_all_parameters_receive_gradient(cell_model); + + ASSERT_TRUE(x_t.grad().defined()); + ASSERT_GT(x_t.grad().abs().sum().item(), 0.f); + ASSERT_TRUE(input_t.grad().defined()); + ASSERT_GT(input_t.grad().abs().sum().item(), 0.f); +} + +INSTANTIATE_TEST_SUITE_P( + TestCellModel, CellModelTestParam, + testing::Combine(testing::Values(2, 8, 16), testing::Values(1, 3, 12), testing::Values(1, 4))); + +/* + * Liquid cell + */ + +TEST_P(LiquidCellTestParam, StepOutputShape) { + const auto [neuron_number, input_size, output_size, unfolding_steps, batch_size, time_steps] = + GetParam(); + + LiquidCell liquid_cell( + neuron_number, input_size, output_size, unfolding_steps, tanh_activation, DELTA_T); + + const auto x_t = torch::randn({batch_size, neuron_number}); + const auto input_t = torch::randn({batch_size, input_size}); + + const auto [output, x_t_next] = liquid_cell.forward_step(x_t, input_t); + + ASSERT_EQ(output.ndimension(), 2); + ASSERT_EQ(output.size(0), batch_size); + ASSERT_EQ(output.size(1), output_size); + + ASSERT_EQ(x_t_next.ndimension(), 2); + ASSERT_EQ(x_t_next.size(0), batch_size); + ASSERT_EQ(x_t_next.size(1), neuron_number); +} + +TEST_P(LiquidCellTestParam, OutputShape) { + const auto [neuron_number, input_size, output_size, unfolding_steps, batch_size, time_steps] = + GetParam(); + + LiquidCell liquid_cell( + neuron_number, input_size, output_size, unfolding_steps, tanh_activation, DELTA_T); + + const auto inputs = torch::randn({batch_size, time_steps, input_size}); + + const auto output = liquid_cell.forward(inputs); + + ASSERT_EQ(output.ndimension(), 3); + ASSERT_EQ(output.size(0), batch_size); + ASSERT_EQ(output.size(1), time_steps); + ASSERT_EQ(output.size(2), output_size); +} + +TEST_P(LiquidCellTestParam, GradientFlows) { + const auto [neuron_number, input_size, output_size, unfolding_steps, batch_size, time_steps] = + GetParam(); + + // LayerNorm over a single feature outputs exactly zero (x - mean(x) == 0), + // which blocks any gradient from reaching the layers before it + if (output_size == 1) GTEST_SKIP() << "LayerNorm({1}) zeroes the signal, no gradient flows"; + + LiquidCell liquid_cell( + neuron_number, input_size, output_size, unfolding_steps, tanh_activation, DELTA_T); + + const auto inputs = torch::randn({batch_size, time_steps, input_size}, torch::requires_grad()); + + liquid_cell.forward(inputs).sum().backward(); + + assert_all_parameters_receive_gradient(liquid_cell); + + ASSERT_TRUE(inputs.grad().defined()); + ASSERT_TRUE(torch::all(torch::isfinite(inputs.grad())).item()); + // every time step must contribute to the loss + const auto per_step_grad = inputs.grad().abs().sum(std::vector{0, 2}); + ASSERT_TRUE(torch::all(torch::gt(per_step_grad, 0.f)).item()) + << "Some time steps received no gradient: " << per_step_grad; +} + +TEST(LiquidCellEdge, RejectsNon3DInput) { + LiquidCell liquid_cell(4, 3, 2, 2, tanh_activation, DELTA_T); + + EXPECT_THROW(liquid_cell.forward(torch::randn({2, 3})), c10::Error); + EXPECT_THROW(liquid_cell.forward(torch::randn({2, 5, 3, 1})), c10::Error); +} + +INSTANTIATE_TEST_SUITE_P( + TestLiquidCell, LiquidCellTestParam, + testing::Combine( + testing::Values(2, 8, 16), testing::Values(1, 3, 12), testing::Values(1, 2, 6), + testing::Values(1, 3), testing::Values(1, 4), testing::Values(1, 2, 5))); diff --git a/arenai_agent/tests/src/tests_networks/tests_q_function.cpp b/arenai_agent/tests/src/tests_networks/tests_q_function.cpp deleted file mode 100644 index 27c726bd..00000000 --- a/arenai_agent/tests/src/tests_networks/tests_q_function.cpp +++ /dev/null @@ -1,77 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#include - -#include - -using namespace arenai; -using namespace arenai::agent; - -TEST_P(QFunctionTestParam, TestQFunctionPerDiscreteAction) { - const auto - [layers, cont_actions_nb, discrete_actions_nb, sensors_nb, sensors_hidden_size, - actions_hidden_size, batch_size] = GetParam(); - - constexpr int input_channels = 3; - constexpr int width = 32; - constexpr int height = 32; - - QFunction q_function( - height, width, sensors_nb, cont_actions_nb, discrete_actions_nb, sensors_hidden_size, - actions_hidden_size, layers, {{input_channels, 4}, {4, 8}}, {2, 4}); - - const auto image = torch::randint( - 255, {batch_size, input_channels, height, width}, - torch::TensorOptions().dtype(torch::kUInt8)); - const auto sensors = torch::randn({batch_size, sensors_nb}); - - const auto continuous_actions = torch::rand({batch_size, cont_actions_nb}) * 2.f - 1.f; - - const auto value = q_function.value_per_discrete_action(image, sensors, continuous_actions); - - ASSERT_EQ(value.ndimension(), 2); - ASSERT_EQ(value.size(0), batch_size); - ASSERT_EQ(value.size(1), discrete_actions_nb); -} - -TEST_P(QFunctionTestParam, TestQFunctionOHE) { - const auto - [layers, cont_actions_nb, discrete_actions_nb, sensors_nb, sensors_hidden_size, - actions_hidden_size, batch_size] = GetParam(); - - constexpr int input_channels = 3; - constexpr int width = 32; - constexpr int height = 32; - - QFunction q_function( - height, width, sensors_nb, cont_actions_nb, discrete_actions_nb, sensors_hidden_size, - actions_hidden_size, layers, {{input_channels, 4}, {4, 8}}, {2, 4}); - - const auto image = torch::randint( - 255, {batch_size, input_channels, height, width}, - torch::TensorOptions().dtype(torch::kUInt8)); - const auto sensors = torch::randn({batch_size, sensors_nb}); - - const auto continuous_actions = torch::rand({batch_size, cont_actions_nb}) * 2.f - 1.f; - const auto chosen_actions_index = - torch::randint(discrete_actions_nb, {batch_size, 1}).to(torch::kLong); - - auto discrete_actions_ohe = torch::zeros({batch_size, discrete_actions_nb}); - discrete_actions_ohe = discrete_actions_ohe.scatter_(1, chosen_actions_index, 1.0); - - const auto value = - q_function.value_ohe(image, sensors, continuous_actions, discrete_actions_ohe); - - ASSERT_EQ(value.ndimension(), 2); - ASSERT_EQ(value.size(0), batch_size); - ASSERT_EQ(value.size(1), 1); -} - -INSTANTIATE_TEST_SUITE_P( - TestQFunction, QFunctionTestParam, - testing::Combine( - testing::Values(HiddenLayers{16, 32}, HiddenLayers{2, 3}), testing::Values(1, 2, 3), - testing::Values(2, 3, 4), testing::Values(1, 2, 3), testing::Values(2, 3, 4), - testing::Values(2, 3, 4), testing::Values(1, 2, 3))); diff --git a/arenai_agent/tests/src/tests_networks/tests_q_function_consistency.cpp b/arenai_agent/tests/src/tests_networks/tests_q_function_consistency.cpp deleted file mode 100644 index 983873df..00000000 --- a/arenai_agent/tests/src/tests_networks/tests_q_function_consistency.cpp +++ /dev/null @@ -1,150 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#include -#include - -#include - -using namespace arenai; -using namespace arenai::agent; - -// ======================================================================== -// value_per_discrete_action must match value_ohe action by action -// ======================================================================== - -TEST_F(QFunctionConsistencyTest, PerDiscreteActionMatchesOhe) { - constexpr int height = 8, width = 8; - constexpr int nb_sensors = 5, nb_cont = 3, nb_disc = 4; - constexpr int batch = 4; - - QFunction q(height, width, nb_sensors, nb_cont, nb_disc, 8, 8, {16}, {{3, 4}}, {2}); - - const auto vision = torch::randint(0, 255, {batch, 3, height, width}, torch::kUInt8); - const auto sensors = torch::randn({batch, nb_sensors}); - const auto cont_actions = torch::randn({batch, nb_cont}); - - torch::NoGradGuard no_grad; - - const auto q_per_action = q.value_per_discrete_action(vision, sensors, cont_actions); - - ASSERT_EQ(q_per_action.size(0), batch); - ASSERT_EQ(q_per_action.size(1), nb_disc); - - const auto one_hots = torch::eye(nb_disc); - for (int a = 0; a < nb_disc; a++) { - const auto ohe = one_hots[a].unsqueeze(0).expand({batch, -1}); - const auto q_a = q.value_ohe(vision, sensors, cont_actions, ohe); - - ASSERT_TRUE(torch::allclose(q_per_action.select(1, a).unsqueeze(1), q_a, 1e-4, 1e-4)) - << "value_per_discrete_action column " << a << " should equal value_ohe(one_hot[" << a - << "])"; - } -} - -TEST_F(QFunctionConsistencyTest, ValueOheOutputFinite) { - constexpr int height = 8, width = 8; - constexpr int nb_sensors = 5, nb_cont = 3, nb_disc = 2; - constexpr int batch = 4; - - QFunction q(height, width, nb_sensors, nb_cont, nb_disc, 8, 8, {16}, {{3, 4}}, {2}); - - const auto vision = torch::randint(0, 255, {batch, 3, height, width}, torch::kUInt8); - const auto sensors = torch::randn({batch, nb_sensors}); - const auto cont_actions = torch::randn({batch, nb_cont}); - const auto disc_ohe = torch::zeros({batch, nb_disc}); - disc_ohe.select(1, 0).fill_(1.0f); - - const auto value = q.value_ohe(vision, sensors, cont_actions, disc_ohe); - - ASSERT_TRUE(torch::all(torch::isfinite(value)).item()) << "Q-value should be finite"; -} - -// ======================================================================== -// Gradient flow tests for networks -// ======================================================================== - -TEST_F(QFunctionGradientTest, GradientFlowsThroughQFunction) { - constexpr int height = 8, width = 8; - constexpr int nb_sensors = 3, nb_cont = 2, nb_disc = 2; - constexpr int batch = 2; - - QFunction q(height, width, nb_sensors, nb_cont, nb_disc, 8, 8, {16}, {{3, 4}}, {2}); - - const auto vision = torch::randint(0, 255, {batch, 3, height, width}, torch::kUInt8); - const auto sensors = torch::randn({batch, nb_sensors}); - const auto cont_actions = torch::randn({batch, nb_cont}); - - const auto value = q.value_per_discrete_action(vision, sensors, cont_actions); - const auto loss = value.sum(); - - loss.backward(); - - bool any_grad = false; - for (const auto &p: q.parameters()) { - if (p.grad().defined() && p.grad().abs().sum().item() > 0) { - any_grad = true; - ASSERT_TRUE(torch::all(torch::isfinite(p.grad())).item()) - << "All gradients should be finite"; - } - } - ASSERT_TRUE(any_grad) << "At least some parameters should receive gradients"; -} - -TEST_F(ActorGradientTest, GradientFlowsThroughActor) { - constexpr int height = 8, width = 8; - constexpr int nb_sensors = 3, nb_cont = 2, nb_disc = 2; - constexpr int batch = 2; - - Actor actor(height, width, nb_sensors, nb_cont, nb_disc, 8, {16}, {{3, 4}}, {2}, 0.1f, 0.2f); - - const auto vision = torch::randint(0, 255, {batch, 3, height, width}, torch::kUInt8); - const auto sensors = torch::randn({batch, nb_sensors}); - - const auto [mu, sigma, discrete] = actor.act(vision, sensors); - const auto loss = mu.sum() + sigma.sum() + discrete.sum(); - - loss.backward(); - - bool any_grad = false; - for (const auto &p: actor.parameters()) { - if (p.grad().defined() && p.grad().abs().sum().item() > 0) { - any_grad = true; - ASSERT_TRUE(torch::all(torch::isfinite(p.grad())).item()) - << "All actor gradients should be finite"; - } - } - ASSERT_TRUE(any_grad) << "At least some actor parameters should receive gradients"; -} - -TEST_F(ActorGradientTest, ActorWeightsChangeAfterOptimStep) { - constexpr int height = 8, width = 8; - constexpr int nb_sensors = 3, nb_cont = 2, nb_disc = 2; - constexpr int batch = 2; - - Actor actor(height, width, nb_sensors, nb_cont, nb_disc, 8, {16}, {{3, 4}}, {2}, 0.1f, 0.2f); - - auto optimizer = torch::optim::Adam(actor.parameters(), 1e-3); - - auto params_before = std::vector(); - for (const auto &p: actor.parameters()) params_before.push_back(p.clone()); - - const auto vision = torch::randint(0, 255, {batch, 3, height, width}, torch::kUInt8); - const auto sensors = torch::randn({batch, nb_sensors}); - - const auto [mu, sigma, discrete] = actor.act(vision, sensors); - const auto loss = -(mu.sum() + sigma.log().sum()); - - optimizer.zero_grad(); - loss.backward(); - optimizer.step(); - - bool any_changed = false; - int i = 0; - for (const auto &p: actor.parameters()) { - if (!torch::equal(p, params_before[i])) any_changed = true; - i++; - } - ASSERT_TRUE(any_changed) << "Some actor parameters should change after optimizer step"; -} diff --git a/arenai_agent/tests/src/tests_networks/tests_vision_impala.cpp b/arenai_agent/tests/src/tests_networks/tests_vision_impala.cpp new file mode 100644 index 00000000..e56c1ee3 --- /dev/null +++ b/arenai_agent/tests/src/tests_networks/tests_vision_impala.cpp @@ -0,0 +1,92 @@ +// +// Created by claude on 20/09/2026. +// + +#include + +#include + +using namespace arenai; +using namespace arenai::agent; + +TEST_P(ImpalaVisionTestParam, TestImpalaVisionForward) { + const auto [width, height, channels, output_conv_channels, batch_size] = GetParam(); + + std::vector> conv_layers; + + int curr_channels = channels; + for (const auto &c_o: output_conv_channels) { + conv_layers.emplace_back(curr_channels, c_o); + curr_channels = c_o; + } + + ImpalaConvolutionNetwork conv(height, width, conv_layers); + + const auto images = torch::randint( + 255, {batch_size, channels, height, width}, torch::TensorOptions().dtype(torch::kUInt8)); + + const auto encoded_images = conv.forward(images); + + ASSERT_EQ(encoded_images.ndimension(), 2); + ASSERT_EQ(encoded_images.size(0), batch_size); + ASSERT_EQ(encoded_images.size(1), conv.get_output_size()); + ASSERT_EQ(encoded_images.size(1) % output_conv_channels.back(), 0); +} + +// Create parametrized tests + +INSTANTIATE_TEST_SUITE_P( + TestImpalaVision, ImpalaVisionTestParam, + testing::Combine( + testing::Values(16, 32), testing::Values(16, 32), testing::Values(1, 2, 3), + testing::Values( + ImpalaOutputConvChannels{4}, ImpalaOutputConvChannels{4, 8}, + ImpalaOutputConvChannels{16, 32, 48}), + testing::Values(1, 2, 3))); + +// Edge cases + +TEST_F(ImpalaVisionEdgeTest, RejectsNonUint8Input) { + ImpalaConvolutionNetwork conv(8, 8, {{3, 4}}); + + const auto float_input = torch::randn({1, 3, 8, 8}); + + ASSERT_THROW(conv.forward(float_input), c10::Error) << "Should throw when input is not UInt8"; +} + +TEST_F(ImpalaVisionEdgeTest, NormalizesToExpectedRange) { + ImpalaConvolutionNetwork conv(8, 8, {{3, 4}}); + + const auto zeros = torch::zeros({1, 3, 8, 8}, torch::kUInt8); + const auto result_zeros = conv.forward(zeros); + + const auto max255 = torch::ones({1, 3, 8, 8}, torch::kUInt8) * 255; + const auto result_255 = conv.forward(max255); + + ASSERT_TRUE(torch::all(torch::isfinite(result_zeros)).item()); + ASSERT_TRUE(torch::all(torch::isfinite(result_255)).item()); +} + +TEST_F(ImpalaVisionEdgeTest, OutputSizeMatchesGetOutputSize) { + const std::vector> channels = {{3, 8}, {8, 16}}; + constexpr int h = 16, w = 16; + constexpr int batch_size = 2; + + ImpalaConvolutionNetwork conv(h, w, channels); + + const auto input = torch::randint(255, {batch_size, 3, h, w}, torch::kUInt8); + const auto output = conv.forward(input); + + ASSERT_EQ(output.size(0), batch_size); + ASSERT_EQ(output.size(1), conv.get_output_size()); +} + +TEST_F(ImpalaVisionEdgeTest, ResidualBlockPreservesShape) { + ResidualBlock block(8); + + const auto input = torch::randn({2, 8, 16, 16}); + const auto output = block.forward(input); + + ASSERT_TRUE(output.sizes() == input.sizes()); + ASSERT_TRUE(torch::all(torch::isfinite(output)).item()); +} diff --git a/arenai_agent/tests/src/tests_networks_utils/tests_init.cpp b/arenai_agent/tests/src/tests_networks_utils/tests_init.cpp index 8e7e3862..3b22409f 100644 --- a/arenai_agent/tests/src/tests_networks_utils/tests_init.cpp +++ b/arenai_agent/tests/src/tests_networks_utils/tests_init.cpp @@ -81,23 +81,24 @@ TEST_F(InitWeightsTest, SigmaOutputIsEqualToWantedOne) { TEST_F(InitWeightsTest, DiscreteOutputWeightsOrthogonal) { torch::nn::Linear linear(32, 6); - init_discrete_output_weights(*linear, 0.f); + init_discrete_output_weights(*linear, std::vector(6, 0.5f)); assert_orthogonal(linear->weight, 0.01f); } -TEST_F(InitWeightsTest, DiscreteOutputFireProbaIsEqualToWantedOne) { - constexpr float wanted_fire_proba = 0.2f; +TEST_F(InitWeightsTest, DiscreteOutputProbasAreEqualToWantedOnes) { + const std::vector wanted_probas{0.2f, 0.35f}; torch::nn::Sequential seq( - torch::nn::Linear(32, model::ENEMY_NB_DISCRETE_ACTION), torch::nn::Softmax(-1)); - seq->apply([](torch::nn::Module &m) { init_discrete_output_weights(m, wanted_fire_proba); }); + torch::nn::Linear(32, model::ENEMY_NB_DISCRETE_ACTION), torch::nn::Sigmoid()); + seq->apply( + [&wanted_probas](torch::nn::Module &m) { init_discrete_output_weights(m, wanted_probas); }); torch::Tensor x = torch::randn({1, 32}); const auto out = seq->forward(x); - ASSERT_NEAR(out[0][0].item(), wanted_fire_proba, 1e-2f); - ASSERT_NEAR(out[0][1].item(), 1.f - wanted_fire_proba, 1e-2f); + ASSERT_NEAR(out[0][0].item(), wanted_probas[0], 1e-2f); + ASSERT_NEAR(out[0][1].item(), wanted_probas[1], 1e-2f); } TEST_F(InitWeightsTest, ValueOutputWeightsOrthogonal) { diff --git a/arenai_agent/tests/src/tests_replay_buffer/create_random_step.cpp b/arenai_agent/tests/src/tests_replay_buffer/create_random_step.cpp deleted file mode 100644 index dd7ceb33..00000000 --- a/arenai_agent/tests/src/tests_replay_buffer/create_random_step.cpp +++ /dev/null @@ -1,27 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#include "./create_random_step.h" - -using namespace arenai; -using namespace arenai::agent; - -// single-tank state: every tensor carries a leading nb_tanks dimension of 1 -TorchState create_random_state(const int width, const int height, const int nb_sensors) { - return { - .vision = torch::randint(255, {1, 3, height, width}, torch::kUInt8), - .proprioception = torch::randn({1, nb_sensors})}; -} - -SacInputStep create_random_step( - const int width, const int height, const int nb_cont_actions, const int nb_discrete_actions, - const int nb_sensors, const bool done) { - return { - .state = create_random_state(width, height, nb_sensors), - .action = - {.continuous_action = torch::rand({1, nb_cont_actions}) * 2.f - 1.f, - .discrete_action = torch::softmax(torch::randn({1, nb_discrete_actions}), -1)}, - .reward = torch::randn({1, 1}), - .done = torch::full({1, 1}, done, torch::kBool)}; -} diff --git a/arenai_agent/tests/src/tests_replay_buffer/create_random_step.h b/arenai_agent/tests/src/tests_replay_buffer/create_random_step.h deleted file mode 100644 index 7af26608..00000000 --- a/arenai_agent/tests/src/tests_replay_buffer/create_random_step.h +++ /dev/null @@ -1,16 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#ifndef ARENAI_ADD_RANDOM_STEP_H -#define ARENAI_ADD_RANDOM_STEP_H - -#include -#include - -arenai::agent::TorchState create_random_state(int width, int height, int nb_sensors); - -arenai::agent::SacInputStep create_random_step( - int width, int height, int nb_cont_actions, int nb_discrete_actions, int nb_sensors, bool done); - -#endif//ARENAI_ADD_RANDOM_STEP_H diff --git a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_add.cpp b/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_add.cpp deleted file mode 100644 index 9eb17ad6..00000000 --- a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_add.cpp +++ /dev/null @@ -1,49 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#include - -#include - -#include "./create_random_step.h" - -using namespace arenai; -using namespace arenai::agent; - -// ======================================================================== -// Normal: steps_to_add <= memory_size -// ======================================================================== - -TEST_P(ReplayBufferAddNormalTestParam, SizeMatchesAdded) { - const auto [memory_size, to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - - for (uint32_t i = 0; i < to_add; ++i) buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - - ASSERT_EQ(buffer.size(), to_add); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferAddNormal, ReplayBufferAddNormalTestParam, - testing::Combine(testing::Values(8, 16, 32), testing::Values(1, 2, 4, 8))); - -// ======================================================================== -// Overflow: steps_to_add > memory_size -// ======================================================================== - -TEST_P(ReplayBufferAddOverflowTestParam, SizeCappedAtMemorySize) { - const auto [memory_size, to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - - for (uint32_t i = 0; i < to_add; ++i) buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - - ASSERT_EQ(buffer.size(), memory_size); - ASSERT_LT(buffer.size(), to_add); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferAddOverflow, ReplayBufferAddOverflowTestParam, - testing::Combine(testing::Values(2, 4, 8), testing::Values(10, 16, 32))); diff --git a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_edge.cpp b/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_edge.cpp deleted file mode 100644 index 72df747d..00000000 --- a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_edge.cpp +++ /dev/null @@ -1,128 +0,0 @@ -// -// Created by claude on 01/07/2026. -// - -#include - -#include - -#include - -#include "./create_random_step.h" - -using namespace arenai; -using namespace arenai::agent; - -TEST_F(ReplayBufferEdgeTest, SampleFromEmptyBufferDoesNotCrash) { - const SacReplayBuffer buffer(10); - - ASSERT_EQ(buffer.size(), 0u); - - // Sampling from empty buffer should either throw or return a valid (if degenerate) result - // but must NOT produce undefined behavior - ASSERT_ANY_THROW(buffer.sample(1, torch::kCPU)) - << "Sampling from empty buffer should throw (randint with upper_bound=0 is invalid)"; -} - -TEST_F(ReplayBufferEdgeTest, SampleFromSingleElement) { - SacReplayBuffer buffer(10); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - buffer.finish_episode(create_random_state(8, 8, 5)); - - ASSERT_EQ(buffer.size(), 2u); - - const auto output = buffer.sample(1, torch::kCPU); - - ASSERT_EQ(output.state.vision.size(0), 1); - ASSERT_EQ(output.action.continuous_action.size(0), 1); -} - -TEST_F(ReplayBufferEdgeTest, SampleBatchLargerThanSingleElement) { - SacReplayBuffer buffer(10); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - buffer.finish_episode(create_random_state(8, 8, 5)); - - const auto output = buffer.sample(5, torch::kCPU); - - ASSERT_EQ(output.state.vision.size(0), 1) - << "Batch size should be clamped to the single available transition"; -} - -namespace { - SacInputStep create_step_with_reward(const float reward) { - SacInputStep step; - step.state.vision = torch::randint(255, {1, 3, 8, 8}, torch::kUInt8); - step.state.proprioception = torch::randn({1, 5}); - step.action.continuous_action = torch::randn({1, 3}); - step.action.discrete_action = torch::zeros({1, 2}); - step.action.discrete_action[0][0] = 1.0f; - step.reward = torch::full({1, 1}, reward); - step.done = torch::zeros({1, 1}); - return step; - } -}// namespace - -TEST_F(ReplayBufferEdgeTest, ConstantRewardsAreNotRescaled) { - SacReplayBuffer buffer(10); - - buffer.add(create_step_with_reward(2.0f)); - buffer.add(create_step_with_reward(2.0f)); - - const auto output = buffer.sample(1, torch::kCPU); - - ASSERT_NEAR(output.reward.item(), 2.0f, 1e-5f) - << "Zero reward variance must keep the scale at 1 (no division by ~0)"; -} - -TEST_F(ReplayBufferEdgeTest, DeadStepIsStoredAsTerminal) { - SacReplayBuffer buffer(10); - - buffer.add(create_random_step(8, 8, 3, 2, 5, true)); - buffer.finish_episode(create_random_state(8, 8, 5)); - - const auto output = buffer.sample(1, torch::kCPU); - - ASSERT_TRUE(output.done.to(torch::kBool).item()) - << "A termination (done) must sample done=true"; -} - -TEST_F(ReplayBufferEdgeTest, SampleWithZeroBatchSize) { - SacReplayBuffer buffer(10); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - - const auto output = buffer.sample(0, torch::kCPU); - - ASSERT_EQ(output.state.vision.size(0), 1) << "Batch size 0 should be clamped to 1"; -} - -TEST_F(ReplayBufferEdgeTest, SampleWithNegativeBatchSize) { - SacReplayBuffer buffer(10); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - buffer.finish_episode(create_random_state(8, 8, 5)); - - const auto output = buffer.sample(-5, torch::kCPU); - - ASSERT_EQ(output.state.vision.size(0), 1) << "Negative batch size should be clamped to 1"; -} - -TEST_F(ReplayBufferEdgeTest, CircularOverwriteKeepsMaxSize) { - SacReplayBuffer buffer(3); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - ASSERT_EQ(buffer.size(), 2u); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - ASSERT_EQ(buffer.size(), 3u); - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - ASSERT_EQ(buffer.size(), 3u) << "Buffer size should not exceed capacity after wraparound"; - - buffer.add(create_random_step(8, 8, 3, 2, 5, false)); - ASSERT_EQ(buffer.size(), 3u); -} diff --git a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_sample.cpp b/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_sample.cpp deleted file mode 100644 index 21db16ec..00000000 --- a/arenai_agent/tests/src/tests_replay_buffer/tests_replay_buffer_sample.cpp +++ /dev/null @@ -1,147 +0,0 @@ -// -// Created by samuel on 30/06/2026. -// - -#include - -#include - -#include "./create_random_step.h" - -using namespace arenai; -using namespace arenai::agent; - -namespace { - constexpr int WIDTH = 16, HEIGHT = 8; - constexpr int CONT_ACTIONS_NB = 3, DISCRETE_ACTIONS_NB = 4; - constexpr int SENSORS_NB = 5; - - void assert_sample_shapes(const SacTrainStep &output, const int expected_batch_size) { - const auto &[state, action, reward, done, next_state] = output; - - ASSERT_EQ(state.vision.ndimension(), 4); - ASSERT_EQ(state.vision.size(0), expected_batch_size); - ASSERT_EQ(state.vision.size(1), 3); - ASSERT_EQ(state.vision.size(2), HEIGHT); - ASSERT_EQ(state.vision.size(3), WIDTH); - - ASSERT_EQ(state.proprioception.ndimension(), 2); - ASSERT_EQ(state.proprioception.size(0), expected_batch_size); - ASSERT_EQ(state.proprioception.size(1), SENSORS_NB); - - ASSERT_EQ(action.continuous_action.ndimension(), 2); - ASSERT_EQ(action.continuous_action.size(0), expected_batch_size); - ASSERT_EQ(action.continuous_action.size(1), CONT_ACTIONS_NB); - - ASSERT_EQ(action.discrete_action.ndimension(), 2); - ASSERT_EQ(action.discrete_action.size(0), expected_batch_size); - ASSERT_EQ(action.discrete_action.size(1), DISCRETE_ACTIONS_NB); - - ASSERT_EQ(reward.ndimension(), 2); - ASSERT_EQ(reward.size(0), expected_batch_size); - ASSERT_EQ(reward.size(1), 1); - - ASSERT_EQ(done.ndimension(), 2); - ASSERT_EQ(done.size(0), expected_batch_size); - ASSERT_EQ(done.size(1), 1); - - ASSERT_EQ(next_state.vision.ndimension(), 4); - ASSERT_EQ(next_state.vision.size(0), expected_batch_size); - ASSERT_EQ(next_state.vision.size(1), 3); - ASSERT_EQ(next_state.vision.size(2), HEIGHT); - ASSERT_EQ(next_state.vision.size(3), WIDTH); - - ASSERT_EQ(next_state.proprioception.ndimension(), 2); - ASSERT_EQ(next_state.proprioception.size(0), expected_batch_size); - ASSERT_EQ(next_state.proprioception.size(1), SENSORS_NB); - } - - void fill_buffer(SacReplayBuffer &buffer, const uint32_t count) { - for (uint32_t i = 0; i < count; ++i) - buffer.add(create_random_step( - WIDTH, HEIGHT, CONT_ACTIONS_NB, DISCRETE_ACTIONS_NB, SENSORS_NB, false)); - - // the final observation closes the episode and makes the last added step sampleable - buffer.finish_episode(create_random_state(WIDTH, HEIGHT, SENSORS_NB)); - } -}// namespace - -// ======================================================================== -// Normal batch size: batch_size <= steps_added <= memory_size -// ======================================================================== - -TEST_P(ReplayBufferSampleNormalTestParam, BatchSizeRespected) { - const auto [memory_size, batch_size, steps_to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - fill_buffer(buffer, steps_to_add); - - const auto output = buffer.sample(static_cast(batch_size), torch::kCPU); - - assert_sample_shapes(output, static_cast(batch_size)); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferSampleNormal, ReplayBufferSampleNormalTestParam, - testing::Combine( - testing::Values(16, 32), testing::Values(1, 4, 8), testing::Values(8, 16, 32))); - -// ======================================================================== -// Overflow batch size: batch_size > memory_size (buffer full) -// ======================================================================== - -TEST_P(ReplayBufferSampleOverflowTestParam, BatchSizeClampedToMemorySize) { - const auto [memory_size, batch_size, steps_to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - fill_buffer(buffer, steps_to_add); - - const auto output = buffer.sample(static_cast(batch_size), torch::kCPU); - - // once the ring has wrapped, the final-observation slot is not sampleable: - // memory_size - 1 transitions remain - assert_sample_shapes(output, static_cast(memory_size) - 1); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferSampleOverflow, ReplayBufferSampleOverflowTestParam, - testing::Combine( - testing::Values(4, 8), testing::Values(16, 32, 64), testing::Values(8, 16, 32))); - -// ======================================================================== -// Underflow batch size: batch_size > steps_added (buffer not full yet) -// ======================================================================== - -TEST_P(ReplayBufferSampleUnderflowTestParam, BatchSizeClampedToBufferSize) { - const auto [memory_size, batch_size, steps_to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - fill_buffer(buffer, steps_to_add); - - const auto output = buffer.sample(static_cast(batch_size), torch::kCPU); - - assert_sample_shapes(output, static_cast(steps_to_add)); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferSampleUnderflow, ReplayBufferSampleUnderflowTestParam, - testing::Combine(testing::Values(32, 64), testing::Values(16, 32), testing::Values(1, 2, 4))); - -// ======================================================================== -// Double overflow: batch_size > memory_size > steps_added -// ======================================================================== - -TEST_P(ReplayBufferSampleDoubleOverflowTestParam, BatchSizeClampedToBufferSize) { - const auto [memory_size, batch_size, steps_to_add] = GetParam(); - - SacReplayBuffer buffer(static_cast(memory_size)); - fill_buffer(buffer, steps_to_add); - - const auto output = buffer.sample(static_cast(batch_size), torch::kCPU); - - assert_sample_shapes(output, static_cast(steps_to_add)); -} - -INSTANTIATE_TEST_SUITE_P( - ReplayBufferSampleDoubleOverflow, ReplayBufferSampleDoubleOverflowTestParam, - testing::Combine(testing::Values(8, 16), testing::Values(32, 64), testing::Values(1, 2, 4))); diff --git a/arenai_agent/tests/src/tests_utils/tests_torch_converter.cpp b/arenai_agent/tests/src/tests_utils/tests_torch_converter.cpp index 97415edb..d427a94f 100644 --- a/arenai_agent/tests/src/tests_utils/tests_torch_converter.cpp +++ b/arenai_agent/tests/src/tests_utils/tests_torch_converter.cpp @@ -17,7 +17,7 @@ using namespace arenai::agent; TEST_F(TorchConverterTest, SingleActionMapping) { const auto continuous = torch::tensor({{0.1f, -0.2f, 0.3f, -0.4f}}); - const auto discrete = torch::tensor({{0.8f, 0.2f}}); + const auto discrete = torch::tensor({{1.f, 0.f}}); const auto actions = tensor_to_actions(continuous, discrete); @@ -27,24 +27,27 @@ TEST_F(TorchConverterTest, SingleActionMapping) { ASSERT_FLOAT_EQ(actions[0].right_joystick.x, 0.3f); ASSERT_FLOAT_EQ(actions[0].right_joystick.y, -0.4f); ASSERT_TRUE(actions[0].fire_button.pressed); + ASSERT_FALSE(actions[0].zoom_button.pressed); } -TEST_F(TorchConverterTest, FireButtonFalseWhenSecondLarger) { +TEST_F(TorchConverterTest, IndependentFireAndZoom) { const auto continuous = torch::zeros({1, 4}); - const auto discrete = torch::tensor({{0.2f, 0.8f}}); + const auto discrete = torch::tensor({{0.f, 1.f}}); const auto actions = tensor_to_actions(continuous, discrete); ASSERT_FALSE(actions[0].fire_button.pressed); + ASSERT_TRUE(actions[0].zoom_button.pressed); } -TEST_F(TorchConverterTest, FireButtonFalseWhenEqual) { +TEST_F(TorchConverterTest, BothActionsEngagedTogether) { const auto continuous = torch::zeros({1, 4}); - const auto discrete = torch::tensor({{0.5f, 0.5f}}); + const auto discrete = torch::tensor({{1.f, 1.f}}); const auto actions = tensor_to_actions(continuous, discrete); - ASSERT_FALSE(actions[0].fire_button.pressed); + ASSERT_TRUE(actions[0].fire_button.pressed); + ASSERT_TRUE(actions[0].zoom_button.pressed); } TEST_F(TorchConverterTest, SingleStateToTensor) { @@ -137,7 +140,7 @@ TEST_P(TensorToActionsParamTest, ContinuousValuesMatchTensor) { } } -TEST_P(TensorToActionsParamTest, DiscreteFireButtonConsistent) { +TEST_P(TensorToActionsParamTest, DiscreteButtonsConsistent) { const auto batch = GetParam(); const auto discrete = torch::rand({batch, 2}); @@ -147,7 +150,8 @@ TEST_P(TensorToActionsParamTest, DiscreteFireButtonConsistent) { auto acc = discrete.accessor(); for (int i = 0; i < batch; ++i) { - ASSERT_EQ(actions[i].fire_button.pressed, acc[i][0] > acc[i][1]); + ASSERT_EQ(actions[i].fire_button.pressed, acc[i][0] > 0.5f); + ASSERT_EQ(actions[i].zoom_button.pressed, acc[i][1] > 0.5f); } } diff --git a/arenai_controller/include/arenai_controller/inputs.h b/arenai_controller/include/arenai_controller/inputs.h index 1cd326df..23460300 100644 --- a/arenai_controller/include/arenai_controller/inputs.h +++ b/arenai_controller/include/arenai_controller/inputs.h @@ -25,6 +25,7 @@ namespace arenai::controller { joystick right_joystick; button fire_button; + button zoom_button; }; }// namespace arenai::controller diff --git a/arenai_core/include/arenai_core/environment.h b/arenai_core/include/arenai_core/environment.h index a1dfca97..9f5fe840 100644 --- a/arenai_core/include/arenai_core/environment.h +++ b/arenai_core/include/arenai_core/environment.h @@ -27,9 +27,9 @@ namespace arenai::core { const std::shared_ptr &file_reader, const std::shared_ptr &graphics_backend, int nb_tanks, float wanted_frequency, int vision_height, int vision_width, int vision_num_threads, - bool vision_thread_sleep); + bool vision_thread_sleep, bool apply_timeout); - virtual std::vector> + virtual std::vector> step(float time_delta, const std::vector &actions); std::vector reset(float spawn_width, float spawn_height); @@ -60,6 +60,8 @@ namespace arenai::core { bool drawing_started_; + bool apply_timeout; + std::shared_ptr graphics_backend; std::shared_ptr gl_context; diff --git a/arenai_core/include/arenai_core/types.h b/arenai_core/include/arenai_core/types.h index 164057e3..9e73f225 100644 --- a/arenai_core/include/arenai_core/types.h +++ b/arenai_core/include/arenai_core/types.h @@ -20,6 +20,10 @@ namespace arenai::core { typedef bool IsDone; + // done from starving out: the episode ends for the tank, but the value target must + // bootstrap instead of cutting the return (truncation, not termination) + typedef bool IsTruncated; + typedef controller::user_input Action; }// namespace arenai::core diff --git a/arenai_core/src/enemy_handler.cpp b/arenai_core/src/enemy_handler.cpp index 30f368f1..2be1e5df 100644 --- a/arenai_core/src/enemy_handler.cpp +++ b/arenai_core/src/enemy_handler.cpp @@ -19,7 +19,8 @@ namespace arenai::core { .right_joystick = {.x = event.right_joystick.x * turret_scale_per_frame, .y = event.right_joystick.y * turret_scale_per_frame}, - .fire_button = event.fire_button}; + .fire_button = event.fire_button, + .zoom_button = event.zoom_button}; return {true, action}; } diff --git a/arenai_core/src/environment.cpp b/arenai_core/src/environment.cpp index bc55b76a..f0448ecc 100644 --- a/arenai_core/src/environment.cpp +++ b/arenai_core/src/environment.cpp @@ -18,7 +18,7 @@ namespace arenai::core { const std::shared_ptr &file_reader, const std::shared_ptr &graphics_backend, const int nb_tanks, float wanted_frequency, const int vision_height, const int vision_width, - const int vision_num_threads, const bool vision_thread_sleep) + const int vision_num_threads, const bool vision_thread_sleep, const bool apply_timeout) : wanted_frequency(wanted_frequency), nb_tanks(nb_tanks), vision_height(vision_height), vision_width(vision_width), vision_num_threads(vision_num_threads), vision_thread_sleep(vision_thread_sleep), @@ -27,10 +27,10 @@ namespace arenai::core { vision_thread_sleep)), physic_engine(model::make_physic_engine(wanted_frequency)), nb_reset_frames(static_cast(4.f / wanted_frequency)), drawing_started_(false), - graphics_backend(graphics_backend), gl_context(graphics_backend->render_context()), - rng(dev()), file_reader(file_reader) {} + apply_timeout(apply_timeout), graphics_backend(graphics_backend), + gl_context(graphics_backend->render_context()), rng(dev()), file_reader(file_reader) {} - std::vector> + std::vector> BaseTanksEnvironment::step(const float time_delta, const std::vector &actions) { // 1. apply action @@ -51,7 +51,7 @@ namespace arenai::core { vision_pool_->loop_wait(); // 4. build State - std::vector> result; + std::vector> result; result.reserve(tanks.size()); for (int i = 0; i < tanks.size(); i++) { @@ -61,7 +61,7 @@ namespace arenai::core { // compute/get state result.emplace_back( State(vision_pool_->read_vision(i), tanks[i]->get_proprioception()), - tanks[i]->get_reward(), tanks[i]->is_dead()); + tanks[i]->get_reward(), tanks[i]->is_dead(), tanks[i]->is_timeout()); } return result; @@ -79,6 +79,15 @@ namespace arenai::core { "height_map", file_reader, "heightmap/heightmap6.png", glm::vec3(0., 40., 0.), glm::vec3(10., 200., 10.)); + // terrain surface lies in [40 - 200/2, 40 + 200/2] = [-60, 140]: a vertical + // segment from 300 to -300 always crosses it inside the map + constexpr float terrain_max_y = 40.f + 200.f / 2.f; + const auto ground_height = [this](const float x, const float z) { + const auto hit = + physic_engine->ray_cast(glm::vec3(x, 300.f, z), glm::vec3(x, -300.f, z)); + return hit.has_value() ? hit->y : terrain_max_y; + }; + std::uniform_real_distribution x_pos_u_dist(-spawn_width / 2, spawn_width / 2); std::uniform_real_distribution y_pos_u_dist(-spawn_height / 2, spawn_height / 2); @@ -86,9 +95,10 @@ namespace arenai::core { // add tanks for (int i = 0; i < nb_tanks; i++) { + const float x = x_pos_u_dist(rng), z = y_pos_u_dist(rng); tanks.push_back(tank_factory->make_enemy_tank( file_reader, "enemy_" + std::to_string(i), - glm::vec3(x_pos_u_dist(rng), 0.f, y_pos_u_dist(rng)))); + glm::vec3(x, ground_height(x, z) + 3.f, z), apply_timeout)); tank_controller_handler.push_back(std::make_unique( wanted_frequency, model::ENEMY_TURRET_RADIAL_VELOCITY)); @@ -98,29 +108,42 @@ namespace arenai::core { } // add basic shapes + constexpr float item_spawn_size = 2000.f; + std::uniform_real_distribution item_x_pos_u_dist(-item_spawn_size / 2, item_spawn_size / 2); + std::uniform_real_distribution item_y_pos_u_dist(-item_spawn_size / 2, item_spawn_size / 2); + std::uniform_real_distribution scale_u_dist(2.5, 10); - constexpr int nb_shapes = 5; + constexpr int nb_shapes = 30; + + // unit meshes scaled by scale.y: spawn the item's half-height (+1 margin) above the ground + const auto item_pos = [&](const float x, const float z, const glm::vec3 &scale) { + return glm::vec3(x, ground_height(x, z) + scale.y + 1.f, z); + }; for (int i = 0; i < nb_shapes; i++) { - glm::vec3 pos(x_pos_u_dist(rng), 0.f, y_pos_u_dist(rng)); + float x = item_x_pos_u_dist(rng), z = item_y_pos_u_dist(rng); glm::vec3 scale(scale_u_dist(rng)); item_factory->make_sphere_item( - "sphere_" + std::to_string(i), file_reader, pos, scale, mass_u_dist(rng)); + "sphere_" + std::to_string(i), file_reader, item_pos(x, z, scale), scale, + mass_u_dist(rng)); - pos = glm::vec3(x_pos_u_dist(rng), 0.f, y_pos_u_dist(rng)); + x = item_x_pos_u_dist(rng), z = item_y_pos_u_dist(rng); scale = glm::vec3(scale_u_dist(rng)); item_factory->make_cube_item( - "cube_" + std::to_string(i), file_reader, pos, scale, mass_u_dist(rng)); + "cube_" + std::to_string(i), file_reader, item_pos(x, z, scale), scale, + mass_u_dist(rng)); - pos = glm::vec3(x_pos_u_dist(rng), 0.f, y_pos_u_dist(rng)); + x = item_x_pos_u_dist(rng), z = item_y_pos_u_dist(rng); scale = glm::vec3(scale_u_dist(rng)); item_factory->make_tetra_item( - "tetra_" + std::to_string(i), file_reader, pos, scale, mass_u_dist(rng)); + "tetra_" + std::to_string(i), file_reader, item_pos(x, z, scale), scale, + mass_u_dist(rng)); - pos = glm::vec3(x_pos_u_dist(rng), 0.f, y_pos_u_dist(rng)); + x = item_x_pos_u_dist(rng), z = item_y_pos_u_dist(rng); scale = glm::vec3(scale_u_dist(rng)); item_factory->make_cylinder_item( - "cylinder_" + std::to_string(i), file_reader, pos, scale, mass_u_dist(rng)); + "cylinder_" + std::to_string(i), file_reader, item_pos(x, z, scale), scale, + mass_u_dist(rng)); } on_reset_physics(physic_engine); diff --git a/arenai_core/tests/include/arenai_core_tests/tests_environment.h b/arenai_core/tests/include/arenai_core_tests/tests_environment.h index 58ef69c8..729cc4d5 100644 --- a/arenai_core/tests/include/arenai_core_tests/tests_environment.h +++ b/arenai_core/tests/include/arenai_core_tests/tests_environment.h @@ -19,6 +19,9 @@ class TestTanksEnvironment final : public core::BaseTanksEnvironment { int reset_physics_call_count = 0; int reset_drawables_call_count = 0; + // item name → position right after spawn, before the settle steps + std::vector> spawn_positions; + protected: void on_draw(const std::vector> &model_matrices) override { draw_call_count++; @@ -26,6 +29,10 @@ class TestTanksEnvironment final : public core::BaseTanksEnvironment { void on_reset_physics(const std::unique_ptr &engine) override { reset_physics_call_count++; + + spawn_positions.clear(); + for (const auto &item: engine->get_items()) + spawn_positions.emplace_back(item->get_name(), glm::vec3(item->get_model_matrix()[3])); } void on_reset_drawables(const std::unique_ptr &engine) override { diff --git a/arenai_core/tests/resources/golden_images/golden_env_reset_tank_0.json b/arenai_core/tests/resources/golden_images/golden_env_reset_tank_0.json index bb387f12..60107d8e 100644 --- a/arenai_core/tests/resources/golden_images/golden_env_reset_tank_0.json +++ b/arenai_core/tests/resources/golden_images/golden_env_reset_tank_0.json @@ -1 +1 @@ -[97,92,88,87,80,76,72,75,73,71,74,73,81,81,77,77,88,78,58,90,83,83,81,79,75,74,72,76,84,84,70,72,89,89,86,78,75,77,76,72,71,76,76,76,79,79,79,71,88,77,76,76,76,75,71,70,80,84,84,85,77,77,81,73,76,75,75,75,75,72,85,85,82,82,84,84,77,81,81,81,73,73,74,79,83,85,85,85,82,82,82,84,80,80,81,81,72,79,79,83,83,83,85,85,82,83,84,84,80,80,80,80,79,79,79,83,83,83,83,88,49,83,83,83,80,80,80,80,79,79,83,83,83,150,41,48,47,64,83,83,80,80,80,69,81,79,82,77,77,151,158,158,158,157,158,83,80,80,69,69,81,81,77,77,77,161,154,154,154,155,162,88,77,69,69,69,81,81,81,77,26,161,154,154,154,161,161,202,77,77,77,77,81,81,81,81,21,113,112,67,67,67,113,137,77,77,77,77,81,81,81,81,16,16,66,66,66,66,66,77,77,77,77,77,81,81,62,68,68,68,68,66,66,66,77,77,77,77,77,61,62,68,68,68,68,68,68,68,66,66,77,77,61,61,61,61,111,107,104,104,97,94,90,93,91,88,93,92,103,103,98,98,103,93,71,112,104,103,101,98,94,93,91,97,107,107,89,91,110,110,107,97,94,96,96,91,91,96,96,96,101,101,101,91,110,97,96,96,95,95,90,90,102,107,108,108,98,98,104,93,95,95,95,95,95,92,109,109,105,105,107,108,98,103,103,103,92,92,94,101,106,109,109,109,105,105,105,107,103,103,103,103,92,101,101,106,106,106,109,109,105,107,107,107,103,103,103,103,101,101,101,106,106,106,106,112,150,107,107,107,102,103,103,103,101,101,106,106,106,125,127,147,144,196,107,107,102,102,103,89,103,101,106,99,99,126,132,132,132,131,131,107,102,102,89,89,103,105,99,99,99,134,129,129,129,129,134,112,99,89,89,89,105,105,105,99,223,134,129,129,129,134,134,170,99,99,100,100,105,105,105,105,175,95,95,86,86,86,96,118,99,99,99,100,105,105,105,105,126,127,86,86,86,86,86,99,99,99,99,99,105,105,80,89,89,89,89,86,86,86,99,99,99,99,99,80,80,89,89,89,89,89,89,89,86,86,99,99,79,79,79,79,160,160,163,167,162,165,160,166,164,161,168,169,183,183,177,178,160,154,134,188,178,178,177,174,170,171,169,176,189,189,167,170,187,187,183,172,169,172,173,168,168,176,176,176,183,183,183,170,187,171,171,172,173,173,168,168,184,191,192,192,180,180,187,174,171,172,172,173,174,172,194,194,189,189,192,192,180,187,187,187,169,170,173,184,191,194,194,194,189,189,189,192,186,187,187,187,170,184,184,191,191,191,194,194,189,192,192,192,186,187,187,187,184,184,185,191,191,191,191,199,125,192,192,192,187,187,187,187,184,184,191,191,191,169,113,123,122,148,192,192,187,187,187,170,188,185,191,182,182,170,175,175,175,174,175,192,187,187,170,170,188,190,182,182,182,177,172,172,172,173,178,200,183,170,170,170,190,190,190,182,124,177,172,172,172,177,178,177,184,184,184,184,190,190,190,190,106,142,142,167,167,167,142,140,184,184,184,184,190,190,190,190,88,88,167,167,167,167,168,184,184,184,184,184,190,190,160,170,171,171,171,167,167,168,184,184,184,184,184,159,160,171,171,171,171,171,171,171,168,168,184,184,159,159,159,159] +[39,39,39,102,87,72,63,66,62,61,60,59,62,61,60,60,39,39,39,39,39,91,70,65,65,61,58,61,59,61,62,60,40,39,39,39,39,39,39,90,76,73,64,60,61,61,58,64,40,40,40,39,39,39,39,39,39,84,76,67,62,64,64,66,40,40,40,40,40,39,39,39,39,39,84,81,67,66,62,66,40,40,40,40,40,40,40,39,39,39,39,39,77,64,59,59,41,40,40,40,40,40,40,40,40,39,86,39,39,39,56,56,41,41,41,40,40,40,40,40,40,40,40,39,39,39,39,39,42,42,41,41,41,40,40,40,40,40,40,50,40,39,39,39,61,61,61,61,61,61,61,61,61,61,61,61,61,61,61,62,62,61,61,61,61,61,64,64,57,57,57,57,62,62,62,62,61,61,61,61,61,61,63,36,61,57,57,57,62,62,62,62,61,61,61,61,61,63,63,36,24,24,33,34,47,46,45,44,126,126,125,142,140,134,24,24,24,24,161,36,39,48,47,46,126,128,128,142,148,153,24,24,24,24,34,37,40,39,58,165,146,146,146,146,146,146,146,146,146,146,146,146,146,146,146,147,53,53,53,110,94,82,75,73,71,68,67,66,69,68,67,67,53,53,53,53,53,101,80,78,78,70,67,68,66,69,71,69,53,53,53,53,53,53,52,97,82,83,76,67,71,68,65,73,53,53,53,53,53,53,53,53,52,92,85,77,71,73,72,75,54,53,53,53,53,53,53,53,53,53,90,88,76,77,69,75,54,54,54,53,53,53,53,53,53,53,53,52,85,74,68,66,54,54,54,54,54,53,53,53,53,53,111,53,53,52,63,65,55,54,54,54,54,54,54,53,53,53,53,53,53,53,53,52,56,56,55,55,54,54,54,54,54,53,53,158,53,53,53,53,193,193,193,193,193,193,193,193,193,193,194,194,194,194,194,195,195,194,194,194,194,195,203,203,181,181,181,181,198,198,195,195,195,195,195,195,195,193,199,117,194,181,181,181,198,195,195,195,194,194,194,194,193,199,199,117,109,109,108,102,60,59,58,58,162,162,158,174,172,166,109,109,109,109,193,115,127,61,60,59,164,166,166,176,180,183,109,109,109,109,102,122,132,128,72,136,122,122,122,122,122,122,122,122,122,122,122,122,122,122,122,122,127,127,126,123,106,95,87,89,86,86,85,82,87,86,85,85,127,127,127,127,126,112,94,94,87,87,83,84,82,84,88,85,127,127,127,127,127,127,126,113,99,94,89,83,82,84,81,88,127,127,127,127,127,127,127,126,126,103,100,89,86,89,85,89,127,127,127,127,127,127,127,127,127,126,106,104,93,89,85,88,127,127,127,127,127,127,127,127,127,127,127,126,101,84,83,84,127,127,127,127,127,127,127,127,127,127,200,127,127,126,79,80,128,128,128,127,127,127,127,127,127,127,127,127,127,127,126,126,128,128,128,128,128,127,127,127,127,127,127,128,127,127,127,127,145,146,146,146,146,146,146,146,146,146,146,146,146,146,146,147,147,146,146,146,146,146,151,151,140,140,140,140,148,148,147,147,146,146,146,146,146,146,149,107,146,140,140,140,148,147,147,147,146,146,146,146,146,149,149,107,114,114,103,96,130,130,129,129,176,176,170,185,183,177,114,114,114,114,204,102,109,130,130,130,177,179,179,186,191,194,114,114,114,114,96,106,111,109,138,180,166,166,166,166,166,166,166,166,166,166,166,166,166,166,166,166] diff --git a/arenai_core/tests/resources/golden_images/golden_env_reset_tank_1.json b/arenai_core/tests/resources/golden_images/golden_env_reset_tank_1.json index e351a9cd..62efc607 100644 --- a/arenai_core/tests/resources/golden_images/golden_env_reset_tank_1.json +++ b/arenai_core/tests/resources/golden_images/golden_env_reset_tank_1.json @@ -1 +1 @@ -[48,47,47,47,42,42,42,42,41,41,41,41,42,42,42,42,43,47,47,47,43,42,42,42,41,41,41,43,42,42,42,42,43,47,47,47,42,42,42,41,41,41,43,42,42,42,42,42,43,47,47,47,42,42,42,41,43,43,42,42,42,42,42,42,44,47,47,47,47,42,42,44,47,47,42,42,42,42,42,42,44,47,47,47,46,42,44,44,47,47,47,47,46,42,42,42,44,47,47,47,51,44,44,44,47,47,46,46,46,46,46,42,42,47,51,51,49,48,44,44,47,46,46,46,46,46,46,46,42,47,51,50,48,48,23,28,28,37,46,46,46,46,46,46,44,46,50,50,147,147,148,147,147,147,147,46,46,46,46,46,44,46,127,90,90,146,128,147,147,175,123,135,46,46,46,46,46,46,125,83,83,83,90,126,123,201,210,121,46,46,46,46,44,125,128,84,83,83,126,128,191,156,178,107,46,46,46,46,38,43,78,84,41,45,126,152,166,167,178,41,46,46,46,46,39,43,43,41,41,45,126,119,166,42,41,41,41,46,46,46,41,38,43,41,41,41,45,103,166,42,41,41,41,41,46,46,11,10,9,9,8,8,7,7,7,6,6,6,6,5,5,5,10,10,9,9,8,8,7,7,7,6,6,6,5,5,5,5,10,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,10,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,4,10,9,9,8,8,7,7,6,6,5,5,5,5,5,4,4,10,9,9,8,8,7,6,6,5,5,5,5,5,5,4,4,10,9,8,8,7,7,6,6,5,5,5,5,5,4,4,4,10,9,8,7,7,6,26,31,31,41,5,5,4,4,4,4,9,8,7,7,182,182,184,183,183,183,183,4,4,4,4,4,9,8,158,115,115,181,159,183,183,216,154,30,4,4,4,4,8,7,157,107,107,107,115,157,154,50,52,27,4,4,4,4,8,157,160,107,107,107,157,160,140,40,45,25,4,4,4,4,8,7,127,136,6,5,157,113,123,123,45,4,4,4,4,4,9,7,7,6,6,5,157,89,122,4,4,4,4,4,4,4,9,7,7,6,6,5,5,78,123,4,4,4,4,4,4,4,93,93,93,92,87,88,87,87,87,87,87,86,88,88,88,88,88,93,93,92,88,87,87,87,87,87,86,88,88,88,88,88,88,93,92,93,88,87,87,87,87,87,88,88,88,88,88,88,88,93,93,93,88,87,87,87,89,88,88,88,88,88,88,88,89,93,93,93,93,87,87,90,94,94,88,88,88,88,88,88,89,93,93,93,93,87,90,90,94,93,93,93,93,88,88,88,89,93,93,93,98,90,90,90,93,93,93,93,93,93,93,88,87,93,98,98,95,95,90,90,93,93,93,93,93,93,93,93,87,93,98,97,95,95,31,35,35,41,93,93,93,93,93,93,90,92,97,97,7,7,7,7,7,7,7,93,93,93,93,93,89,92,7,6,6,6,6,6,6,7,7,88,93,93,93,93,92,92,7,6,6,6,6,6,6,23,24,83,93,93,93,93,89,7,7,6,6,6,6,6,175,20,22,77,93,93,93,93,83,89,196,205,87,92,6,152,161,161,22,87,93,93,93,93,83,89,89,87,87,92,6,133,161,88,87,87,87,93,93,93,86,83,89,87,87,87,92,123,161,88,87,87,87,87,93,93] +[76,76,72,72,72,70,72,71,72,72,76,77,77,77,79,78,74,74,69,69,69,71,71,73,73,75,75,76,77,78,78,78,72,70,70,68,67,69,69,75,73,73,75,75,73,76,77,78,68,68,68,66,68,69,76,76,75,70,72,72,72,72,72,71,68,68,68,68,68,72,72,72,74,71,69,69,73,72,72,72,69,67,67,67,67,65,65,63,72,73,72,72,73,73,73,72,68,69,67,67,58,57,57,74,73,73,73,74,74,75,73,71,68,68,68,57,57,57,57,73,73,74,74,73,74,75,75,75,68,68,68,53,53,52,73,73,73,73,74,74,70,71,74,74,68,68,53,53,52,52,69,72,72,72,73,74,70,70,70,70,70,53,52,52,52,69,69,69,72,72,72,72,63,67,67,67,70,52,52,52,52,69,23,30,30,36,72,72,63,63,63,63,52,52,158,140,140,140,141,125,125,125,72,72,63,63,63,63,52,52,124,118,139,166,134,123,135,125,125,69,63,63,63,63,52,51,100,119,126,126,164,202,124,125,115,140,148,52,52,55,51,100,142,189,187,125,179,174,211,82,116,116,52,52,52,52,25,25,24,24,23,22,21,21,20,20,20,20,20,21,21,22,22,22,21,20,19,19,19,18,18,18,18,18,19,19,19,20,20,20,19,19,18,17,16,16,17,16,16,17,17,17,18,18,17,17,16,15,15,14,14,15,15,15,15,15,15,15,15,15,14,14,14,13,13,12,12,13,13,13,13,13,13,13,13,13,13,12,12,12,11,11,11,10,10,11,11,11,12,12,12,12,11,11,11,10,10,9,9,8,9,9,10,10,10,11,11,11,10,10,10,9,8,8,7,7,8,8,8,9,9,10,10,10,9,9,9,8,7,7,7,7,7,7,7,8,8,8,9,9,9,8,8,7,7,6,6,6,6,7,7,7,7,7,7,8,8,8,7,7,6,6,6,6,6,6,6,6,6,6,7,7,8,7,7,6,6,6,26,33,33,41,6,6,6,6,6,6,7,6,196,174,175,174,176,156,156,156,5,6,6,6,6,6,6,6,156,148,173,205,168,154,168,156,156,5,5,5,5,5,6,5,126,149,158,158,42,51,156,156,26,30,184,5,5,5,5,126,106,139,137,156,45,44,53,19,26,26,4,4,4,4,123,124,120,120,120,117,120,119,120,120,125,125,126,126,128,127,122,123,116,116,116,119,119,121,122,124,125,125,127,127,128,128,120,118,118,115,115,118,118,125,122,122,124,124,122,125,127,127,116,116,116,115,117,117,126,126,125,119,121,121,121,122,122,120,116,117,117,117,117,122,122,122,124,121,118,118,123,121,122,122,118,116,116,116,116,114,114,112,123,123,122,122,123,123,123,121,118,118,116,116,105,105,105,125,124,124,124,124,125,126,123,121,118,118,118,105,105,105,105,124,124,125,125,124,125,126,125,126,118,118,118,100,100,100,124,124,124,124,125,125,120,122,126,125,118,118,100,100,100,100,120,123,123,124,124,125,120,120,121,121,120,100,100,100,100,120,120,120,123,123,124,124,113,117,117,117,120,100,100,99,99,120,31,36,36,41,124,124,113,113,113,113,100,100,7,7,7,7,7,7,7,7,123,124,113,113,113,113,99,99,7,6,7,7,7,6,7,7,7,120,112,112,113,113,99,99,6,6,6,6,21,23,7,7,80,90,7,100,100,103,99,6,147,175,173,7,21,21,24,66,80,80,99,99,99,99] diff --git a/arenai_core/tests/src/tests_environment.cpp b/arenai_core/tests/src/tests_environment.cpp index 21533823..a2148c09 100644 --- a/arenai_core/tests/src/tests_environment.cpp +++ b/arenai_core/tests/src/tests_environment.cpp @@ -11,6 +11,7 @@ #include #include +#include using namespace arenai; using namespace arenai::core; @@ -26,7 +27,7 @@ TEST_F(EnvironmentTest, ResetReturnsCorrectNumberOfStates) { constexpr int vision_w = 16; TestTanksEnvironment env( - file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false); + file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false, false); const auto states = env.reset(100.f, 100.f); @@ -46,7 +47,7 @@ TEST_F(EnvironmentTest, ResetInitialVisionIsNotBlack) { constexpr int vision_w = 16; TestTanksEnvironment env( - file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false); + file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false, false); for (const auto states = env.reset(100.f, 100.f); const auto &[vision, proprioception]: states) { @@ -83,7 +84,7 @@ TEST_F(EnvironmentTest, ResetGoldenImage) { constexpr int vision_w = 16; TestTanksEnvironment env( - file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false); + file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false, false); env.seed(42); @@ -147,7 +148,8 @@ TEST_F(EnvironmentTest, ResetProprioceptionNonEmpty) { constexpr int nb_tanks = 2; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 8, 8, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 8, 8, 1, false, false); for (const auto states = env.reset(100.f, 100.f); const auto &[vision, proprioception]: states) { @@ -165,7 +167,8 @@ TEST_F(EnvironmentTest, ResetCallsOnResetPhysics) { constexpr int nb_tanks = 1; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 8, 8, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 8, 8, 1, false, false); env.reset(100.f, 100.f); @@ -185,7 +188,7 @@ TEST_F(EnvironmentTest, StepReturnsCorrectNumberOfTuples) { constexpr int vision_w = 16; TestTanksEnvironment env( - file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false); + file_reader, graphics_backend, nb_tanks, frequency, vision_h, vision_w, 1, false, false); env.reset(100.f, 100.f); @@ -209,7 +212,8 @@ TEST_F(EnvironmentTest, StepRewardAndDoneAreValid) { constexpr int nb_tanks = 2; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); env.reset(100.f, 100.f); @@ -219,7 +223,7 @@ TEST_F(EnvironmentTest, StepRewardAndDoneAreValid) { .fire_button = {false}}); for (const auto results = env.step(frequency, actions); - const auto &[state, reward, is_done]: results) { + const auto &[state, reward, is_done, is_truncated]: results) { ASSERT_FALSE(std::isnan(reward)) << "reward should not be NaN"; ASSERT_FALSE(std::isinf(reward)) << "reward should not be Inf"; } @@ -235,7 +239,8 @@ TEST_F(EnvironmentTest, StepCallsOnDraw) { constexpr int nb_tanks = 1; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); env.reset(100.f, 100.f); @@ -263,7 +268,8 @@ TEST_F(EnvironmentTest, MultipleStepsDoNotCrash) { constexpr int nb_tanks = 2; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); env.reset(100.f, 100.f); @@ -288,7 +294,8 @@ TEST_F(EnvironmentTest, ResetCallsOnResetDrawables) { constexpr int nb_tanks = 1; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); env.reset(100.f, 100.f); @@ -305,7 +312,8 @@ TEST_F(EnvironmentTest, FullLifecycle) { constexpr int nb_tanks = 2; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); // First episode const auto initial_states = env.reset(100.f, 100.f); @@ -334,10 +342,52 @@ TEST_F(EnvironmentTest, StopDrawingDoubleCallDoesNotCrash) { constexpr int nb_tanks = 1; constexpr float frequency = 1.f / 60.f; - TestTanksEnvironment env(file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false); + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); env.reset(100.f, 100.f); env.stop_drawing(); ASSERT_NO_THROW(env.stop_drawing()); } + +// ======================================================================== +// reset — every body spawns above the terrain surface +// ======================================================================== + +TEST_F(EnvironmentTest, SpawnedBodiesStartAboveTerrain) { + constexpr int nb_tanks = 4; + constexpr int nb_shapes_per_kind = 30; + constexpr float frequency = 1.f / 60.f; + + TestTanksEnvironment env( + file_reader, graphics_backend, nb_tanks, frequency, 16, 16, 1, false, false); + + // wide spawn zone so tanks land on varied terrain heights + env.reset(2000.f, 2000.f); + env.stop_drawing(); + + // terrain-only world (the fixture's engine) to measure the ground height at each (x, z) + engine->get_item_factory()->make_height_map_item( + "height_map", file_reader, "heightmap/heightmap6.png", glm::vec3(0., 40., 0.), + glm::vec3(10., 200., 10.)); + + ASSERT_FALSE(env.spawn_positions.empty()); + + int checked = 0; + for (const auto &[name, pos]: env.spawn_positions) { + // only bodies created at the spawn (x, z) itself: props and tank chassis; + // wheels/turret/canon are laterally offset so the local ground height does not apply + const bool is_prop = name.starts_with("sphere_") || name.starts_with("cube_") + || name.starts_with("tetra_") || name.starts_with("cylinder_"); + if (!is_prop && !name.ends_with("_chassis")) continue; + + const auto ground = + engine->ray_cast(glm::vec3(pos.x, 300.f, pos.z), glm::vec3(pos.x, -300.f, pos.z)); + ASSERT_TRUE(ground.has_value()) << name; + EXPECT_GT(pos.y, ground->y) << name; + checked++; + } + + EXPECT_EQ(checked, nb_tanks + 4 * nb_shapes_per_kind); +} diff --git a/arenai_desktop/src/controller/bindings.h b/arenai_desktop/src/controller/bindings.h index f5e8e69a..a484c78c 100644 --- a/arenai_desktop/src/controller/bindings.h +++ b/arenai_desktop/src/controller/bindings.h @@ -25,9 +25,9 @@ namespace arenai::desktop { std::optional turn_left = controller::Key::A; std::optional turn_right = controller::Key::D; std::optional fire = controller::MouseButton::Left; + std::optional zoom = controller::MouseButton::Right; }; - // the six analog channels a pad exposes through the window's callbacks enum class GamepadAxis { LeftStickX, LeftStickY, @@ -39,9 +39,6 @@ namespace arenai::desktop { inline constexpr int NB_GAMEPAD_AXES = 6; - // One analog slot. `sign` keeps the direction captured for the one-way - // actions (accelerate / reverse read max(0, sign * value)); the two-way - // actions (steer, aim) ignore it and stay at +1. struct GamepadAxisBinding { GamepadAxis axis; float sign = 1.f; @@ -51,6 +48,7 @@ namespace arenai::desktop { struct GamepadBindings { std::optional fire = controller::GamepadButton::RB; + std::optional zoom = controller::GamepadButton::LB; std::optional steer = GamepadAxisBinding{.axis = GamepadAxis::LeftStickX}; std::optional aim_x = @@ -61,9 +59,7 @@ namespace arenai::desktop { GamepadAxisBinding{.axis = GamepadAxis::RightTrigger}; std::optional reverse = GamepadAxisBinding{.axis = GamepadAxis::LeftTrigger}; - // preferred pad across reconnections; empty = first connected std::string device_guid; - // display only, shown while the preferred pad is unplugged std::string device_name; }; @@ -72,7 +68,6 @@ namespace arenai::desktop { GamepadBindings gamepad; }; - // canonical names for persistence ("" / nullopt on unknown) const char *to_string(GamepadAxis axis); std::optional gamepad_axis_from_string(std::string_view name); diff --git a/arenai_desktop/src/controller/gamepad.cpp b/arenai_desktop/src/controller/gamepad.cpp index ca03e17c..47936e67 100644 --- a/arenai_desktop/src/controller/gamepad.cpp +++ b/arenai_desktop/src/controller/gamepad.cpp @@ -92,9 +92,14 @@ namespace arenai::desktop { float canon_rotation = 0.f; if (event.button.has_value()) { - if (const auto &[button, action] = event.button.value(); - action == controller::InputAction::Press && button == bindings.fire) + const auto &[button, action] = event.button.value(); + if (action == controller::InputAction::Press && button == bindings.fire) need_fire = true; + + if (button == bindings.zoom) { + if (action == controller::InputAction::Press) zoom_held = true; + else if (action == controller::InputAction::Release) zoom_held = false; + } } else { // per-frame tick: controllers consume rad/frame deltas, so the stick // deflection is scaled into radians here (like the mouse handler) @@ -113,7 +118,8 @@ namespace arenai::desktop { true, {.left_joystick = {.x = direction, .y = speed}, .right_joystick = {.x = turret_rotation, .y = canon_rotation}, - .fire_button = {need_fire}}}; + .fire_button = {need_fire}, + .zoom_button = {zoom_held}}}; } }// namespace arenai::desktop diff --git a/arenai_desktop/src/controller/gamepad.h b/arenai_desktop/src/controller/gamepad.h index 3de95941..f2bbf68f 100644 --- a/arenai_desktop/src/controller/gamepad.h +++ b/arenai_desktop/src/controller/gamepad.h @@ -42,13 +42,14 @@ namespace arenai::desktop { PlayerGamepadInput state; + // held state: true between the press and the release of the zoom button + bool zoom_held = false; + static float apply_dead_zone(double value); - // deflection of a two-way slot (steer, aim), 0 when unbound float axis_value(const std::optional &slot, const PlayerGamepadInput &event); - // deflection of a one-way slot (accelerate, reverse): the captured - // direction reads positive, the other way is ignored + float one_way_axis_value( const std::optional &slot, const PlayerGamepadInput &event); }; diff --git a/arenai_desktop/src/controller/mouse_keyboard.cpp b/arenai_desktop/src/controller/mouse_keyboard.cpp index 72623a1c..1005ba54 100644 --- a/arenai_desktop/src/controller/mouse_keyboard.cpp +++ b/arenai_desktop/src/controller/mouse_keyboard.cpp @@ -13,7 +13,7 @@ namespace arenai::desktop { const KeyboardBindings &bindings) : window(std::move(window)), renderer(renderer), bindings(bindings), last_mouse_x(0.), last_mouse_y(0.), current_dir(0.f), current_speed(0.f), current_turret_rotation(0.f), - current_canon_rotation(0.f), cursor_captured(true) { + current_canon_rotation(0.f), current_zoom(false), cursor_captured(true) { const auto center_x = static_cast(renderer.get_width()) / 2., center_y = static_cast(renderer.get_height()) / 2.; @@ -57,9 +57,11 @@ namespace arenai::desktop { else if (input == bindings.turn_left) current_dir = -1.f; else if (input == bindings.turn_right) current_dir = 1.f; else if (input == bindings.fire) need_fire = true; + else if (input == bindings.zoom) current_zoom = true; } else if (action == controller::InputAction::Release) { if (input == bindings.forward || input == bindings.backward) current_speed = 0.f; if (input == bindings.turn_left || input == bindings.turn_right) current_dir = 0.f; + if (input == bindings.zoom) current_zoom = false; } } @@ -118,7 +120,8 @@ namespace arenai::desktop { true, {.left_joystick = {.x = current_dir, .y = current_speed}, .right_joystick = {.x = current_turret_rotation, .y = current_canon_rotation}, - .fire_button = {need_fire}}}; + .fire_button = {need_fire}, + .zoom_button = {current_zoom}}}; } }// namespace arenai::desktop diff --git a/arenai_desktop/src/controller/mouse_keyboard.h b/arenai_desktop/src/controller/mouse_keyboard.h index 43275f8f..0a7ae056 100644 --- a/arenai_desktop/src/controller/mouse_keyboard.h +++ b/arenai_desktop/src/controller/mouse_keyboard.h @@ -61,6 +61,9 @@ namespace arenai::desktop { float current_turret_rotation; float current_canon_rotation; + // held state: true between the press and the release of the zoom slot + bool current_zoom; + bool cursor_captured; }; diff --git a/arenai_desktop/src/core/agent_loading_checker.cpp b/arenai_desktop/src/core/agent_loading_checker.cpp index ca069915..0e244d81 100644 --- a/arenai_desktop/src/core/agent_loading_checker.cpp +++ b/arenai_desktop/src/core/agent_loading_checker.cpp @@ -4,32 +4,64 @@ #include "./agent_loading_checker.h" -#include +#include + +#include + #include -#include namespace arenai::desktop { - std::optional - check_agent_folder(const ModelOptions &model_options, const std::filesystem::path &folder) { + agent::AgentAlgorithm to_agent_algorithm(const gui::AiAlgorithm algorithm) { + switch (algorithm) { + case gui::AiAlgorithm::Ppo: return agent::PPO; + default: return agent::PPO_LIQUID; + } + } + + std::optional + resolve_agent_config(const gui::AgentSelection &selection) { + if (!selection.config.empty()) return selection.config; + + for (const auto &candidate: + {selection.folder / "config.json", selection.folder.parent_path() / "config.json"}) + if (std::filesystem::is_regular_file(candidate)) return candidate; + + return std::nullopt; + } + + std::optional check_agent(const gui::AgentSelection &selection) { + const auto config_path = resolve_agent_config(selection); + if (!config_path) + return "config.json not found in " + selection.folder.string() + " or its parent"; + try { + std::ifstream config_stream(*config_path); + if (!config_stream.is_open()) + throw std::runtime_error("cannot open " + config_path->string()); + const auto config = nlohmann::json::parse(config_stream); - const auto agent = - agent::ActorAgentFactory(model_options.hyper_parameters) - .get_agent( - model_options.vision_height, model_options.vision_width, - model::ENEMY_PROPRIOCEPTION_SIZE, model::ENEMY_NB_CONTINUOUS_ACTION, - model::ENEMY_NB_DISCRETE_ACTION); + agent::AgentFactory factory(config); // stays on CPU: the point is only to prove the state dicts load - agent->load(folder); + // and the networks answer a forward pass + const auto agent = factory.get_agent( + to_agent_algorithm(selection.algorithm), model::ENEMY_PROPRIOCEPTION_SIZE, + model::ENEMY_NB_CONTINUOUS_ACTION, model::ENEMY_NB_DISCRETE_ACTION, false); + + agent->load(selection.folder); + + const core::State blank_state = { + .vision = + {.pixels = std::vector( + static_cast(3) * factory.get_vision_height() + * factory.get_vision_width(), + 0)}, + .proprioception = std::vector(model::ENEMY_PROPRIOCEPTION_SIZE, 0.f)}; + agent->act({blank_state}, factory.get_vision_height(), factory.get_vision_width()); return std::nullopt; - } catch (utils::FileDoesNotExistException &e) { - return "Missing file: " + e.missing_file().filename().string(); - } catch (utils::ModelLoadException &e) { - return "Error while loading file: " + e.wrong_state_dict_file().filename().string(); - } catch (const std::exception &_) { return "Unknow error while loading Model"; } + } catch (const std::exception &e) { return e.what(); } } }// namespace arenai::desktop diff --git a/arenai_desktop/src/core/agent_loading_checker.h b/arenai_desktop/src/core/agent_loading_checker.h index 3cf6560f..752440b7 100644 --- a/arenai_desktop/src/core/agent_loading_checker.h +++ b/arenai_desktop/src/core/agent_loading_checker.h @@ -9,15 +9,24 @@ #include #include -#include "../game.h" +#include + +#include "../gui/menu.h" namespace arenai::desktop { - // Dry-run load of the SAC agent: builds the networks and loads every - // state dict on CPU, then throws the agent away. Returns std::nullopt - // when the folder is a valid checkpoint, else the message to display. - std::optional - check_agent_folder(const ModelOptions &model_options, const std::filesystem::path &folder); + agent::AgentAlgorithm to_agent_algorithm(gui::AiAlgorithm algorithm); + + // the config.json a selection uses: the explicitly picked file when set, + // else the one sitting in the state-dict folder, else in its parent (the + // training layout: train_NNN/config.json next to train_NNN/save_K/) + std::optional resolve_agent_config(const gui::AgentSelection &selection); + + // Dry-run load of the selected agent: builds the networks from the + // config.json, loads every state dict on CPU and answers one forward pass + // on a blank state, then throws the agent away. Returns std::nullopt when + // the selection is a working model, else the exception message to display. + std::optional check_agent(const gui::AgentSelection &selection); }// namespace arenai::desktop diff --git a/arenai_desktop/src/core/game_environment.cpp b/arenai_desktop/src/core/game_environment.cpp index 57c09e13..bac3f2dd 100644 --- a/arenai_desktop/src/core/game_environment.cpp +++ b/arenai_desktop/src/core/game_environment.cpp @@ -9,6 +9,7 @@ #include #include +#include #include #include #include @@ -25,7 +26,8 @@ namespace arenai::desktop { : BaseTanksEnvironment( std::make_shared(asset_folder_path), view::make_vulkan_backend(settings.vision_gpu), settings.nb_tanks, wanted_frequency, - vision_height, vision_width, 8, true), + // no starving timeout in the playable game + vision_height, vision_width, 8, true, false), windowed_backend(graphics_backend), asset_file_reader(std::make_shared(asset_folder_path)), player_tank(std::nullptr_t()), player_renderer(std::nullptr_t()), @@ -73,7 +75,7 @@ namespace arenai::desktop { const glm::vec3 origin = canon_matrix * glm::vec4(0.f, 0.f, 0.f, 1.f); const glm::vec3 forward = glm::normalize(glm::mat3(canon_matrix) * glm::vec3(0.f, 0.f, 1.f)); - const glm::vec3 aim = origin + forward * AIM_DISTANCE; + const glm::vec3 aim = origin + forward * model::CANON_AIM_DISTANCE; const glm::vec4 clip = player_renderer->last_view_projection() * glm::vec4(aim, 1.f); if (clip.w <= 0.f) return std::nullopt; diff --git a/arenai_desktop/src/core/game_environment.h b/arenai_desktop/src/core/game_environment.h index d052c8da..092c8e81 100644 --- a/arenai_desktop/src/core/game_environment.h +++ b/arenai_desktop/src/core/game_environment.h @@ -43,8 +43,6 @@ namespace arenai::desktop { model::PlayerHits consume_player_hits() const; std::vector consume_damage_screen_angles() const; - static constexpr float AIM_DISTANCE = 100.f; - std::optional aim_point_on_screen() const; protected: diff --git a/arenai_desktop/src/core/user_preferences.cpp b/arenai_desktop/src/core/user_preferences.cpp index 6d0e0a3e..9f44781b 100644 --- a/arenai_desktop/src/core/user_preferences.cpp +++ b/arenai_desktop/src/core/user_preferences.cpp @@ -77,6 +77,16 @@ namespace arenai::desktop { slot = GamepadAxisBinding{.axis = *axis, .sign = sign}; } + void load_gamepad_button_binding( + const nlohmann::json &json, const char *field, + std::optional &slot) { + if (!json.contains(field)) return; + const auto name = json.value(field, std::string()); + if (name.empty()) slot = std::nullopt; + else if (const auto button = controller::gamepad_button_from_string(name)) + slot = *button; + } + void load_bindings(const nlohmann::json &json, ControlBindings &bindings) { if (const auto keyboard = json.value("keyboard", nlohmann::json::object()); keyboard.is_object()) { @@ -85,16 +95,13 @@ namespace arenai::desktop { load_keyboard_binding(keyboard, "turn_left", bindings.keyboard.turn_left); load_keyboard_binding(keyboard, "turn_right", bindings.keyboard.turn_right); load_keyboard_binding(keyboard, "fire", bindings.keyboard.fire); + load_keyboard_binding(keyboard, "zoom", bindings.keyboard.zoom); } if (const auto gamepad = json.value("gamepad", nlohmann::json::object()); gamepad.is_object()) { - if (gamepad.contains("fire")) { - const auto name = gamepad.value("fire", std::string()); - if (name.empty()) bindings.gamepad.fire = std::nullopt; - else if (const auto button = controller::gamepad_button_from_string(name)) - bindings.gamepad.fire = *button; - } + load_gamepad_button_binding(gamepad, "fire", bindings.gamepad.fire); + load_gamepad_button_binding(gamepad, "zoom", bindings.gamepad.zoom); load_axis_binding(gamepad, "steer", bindings.gamepad.steer); load_axis_binding(gamepad, "aim_x", bindings.gamepad.aim_x); load_axis_binding(gamepad, "aim_y", bindings.gamepad.aim_y); @@ -114,10 +121,13 @@ namespace arenai::desktop { {"backward", keyboard_binding_to_string(bindings.keyboard.backward)}, {"turn_left", keyboard_binding_to_string(bindings.keyboard.turn_left)}, {"turn_right", keyboard_binding_to_string(bindings.keyboard.turn_right)}, - {"fire", keyboard_binding_to_string(bindings.keyboard.fire)}}}, + {"fire", keyboard_binding_to_string(bindings.keyboard.fire)}, + {"zoom", keyboard_binding_to_string(bindings.keyboard.zoom)}}}, {"gamepad", {{"fire", bindings.gamepad.fire ? controller::to_string(*bindings.gamepad.fire) : ""}, + {"zoom", + bindings.gamepad.zoom ? controller::to_string(*bindings.gamepad.zoom) : ""}, {"steer", axis_binding_to_string(bindings.gamepad.steer)}, {"aim_x", axis_binding_to_string(bindings.gamepad.aim_x)}, {"aim_y", axis_binding_to_string(bindings.gamepad.aim_y)}, @@ -182,11 +192,19 @@ namespace arenai::desktop { bindings.is_object()) load_bindings(bindings, settings.bindings); - // a stale folder (moved, deleted, unplugged drive) falls back to + // a stale path (moved, deleted, unplugged drive) falls back to // the default so the menu never starts on an unplayable selection - if (const std::filesystem::path sac_folder = json.value("sac_folder", std::string()); - !sac_folder.empty() && std::filesystem::is_directory(sac_folder)) - settings.sac_folder = sac_folder; + if (const std::filesystem::path agent_folder = + json.value("agent_folder", std::string()); + !agent_folder.empty() && std::filesystem::is_directory(agent_folder)) + settings.agent_folder = agent_folder; + if (const std::filesystem::path agent_config = + json.value("agent_config", std::string()); + !agent_config.empty() && std::filesystem::is_regular_file(agent_config)) + settings.agent_config = agent_config; + if (const auto algorithm = + gui::ai_algorithm_from_string(json.value("algorithm", std::string()))) + settings.agent_algorithm = *algorithm; } catch (const std::exception &e) { std::cerr << "Cannot load preferences " << path << ": " << e.what() << std::endl; return defaults; @@ -210,7 +228,9 @@ namespace arenai::desktop { {"window_gpu", settings.window_gpu}, {"vision_gpu", settings.vision_gpu}, {"bindings", bindings_to_json(settings.bindings)}, - {"sac_folder", settings.sac_folder.string()}, + {"agent_folder", settings.agent_folder.string()}, + {"agent_config", settings.agent_config.string()}, + {"algorithm", gui::to_string(settings.agent_algorithm)}, }; std::filesystem::create_directories(path.parent_path()); diff --git a/arenai_desktop/src/game.cpp b/arenai_desktop/src/game.cpp index 5960e7b3..d3c6caa8 100644 --- a/arenai_desktop/src/game.cpp +++ b/arenai_desktop/src/game.cpp @@ -5,9 +5,12 @@ #include "./game.h" #include +#include #include -#include +#include + +#include #include #include #include @@ -29,14 +32,30 @@ namespace arenai::desktop { const std::unique_ptr &gui) { const auto window = graphics_backend->get_window(); - const std::shared_ptr sac_agent = - agent::ActorAgentFactory(model_options.hyper_parameters) - .get_agent( - model_options.vision_height, model_options.vision_width, - model::ENEMY_PROPRIOCEPTION_SIZE, model::ENEMY_NB_CONTINUOUS_ACTION, - model::ENEMY_NB_DISCRETE_ACTION); + // the menu only lets Play through once check_agent() validated the + // selection, so the config.json is there and the load succeeds + const gui::AgentSelection selection = { + .config = settings.agent_config, + .folder = settings.agent_folder, + .algorithm = settings.agent_algorithm}; + const auto config_path = resolve_agent_config(selection); + if (!config_path) + throw std::runtime_error( + "config.json not found in " + selection.folder.string() + " or its parent"); + + std::ifstream config_stream(*config_path); + const auto config = nlohmann::json::parse(config_stream); + + agent::AgentFactory factory(config); + const int vision_height = factory.get_vision_height(); + const int vision_width = factory.get_vision_width(); + const float wanted_frequency = factory.get_wanted_frequency(); + + const std::shared_ptr enemy_agent = factory.get_agent( + to_agent_algorithm(settings.agent_algorithm), model::ENEMY_PROPRIOCEPTION_SIZE, + model::ENEMY_NB_CONTINUOUS_ACTION, model::ENEMY_NB_DISCRETE_ACTION, model_options.cuda); - sac_agent->load(settings.sac_folder); + enemy_agent->load(settings.agent_folder); // route the pad input to the configured device when it is connected // (also covers runs that skip the menu, e.g. ARENAI_DEBUG_AUTOPLAY) @@ -49,8 +68,8 @@ namespace arenai::desktop { } const auto env = std::make_shared( - game_options.resources_folder, graphics_backend, settings, model_options.vision_height, - model_options.vision_width, game_options.wanted_frequency); + game_options.resources_folder, graphics_backend, settings, vision_height, vision_width, + wanted_frequency); auto states = env->reset( static_cast(settings.spawn_side), static_cast(settings.spawn_side)); @@ -98,7 +117,7 @@ namespace arenai::desktop { auto outcome = InGameOutcome::ExitGame; const auto frame_dt = - std::chrono::milliseconds(static_cast(game_options.wanted_frequency * 1000.f)); + std::chrono::milliseconds(static_cast(wanted_frequency * 1000.f)); while (!window->should_close()) { window->poll_events(); @@ -129,10 +148,9 @@ namespace arenai::desktop { auto last_time = std::chrono::steady_clock::now(); - const auto action = - sac_agent->act(states, model_options.vision_height, model_options.vision_width); + const auto action = enemy_agent->act(states, vision_height, vision_width); - const auto steps = env->step(game_options.wanted_frequency, action); + const auto steps = env->step(wanted_frequency, action); if (const auto [hits, kills] = env->consume_player_hits(); kills > 0) gui->notify_hit(gui::HitKind::Kill); @@ -149,7 +167,7 @@ namespace arenai::desktop { states.clear(); - for (const auto &[state, reward, done]: steps) states.push_back(state); + for (const auto &[state, reward, done, truncated]: steps) states.push_back(state); auto now = std::chrono::steady_clock::now(); auto dt = now - last_time; @@ -171,8 +189,9 @@ namespace arenai::desktop { void run_gui(const GameOptions &game_options, const ModelOptions &model_options) { // loaded before the backend: the window GPU choice only applies at // device creation, i.e. here - const auto initial_settings = - load_preferences({.sac_folder = model_options.state_dict_folder}); + const auto initial_settings = load_preferences( + {.agent_folder = model_options.state_dict_folder, + .agent_config = model_options.config_json}); const std::shared_ptr graphics_backend = view::make_glfw_vulkan_backend( game_options.window_width, game_options.window_height, "ArenAI", @@ -187,9 +206,7 @@ namespace arenai::desktop { const auto gui = gui::make_gui( graphics_backend, asset_reader, initial_settings, view::list_vulkan_gpus(), game_options.window_width, game_options.window_height, - [&model_options](const std::filesystem::path &folder) { - return check_agent_folder(model_options, folder); - }); + [](const gui::AgentSelection &selection) { return check_agent(selection); }); window->set_resize_callback( [&gui](const int width, const int height) { gui->on_window_resized(width, height); }); diff --git a/arenai_desktop/src/game.h b/arenai_desktop/src/game.h index 96eacc5f..b035734e 100644 --- a/arenai_desktop/src/game.h +++ b/arenai_desktop/src/game.h @@ -6,23 +6,20 @@ #define ARENAI_DESKTOP_GAME_H #include -#include -#include #include "./gui/menu.h" namespace arenai::desktop { + // everything else about the model (vision size, control frequency, + // hyper-parameters) comes from the selected training run's config.json struct ModelOptions { - int vision_height; - int vision_width; - std::map hyper_parameters; std::filesystem::path state_dict_folder; + std::filesystem::path config_json; bool cuda; }; struct GameOptions { - float wanted_frequency; int window_width; int window_height; std::filesystem::path resources_folder; diff --git a/arenai_desktop/src/gui/menu.h b/arenai_desktop/src/gui/menu.h index 7953c6cc..6af21a5d 100644 --- a/arenai_desktop/src/gui/menu.h +++ b/arenai_desktop/src/gui/menu.h @@ -24,7 +24,7 @@ // The gui/ folder is a hexagon of its own: this header is its only public // port, and it exposes no RmlUi type — the library stays an implementation -// detail of rml_menu.cpp and the rml/ subfolder, exactly like GL stays +// detail of the rml/ subfolder, exactly like GL stays // inside arenai_view. namespace arenai::desktop::gui { @@ -61,6 +61,24 @@ namespace arenai::desktop::gui { } } + // the RL algorithms the enemy agent can be built from; mirrors + // agent::AgentAlgorithm without leaking arenai_agent into the gui port + enum class AiAlgorithm { Ppo, PpoLiquid }; + + // canonical names, shared by the JSON preferences and the menu bindings + constexpr const char *to_string(const AiAlgorithm algorithm) { + switch (algorithm) { + case AiAlgorithm::Ppo: return "ppo"; + default: return "ppo_liquid"; + } + } + + constexpr std::optional ai_algorithm_from_string(std::string_view name) { + for (const auto algorithm: {AiAlgorithm::Ppo, AiAlgorithm::PpoLiquid}) + if (name == to_string(algorithm)) return algorithm; + return std::nullopt; + } + // what the player can tune in the menu before launching a game struct GameSettings { int nb_tanks = 16; @@ -77,7 +95,12 @@ namespace arenai::desktop::gui { std::string window_gpu; std::string vision_gpu; - std::filesystem::path sac_folder; + // the trained enemy agent: its state-dict folder, the config.json of + // its training run (empty = looked up next to the state dicts) and + // the algorithm to rebuild the networks with + std::filesystem::path agent_folder; + std::filesystem::path agent_config; + AiAlgorithm agent_algorithm = AiAlgorithm::Ppo; }; enum class MenuOutcome { Play, Quit }; @@ -115,14 +138,21 @@ namespace arenai::desktop::gui { virtual void on_window_resized(int width, int height) = 0; }; - using SacFolderValidator = - std::function(const std::filesystem::path &)>; + // what the AI page hands to the dry-run loader: nullopt = the model + // loaded and answered a forward pass, else the message to display + struct AgentSelection { + std::filesystem::path config;// empty = auto-resolved config.json + std::filesystem::path folder; + AiAlgorithm algorithm; + }; + + using AgentValidator = std::function(const AgentSelection &)>; std::unique_ptr make_gui( const std::shared_ptr &backend, const std::shared_ptr &asset_reader, const GameSettings &initial_settings, const std::vector &gpus, - int window_width, int window_height, SacFolderValidator sac_validator); + int window_width, int window_height, AgentValidator agent_validator); }// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/corner_reticle.cpp b/arenai_desktop/src/gui/rml/corner_reticle.cpp new file mode 100644 index 00000000..6c86e063 --- /dev/null +++ b/arenai_desktop/src/gui/rml/corner_reticle.cpp @@ -0,0 +1,95 @@ +// +// Created by samuel on 02/08/2026. +// + +#include "./corner_reticle.h" + +namespace arenai::desktop::gui { + + CornerReticleDecorator::CornerReticleDecorator( + const Rml::Colourb tick_color, const Rml::Colourb fill_color, + const Rml::NumericValue tick_length, const Rml::NumericValue thickness, + const Rml::NumericValue inset, const Rml::NumericValue fill_inset) + : tick_color_(tick_color), fill_color_(fill_color), tick_length_(tick_length), + thickness_(thickness), inset_(inset), fill_inset_(fill_inset) {} + + Rml::DecoratorDataHandle + CornerReticleDecorator::GenerateElementData(Rml::Element *element, Rml::BoxArea) const { + const float opacity = element->GetComputedValues().opacity(); + const Rml::Vector2f size = element->GetBox().GetSize(Rml::BoxArea::Border); + const float length = element->ResolveLength(tick_length_); + const float thickness = element->ResolveLength(thickness_); + const float inset = element->ResolveLength(inset_); + const Rml::ColourbPremultiplied tick = tick_color_.ToPremultiplied(opacity); + + Rml::Mesh mesh; + + // each corner is an L: the horizontal tick, then the vertical + // remainder below/above it (no overlap, so translucent tick colors + // blend once) + for (const float sx: {1.f, -1.f}) + for (const float sy: {1.f, -1.f}) { + const float hx = sx > 0.f ? inset : size.x - inset - length; + const float hy = sy > 0.f ? inset : size.y - inset - thickness; + Rml::MeshUtilities::GenerateQuad(mesh, {hx, hy}, {length, thickness}, tick); + if (length > thickness) { + const float vx = sx > 0.f ? inset : size.x - inset - thickness; + const float vy = sy > 0.f ? inset + thickness : size.y - inset - length; + Rml::MeshUtilities::GenerateQuad( + mesh, {vx, vy}, {thickness, length - thickness}, tick); + } + } + + // centered fill (slider knob): inset from the edges, the element's + // own border-radius, skipped when left transparent + if (fill_color_.alpha > 0) { + const float fill_inset = element->ResolveLength(fill_inset_); + const float radius = element->GetComputedValues().border_top_left_radius(); + const Rml::RenderBox fill_box( + {size.x - 2.f * fill_inset, size.y - 2.f * fill_inset}, {fill_inset, fill_inset}, + {0.f, 0.f, 0.f, 0.f}, {radius, radius, radius, radius}); + Rml::MeshUtilities::GenerateBackground( + mesh, fill_box, fill_color_.ToPremultiplied(opacity)); + } + + auto *geometry = + new Rml::Geometry(element->GetRenderManager()->MakeGeometry(std::move(mesh))); + return reinterpret_cast(geometry); + } + + void + CornerReticleDecorator::ReleaseElementData(const Rml::DecoratorDataHandle element_data) const { + delete reinterpret_cast(element_data); + } + + void CornerReticleDecorator::RenderElement( + Rml::Element *element, const Rml::DecoratorDataHandle element_data) const { + reinterpret_cast(element_data) + ->Render(element->GetAbsoluteOffset(Rml::BoxArea::Border)); + } + + CornerReticleDecoratorInstancer::CornerReticleDecoratorInstancer() { + id_tick_color_ = RegisterProperty("tick-color", "#00A6FB").AddParser("color").GetId(); + id_fill_color_ = RegisterProperty("fill-color", "transparent").AddParser("color").GetId(); + id_tick_length_ = RegisterProperty("tick-length", "7dp").AddParser("length").GetId(); + id_thickness_ = RegisterProperty("thickness", "2dp").AddParser("length").GetId(); + id_inset_ = RegisterProperty("inset", "3dp").AddParser("length").GetId(); + id_fill_inset_ = RegisterProperty("fill-inset", "0dp").AddParser("length").GetId(); + RegisterShorthand( + "decorator", "tick-color, fill-color, tick-length, thickness, inset, fill-inset", + Rml::ShorthandType::FallThrough); + } + + Rml::SharedPtr CornerReticleDecoratorInstancer::InstanceDecorator( + const Rml::String &, const Rml::PropertyDictionary &properties, + const Rml::DecoratorInstancerInterface &) { + return Rml::MakeShared( + properties.GetProperty(id_tick_color_)->Get(), + properties.GetProperty(id_fill_color_)->Get(), + properties.GetProperty(id_tick_length_)->GetNumericValue(), + properties.GetProperty(id_thickness_)->GetNumericValue(), + properties.GetProperty(id_inset_)->GetNumericValue(), + properties.GetProperty(id_fill_inset_)->GetNumericValue()); + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/corner_reticle.h b/arenai_desktop/src/gui/rml/corner_reticle.h new file mode 100644 index 00000000..f9c03a05 --- /dev/null +++ b/arenai_desktop/src/gui/rml/corner_reticle.h @@ -0,0 +1,59 @@ +// +// Created by samuel on 02/08/2026. +// + +#ifndef ARENAI_DESKTOP_GUI_RML_CORNER_RETICLE_H +#define ARENAI_DESKTOP_GUI_RML_CORNER_RETICLE_H + +#include + +namespace arenai::desktop::gui { + + // The menu's single cursor mark: four L-shaped corner ticks framing the + // control under the cursor (mouse hover or gamepad focus). The optional + // centered fill exists for the slider knob, a widget-internal element + // (sliderbar) that cannot host children — plain triangle geometry, the + // only thing the Vulkan backend renders (no box-shadow, no textures). + // RCSS: decorator: corner-reticle( + // ); + class CornerReticleDecorator final : public Rml::Decorator { + public: + CornerReticleDecorator( + Rml::Colourb tick_color, Rml::Colourb fill_color, Rml::NumericValue tick_length, + Rml::NumericValue thickness, Rml::NumericValue inset, Rml::NumericValue fill_inset); + + Rml::DecoratorDataHandle + GenerateElementData(Rml::Element *element, Rml::BoxArea) const override; + void ReleaseElementData(Rml::DecoratorDataHandle element_data) const override; + void + RenderElement(Rml::Element *element, Rml::DecoratorDataHandle element_data) const override; + + private: + Rml::Colourb tick_color_; + Rml::Colourb fill_color_; + Rml::NumericValue tick_length_; + Rml::NumericValue thickness_; + Rml::NumericValue inset_; + Rml::NumericValue fill_inset_; + }; + + class CornerReticleDecoratorInstancer final : public Rml::DecoratorInstancer { + public: + CornerReticleDecoratorInstancer(); + + Rml::SharedPtr InstanceDecorator( + const Rml::String &, const Rml::PropertyDictionary &properties, + const Rml::DecoratorInstancerInterface &) override; + + private: + Rml::PropertyId id_tick_color_{}; + Rml::PropertyId id_fill_color_{}; + Rml::PropertyId id_tick_length_{}; + Rml::PropertyId id_thickness_{}; + Rml::PropertyId id_inset_{}; + Rml::PropertyId id_fill_inset_{}; + }; + +}// namespace arenai::desktop::gui + +#endif// ARENAI_DESKTOP_GUI_RML_CORNER_RETICLE_H diff --git a/arenai_desktop/src/gui/rml/cursor_ring.cpp b/arenai_desktop/src/gui/rml/cursor_ring.cpp deleted file mode 100644 index 6015e43f..00000000 --- a/arenai_desktop/src/gui/rml/cursor_ring.cpp +++ /dev/null @@ -1,84 +0,0 @@ -// -// Created by samuel on 02/08/2026. -// - -#include "./cursor_ring.h" - -namespace arenai::desktop::gui { - - CursorRingDecorator::CursorRingDecorator( - const Rml::Colourb ring_color, const Rml::Colourb fill_color, - const Rml::NumericValue ring_width, const Rml::NumericValue gap) - : ring_color_(ring_color), fill_color_(fill_color), ring_width_(ring_width), gap_(gap) {} - - Rml::DecoratorDataHandle - CursorRingDecorator::GenerateElementData(Rml::Element *element, Rml::BoxArea) const { - const float opacity = element->GetComputedValues().opacity(); - const Rml::Vector2f size = element->GetBox().GetSize(Rml::BoxArea::Border); - const float ring_width = element->ResolveLength(ring_width_); - const float inset = ring_width + element->ResolveLength(gap_); - // ring and fill both adopt the element's border-radius: the - // ring must read as the same shape as the knob it surrounds - // (a concentric radius+inset ring looks round next to a - // square knob), so no radius inflation on the outer box - const float ring_radius = element->GetComputedValues().border_top_left_radius(); - const float fill_radius = ring_radius; - - Rml::Mesh mesh; - - // hollow ring at the outer edge (zero-alpha background: only - // the border geometry is emitted) - const Rml::RenderBox ring_box( - {size.x - 2.f * ring_width, size.y - 2.f * ring_width}, {0.f, 0.f}, - {ring_width, ring_width, ring_width, ring_width}, - {ring_radius, ring_radius, ring_radius, ring_radius}); - const Rml::ColourbPremultiplied ring = ring_color_.ToPremultiplied(opacity); - const Rml::ColourbPremultiplied ring_colors[4] = {ring, ring, ring, ring}; - Rml::MeshUtilities::GenerateBackgroundBorder( - mesh, ring_box, Rml::ColourbPremultiplied(0, 0, 0, 0), ring_colors); - - // fill disc, inset past the ring and the gap - const Rml::Vector2f fill_size = {size.x - 2.f * inset, size.y - 2.f * inset}; - const Rml::RenderBox fill_box( - fill_size, {inset, inset}, {0.f, 0.f, 0.f, 0.f}, - {fill_radius, fill_radius, fill_radius, fill_radius}); - Rml::MeshUtilities::GenerateBackground( - mesh, fill_box, fill_color_.ToPremultiplied(opacity)); - - auto *geometry = - new Rml::Geometry(element->GetRenderManager()->MakeGeometry(std::move(mesh))); - return reinterpret_cast(geometry); - } - - void - CursorRingDecorator::ReleaseElementData(const Rml::DecoratorDataHandle element_data) const { - delete reinterpret_cast(element_data); - } - - void CursorRingDecorator::RenderElement( - Rml::Element *element, const Rml::DecoratorDataHandle element_data) const { - reinterpret_cast(element_data) - ->Render(element->GetAbsoluteOffset(Rml::BoxArea::Border)); - } - - CursorRingDecoratorInstancer::CursorRingDecoratorInstancer() { - id_ring_color_ = RegisterProperty("ring-color", "#EAF7FF").AddParser("color").GetId(); - id_fill_color_ = RegisterProperty("fill-color", "#00A6FB").AddParser("color").GetId(); - id_ring_width_ = RegisterProperty("ring-width", "2dp").AddParser("length").GetId(); - id_gap_ = RegisterProperty("gap", "2dp").AddParser("length").GetId(); - RegisterShorthand( - "decorator", "ring-color, fill-color, ring-width, gap", - Rml::ShorthandType::FallThrough); - } - - Rml::SharedPtr CursorRingDecoratorInstancer::InstanceDecorator( - const Rml::String &, const Rml::PropertyDictionary &properties, - const Rml::DecoratorInstancerInterface &) { - return Rml::MakeShared( - properties.GetProperty(id_ring_color_)->Get(), - properties.GetProperty(id_fill_color_)->Get(), - properties.GetProperty(id_ring_width_)->GetNumericValue(), - properties.GetProperty(id_gap_)->GetNumericValue()); - } - -}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/cursor_ring.h b/arenai_desktop/src/gui/rml/cursor_ring.h deleted file mode 100644 index baa83c24..00000000 --- a/arenai_desktop/src/gui/rml/cursor_ring.h +++ /dev/null @@ -1,54 +0,0 @@ -// -// Created by samuel on 02/08/2026. -// - -#ifndef ARENAI_DESKTOP_GUI_RML_CURSOR_RING_H -#define ARENAI_DESKTOP_GUI_RML_CURSOR_RING_H - -#include - -namespace arenai::desktop::gui { - - // The slider knob is a widget-internal element (sliderbar) that cannot - // be wrapped in RML like the buttons, so its detached cursor ring is a - // gui-local decorator instead: an ink ring hugging the element's edge, - // a transparent gap, and the fill disc — plain triangle geometry, the - // only thing the Vulkan backend renders (no box-shadow, no textures). - // RCSS: decorator: cursor-ring( ); - class CursorRingDecorator final : public Rml::Decorator { - public: - CursorRingDecorator( - Rml::Colourb ring_color, Rml::Colourb fill_color, Rml::NumericValue ring_width, - Rml::NumericValue gap); - - Rml::DecoratorDataHandle - GenerateElementData(Rml::Element *element, Rml::BoxArea) const override; - void ReleaseElementData(Rml::DecoratorDataHandle element_data) const override; - void - RenderElement(Rml::Element *element, Rml::DecoratorDataHandle element_data) const override; - - private: - Rml::Colourb ring_color_; - Rml::Colourb fill_color_; - Rml::NumericValue ring_width_; - Rml::NumericValue gap_; - }; - - class CursorRingDecoratorInstancer final : public Rml::DecoratorInstancer { - public: - CursorRingDecoratorInstancer(); - - Rml::SharedPtr InstanceDecorator( - const Rml::String &, const Rml::PropertyDictionary &properties, - const Rml::DecoratorInstancerInterface &) override; - - private: - Rml::PropertyId id_ring_color_{}; - Rml::PropertyId id_fill_color_{}; - Rml::PropertyId id_ring_width_{}; - Rml::PropertyId id_gap_{}; - }; - -}// namespace arenai::desktop::gui - -#endif// ARENAI_DESKTOP_GUI_RML_CURSOR_RING_H diff --git a/arenai_desktop/src/gui/rml/gui.cpp b/arenai_desktop/src/gui/rml/gui.cpp new file mode 100644 index 00000000..82721c58 --- /dev/null +++ b/arenai_desktop/src/gui/rml/gui.cpp @@ -0,0 +1,466 @@ +// +// Created by samuel on 17/07/2026. +// + +#include +#include +#include +#include + +#include + +#include +#include + +#include "../menu.h" +#include "./adapters.h" +#include "./corner_reticle.h" +#include "./damage_arc.h" +#include "./hit_marker.h" +#include "./hud.h" +#include "./input.h" +#include "./pages/ai_page.h" +#include "./pages/controls_page.h" +#include "./pages/graphics_page.h" +#include "./reticle.h" + +namespace arenai::desktop::gui { + + namespace { + + // Orchestrates the RmlUi machinery: library lifecycle, the loaded + // documents and the navigation between them, the main-menu loop and + // the pause / game-over popups. Everything page-specific lives in the + // pages/ classes, the in-game overlay in Hud. + class RmlGui final : public AbstractGui { + public: + RmlGui( + const std::shared_ptr &backend, + const std::shared_ptr &asset_reader, + const GameSettings &initial_settings, const std::vector &gpus, + const int window_width, const int window_height, AgentValidator agent_validator) + : backend_(backend), window_(backend->get_window()), settings_(initial_settings), + width_(window_width), height_(window_height), file_interface_(asset_reader), + graphics_page_(settings_, window_, gpus), controls_page_(settings_, window_), + ai_page_(settings_, std::move(agent_validator)) { + Rml::SetSystemInterface(&system_interface_); + Rml::SetFileInterface(&file_interface_); + Rml::SetRenderInterface(&backend_->ui_render_interface()); + Rml::Initialise(); + + // menu.rcss frames the focused control (and paints the slider + // knob) with this gui-local decorator. Built only now: its + // property registration needs the style-sheet specification + // that Rml::Initialise() just created, so it cannot be a plain + // member (members are constructed before this body runs). + corner_reticle_instancer_ = std::make_unique(); + Rml::Factory::RegisterDecoratorInstancer( + "corner-reticle", corner_reticle_instancer_.get()); + hit_marker_instancer_ = std::make_unique(); + Rml::Factory::RegisterDecoratorInstancer("hit-marker", hit_marker_instancer_.get()); + reticle_instancer_ = std::make_unique(); + Rml::Factory::RegisterDecoratorInstancer("reticle", reticle_instancer_.get()); + damage_arc_instancer_ = std::make_unique(); + Rml::Factory::RegisterDecoratorInstancer("damage-arc", damage_arc_instancer_.get()); + + load_fonts(asset_reader); + + context_ = Rml::CreateContext("menu", Rml::Vector2i(width_, height_)); + if (!context_) throw std::runtime_error("RmlUi context creation failed"); + update_dp_ratio(); + + bind_data_model(); + ai_page_.refresh_explorer(); + + main_document_ = context_->LoadDocument("menu/main_menu.rml"); + params_document_ = context_->LoadDocument("menu/parameters.rml"); + controls_document_ = context_->LoadDocument("menu/controls.rml"); + graphics_document_ = context_->LoadDocument("menu/graphics.rml"); + ai_document_ = context_->LoadDocument("menu/ai.rml"); + pause_document_ = context_->LoadDocument("menu/pause.rml"); + game_over_document_ = context_->LoadDocument("menu/game_over.rml"); + hud_document_ = context_->LoadDocument("menu/hud.rml"); + if (!main_document_ || !params_document_ || !controls_document_ + || !graphics_document_ || !ai_document_ || !pause_document_ + || !game_over_document_ || !hud_document_) + throw std::runtime_error("RmlUi menu documents failed to load"); + + hud_ = std::make_unique(*hud_document_); + hud_document_->Show(Rml::ModalFlag::None, Rml::FocusFlag::None); + + // D-pad bridge across the file explorer's scroll container + ai_page_.attach(*ai_document_); + + input_adapter_ = std::make_shared( + context_, + [this] { + // Escape / gamepad B back out of the controls and + // parameters screens; while paused B resumes the game + // (the application intercepts Escape itself before + // this adapter); the game-over popup cannot be backed + // out of + if (game_over_document_->IsVisible()) return; + if (pause_document_->IsVisible()) + pending_pause_action_ = PauseAction::Continue; + else if (controls_document_->IsVisible()) close_controls(); + else if (graphics_document_->IsVisible()) close_graphics(); + else if (ai_document_->IsVisible()) close_ai(); + else if (params_document_->IsVisible()) close_params(); + }, + [this](const bool gamepad) { + // menu.rcss shows the :focus highlight only under + // .gamepad-nav, so the mouse hover and the gamepad + // cursor are never visible together + for (auto *document: + {main_document_, params_document_, controls_document_, + graphics_document_, ai_document_, pause_document_, + game_over_document_}) + document->SetClass("gamepad-nav", gamepad); + }); + controls_page_.set_input_adapter(input_adapter_); + + // route the pad input to the pad persisted from the previous + // session, when it is connected + controls_page_.on_open(); + } + + MenuOutcome run_main_menu() override { + play_clicked_ = false; + quit_clicked_ = false; + + window_->set_keyboard_callback(input_adapter_); + window_->set_gamepad_callback(input_adapter_); + window_->set_cursor_mode(controller::CursorMode::Normal); + main_document_->Show(); + + while (!window_->should_close() && !play_clicked_ && !quit_clicked_) { + window_->poll_events(); + + controls_page_.update(controls_document_->IsVisible()); + + context_->Update(); + + ai_page_.update_after_context(); + + backend_->begin_ui_frame(width_, height_); + context_->Render(); + backend_->end_ui_frame(); + backend_->present(); + } + + main_document_->Hide(); + params_document_->Hide(); + graphics_document_->Hide(); + ai_document_->Hide(); + controls_page_.on_close(); + controls_document_->Hide(); + window_->set_keyboard_callback(nullptr); + window_->set_gamepad_callback(nullptr); + + return play_clicked_ ? MenuOutcome::Play : MenuOutcome::Quit; + } + + GameSettings settings() const override { return settings_; } + + void open_pause(const int score) override { + pending_pause_action_ = PauseAction::None; + score_ = score; + model_handle_.DirtyVariable("score"); + pause_document_->Show(); + } + + void close_pause() override { pause_document_->Hide(); } + + void open_game_over(const int score) override { + pending_pause_action_ = PauseAction::None; + score_ = score; + model_handle_.DirtyVariable("score"); + game_over_document_->Show(); + } + + void close_game_over() override { game_over_document_->Hide(); } + + void render_pause_overlay() override { + context_->Update(); + + backend_->begin_ui_overlay(width_, height_); + context_->Render(); + backend_->end_ui_frame(); + } + + PauseAction poll_pause_action() override { + return std::exchange(pending_pause_action_, PauseAction::None); + } + + void notify_hit(const HitKind kind) override { hud_->notify_hit(kind); } + + void notify_damage(const float screen_angle) override { + hud_->notify_damage(screen_angle); + } + + void set_aim_point(const std::optional normalized) override { + hud_->set_aim_point(normalized, width_, height_); + } + + void render_hud_overlay() override { + context_->Update(); + + backend_->begin_ui_overlay(width_, height_); + context_->Render(); + backend_->end_ui_frame(); + } + + std::shared_ptr pause_input() override { + return input_adapter_; + } + + std::shared_ptr pause_gamepad_input() override { + return input_adapter_; + } + + void on_window_resized(const int width, const int height) override { + width_ = width; + height_ = height; + context_->SetDimensions(Rml::Vector2i(width_, height_)); + update_dp_ratio(); + } + + ~RmlGui() override { + // nothing may keep pointing at this object through the window + window_->set_keyboard_callback(nullptr); + window_->set_gamepad_callback(nullptr); + window_->set_resize_callback(nullptr); + + // releases the GL resources through the backend's render + // interface, whose context is still current on this thread + Rml::Shutdown(); + } + + private: + // Every dp length in menu.rcss is mapped to pixels relative to a + // 1080p design baseline, measured against the monitor the window + // sits on — not the window itself — so the menu keeps the same + // physical size on the display whether the game is fullscreen or + // in a small window (a 4K TV renders it twice as large either + // way). The min of both axes keeps the design fitting on unusual + // ratios. + void update_dp_ratio() const { + const auto [screen_width, screen_height] = window_->screen_size(); + context_->SetDensityIndependentPixelRatio(std::max( + 0.5f, std::min( + static_cast(screen_width) / 1920.0f, + static_cast(screen_height) / 1080.0f))); + } + + // Registered with an explicit family/weight (the static TTFs carry + // per-weight legacy family names that would not match the RCSS + // font-family otherwise). The buffers must outlive Rml::Shutdown(). + void + load_fonts(const std::shared_ptr &asset_reader) { + struct FontSpec { + const char *path; + const char *family; + int weight; + }; + constexpr FontSpec MENU_FONTS[] = { + {.path = "font/Sora-Regular.ttf", .family = "Sora", .weight = 400}, + {.path = "font/Sora-SemiBold.ttf", .family = "Sora", .weight = 600}, + {.path = "font/Sora-Bold.ttf", .family = "Sora", .weight = 700}, + {.path = "font/IBMPlexMono-Regular.ttf", + .family = "IBM Plex Mono", + .weight = 400}, + {.path = "font/IBMPlexMono-Medium.ttf", + .family = "IBM Plex Mono", + .weight = 500}, + {.path = "font/IBMPlexMono-SemiBold.ttf", + .family = "IBM Plex Mono", + .weight = 600}, + }; + + font_buffers_.reserve(std::size(MENU_FONTS)); + for (const auto &[path, family, weight]: MENU_FONTS) { + font_buffers_.push_back(asset_reader->read_text(path)); + const auto &buffer = font_buffers_.back(); + Rml::LoadFontFace( + Rml::Span( + reinterpret_cast(buffer.data()), buffer.size()), + family, Rml::Style::FontStyle::Normal, + static_cast(weight)); + } + } + + void bind_data_model() { + Rml::DataModelConstructor constructor = context_->CreateDataModel("settings"); + if (!constructor) throw std::runtime_error("RmlUi data model creation failed"); + + constructor.RegisterArray>(); + + constructor.Bind("nb_tanks", &settings_.nb_tanks); + constructor.Bind("spawn_side", &settings_.spawn_side); + constructor.Bind("score", &score_); + + graphics_page_.bind(constructor); + controls_page_.bind(constructor); + ai_page_.bind(constructor); + + constructor.BindEventCallback( + "play", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + if (ai_page_.can_play()) play_clicked_ = true; + }); + constructor.BindEventCallback( + "exit", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + quit_clicked_ = true; + }); + constructor.BindEventCallback( + "open_params", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + main_document_->Hide(); + params_document_->Show(); + }); + constructor.BindEventCallback( + "back", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + close_params(); + }); + constructor.BindEventCallback( + "open_graphics", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + params_document_->Hide(); + graphics_document_->Show(); + }); + constructor.BindEventCallback( + "graphics_back", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + close_graphics(); + }); + constructor.BindEventCallback( + "open_controls", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + params_document_->Hide(); + controls_page_.on_open(); + controls_document_->Show(); + }); + constructor.BindEventCallback( + "controls_back", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + close_controls(); + }); + constructor.BindEventCallback( + "open_ai", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + params_document_->Hide(); + ai_document_->Show(); + }); + constructor.BindEventCallback( + "ai_back", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + close_ai(); + }); + + constructor.BindEventCallback( + "pause_continue", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + pending_pause_action_ = PauseAction::Continue; + }); + constructor.BindEventCallback( + "pause_main_menu", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + pending_pause_action_ = PauseAction::MainMenu; + }); + constructor.BindEventCallback( + "pause_exit", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + pending_pause_action_ = PauseAction::ExitGame; + }); + constructor.BindEventCallback( + "game_over_retry", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + pending_pause_action_ = PauseAction::Retry; + }); + + model_handle_ = constructor.GetModelHandle(); + graphics_page_.set_model_handle(model_handle_); + controls_page_.set_model_handle(model_handle_); + ai_page_.set_model_handle(model_handle_); + } + + void close_params() const { + params_document_->Hide(); + main_document_->Show(); + } + + void close_ai() const { + ai_document_->Hide(); + params_document_->Show(); + } + + void close_graphics() const { + graphics_document_->Hide(); + params_document_->Show(); + } + + void close_controls() { + controls_page_.on_close(); + controls_document_->Hide(); + params_document_->Show(); + } + + std::shared_ptr backend_; + std::shared_ptr window_; + + GameSettings settings_; + int width_; + int height_; + + MenuSystemInterface system_interface_; + ReaderBackedFileInterface file_interface_; + + // the pages register their slice of the "settings" data model and + // write straight into settings_ + GraphicsPage graphics_page_; + ControlsPage controls_page_; + AiPage ai_page_; + + // unique_ptr: created after Rml::Initialise(), and member + // destruction keeps it alive until after Rml::Shutdown() as + // RmlUi requires of registered instancers + std::unique_ptr corner_reticle_instancer_; + std::unique_ptr hit_marker_instancer_; + std::unique_ptr reticle_instancer_; + std::unique_ptr damage_arc_instancer_; + std::vector font_buffers_; + + Rml::Context *context_ = nullptr; + Rml::ElementDocument *main_document_ = nullptr; + Rml::ElementDocument *params_document_ = nullptr; + Rml::ElementDocument *controls_document_ = nullptr; + Rml::ElementDocument *graphics_document_ = nullptr; + Rml::ElementDocument *ai_document_ = nullptr; + Rml::ElementDocument *pause_document_ = nullptr; + Rml::ElementDocument *game_over_document_ = nullptr; + Rml::ElementDocument *hud_document_ = nullptr; + // built once hud_document_ is loaded + std::unique_ptr hud_; + Rml::DataModelHandle model_handle_; + + std::shared_ptr input_adapter_; + + // score shown by the pause and game-over popups + int score_ = 0; + bool play_clicked_ = false; + bool quit_clicked_ = false; + PauseAction pending_pause_action_ = PauseAction::None; + }; + + }// namespace + + std::unique_ptr make_gui( + const std::shared_ptr &backend, + const std::shared_ptr &asset_reader, + const GameSettings &initial_settings, const std::vector &gpus, + const int window_width, const int window_height, AgentValidator agent_validator) { + return std::make_unique( + backend, asset_reader, initial_settings, gpus, window_width, window_height, + std::move(agent_validator)); + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/hud.cpp b/arenai_desktop/src/gui/rml/hud.cpp new file mode 100644 index 00000000..264f299c --- /dev/null +++ b/arenai_desktop/src/gui/rml/hud.cpp @@ -0,0 +1,83 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./hud.h" + +#include +#include + +namespace arenai::desktop::gui { + + Hud::Hud(Rml::ElementDocument &document) { + reticle_ = document.GetElementById("reticle"); + hit_marker_ = document.GetElementById("hit-marker"); + if (reticle_ == nullptr || hit_marker_ == nullptr) + throw std::runtime_error("hud.rml misses its #reticle / #hit-marker elements"); + + if (Rml::Element *arcs = document.GetElementById("damage-arcs")) + for (int i = 0; i < arcs->GetNumChildren(); i++) + damage_arcs_.push_back(arcs->GetChild(i)); + if (damage_arcs_.empty()) + throw std::runtime_error("hud.rml misses its #damage-arcs elements"); + } + + void Hud::notify_hit(const HitKind kind) const { + const bool kill = kind == HitKind::Kill; + hit_marker_->SetClass("kill", kill); + + // Battlefield-style feedback: the ticks spread outward while + // fading; a kill starts bigger, flares wider and lasts longer + const float duration = kill ? KILL_MARKER_FADE_SECONDS : HIT_MARKER_FADE_SECONDS; + const Rml::Tween tween(Rml::Tween::Quadratic, Rml::Tween::Out); + + const Rml::Property opaque(1.f, Rml::Unit::NUMBER); + hit_marker_->Animate( + "opacity", Rml::Property(0.f, Rml::Unit::NUMBER), duration, tween, 1, false, 0.f, + &opaque); + + const Rml::Property start_scale = Rml::Transform::MakeProperty( + {Rml::Transforms::Scale2D(kill ? KILL_MARKER_START_SCALE : 1.f)}); + hit_marker_->Animate( + "transform", + Rml::Transform::MakeProperty( + {Rml::Transforms::Scale2D(kill ? KILL_MARKER_END_SCALE : HIT_MARKER_END_SCALE)}), + duration, tween, 1, false, 0.f, &start_scale); + } + + void Hud::notify_damage(const float screen_angle) { + // oldest-slot reuse: a burst of impacts shows as many arcs as + // the pool holds, each rotated toward its own shooter + Rml::Element *arc = damage_arcs_[next_damage_arc_]; + next_damage_arc_ = (next_damage_arc_ + 1) % damage_arcs_.size(); + + constexpr float rad_to_deg = 180.f / std::numbers::pi_v; + arc->SetProperty( + Rml::PropertyId::Transform, + Rml::Transform::MakeProperty({Rml::Transforms::Rotate2D(screen_angle * rad_to_deg)})); + + const Rml::Tween tween(Rml::Tween::Quadratic, Rml::Tween::Out); + const Rml::Property opaque(1.f, Rml::Unit::NUMBER); + arc->Animate( + "opacity", Rml::Property(0.f, Rml::Unit::NUMBER), DAMAGE_ARC_FADE_SECONDS, tween, 1, + false, 0.f, &opaque); + } + + void Hud::set_aim_point( + const std::optional normalized, const int width, const int height) const { + if (!normalized) { + reticle_->SetProperty( + Rml::PropertyId::Visibility, Rml::Property(Rml::Style::Visibility::Hidden)); + return; + } + reticle_->SetProperty( + Rml::PropertyId::Visibility, Rml::Property(Rml::Style::Visibility::Visible)); + reticle_->SetProperty( + Rml::PropertyId::Left, + Rml::Property(normalized->x * static_cast(width), Rml::Unit::PX)); + reticle_->SetProperty( + Rml::PropertyId::Top, + Rml::Property(normalized->y * static_cast(height), Rml::Unit::PX)); + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/hud.h b/arenai_desktop/src/gui/rml/hud.h new file mode 100644 index 00000000..933fcd5a --- /dev/null +++ b/arenai_desktop/src/gui/rml/hud.h @@ -0,0 +1,46 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_DESKTOP_GUI_RML_HUD_H +#define ARENAI_DESKTOP_GUI_RML_HUD_H + +#include +#include +#include + +#include +#include + +#include "../menu.h" + +namespace arenai::desktop::gui { + + // The in-game overlay elements of hud.rml: the aim reticle, the + // Battlefield-style hit marker and the incoming-damage arc pool. + class Hud { + public: + // grabs the overlay elements from the loaded hud document + explicit Hud(Rml::ElementDocument &document); + + void notify_hit(HitKind kind) const; + void notify_damage(float screen_angle); + void set_aim_point(std::optional normalized, int width, int height) const; + + private: + Rml::Element *reticle_ = nullptr; + Rml::Element *hit_marker_ = nullptr; + std::vector damage_arcs_; + std::size_t next_damage_arc_ = 0; + + static constexpr float DAMAGE_ARC_FADE_SECONDS = 0.8f; + static constexpr float HIT_MARKER_FADE_SECONDS = 0.45f; + static constexpr float HIT_MARKER_END_SCALE = 1.3f; + static constexpr float KILL_MARKER_FADE_SECONDS = 0.6f; + static constexpr float KILL_MARKER_START_SCALE = 1.1f; + static constexpr float KILL_MARKER_END_SCALE = 1.55f; + }; + +}// namespace arenai::desktop::gui + +#endif// ARENAI_DESKTOP_GUI_RML_HUD_H diff --git a/arenai_desktop/src/gui/rml/input.cpp b/arenai_desktop/src/gui/rml/input.cpp index 51eb548a..f1e2b19a 100644 --- a/arenai_desktop/src/gui/rml/input.cpp +++ b/arenai_desktop/src/gui/rml/input.cpp @@ -27,7 +27,7 @@ namespace arenai::desktop::gui { Rml::ElementDocument *document = focus->GetOwnerDocument(); if (document == nullptr) return; const Rml::Element *list = document->GetElementById("file-list"); - Rml::Element *above = document->GetElementById("graphics-configure"); + Rml::Element *above = document->GetElementById("algorithm-toggles"); Rml::Element *below = document->GetElementById("use-folder"); if (list == nullptr || above == nullptr || below == nullptr) return; @@ -36,9 +36,13 @@ namespace arenai::desktop::gui { Rml::Element *target = nullptr; if (focus->GetParentNode() == list) { - if (key == Rml::Input::KI_UP && focus == entries.front()) target = above; - else if (key == Rml::Input::KI_DOWN && focus == entries.back()) target = below; - } else if (key == Rml::Input::KI_DOWN && focus == above) target = entries.front(); + if (key == Rml::Input::KI_UP && focus == entries.front()) { + // land on the algorithm currently in use, not a corner of the row + target = above->QuerySelector(".toggle.selected"); + if (target == nullptr && above->GetNumChildren() > 0) target = above->GetChild(0); + } else if (key == Rml::Input::KI_DOWN && focus == entries.back()) target = below; + } else if (key == Rml::Input::KI_DOWN && focus->GetParentNode() == above) + target = entries.front(); else if (key == Rml::Input::KI_UP && focus == below) target = entries.back(); if (target != nullptr && target->Focus(true)) { diff --git a/arenai_desktop/src/gui/rml/pages/ai_page.cpp b/arenai_desktop/src/gui/rml/pages/ai_page.cpp new file mode 100644 index 00000000..b8564210 --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/ai_page.cpp @@ -0,0 +1,161 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./ai_page.h" + +#include +#include +#include + +namespace arenai::desktop::gui { + + AiPage::AiPage(GameSettings &settings, AgentValidator validator) + : settings_(settings), agent_validator_(std::move(validator)), + algorithm_display_(to_string(settings.agent_algorithm)) { + current_dir_ = std::filesystem::exists(settings_.agent_folder) + ? std::filesystem::canonical(settings_.agent_folder) + : std::filesystem::current_path(); + + // the selection persisted from the previous run gets the same + // dry-run check as a freshly picked one + validate_agent(); + } + + void AiPage::bind(Rml::DataModelConstructor &constructor) { + constructor.Bind("algorithm", &algorithm_display_); + constructor.Bind("agent_folder", &agent_folder_display_); + constructor.Bind("agent_config", &agent_config_display_); + constructor.Bind("agent_status", &agent_status_); + constructor.Bind("agent_valid", &agent_valid_); + constructor.Bind("current_dir", ¤t_dir_display_); + constructor.Bind("entries", &entries_); + constructor.Bind("can_play", &can_play_); + + constructor.BindEventCallback( + "select_entry", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + const auto index = static_cast(arguments[0].Get()); + if (index >= entries_.size()) return; + + const std::string &entry = entries_[index]; + if (entry == "..") current_dir_ = current_dir_.parent_path(); + else if (const auto path = current_dir_ / entry; + std::filesystem::is_directory(path)) + current_dir_ = path; + else { + // a listed file is a config.json candidate + settings_.agent_config = path; + validate_agent(); + refresh_explorer(); + return; + } + refresh_explorer(); + focus_explorer_pending_ = true; + }); + constructor.BindEventCallback( + "set_algorithm", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + const auto algorithm = ai_algorithm_from_string(arguments[0].Get()); + if (!algorithm) return; + settings_.agent_algorithm = *algorithm; + algorithm_display_ = to_string(*algorithm); + model_handle_.DirtyVariable("algorithm"); + // the networks are rebuilt differently: re-run the check + validate_agent(); + refresh_explorer(); + }); + constructor.BindEventCallback( + "select_folder", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + settings_.agent_folder = current_dir_; + validate_agent(); + refresh_explorer(); + }); + } + + void AiPage::set_model_handle(const Rml::DataModelHandle handle) { model_handle_ = handle; } + + void AiPage::attach(Rml::ElementDocument &document) { + document_ = &document; + document.AddEventListener(Rml::EventId::Keydown, &explorer_nav_listener_); + } + + void AiPage::refresh_explorer() { + entries_.clear(); + if (current_dir_.has_parent_path() && current_dir_ != current_dir_.root_path()) + entries_.emplace_back(".."); + + std::vector config_files; + std::error_code list_error; + for (const auto &entry: std::filesystem::directory_iterator(current_dir_, list_error)) + if (std::error_code type_error; entry.is_directory(type_error)) + entries_.push_back(entry.path().filename().string()); + else if (entry.path().extension() == ".json") + config_files.push_back(entry.path().filename().string()); + if (list_error) + std::cerr << "Cannot list " << current_dir_ << ": " << list_error.message() + << std::endl; + + // keep ".." pinned first, then the directories, then the + // config.json candidates, each block sorted + const auto first_dir = + entries_.begin() + (!entries_.empty() && entries_[0] == ".." ? 1 : 0); + std::sort(first_dir, entries_.end()); + std::sort(config_files.begin(), config_files.end()); + entries_.insert(entries_.end(), config_files.begin(), config_files.end()); + + current_dir_display_ = current_dir_.string(); + agent_folder_display_ = settings_.agent_folder.string(); + agent_config_display_ = settings_.agent_config.empty() + ? "auto (config.json next to the state dicts)" + : settings_.agent_config.string(); + + if (model_handle_) { + model_handle_.DirtyVariable("entries"); + model_handle_.DirtyVariable("current_dir"); + model_handle_.DirtyVariable("agent_folder"); + model_handle_.DirtyVariable("agent_config"); + model_handle_.DirtyVariable("agent_status"); + model_handle_.DirtyVariable("agent_valid"); + model_handle_.DirtyVariable("can_play"); + } + } + + void AiPage::update_after_context() { + // entering a directory rebuilt the entry clones during Update + // (dropping the focused one): put the cursor back on the first entry + // of the fresh listing so the gamepad walk resumes there — invisible + // for the mouse, the :focus highlight only shows under .gamepad-nav + if (std::exchange(focus_explorer_pending_, false)) focus_first_entry(); + } + + // runs the injected dry-run load and turns its outcome into the + // tri-state the AI screen displays (nothing chosen yet / model + // loaded / error message); can_play_ follows the real load + void AiPage::validate_agent() { + if (settings_.agent_folder.empty()) { + agent_valid_ = false; + agent_status_ = ""; + } else { + const auto error = agent_validator_( + {.config = settings_.agent_config, + .folder = settings_.agent_folder, + .algorithm = settings_.agent_algorithm}); + agent_valid_ = !error.has_value(); + agent_status_ = error.value_or("AI model loaded"); + } + can_play_ = agent_valid_; + } + + void AiPage::focus_first_entry() const { + const Rml::Element *list = document_->GetElementById("file-list"); + if (list == nullptr) return; + const auto entries = ExplorerNavListener::visible_file_entries(list); + if (entries.empty()) return; + if (entries.front()->Focus(true)) + entries.front()->ScrollIntoView(Rml::ScrollAlignment::Nearest); + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/pages/ai_page.h b/arenai_desktop/src/gui/rml/pages/ai_page.h new file mode 100644 index 00000000..80d2eb0a --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/ai_page.h @@ -0,0 +1,71 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_DESKTOP_GUI_RML_PAGES_AI_PAGE_H +#define ARENAI_DESKTOP_GUI_RML_PAGES_AI_PAGE_H + +#include +#include + +#include + +#include "../../menu.h" +#include "../input.h" + +namespace arenai::desktop::gui { + + // The AI screen: the algorithm toggle and the file explorer picking the + // trained model's folder, dry-run validated through the injected + // AgentValidator. Registers its slice of the shared "settings" data model + // and writes straight into the GameSettings it was built around. + class AiPage { + public: + AiPage(GameSettings &settings, AgentValidator validator); + + // registers the page's variables and callbacks into the shared model + void bind(Rml::DataModelConstructor &constructor); + void set_model_handle(Rml::DataModelHandle handle); + + // hooks the D-pad bridge across the explorer's scroll container and + // remembers the document for the post-Update focus restore + void attach(Rml::ElementDocument &document); + + void refresh_explorer(); + + // after the context Update rebuilt the entry clones, put the gamepad + // cursor back on the first entry of the fresh listing + void update_after_context(); + + // the menu's Play gate: true once the dry-run load succeeded + bool can_play() const { return can_play_; } + + private: + void validate_agent(); + void focus_first_entry() const; + + GameSettings &settings_; + AgentValidator agent_validator_; + Rml::DataModelHandle model_handle_; + + Rml::ElementDocument *document_ = nullptr; + // removed from the document when Rml::Shutdown() destroys it, before + // the members are torn down + ExplorerNavListener explorer_nav_listener_; + bool focus_explorer_pending_ = false; + + std::filesystem::path current_dir_; + Rml::String algorithm_display_; + Rml::String current_dir_display_; + Rml::String agent_folder_display_; + Rml::String agent_config_display_; + // empty while no folder is chosen; otherwise success or error text + Rml::String agent_status_; + std::vector entries_; + bool agent_valid_ = false; + bool can_play_ = false; + }; + +}// namespace arenai::desktop::gui + +#endif// ARENAI_DESKTOP_GUI_RML_PAGES_AI_PAGE_H diff --git a/arenai_desktop/src/gui/rml/pages/controls_page.cpp b/arenai_desktop/src/gui/rml/pages/controls_page.cpp new file mode 100644 index 00000000..4547ca1a --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/controls_page.cpp @@ -0,0 +1,385 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./controls_page.h" + +#include +#include +#include +#include + +#include + +namespace arenai::desktop::gui { + + namespace { + + constexpr const char *KB_SLOT_LABELS[] = {"FORWARD", "BACKWARD", "TURN LEFT", + "TURN RIGHT", "FIRE", "ZOOM"}; + constexpr const char *GP_SLOT_LABELS[] = {"FIRE", "ZOOM", "STEER", "AIM X", + "AIM Y", "ACCELERATE", "REVERSE"}; + constexpr int NB_KB_SLOTS = 6; + constexpr int NB_GP_SLOTS = 7; + constexpr int NB_GP_BUTTON_SLOTS = 2; + + constexpr bool gp_slot_is_one_way(const int slot) { return slot >= 5; } + + constexpr double CAPTURE_ENGAGE = 0.6, CAPTURE_REST = 0.3; + + constexpr auto CAPTURE_TIMEOUT = std::chrono::seconds(15); + + constexpr const char *GAMEPAD_AXIS_LABELS[] = {"L-STICK X", "L-STICK Y", "R-STICK X", + "R-STICK Y", "LT", "RT"}; + + }// namespace + + ControlsPage::ControlsPage(GameSettings &settings, std::shared_ptr window) + : settings_(settings), window_(std::move(window)), + controller_display_( + settings.controller_kind == ControllerKind::Gamepad ? "gamepad" : "keyboard") {} + + void ControlsPage::bind(Rml::DataModelConstructor &constructor) { + if (auto row_handle = constructor.RegisterStruct()) { + row_handle.RegisterMember("label", &BindingRow::label); + row_handle.RegisterMember("binding", &BindingRow::binding); + row_handle.RegisterMember("listening", &BindingRow::listening); + row_handle.RegisterMember("bound", &BindingRow::bound); + } + constructor.RegisterArray>(); + + constructor.Bind("kb_rows", &kb_rows_); + constructor.Bind("gp_rows", &gp_rows_); + constructor.Bind("gamepads", &gamepad_names_); + constructor.Bind("selected_gamepad", &selected_gamepad_); + constructor.Bind("bind_status", &bind_status_); + constructor.Bind("bind_warning", &bind_warning_); + constructor.Bind("controller", &controller_display_); + + constructor.BindEventCallback( + "set_controller", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + controller_display_ = arguments[0].Get(); + settings_.controller_kind = controller_display_ == "gamepad" + ? ControllerKind::Gamepad + : ControllerKind::Keyboard; + model_handle_.DirtyVariable("controller"); + // the toggle lives on the controls page: switching the + // kind swaps the visible bindings, so a running capture + // and the status text follow along + end_capture(); + set_default_bind_status(); + rebuild_binding_rows(); + }); + constructor.BindEventCallback( + "capture_kb", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (!arguments.empty()) begin_capture(true, arguments[0].Get()); + }); + constructor.BindEventCallback( + "capture_gp", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (!arguments.empty()) begin_capture(false, arguments[0].Get()); + }); + constructor.BindEventCallback( + "select_pad", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + const auto index = static_cast(arguments[0].Get()); + if (index >= gamepads_.size()) return; + + settings_.bindings.gamepad.device_guid = gamepads_[index].guid; + settings_.bindings.gamepad.device_name = gamepads_[index].name; + window_->select_gamepad(gamepads_[index].id); + selected_gamepad_ = static_cast(index); + model_handle_.DirtyVariable("selected_gamepad"); + }); + constructor.BindEventCallback( + "reset_bindings", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { + end_capture(); + if (settings_.controller_kind == ControllerKind::Keyboard) + settings_.bindings.keyboard = {}; + else { + // the device choice is not a binding: keep it + auto gamepad = GamepadBindings{}; + gamepad.device_guid = std::move(settings_.bindings.gamepad.device_guid); + gamepad.device_name = std::move(settings_.bindings.gamepad.device_name); + settings_.bindings.gamepad = std::move(gamepad); + } + set_default_bind_status(); + rebuild_binding_rows(); + }); + } + + void ControlsPage::set_model_handle(const Rml::DataModelHandle handle) { + model_handle_ = handle; + } + + void ControlsPage::set_input_adapter(std::shared_ptr adapter) { + input_adapter_ = std::move(adapter); + } + + void ControlsPage::on_open() { + refresh_gamepad_list(); + set_default_bind_status(); + rebuild_binding_rows(); + } + + void ControlsPage::on_close() { + end_capture(); + rebuild_binding_rows(); + } + + void ControlsPage::update(const bool page_visible) { + // pads can be (un)plugged while the device list is on screen; the + // refresh is a no-op while nothing changed + if (page_visible && settings_.controller_kind == ControllerKind::Gamepad) + refresh_gamepad_list(); + + // a forgotten capture ends on its own (a pad-only player has no + // Escape at hand) + if (capture_slot_ >= 0 && std::chrono::steady_clock::now() >= capture_deadline_) { + end_capture(); + rebuild_binding_rows(); + } + } + + std::array *, 6> ControlsPage::kb_slots() { + auto &keyboard = settings_.bindings.keyboard; + return {&keyboard.forward, &keyboard.backward, &keyboard.turn_left, + &keyboard.turn_right, &keyboard.fire, &keyboard.zoom}; + } + + std::array *, 2> ControlsPage::gp_button_slots() { + auto &gamepad = settings_.bindings.gamepad; + return {&gamepad.fire, &gamepad.zoom}; + } + + std::array *, 5> ControlsPage::gp_axis_slots() { + auto &gamepad = settings_.bindings.gamepad; + return { + &gamepad.steer, &gamepad.aim_x, &gamepad.aim_y, &gamepad.accelerate, &gamepad.reverse}; + } + + Rml::String + ControlsPage::keyboard_slot_label(const std::optional &slot) const { + if (!slot) return "UNBOUND"; + if (const auto *button = std::get_if(&*slot)) switch (*button) { + case controller::MouseButton::Left: return "MOUSE L"; + case controller::MouseButton::Right: return "MOUSE R"; + case controller::MouseButton::Middle: return "MOUSE M"; + } + // layout-aware label (Key::Q reads "A" on AZERTY) + const auto label = window_->key_label(std::get(*slot)); + return label.empty() ? "?" : Rml::String(label); + } + + Rml::String ControlsPage::axis_slot_label( + const std::optional &slot, const bool one_way) { + if (!slot) return "UNBOUND"; + Rml::String label = GAMEPAD_AXIS_LABELS[static_cast(slot->axis)]; + // a one-way action bound to a stick reads a single direction + const bool on_stick = + slot->axis != GamepadAxis::LeftTrigger && slot->axis != GamepadAxis::RightTrigger; + if (one_way && on_stick) label += slot->sign > 0.f ? "+" : "-"; + return label; + } + + void ControlsPage::rebuild_binding_rows() { + kb_rows_.resize(NB_KB_SLOTS); + const auto keyboard_slots = kb_slots(); + for (int i = 0; i < NB_KB_SLOTS; i++) { + const bool listening = capture_keyboard_page_ && capture_slot_ == i; + kb_rows_[i] = { + .label = KB_SLOT_LABELS[i], + .binding = listening ? "PRESS A KEY..." : keyboard_slot_label(*keyboard_slots[i]), + .listening = listening, + .bound = keyboard_slots[i]->has_value()}; + } + + gp_rows_.resize(NB_GP_SLOTS); + const auto button_slots = gp_button_slots(); + const auto axis_slots = gp_axis_slots(); + for (int i = 0; i < NB_GP_SLOTS; i++) { + const bool listening = !capture_keyboard_page_ && capture_slot_ == i; + const bool button_slot = i < NB_GP_BUTTON_SLOTS; + Rml::String binding; + if (listening) binding = button_slot ? "PRESS A BUTTON..." : "MOVE AN AXIS..."; + else if (button_slot) + binding = *button_slots[i] ? controller::to_string(**button_slots[i]) : "UNBOUND"; + else + binding = + axis_slot_label(*axis_slots[i - NB_GP_BUTTON_SLOTS], gp_slot_is_one_way(i)); + gp_rows_[i] = { + .label = GP_SLOT_LABELS[i], + .binding = std::move(binding), + .listening = listening, + .bound = button_slot ? button_slots[i]->has_value() + : axis_slots[i - NB_GP_BUTTON_SLOTS]->has_value()}; + } + + if (model_handle_) { + model_handle_.DirtyVariable("kb_rows"); + model_handle_.DirtyVariable("gp_rows"); + model_handle_.DirtyVariable("bind_status"); + model_handle_.DirtyVariable("bind_warning"); + } + } + + void ControlsPage::set_default_bind_status() { + bind_warning_ = false; + bind_status_ = settings_.controller_kind == ControllerKind::Keyboard + ? "Click a slot, then press the new key or mouse button. " + "Escape cancels." + : "Click a slot, then press a button or move an axis. " + "Escape or 15s of inactivity cancels — Start stays the " + "pause toggle."; + } + + void ControlsPage::begin_capture(const bool keyboard_page, const int slot) { + const int nb_slots = keyboard_page ? NB_KB_SLOTS : NB_GP_SLOTS; + if (slot < 0 || slot >= nb_slots) return; + + capture_keyboard_page_ = keyboard_page; + capture_slot_ = slot; + capture_deadline_ = std::chrono::steady_clock::now() + CAPTURE_TIMEOUT; + // every axis must return to rest once before it can bind (the + // stick may still be deflected from navigating the menu) + axis_armed_.fill(false); + set_default_bind_status(); + + input_adapter_->set_capture_sink( + [this](const RawMenuInput &input) { on_capture_input(input); }); + rebuild_binding_rows(); + } + + void ControlsPage::end_capture() { + capture_slot_ = -1; + if (input_adapter_) input_adapter_->set_capture_sink(nullptr); + } + + void ControlsPage::on_capture_input(const RawMenuInput &input) { + // Escape is the only cancel input (every pad button stays + // bindable); a pad-only player relies on the capture timeout + if (const auto *key = std::get_if(&input); + key != nullptr && *key == controller::Key::Escape) { + end_capture(); + rebuild_binding_rows(); + return; + } + + if (capture_keyboard_page_) { + if (const auto *key = std::get_if(&input)) + assign_keyboard(KeyboardBinding(*key)); + else if (const auto *button = std::get_if(&input)) + assign_keyboard(KeyboardBinding(*button)); + // pad input has no meaning on the keyboard page + return; + } + + if (capture_slot_ < NB_GP_BUTTON_SLOTS) { + // Start stays the in-game pause toggle, never a binding + if (const auto *button = std::get_if(&input); + button != nullptr && *button != controller::GamepadButton::Start) + assign_gamepad_button(*button); + return; + } + + if (const auto *motion = std::get_if>(&input)) { + const auto &[axis, value] = *motion; + auto &armed = axis_armed_[static_cast(axis)]; + if (std::abs(value) < CAPTURE_REST) armed = true; + else if (armed && std::abs(value) > CAPTURE_ENGAGE) assign_gamepad_axis(axis, value); + } + } + + void ControlsPage::unbind_conflict(const Rml::String &new_label, const Rml::String &old_label) { + bind_warning_ = true; + bind_status_ = + new_label + " was bound to " + old_label + " — " + old_label + " is now unbound."; + } + + void ControlsPage::assign_keyboard(const KeyboardBinding &binding) { + const auto slots = kb_slots(); + *slots[capture_slot_] = binding; + + for (int i = 0; i < NB_KB_SLOTS; i++) + if (i != capture_slot_ && *slots[i] == binding) { + *slots[i] = std::nullopt; + unbind_conflict(keyboard_slot_label(binding), KB_SLOT_LABELS[i]); + } + + end_capture(); + rebuild_binding_rows(); + } + + void ControlsPage::assign_gamepad_button(const controller::GamepadButton button) { + const auto slots = gp_button_slots(); + *slots[capture_slot_] = button; + + // the other button slot loses a conflicting assignment + for (int slot = 0; slot < NB_GP_BUTTON_SLOTS; slot++) + if (slot != capture_slot_ && *slots[slot] == button) { + *slots[slot] = std::nullopt; + unbind_conflict(controller::to_string(button), GP_SLOT_LABELS[slot]); + } + + end_capture(); + rebuild_binding_rows(); + } + + void ControlsPage::assign_gamepad_axis(const GamepadAxis axis, const double value) { + const bool one_way = gp_slot_is_one_way(capture_slot_); + const auto slots = gp_axis_slots(); + const GamepadAxisBinding binding{.axis = axis, .sign = one_way && value < 0. ? -1.f : 1.f}; + *slots[capture_slot_ - NB_GP_BUTTON_SLOTS] = binding; + + // two slots clash when they read the same range of an axis: a + // two-way slot owns the whole axis, one-way slots only their + // captured side + for (int slot = NB_GP_BUTTON_SLOTS; slot < NB_GP_SLOTS; slot++) { + if (slot == capture_slot_) continue; + auto &other = *slots[slot - NB_GP_BUTTON_SLOTS]; + if (!other || other->axis != axis) continue; + if (one_way && gp_slot_is_one_way(slot) && other->sign != binding.sign) continue; + other = std::nullopt; + unbind_conflict(axis_slot_label(binding, one_way), GP_SLOT_LABELS[slot]); + } + + end_capture(); + rebuild_binding_rows(); + } + + void ControlsPage::refresh_gamepad_list() { + auto gamepads = window_->list_gamepads(); + const bool changed = gamepads.size() != gamepads_.size() + || !std::equal( + gamepads.begin(), gamepads.end(), gamepads_.begin(), + [](const view::GamepadInfo &a, const view::GamepadInfo &b) { + return a.id == b.id && a.guid == b.guid && a.name == b.name; + }); + if (!changed) return; + + gamepads_ = std::move(gamepads); + gamepad_names_.clear(); + for (const auto &pad: gamepads_) gamepad_names_.emplace_back(pad.name); + + // the preferred pad when connected, else the first one (which + // is what the window falls back to) + selected_gamepad_ = gamepads_.empty() ? -1 : 0; + for (size_t i = 0; i < gamepads_.size(); i++) + if (!settings_.bindings.gamepad.device_guid.empty() + && gamepads_[i].guid == settings_.bindings.gamepad.device_guid) { + selected_gamepad_ = static_cast(i); + break; + } + window_->select_gamepad(selected_gamepad_ >= 0 ? gamepads_[selected_gamepad_].id : -1); + + if (model_handle_) { + model_handle_.DirtyVariable("gamepads"); + model_handle_.DirtyVariable("selected_gamepad"); + } + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/pages/controls_page.h b/arenai_desktop/src/gui/rml/pages/controls_page.h new file mode 100644 index 00000000..b4cb38ea --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/controls_page.h @@ -0,0 +1,91 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_DESKTOP_GUI_RML_PAGES_CONTROLS_PAGE_H +#define ARENAI_DESKTOP_GUI_RML_PAGES_CONTROLS_PAGE_H + +#include +#include +#include +#include +#include + +#include + +#include + +#include "../../../controller/bindings.h" +#include "../../menu.h" +#include "../input.h" + +namespace arenai::desktop::gui { + + // one row of the controls page: an action and its current binding + struct BindingRow { + Rml::String label; + Rml::String binding; + bool listening = false; + bool bound = true; + }; + + class ControlsPage { + public: + ControlsPage(GameSettings &settings, std::shared_ptr window); + + void bind(Rml::DataModelConstructor &constructor); + void set_model_handle(Rml::DataModelHandle handle); + void set_input_adapter(std::shared_ptr adapter); + + void on_open(); + void on_close(); + + void update(bool page_visible); + + private: + std::array *, 6> kb_slots(); + std::array *, 2> gp_button_slots(); + std::array *, 5> gp_axis_slots(); + + Rml::String keyboard_slot_label(const std::optional &slot) const; + static Rml::String + axis_slot_label(const std::optional &slot, bool one_way); + + void rebuild_binding_rows(); + void set_default_bind_status(); + + void begin_capture(bool keyboard_page, int slot); + void end_capture(); + void on_capture_input(const RawMenuInput &input); + + void unbind_conflict(const Rml::String &new_label, const Rml::String &old_label); + void assign_keyboard(const KeyboardBinding &binding); + void assign_gamepad_button(controller::GamepadButton button); + void assign_gamepad_axis(GamepadAxis axis, double value); + + void refresh_gamepad_list(); + + GameSettings &settings_; + std::shared_ptr window_; + std::shared_ptr input_adapter_; + Rml::DataModelHandle model_handle_; + + Rml::String controller_display_; + std::vector kb_rows_; + std::vector gp_rows_; + std::vector gamepads_; + std::vector gamepad_names_; + + int selected_gamepad_ = -1; + Rml::String bind_status_; + bool bind_warning_ = false; + + int capture_slot_ = -1; + bool capture_keyboard_page_ = false; + std::chrono::steady_clock::time_point capture_deadline_; + std::array axis_armed_{}; + }; + +}// namespace arenai::desktop::gui + +#endif// ARENAI_DESKTOP_GUI_RML_PAGES_CONTROLS_PAGE_H diff --git a/arenai_desktop/src/gui/rml/pages/graphics_page.cpp b/arenai_desktop/src/gui/rml/pages/graphics_page.cpp new file mode 100644 index 00000000..8d90cfe7 --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/graphics_page.cpp @@ -0,0 +1,106 @@ +// +// Created by samuel on 06/09/2026. +// + +#include "./graphics_page.h" + +#include +#include + +namespace arenai::desktop::gui { + + GraphicsPage::GraphicsPage( + GameSettings &settings, std::shared_ptr window, + const std::vector &gpus) + : settings_(settings), window_(std::move(window)), + display_display_(settings.fullscreen ? "fullscreen" : "windowed"), + shadow_display_(to_string(settings.shadow_quality)) { + gpu_names_.emplace_back("Auto"); + for (const auto &gpu: gpus) gpu_names_.emplace_back(gpu); + + const auto gpu_index = [this](const std::string &name) -> int { + for (size_t i = 1; i < gpu_names_.size(); i++) + if (!name.empty() && gpu_names_[i] == name) return static_cast(i); + return 0; + }; + selected_window_gpu_ = gpu_index(settings_.window_gpu); + selected_vision_gpu_ = gpu_index(settings_.vision_gpu); + // a saved GPU that disappeared falls back to Auto, re-saved as such + if (selected_window_gpu_ == 0) settings_.window_gpu.clear(); + if (selected_vision_gpu_ == 0) settings_.vision_gpu.clear(); + + window_gpu_env_override_ = std::getenv("ARENAI_VK_DEVICE_WINDOW") != nullptr; + vision_gpu_env_override_ = std::getenv("ARENAI_VK_DEVICE") != nullptr; + } + + void GraphicsPage::bind(Rml::DataModelConstructor &constructor) { + constructor.Bind("display", &display_display_); + constructor.Bind("shadow_quality", &shadow_display_); + constructor.Bind("msaa", &settings_.msaa_samples); + constructor.Bind("gpus", &gpu_names_); + constructor.Bind("selected_window_gpu", &selected_window_gpu_); + constructor.Bind("selected_vision_gpu", &selected_vision_gpu_); + constructor.Bind("window_gpu_env", &window_gpu_env_override_); + constructor.Bind("vision_gpu_env", &vision_gpu_env_override_); + + constructor.BindEventCallback( + "set_display", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + display_display_ = arguments[0].Get(); + settings_.fullscreen = display_display_ == "fullscreen"; + // applied immediately; the window reports its new size + // through the resize callback (dp-ratio included) + window_->set_fullscreen(settings_.fullscreen); + model_handle_.DirtyVariable("display"); + }); + constructor.BindEventCallback( + "set_shadow_quality", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + const auto quality = shadow_quality_from_string(arguments[0].Get()); + if (!quality) return; + settings_.shadow_quality = *quality; + shadow_display_ = to_string(*quality); + model_handle_.DirtyVariable("shadow_quality"); + }); + constructor.BindEventCallback( + "set_msaa", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + const int samples = arguments[0].Get(); + if (samples != 1 && samples != 2 && samples != 4 && samples != 8) return; + settings_.msaa_samples = samples; + model_handle_.DirtyVariable("msaa"); + }); + constructor.BindEventCallback( + "set_window_gpu", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + select_gpu( + arguments[0].Get(), selected_window_gpu_, settings_.window_gpu, + "selected_window_gpu"); + }); + constructor.BindEventCallback( + "set_vision_gpu", + [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { + if (arguments.empty()) return; + select_gpu( + arguments[0].Get(), selected_vision_gpu_, settings_.vision_gpu, + "selected_vision_gpu"); + }); + } + + void GraphicsPage::set_model_handle(const Rml::DataModelHandle handle) { + model_handle_ = handle; + } + + void GraphicsPage::select_gpu( + const int index, int &selected, std::string &setting, const char *variable) { + if (index < 0 || index >= static_cast(gpu_names_.size())) return; + selected = index; + setting = index == 0 ? std::string() : std::string(gpu_names_[index]); + model_handle_.DirtyVariable(variable); + } + +}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/gui/rml/pages/graphics_page.h b/arenai_desktop/src/gui/rml/pages/graphics_page.h new file mode 100644 index 00000000..6a662ad9 --- /dev/null +++ b/arenai_desktop/src/gui/rml/pages/graphics_page.h @@ -0,0 +1,53 @@ +// +// Created by samuel on 06/09/2026. +// + +#ifndef ARENAI_DESKTOP_GUI_RML_PAGES_GRAPHICS_PAGE_H +#define ARENAI_DESKTOP_GUI_RML_PAGES_GRAPHICS_PAGE_H + +#include +#include +#include + +#include + +#include + +#include "../../menu.h" + +namespace arenai::desktop::gui { + + // The Graphics screen: display mode, shadow quality, MSAA and the two GPU + // pickers. Registers its slice of the shared "settings" data model and + // writes straight into the GameSettings it was built around. + class GraphicsPage { + public: + GraphicsPage( + GameSettings &settings, std::shared_ptr window, + const std::vector &gpus); + + // registers the page's variables and callbacks into the shared model + void bind(Rml::DataModelConstructor &constructor); + void set_model_handle(Rml::DataModelHandle handle); + + private: + // "Auto" rides index 0 of gpu_names_, the actual devices follow + void select_gpu(int index, int &selected, std::string &setting, const char *variable); + + GameSettings &settings_; + std::shared_ptr window_; + + Rml::String display_display_; + Rml::String shadow_display_; + std::vector gpu_names_; + int selected_window_gpu_ = 0; + int selected_vision_gpu_ = 0; + bool window_gpu_env_override_ = false; + bool vision_gpu_env_override_ = false; + + Rml::DataModelHandle model_handle_; + }; + +}// namespace arenai::desktop::gui + +#endif// ARENAI_DESKTOP_GUI_RML_PAGES_GRAPHICS_PAGE_H diff --git a/arenai_desktop/src/gui/rml_menu.cpp b/arenai_desktop/src/gui/rml_menu.cpp deleted file mode 100644 index 397e37df..00000000 --- a/arenai_desktop/src/gui/rml_menu.cpp +++ /dev/null @@ -1,1066 +0,0 @@ -// -// Created by samuel on 17/07/2026. -// - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include "./menu.h" -#include "./rml/adapters.h" -#include "./rml/cursor_ring.h" -#include "./rml/damage_arc.h" -#include "./rml/hit_marker.h" -#include "./rml/input.h" -#include "./rml/reticle.h" - -namespace arenai::desktop::gui { - - namespace { - - // one row of the controls page: an action and its current binding - struct BindingRow { - Rml::String label; - Rml::String binding; - bool listening = false; - bool bound = true; - }; - - // slot order of the two pages; the gamepad page puts its button slot - // (fire) first, then the five axis slots - constexpr const char *KB_SLOT_LABELS[] = { - "FORWARD", "BACKWARD", "TURN LEFT", "TURN RIGHT", "FIRE"}; - constexpr const char *GP_SLOT_LABELS[] = {"FIRE", "STEER", "AIM X", - "AIM Y", "ACCELERATE", "REVERSE"}; - constexpr int NB_KB_SLOTS = 5; - constexpr int NB_GP_SLOTS = 6; - // accelerate / reverse read a single direction of their axis - constexpr bool gp_slot_is_one_way(const int slot) { return slot >= 4; } - - // an axis must come back to rest before it can be captured: without - // this the stick still deflected from navigating the menu (or the A - // press bound to a trigger) would bind itself instantly - constexpr double CAPTURE_ENGAGE = 0.6, CAPTURE_REST = 0.3; - - // a capture nobody feeds cancels itself: Escape is the only cancel - // input (every pad button stays bindable), so a pad-only player - // needs the timeout - constexpr auto CAPTURE_TIMEOUT = std::chrono::seconds(15); - - constexpr const char *GAMEPAD_AXIS_LABELS[] = {"L-STICK X", "L-STICK Y", "R-STICK X", - "R-STICK Y", "LT", "RT"}; - - class RmlGui final : public AbstractGui { - public: - RmlGui( - const std::shared_ptr &backend, - const std::shared_ptr &asset_reader, - const GameSettings &initial_settings, const std::vector &gpus, - const int window_width, const int window_height, SacFolderValidator sac_validator) - : backend_(backend), window_(backend->get_window()), settings_(initial_settings), - width_(window_width), height_(window_height), file_interface_(asset_reader), - sac_validator_(std::move(sac_validator)), - controller_display_( - initial_settings.controller_kind == ControllerKind::Gamepad ? "gamepad" - : "keyboard"), - display_display_(initial_settings.fullscreen ? "fullscreen" : "windowed"), - shadow_display_(to_string(initial_settings.shadow_quality)) { - gpu_names_.emplace_back("Auto"); - for (const auto &gpu: gpus) gpu_names_.emplace_back(gpu); - - const auto gpu_index = [this](const std::string &name) -> int { - for (size_t i = 1; i < gpu_names_.size(); i++) - if (!name.empty() && gpu_names_[i] == name) return static_cast(i); - return 0; - }; - selected_window_gpu_ = gpu_index(settings_.window_gpu); - selected_vision_gpu_ = gpu_index(settings_.vision_gpu); - // a saved GPU that disappeared falls back to Auto, re-saved as such - if (selected_window_gpu_ == 0) settings_.window_gpu.clear(); - if (selected_vision_gpu_ == 0) settings_.vision_gpu.clear(); - - window_gpu_env_override_ = std::getenv("ARENAI_VK_DEVICE_WINDOW") != nullptr; - vision_gpu_env_override_ = std::getenv("ARENAI_VK_DEVICE") != nullptr; - - Rml::SetSystemInterface(&system_interface_); - Rml::SetFileInterface(&file_interface_); - Rml::SetRenderInterface(&backend_->ui_render_interface()); - Rml::Initialise(); - - // menu.rcss draws the slider knob's detached cursor ring with - // this gui-local decorator. Built only now: its property - // registration needs the style-sheet specification that - // Rml::Initialise() just created, so it cannot be a plain - // member (members are constructed before this body runs). - cursor_ring_instancer_ = std::make_unique(); - Rml::Factory::RegisterDecoratorInstancer( - "cursor-ring", cursor_ring_instancer_.get()); - hit_marker_instancer_ = std::make_unique(); - Rml::Factory::RegisterDecoratorInstancer("hit-marker", hit_marker_instancer_.get()); - reticle_instancer_ = std::make_unique(); - Rml::Factory::RegisterDecoratorInstancer("reticle", reticle_instancer_.get()); - damage_arc_instancer_ = std::make_unique(); - Rml::Factory::RegisterDecoratorInstancer("damage-arc", damage_arc_instancer_.get()); - - load_fonts(asset_reader); - - context_ = Rml::CreateContext("menu", Rml::Vector2i(width_, height_)); - if (!context_) throw std::runtime_error("RmlUi context creation failed"); - update_dp_ratio(); - - current_dir_ = std::filesystem::exists(settings_.sac_folder) - ? std::filesystem::canonical(settings_.sac_folder) - : std::filesystem::current_path(); - - // the folder persisted from the previous run gets the same - // dry-run check as a freshly picked one - validate_sac_folder(); - - bind_data_model(); - refresh_explorer(); - - main_document_ = context_->LoadDocument("menu/main_menu.rml"); - params_document_ = context_->LoadDocument("menu/parameters.rml"); - controls_document_ = context_->LoadDocument("menu/controls.rml"); - graphics_document_ = context_->LoadDocument("menu/graphics.rml"); - pause_document_ = context_->LoadDocument("menu/pause.rml"); - game_over_document_ = context_->LoadDocument("menu/game_over.rml"); - hud_document_ = context_->LoadDocument("menu/hud.rml"); - if (!main_document_ || !params_document_ || !controls_document_ - || !graphics_document_ || !pause_document_ || !game_over_document_ - || !hud_document_) - throw std::runtime_error("RmlUi menu documents failed to load"); - - // D-pad bridge across the file explorer's scroll container - reticle_ = hud_document_->GetElementById("reticle"); - hit_marker_ = hud_document_->GetElementById("hit-marker"); - if (reticle_ == nullptr || hit_marker_ == nullptr) - throw std::runtime_error("hud.rml misses its #reticle / #hit-marker elements"); - - if (Rml::Element *arcs = hud_document_->GetElementById("damage-arcs")) - for (int i = 0; i < arcs->GetNumChildren(); i++) - damage_arcs_.push_back(arcs->GetChild(i)); - if (damage_arcs_.empty()) - throw std::runtime_error("hud.rml misses its #damage-arcs elements"); - - hud_document_->Show(Rml::ModalFlag::None, Rml::FocusFlag::None); - - params_document_->AddEventListener(Rml::EventId::Keydown, &explorer_nav_listener_); - - input_adapter_ = std::make_shared( - context_, - [this] { - // Escape / gamepad B back out of the controls and - // parameters screens; while paused B resumes the game - // (the application intercepts Escape itself before - // this adapter); the game-over popup cannot be backed - // out of - if (game_over_document_->IsVisible()) return; - if (pause_document_->IsVisible()) - pending_pause_action_ = PauseAction::Continue; - else if (controls_document_->IsVisible()) close_controls(); - else if (graphics_document_->IsVisible()) close_graphics(); - else if (params_document_->IsVisible()) close_params(); - }, - [this](const bool gamepad) { - // menu.rcss shows the :focus highlight only under - // .gamepad-nav, so the mouse hover and the gamepad - // cursor are never visible together - for (auto *document: - {main_document_, params_document_, controls_document_, - graphics_document_, pause_document_, game_over_document_}) - document->SetClass("gamepad-nav", gamepad); - }); - - // route the pad input to the pad persisted from the previous - // session, when it is connected - refresh_gamepad_list(); - rebuild_binding_rows(); - } - - MenuOutcome run_main_menu() override { - play_clicked_ = false; - quit_clicked_ = false; - - window_->set_keyboard_callback(input_adapter_); - window_->set_gamepad_callback(input_adapter_); - window_->set_cursor_mode(controller::CursorMode::Normal); - main_document_->Show(); - - while (!window_->should_close() && !play_clicked_ && !quit_clicked_) { - window_->poll_events(); - - // pads can be (un)plugged while the device list is on - // screen; the refresh is a no-op while nothing changed - if (controls_document_->IsVisible() - && settings_.controller_kind == ControllerKind::Gamepad) - refresh_gamepad_list(); - - // a forgotten capture ends on its own (a pad-only player - // has no Escape at hand) - if (capture_slot_ >= 0 - && std::chrono::steady_clock::now() >= capture_deadline_) { - end_capture(); - rebuild_binding_rows(); - } - - context_->Update(); - - // entering a directory rebuilt the entry clones during - // Update (dropping the focused one): put the cursor back - // on the first entry of the fresh listing so the gamepad - // walk resumes there — invisible for the mouse, the - // :focus highlight only shows under .gamepad-nav - if (std::exchange(focus_explorer_pending_, false)) focus_first_entry(); - - backend_->begin_ui_frame(width_, height_); - context_->Render(); - backend_->end_ui_frame(); - backend_->present(); - } - - main_document_->Hide(); - params_document_->Hide(); - graphics_document_->Hide(); - end_capture(); - controls_document_->Hide(); - window_->set_keyboard_callback(nullptr); - window_->set_gamepad_callback(nullptr); - - return play_clicked_ ? MenuOutcome::Play : MenuOutcome::Quit; - } - - GameSettings settings() const override { return settings_; } - - void open_pause(const int score) override { - pending_pause_action_ = PauseAction::None; - score_ = score; - model_handle_.DirtyVariable("score"); - pause_document_->Show(); - } - - void close_pause() override { pause_document_->Hide(); } - - void open_game_over(const int score) override { - pending_pause_action_ = PauseAction::None; - score_ = score; - model_handle_.DirtyVariable("score"); - game_over_document_->Show(); - } - - void close_game_over() override { game_over_document_->Hide(); } - - void render_pause_overlay() override { - context_->Update(); - - backend_->begin_ui_overlay(width_, height_); - context_->Render(); - backend_->end_ui_frame(); - } - - PauseAction poll_pause_action() override { - return std::exchange(pending_pause_action_, PauseAction::None); - } - - void notify_hit(const HitKind kind) override { - const bool kill = kind == HitKind::Kill; - hit_marker_->SetClass("kill", kill); - - // Battlefield-style feedback: the ticks spread outward while - // fading; a kill starts bigger, flares wider and lasts longer - const float duration = kill ? KILL_MARKER_FADE_SECONDS : HIT_MARKER_FADE_SECONDS; - const Rml::Tween tween(Rml::Tween::Quadratic, Rml::Tween::Out); - - const Rml::Property opaque(1.f, Rml::Unit::NUMBER); - hit_marker_->Animate( - "opacity", Rml::Property(0.f, Rml::Unit::NUMBER), duration, tween, 1, false, - 0.f, &opaque); - - const Rml::Property start_scale = Rml::Transform::MakeProperty( - {Rml::Transforms::Scale2D(kill ? KILL_MARKER_START_SCALE : 1.f)}); - hit_marker_->Animate( - "transform", - Rml::Transform::MakeProperty({Rml::Transforms::Scale2D( - kill ? KILL_MARKER_END_SCALE : HIT_MARKER_END_SCALE)}), - duration, tween, 1, false, 0.f, &start_scale); - } - - void notify_damage(const float screen_angle) override { - // oldest-slot reuse: a burst of impacts shows as many arcs as - // the pool holds, each rotated toward its own shooter - Rml::Element *arc = damage_arcs_[next_damage_arc_]; - next_damage_arc_ = (next_damage_arc_ + 1) % damage_arcs_.size(); - - constexpr float rad_to_deg = 180.f / std::numbers::pi_v; - arc->SetProperty( - Rml::PropertyId::Transform, - Rml::Transform::MakeProperty( - {Rml::Transforms::Rotate2D(screen_angle * rad_to_deg)})); - - const Rml::Tween tween(Rml::Tween::Quadratic, Rml::Tween::Out); - const Rml::Property opaque(1.f, Rml::Unit::NUMBER); - arc->Animate( - "opacity", Rml::Property(0.f, Rml::Unit::NUMBER), DAMAGE_ARC_FADE_SECONDS, - tween, 1, false, 0.f, &opaque); - } - - void set_aim_point(const std::optional normalized) override { - if (!normalized) { - reticle_->SetProperty( - Rml::PropertyId::Visibility, Rml::Property(Rml::Style::Visibility::Hidden)); - return; - } - reticle_->SetProperty( - Rml::PropertyId::Visibility, Rml::Property(Rml::Style::Visibility::Visible)); - reticle_->SetProperty( - Rml::PropertyId::Left, - Rml::Property(normalized->x * static_cast(width_), Rml::Unit::PX)); - reticle_->SetProperty( - Rml::PropertyId::Top, - Rml::Property(normalized->y * static_cast(height_), Rml::Unit::PX)); - } - - void render_hud_overlay() override { - context_->Update(); - - backend_->begin_ui_overlay(width_, height_); - context_->Render(); - backend_->end_ui_frame(); - } - - std::shared_ptr pause_input() override { - return input_adapter_; - } - - std::shared_ptr pause_gamepad_input() override { - return input_adapter_; - } - - void on_window_resized(const int width, const int height) override { - width_ = width; - height_ = height; - context_->SetDimensions(Rml::Vector2i(width_, height_)); - update_dp_ratio(); - } - - ~RmlGui() override { - // nothing may keep pointing at this object through the window - window_->set_keyboard_callback(nullptr); - window_->set_gamepad_callback(nullptr); - window_->set_resize_callback(nullptr); - - // releases the GL resources through the backend's render - // interface, whose context is still current on this thread - Rml::Shutdown(); - } - - private: - // Every dp length in menu.rcss is mapped to pixels relative to a - // 1080p design baseline, measured against the monitor the window - // sits on — not the window itself — so the menu keeps the same - // physical size on the display whether the game is fullscreen or - // in a small window (a 4K TV renders it twice as large either - // way). The min of both axes keeps the design fitting on unusual - // ratios. - void update_dp_ratio() const { - const auto [screen_width, screen_height] = window_->screen_size(); - context_->SetDensityIndependentPixelRatio(std::max( - 0.5f, std::min( - static_cast(screen_width) / 1920.0f, - static_cast(screen_height) / 1080.0f))); - } - - // Registered with an explicit family/weight (the static TTFs carry - // per-weight legacy family names that would not match the RCSS - // font-family otherwise). The buffers must outlive Rml::Shutdown(). - void - load_fonts(const std::shared_ptr &asset_reader) { - struct FontSpec { - const char *path; - const char *family; - int weight; - }; - constexpr FontSpec MENU_FONTS[] = { - {.path = "font/Sora-Regular.ttf", .family = "Sora", .weight = 400}, - {.path = "font/Sora-SemiBold.ttf", .family = "Sora", .weight = 600}, - {.path = "font/Sora-Bold.ttf", .family = "Sora", .weight = 700}, - {.path = "font/IBMPlexMono-Regular.ttf", - .family = "IBM Plex Mono", - .weight = 400}, - {.path = "font/IBMPlexMono-Medium.ttf", - .family = "IBM Plex Mono", - .weight = 500}, - {.path = "font/IBMPlexMono-SemiBold.ttf", - .family = "IBM Plex Mono", - .weight = 600}, - }; - - font_buffers_.reserve(std::size(MENU_FONTS)); - for (const auto &[path, family, weight]: MENU_FONTS) { - font_buffers_.push_back(asset_reader->read_text(path)); - const auto &buffer = font_buffers_.back(); - Rml::LoadFontFace( - Rml::Span( - reinterpret_cast(buffer.data()), buffer.size()), - family, Rml::Style::FontStyle::Normal, - static_cast(weight)); - } - } - - void bind_data_model() { - Rml::DataModelConstructor constructor = context_->CreateDataModel("settings"); - if (!constructor) throw std::runtime_error("RmlUi data model creation failed"); - - constructor.RegisterArray>(); - - if (auto row_handle = constructor.RegisterStruct()) { - row_handle.RegisterMember("label", &BindingRow::label); - row_handle.RegisterMember("binding", &BindingRow::binding); - row_handle.RegisterMember("listening", &BindingRow::listening); - row_handle.RegisterMember("bound", &BindingRow::bound); - } - constructor.RegisterArray>(); - - constructor.Bind("kb_rows", &kb_rows_); - constructor.Bind("gp_rows", &gp_rows_); - constructor.Bind("gamepads", &gamepad_names_); - constructor.Bind("selected_gamepad", &selected_gamepad_); - constructor.Bind("bind_status", &bind_status_); - constructor.Bind("bind_warning", &bind_warning_); - - constructor.Bind("nb_tanks", &settings_.nb_tanks); - constructor.Bind("spawn_side", &settings_.spawn_side); - constructor.Bind("controller", &controller_display_); - constructor.Bind("display", &display_display_); - constructor.Bind("shadow_quality", &shadow_display_); - constructor.Bind("msaa", &settings_.msaa_samples); - constructor.Bind("gpus", &gpu_names_); - constructor.Bind("selected_window_gpu", &selected_window_gpu_); - constructor.Bind("selected_vision_gpu", &selected_vision_gpu_); - constructor.Bind("window_gpu_env", &window_gpu_env_override_); - constructor.Bind("vision_gpu_env", &vision_gpu_env_override_); - constructor.Bind("sac_folder", &sac_folder_display_); - constructor.Bind("sac_status", &sac_status_); - constructor.Bind("sac_valid", &sac_valid_); - constructor.Bind("current_dir", ¤t_dir_display_); - constructor.Bind("entries", &entries_); - constructor.Bind("can_play", &can_play_); - constructor.Bind("score", &score_); - - constructor.BindEventCallback( - "play", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - if (can_play_) play_clicked_ = true; - }); - constructor.BindEventCallback( - "exit", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - quit_clicked_ = true; - }); - constructor.BindEventCallback( - "open_params", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - main_document_->Hide(); - params_document_->Show(); - }); - constructor.BindEventCallback( - "back", [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - close_params(); - }); - constructor.BindEventCallback( - "enter_dir", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - const auto index = static_cast(arguments[0].Get()); - if (index >= entries_.size()) return; - - const std::string &entry = entries_[index]; - current_dir_ = - entry == ".." ? current_dir_.parent_path() : current_dir_ / entry; - refresh_explorer(); - focus_explorer_pending_ = true; - }); - constructor.BindEventCallback( - "set_controller", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - controller_display_ = arguments[0].Get(); - settings_.controller_kind = controller_display_ == "gamepad" - ? ControllerKind::Gamepad - : ControllerKind::Keyboard; - model_handle_.DirtyVariable("controller"); - }); - constructor.BindEventCallback( - "set_display", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - display_display_ = arguments[0].Get(); - settings_.fullscreen = display_display_ == "fullscreen"; - // applied immediately; the window reports its new size - // through the resize callback (dp-ratio included) - window_->set_fullscreen(settings_.fullscreen); - model_handle_.DirtyVariable("display"); - }); - constructor.BindEventCallback( - "open_graphics", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - params_document_->Hide(); - graphics_document_->Show(); - }); - constructor.BindEventCallback( - "graphics_back", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - close_graphics(); - }); - constructor.BindEventCallback( - "set_shadow_quality", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - const auto quality = - shadow_quality_from_string(arguments[0].Get()); - if (!quality) return; - settings_.shadow_quality = *quality; - shadow_display_ = to_string(*quality); - model_handle_.DirtyVariable("shadow_quality"); - }); - constructor.BindEventCallback( - "set_msaa", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - const int samples = arguments[0].Get(); - if (samples != 1 && samples != 2 && samples != 4 && samples != 8) return; - settings_.msaa_samples = samples; - model_handle_.DirtyVariable("msaa"); - }); - constructor.BindEventCallback( - "set_window_gpu", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - select_gpu( - arguments[0].Get(), selected_window_gpu_, settings_.window_gpu, - "selected_window_gpu"); - }); - constructor.BindEventCallback( - "set_vision_gpu", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - select_gpu( - arguments[0].Get(), selected_vision_gpu_, settings_.vision_gpu, - "selected_vision_gpu"); - }); - constructor.BindEventCallback( - "open_controls", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - open_controls(); - }); - constructor.BindEventCallback( - "controls_back", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - close_controls(); - }); - constructor.BindEventCallback( - "capture_kb", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (!arguments.empty()) begin_capture(true, arguments[0].Get()); - }); - constructor.BindEventCallback( - "capture_gp", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (!arguments.empty()) begin_capture(false, arguments[0].Get()); - }); - constructor.BindEventCallback( - "select_pad", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &arguments) { - if (arguments.empty()) return; - const auto index = static_cast(arguments[0].Get()); - if (index >= gamepads_.size()) return; - - settings_.bindings.gamepad.device_guid = gamepads_[index].guid; - settings_.bindings.gamepad.device_name = gamepads_[index].name; - window_->select_gamepad(gamepads_[index].id); - selected_gamepad_ = static_cast(index); - model_handle_.DirtyVariable("selected_gamepad"); - }); - constructor.BindEventCallback( - "reset_bindings", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - end_capture(); - if (settings_.controller_kind == ControllerKind::Keyboard) - settings_.bindings.keyboard = {}; - else { - // the device choice is not a binding: keep it - auto gamepad = GamepadBindings{}; - gamepad.device_guid = std::move(settings_.bindings.gamepad.device_guid); - gamepad.device_name = std::move(settings_.bindings.gamepad.device_name); - settings_.bindings.gamepad = std::move(gamepad); - } - set_default_bind_status(); - rebuild_binding_rows(); - }); - constructor.BindEventCallback( - "select_folder", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - settings_.sac_folder = current_dir_; - validate_sac_folder(); - refresh_explorer(); - }); - - constructor.BindEventCallback( - "pause_continue", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - pending_pause_action_ = PauseAction::Continue; - }); - constructor.BindEventCallback( - "pause_main_menu", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - pending_pause_action_ = PauseAction::MainMenu; - }); - constructor.BindEventCallback( - "pause_exit", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - pending_pause_action_ = PauseAction::ExitGame; - }); - constructor.BindEventCallback( - "game_over_retry", - [this](Rml::DataModelHandle, Rml::Event &, const Rml::VariantList &) { - pending_pause_action_ = PauseAction::Retry; - }); - - model_handle_ = constructor.GetModelHandle(); - } - - void refresh_explorer() { - entries_.clear(); - if (current_dir_.has_parent_path() && current_dir_ != current_dir_.root_path()) - entries_.emplace_back(".."); - - std::error_code list_error; - for (const auto &entry: - std::filesystem::directory_iterator(current_dir_, list_error)) - if (std::error_code type_error; entry.is_directory(type_error)) - entries_.push_back(entry.path().filename().string()); - if (list_error) - std::cerr << "Cannot list " << current_dir_ << ": " << list_error.message() - << std::endl; - - // keep ".." pinned first, sort the actual directories - const auto first_dir = - entries_.begin() + (!entries_.empty() && entries_[0] == ".." ? 1 : 0); - std::sort(first_dir, entries_.end()); - - current_dir_display_ = current_dir_.string(); - sac_folder_display_ = settings_.sac_folder.string(); - - if (model_handle_) { - model_handle_.DirtyVariable("entries"); - model_handle_.DirtyVariable("current_dir"); - model_handle_.DirtyVariable("sac_folder"); - model_handle_.DirtyVariable("sac_status"); - model_handle_.DirtyVariable("sac_valid"); - model_handle_.DirtyVariable("can_play"); - } - } - - // runs the injected dry-run load and turns its outcome into the - // tri-state the parameters screen displays (nothing chosen yet / - // model loaded / error message); can_play_ follows the real load - void validate_sac_folder() { - if (settings_.sac_folder.empty()) { - sac_valid_ = false; - sac_status_ = ""; - } else { - const auto error = sac_validator_(settings_.sac_folder); - sac_valid_ = !error.has_value(); - sac_status_ = error.value_or("AI model loaded"); - } - can_play_ = sac_valid_; - } - - void close_params() const { - params_document_->Hide(); - main_document_->Show(); - } - - // ---- graphics page ------------------------------------------ - - void close_graphics() const { - graphics_document_->Hide(); - params_document_->Show(); - } - - // "Auto" rides index 0 of gpu_names_, the actual devices follow - void - select_gpu(const int index, int &selected, std::string &setting, const char *variable) { - if (index < 0 || index >= static_cast(gpu_names_.size())) return; - selected = index; - setting = index == 0 ? std::string() : std::string(gpu_names_[index]); - model_handle_.DirtyVariable(variable); - } - - // ---- controls page ------------------------------------------ - - void open_controls() { - params_document_->Hide(); - refresh_gamepad_list(); - set_default_bind_status(); - rebuild_binding_rows(); - controls_document_->Show(); - } - - void close_controls() { - end_capture(); - rebuild_binding_rows(); - controls_document_->Hide(); - params_document_->Show(); - } - - std::array *, NB_KB_SLOTS> kb_slots() { - auto &keyboard = settings_.bindings.keyboard; - return { - &keyboard.forward, &keyboard.backward, &keyboard.turn_left, - &keyboard.turn_right, &keyboard.fire}; - } - - // gamepad slots 1..5 (0 is the fire button) - std::array *, NB_GP_SLOTS - 1> gp_axis_slots() { - auto &gamepad = settings_.bindings.gamepad; - return { - &gamepad.steer, &gamepad.aim_x, &gamepad.aim_y, &gamepad.accelerate, - &gamepad.reverse}; - } - - Rml::String keyboard_slot_label(const std::optional &slot) const { - if (!slot) return "UNBOUND"; - if (const auto *button = std::get_if(&*slot)) - switch (*button) { - case controller::MouseButton::Left: return "MOUSE L"; - case controller::MouseButton::Right: return "MOUSE R"; - case controller::MouseButton::Middle: return "MOUSE M"; - } - // layout-aware label (Key::Q reads "A" on AZERTY) - const auto label = window_->key_label(std::get(*slot)); - return label.empty() ? "?" : Rml::String(label); - } - - static Rml::String - axis_slot_label(const std::optional &slot, const bool one_way) { - if (!slot) return "UNBOUND"; - Rml::String label = GAMEPAD_AXIS_LABELS[static_cast(slot->axis)]; - // a one-way action bound to a stick reads a single direction - const bool on_stick = slot->axis != GamepadAxis::LeftTrigger - && slot->axis != GamepadAxis::RightTrigger; - if (one_way && on_stick) label += slot->sign > 0.f ? "+" : "-"; - return label; - } - - void rebuild_binding_rows() { - kb_rows_.resize(NB_KB_SLOTS); - const auto keyboard_slots = kb_slots(); - for (int i = 0; i < NB_KB_SLOTS; i++) { - const bool listening = capture_keyboard_page_ && capture_slot_ == i; - kb_rows_[i] = { - .label = KB_SLOT_LABELS[i], - .binding = - listening ? "PRESS A KEY..." : keyboard_slot_label(*keyboard_slots[i]), - .listening = listening, - .bound = keyboard_slots[i]->has_value()}; - } - - gp_rows_.resize(NB_GP_SLOTS); - const auto axis_slots = gp_axis_slots(); - for (int i = 0; i < NB_GP_SLOTS; i++) { - const bool listening = !capture_keyboard_page_ && capture_slot_ == i; - const bool fire_slot = i == 0; - Rml::String binding; - if (listening) binding = fire_slot ? "PRESS A BUTTON..." : "MOVE AN AXIS..."; - else if (fire_slot) - binding = settings_.bindings.gamepad.fire - ? controller::to_string(*settings_.bindings.gamepad.fire) - : "UNBOUND"; - else binding = axis_slot_label(*axis_slots[i - 1], gp_slot_is_one_way(i)); - gp_rows_[i] = { - .label = GP_SLOT_LABELS[i], - .binding = std::move(binding), - .listening = listening, - .bound = fire_slot ? settings_.bindings.gamepad.fire.has_value() - : axis_slots[i - 1]->has_value()}; - } - - if (model_handle_) { - model_handle_.DirtyVariable("kb_rows"); - model_handle_.DirtyVariable("gp_rows"); - model_handle_.DirtyVariable("bind_status"); - model_handle_.DirtyVariable("bind_warning"); - } - } - - void set_default_bind_status() { - bind_warning_ = false; - bind_status_ = settings_.controller_kind == ControllerKind::Keyboard - ? "Click a slot, then press the new key or mouse button. " - "Escape cancels." - : "Click a slot, then press a button or move an axis. " - "Escape or 15s of inactivity cancels — Start stays the " - "pause toggle."; - } - - void begin_capture(const bool keyboard_page, const int slot) { - const int nb_slots = keyboard_page ? NB_KB_SLOTS : NB_GP_SLOTS; - if (slot < 0 || slot >= nb_slots) return; - - capture_keyboard_page_ = keyboard_page; - capture_slot_ = slot; - capture_deadline_ = std::chrono::steady_clock::now() + CAPTURE_TIMEOUT; - // every axis must return to rest once before it can bind (the - // stick may still be deflected from navigating the menu) - axis_armed_.fill(false); - set_default_bind_status(); - - input_adapter_->set_capture_sink( - [this](const RawMenuInput &input) { on_capture_input(input); }); - rebuild_binding_rows(); - } - - void end_capture() { - capture_slot_ = -1; - input_adapter_->set_capture_sink(nullptr); - } - - void on_capture_input(const RawMenuInput &input) { - // Escape is the only cancel input (every pad button stays - // bindable); a pad-only player relies on the capture timeout - if (const auto *key = std::get_if(&input); - key != nullptr && *key == controller::Key::Escape) { - end_capture(); - rebuild_binding_rows(); - return; - } - - if (capture_keyboard_page_) { - if (const auto *key = std::get_if(&input)) - assign_keyboard(KeyboardBinding(*key)); - else if (const auto *button = std::get_if(&input)) - assign_keyboard(KeyboardBinding(*button)); - // pad input has no meaning on the keyboard page - return; - } - - if (capture_slot_ == 0) { - // Start stays the in-game pause toggle, never a binding - if (const auto *button = std::get_if(&input); - button != nullptr && *button != controller::GamepadButton::Start) - assign_gamepad_fire(*button); - return; - } - - if (const auto *motion = std::get_if>(&input)) { - const auto &[axis, value] = *motion; - auto &armed = axis_armed_[static_cast(axis)]; - if (std::abs(value) < CAPTURE_REST) armed = true; - else if (armed && std::abs(value) > CAPTURE_ENGAGE) - assign_gamepad_axis(axis, value); - } - } - - void unbind_conflict(const Rml::String &new_label, const Rml::String &old_label) { - bind_warning_ = true; - bind_status_ = new_label + " was bound to " + old_label + " — " + old_label - + " is now unbound."; - } - - void assign_keyboard(const KeyboardBinding &binding) { - const auto slots = kb_slots(); - *slots[capture_slot_] = binding; - - for (int i = 0; i < NB_KB_SLOTS; i++) - if (i != capture_slot_ && *slots[i] == binding) { - *slots[i] = std::nullopt; - unbind_conflict(keyboard_slot_label(binding), KB_SLOT_LABELS[i]); - } - - end_capture(); - rebuild_binding_rows(); - } - - void assign_gamepad_fire(const controller::GamepadButton button) { - // single button slot: no conflict possible - settings_.bindings.gamepad.fire = button; - end_capture(); - rebuild_binding_rows(); - } - - void assign_gamepad_axis(const GamepadAxis axis, const double value) { - const bool one_way = gp_slot_is_one_way(capture_slot_); - const auto slots = gp_axis_slots(); - const GamepadAxisBinding binding{ - .axis = axis, .sign = one_way && value < 0. ? -1.f : 1.f}; - *slots[capture_slot_ - 1] = binding; - - // two slots clash when they read the same range of an axis: a - // two-way slot owns the whole axis, one-way slots only their - // captured side - for (int slot = 1; slot < NB_GP_SLOTS; slot++) { - if (slot == capture_slot_) continue; - auto &other = *slots[slot - 1]; - if (!other || other->axis != axis) continue; - if (one_way && gp_slot_is_one_way(slot) && other->sign != binding.sign) - continue; - other = std::nullopt; - unbind_conflict(axis_slot_label(binding, one_way), GP_SLOT_LABELS[slot]); - } - - end_capture(); - rebuild_binding_rows(); - } - - void refresh_gamepad_list() { - auto gamepads = window_->list_gamepads(); - const bool changed = - gamepads.size() != gamepads_.size() - || !std::equal( - gamepads.begin(), gamepads.end(), gamepads_.begin(), - [](const view::GamepadInfo &a, const view::GamepadInfo &b) { - return a.id == b.id && a.guid == b.guid && a.name == b.name; - }); - if (!changed) return; - - gamepads_ = std::move(gamepads); - gamepad_names_.clear(); - for (const auto &pad: gamepads_) gamepad_names_.emplace_back(pad.name); - - // the preferred pad when connected, else the first one (which - // is what the window falls back to) - selected_gamepad_ = gamepads_.empty() ? -1 : 0; - for (size_t i = 0; i < gamepads_.size(); i++) - if (!settings_.bindings.gamepad.device_guid.empty() - && gamepads_[i].guid == settings_.bindings.gamepad.device_guid) { - selected_gamepad_ = static_cast(i); - break; - } - window_->select_gamepad( - selected_gamepad_ >= 0 ? gamepads_[selected_gamepad_].id : -1); - - if (model_handle_) { - model_handle_.DirtyVariable("gamepads"); - model_handle_.DirtyVariable("selected_gamepad"); - } - } - - void focus_first_entry() const { - const Rml::Element *list = params_document_->GetElementById("file-list"); - if (list == nullptr) return; - const auto entries = ExplorerNavListener::visible_file_entries(list); - if (entries.empty()) return; - if (entries.front()->Focus(true)) - entries.front()->ScrollIntoView(Rml::ScrollAlignment::Nearest); - } - - std::shared_ptr backend_; - std::shared_ptr window_; - - GameSettings settings_; - int width_; - int height_; - - MenuSystemInterface system_interface_; - ReaderBackedFileInterface file_interface_; - // unique_ptr: created after Rml::Initialise(), and member - // destruction keeps it alive until after Rml::Shutdown() as - // RmlUi requires of registered instancers - - std::unique_ptr cursor_ring_instancer_; - std::unique_ptr hit_marker_instancer_; - std::unique_ptr reticle_instancer_; - std::unique_ptr damage_arc_instancer_; - std::vector font_buffers_; - - Rml::Context *context_ = nullptr; - Rml::ElementDocument *main_document_ = nullptr; - Rml::ElementDocument *params_document_ = nullptr; - Rml::ElementDocument *controls_document_ = nullptr; - Rml::ElementDocument *graphics_document_ = nullptr; - Rml::ElementDocument *pause_document_ = nullptr; - Rml::ElementDocument *game_over_document_ = nullptr; - Rml::ElementDocument *hud_document_ = nullptr; - Rml::Element *reticle_ = nullptr; - Rml::Element *hit_marker_ = nullptr; - std::vector damage_arcs_; - std::size_t next_damage_arc_ = 0; - static constexpr float DAMAGE_ARC_FADE_SECONDS = 0.8f; - static constexpr float HIT_MARKER_FADE_SECONDS = 0.45f; - static constexpr float HIT_MARKER_END_SCALE = 1.3f; - static constexpr float KILL_MARKER_FADE_SECONDS = 0.6f; - static constexpr float KILL_MARKER_START_SCALE = 1.1f; - static constexpr float KILL_MARKER_END_SCALE = 1.55f; - Rml::DataModelHandle model_handle_; - - std::shared_ptr input_adapter_; - // removed from the document when Rml::Shutdown() destroys it in - // the destructor body, before the members are torn down - ExplorerNavListener explorer_nav_listener_; - bool focus_explorer_pending_ = false; - - std::filesystem::path current_dir_; - SacFolderValidator sac_validator_; - Rml::String controller_display_; - Rml::String display_display_; - // graphics page state - Rml::String shadow_display_; - std::vector gpu_names_; - int selected_window_gpu_ = 0; - int selected_vision_gpu_ = 0; - bool window_gpu_env_override_ = false; - bool vision_gpu_env_override_ = false; - Rml::String current_dir_display_; - Rml::String sac_folder_display_; - // empty while no folder is chosen; otherwise success or error text - Rml::String sac_status_; - std::vector entries_; - bool sac_valid_ = false; - bool can_play_ = false; - - // controls page state - std::vector kb_rows_; - std::vector gp_rows_; - std::vector gamepads_; - std::vector gamepad_names_; - // index into gamepads_ of the pad feeding the game, -1 when none - // is connected - int selected_gamepad_ = -1; - Rml::String bind_status_; - bool bind_warning_ = false; - // slot being captured (-1 = idle) and which page it belongs to - int capture_slot_ = -1; - bool capture_keyboard_page_ = false; - std::chrono::steady_clock::time_point capture_deadline_; - std::array axis_armed_{}; - // score shown by the pause and game-over popups - int score_ = 0; - bool play_clicked_ = false; - bool quit_clicked_ = false; - PauseAction pending_pause_action_ = PauseAction::None; - }; - - }// namespace - - std::unique_ptr make_gui( - const std::shared_ptr &backend, - const std::shared_ptr &asset_reader, - const GameSettings &initial_settings, const std::vector &gpus, - const int window_width, const int window_height, SacFolderValidator sac_validator) { - return std::make_unique( - backend, asset_reader, initial_settings, gpus, window_width, window_height, - std::move(sac_validator)); - } - -}// namespace arenai::desktop::gui diff --git a/arenai_desktop/src/main.cpp b/arenai_desktop/src/main.cpp index 8e44c5f6..ff9ad56d 100644 --- a/arenai_desktop/src/main.cpp +++ b/arenai_desktop/src/main.cpp @@ -3,7 +3,6 @@ // #include -#include #include @@ -29,63 +28,26 @@ static std::filesystem::path executable_dir(const char *argv0) { return std::filesystem::absolute(argv0).parent_path(); } -static void trim_inplace(std::string &s) { - auto not_space = [](const unsigned char c) { return !std::isspace(c); }; - s.erase(s.begin(), std::ranges::find_if(s, not_space)); - s.erase(std::find_if(s.rbegin(), s.rend(), not_space).base(), s.end()); -} - -static std::tuple parse_key_value(const std::string &value) { - const std::regex regex_match(R"(^ *([^=]+)=\"?([^"]+)\"? *$)"); - - std::smatch match; - - if (!std::regex_match(value, match, regex_match)) - throw std::invalid_argument("invalid pairs format, usage : key=value"); - - std::string s1 = match[1].str(); - std::string s2 = match[2].str(); - trim_inplace(s2); - - return {s1, s2}; -} - -typedef std::vector> hyper_params_vector; - int main(const int argc, char **argv) { argparse::ArgumentParser parser("arenai game"); - // Game options (nb_tanks, spawn_side and controller_kind are set in the menu) - parser.add_argument("--wanted_frequency").scan<'g', float>().default_value(1.f / 30.f); - parser.add_argument("--vision_height").scan<'i', int>().default_value(128); - parser.add_argument("--vision_width").scan<'i', int>().default_value(256); + // everything about the AI model (vision size, frequency, hyper-parameters) + // comes from the config.json selected in the menu; the command line only + // keeps what belongs to the player's machine parser.add_argument("--window_width").scan<'i', int>().default_value(1920); parser.add_argument("--window_height").scan<'i', int>().default_value(1080); - - // Model options - parser.add_argument("--hp", "--hyper_parameters") - .action(parse_key_value) - .append() - .default_value({}); parser.add_argument("--cuda").implicit_value(true).default_value(false); parser.parse_args(argc, argv); - std::map hyper_params; - for (const auto &[key, value]: parser.get("--hyper_parameters")) - hyper_params[key] = value; - const auto resources_folder = executable_dir(argv[0]) / "resources"; run_gui( - {.wanted_frequency = parser.get("--wanted_frequency"), - .window_width = parser.get("--window_width"), + {.window_width = parser.get("--window_width"), .window_height = parser.get("--window_height"), .resources_folder = resources_folder}, - {.vision_height = parser.get("--vision_height"), - .vision_width = parser.get("--vision_width"), - .hyper_parameters = hyper_params, - .state_dict_folder = resources_folder / "trained_models" / "ppo_run_386_save_66", + {.state_dict_folder = resources_folder / "trained_models" / "ppo_liquid_409" / "save_57", + .config_json = resources_folder / "trained_models" / "ppo_train_409" / "config.json", .cuda = parser.get("--cuda")}); return 0; diff --git a/arenai_desktop/tests/include/arenai_desktop_tests/interceptor_agent.h b/arenai_desktop/tests/include/arenai_desktop_tests/interceptor_agent.h index addbd5c6..0c824f7d 100644 --- a/arenai_desktop/tests/include/arenai_desktop_tests/interceptor_agent.h +++ b/arenai_desktop/tests/include/arenai_desktop_tests/interceptor_agent.h @@ -10,7 +10,7 @@ #include // Records every batch of states the game loop hands to the agent and answers -// with neutral actions: what act() received is exactly what the SAC agent +// with neutral actions: what act() received is exactly what the enemy agent // would have seen in run_game(). class InterceptorAgent final : public agent::AbstractAgent { public: diff --git a/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_0.json b/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_0.json index bb387f12..340ebc40 100644 --- a/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_0.json +++ b/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_0.json @@ -1 +1 @@ -[97,92,88,87,80,76,72,75,73,71,74,73,81,81,77,77,88,78,58,90,83,83,81,79,75,74,72,76,84,84,70,72,89,89,86,78,75,77,76,72,71,76,76,76,79,79,79,71,88,77,76,76,76,75,71,70,80,84,84,85,77,77,81,73,76,75,75,75,75,72,85,85,82,82,84,84,77,81,81,81,73,73,74,79,83,85,85,85,82,82,82,84,80,80,81,81,72,79,79,83,83,83,85,85,82,83,84,84,80,80,80,80,79,79,79,83,83,83,83,88,49,83,83,83,80,80,80,80,79,79,83,83,83,150,41,48,47,64,83,83,80,80,80,69,81,79,82,77,77,151,158,158,158,157,158,83,80,80,69,69,81,81,77,77,77,161,154,154,154,155,162,88,77,69,69,69,81,81,81,77,26,161,154,154,154,161,161,202,77,77,77,77,81,81,81,81,21,113,112,67,67,67,113,137,77,77,77,77,81,81,81,81,16,16,66,66,66,66,66,77,77,77,77,77,81,81,62,68,68,68,68,66,66,66,77,77,77,77,77,61,62,68,68,68,68,68,68,68,66,66,77,77,61,61,61,61,111,107,104,104,97,94,90,93,91,88,93,92,103,103,98,98,103,93,71,112,104,103,101,98,94,93,91,97,107,107,89,91,110,110,107,97,94,96,96,91,91,96,96,96,101,101,101,91,110,97,96,96,95,95,90,90,102,107,108,108,98,98,104,93,95,95,95,95,95,92,109,109,105,105,107,108,98,103,103,103,92,92,94,101,106,109,109,109,105,105,105,107,103,103,103,103,92,101,101,106,106,106,109,109,105,107,107,107,103,103,103,103,101,101,101,106,106,106,106,112,150,107,107,107,102,103,103,103,101,101,106,106,106,125,127,147,144,196,107,107,102,102,103,89,103,101,106,99,99,126,132,132,132,131,131,107,102,102,89,89,103,105,99,99,99,134,129,129,129,129,134,112,99,89,89,89,105,105,105,99,223,134,129,129,129,134,134,170,99,99,100,100,105,105,105,105,175,95,95,86,86,86,96,118,99,99,99,100,105,105,105,105,126,127,86,86,86,86,86,99,99,99,99,99,105,105,80,89,89,89,89,86,86,86,99,99,99,99,99,80,80,89,89,89,89,89,89,89,86,86,99,99,79,79,79,79,160,160,163,167,162,165,160,166,164,161,168,169,183,183,177,178,160,154,134,188,178,178,177,174,170,171,169,176,189,189,167,170,187,187,183,172,169,172,173,168,168,176,176,176,183,183,183,170,187,171,171,172,173,173,168,168,184,191,192,192,180,180,187,174,171,172,172,173,174,172,194,194,189,189,192,192,180,187,187,187,169,170,173,184,191,194,194,194,189,189,189,192,186,187,187,187,170,184,184,191,191,191,194,194,189,192,192,192,186,187,187,187,184,184,185,191,191,191,191,199,125,192,192,192,187,187,187,187,184,184,191,191,191,169,113,123,122,148,192,192,187,187,187,170,188,185,191,182,182,170,175,175,175,174,175,192,187,187,170,170,188,190,182,182,182,177,172,172,172,173,178,200,183,170,170,170,190,190,190,182,124,177,172,172,172,177,178,177,184,184,184,184,190,190,190,190,106,142,142,167,167,167,142,140,184,184,184,184,190,190,190,190,88,88,167,167,167,167,168,184,184,184,184,184,190,190,160,170,171,171,171,167,167,168,184,184,184,184,184,159,160,171,171,171,171,171,171,171,168,168,184,184,159,159,159,159] +[60,61,61,65,64,67,65,63,57,58,59,60,59,57,56,53,69,66,63,60,60,59,60,56,59,56,58,58,61,58,55,51,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,41,41,40,40,40,40,40,40,40,40,40,40,40,43,50,40,41,41,41,50,41,40,40,40,40,40,40,40,33,40,40,40,42,42,42,50,42,42,42,41,40,40,40,40,33,40,40,40,42,42,42,24,24,24,24,24,24,24,24,24,41,41,41,41,42,42,39,24,24,24,24,24,24,24,24,24,42,41,41,41,42,42,42,42,24,24,24,24,24,24,24,24,42,42,42,42,124,119,125,138,24,24,24,24,24,24,24,24,125,122,126,138,114,118,137,139,24,24,24,24,24,24,24,24,130,132,132,134,127,155,155,141,24,24,24,24,24,24,24,24,142,147,157,154,147,136,133,135,24,24,24,24,24,24,24,132,133,137,142,148,67,68,70,75,73,76,75,72,64,67,69,70,68,66,65,61,76,73,73,68,69,70,68,66,66,65,67,70,73,65,66,61,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,54,55,54,54,54,53,53,53,53,53,53,53,53,137,161,53,55,55,55,160,55,54,54,54,54,54,54,53,108,53,53,53,55,55,55,160,55,55,55,54,54,54,54,54,108,54,54,54,56,56,56,109,109,109,109,109,109,109,109,109,54,54,54,54,56,56,127,109,109,109,109,109,109,109,109,109,56,55,55,55,56,56,56,56,109,109,109,109,109,109,109,109,56,56,56,56,148,142,144,157,109,109,109,109,109,109,109,109,151,148,149,160,137,137,156,159,109,109,109,109,109,109,109,109,158,157,158,160,149,172,174,159,109,109,109,109,109,109,109,109,165,169,175,176,165,154,151,151,109,109,109,109,109,109,109,154,156,159,164,170,83,84,85,86,88,91,88,87,80,81,83,82,83,81,80,76,92,89,84,81,83,84,81,78,82,82,79,79,87,81,78,76,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,128,129,127,127,127,127,127,127,127,127,127,127,127,118,129,127,129,129,129,129,129,127,127,127,127,127,127,127,103,127,127,127,129,129,129,129,129,129,129,127,127,127,127,127,103,127,127,127,129,129,129,114,114,114,114,114,114,114,114,114,128,127,127,127,129,129,112,114,114,114,114,114,114,114,114,114,129,128,128,128,128,128,128,129,114,114,114,114,114,114,114,114,129,129,129,129,163,156,158,171,114,114,114,114,114,114,114,114,164,161,163,174,151,152,170,170,114,114,114,114,114,114,114,114,172,171,171,173,163,182,181,171,114,114,114,114,114,114,114,114,179,183,189,189,177,166,165,164,114,114,114,114,114,114,114,168,172,174,178,183] diff --git a/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_1.json b/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_1.json index e351a9cd..62efc607 100644 --- a/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_1.json +++ b/arenai_desktop/tests/resources/golden_images/golden_desktop_reset_tank_1.json @@ -1 +1 @@ -[48,47,47,47,42,42,42,42,41,41,41,41,42,42,42,42,43,47,47,47,43,42,42,42,41,41,41,43,42,42,42,42,43,47,47,47,42,42,42,41,41,41,43,42,42,42,42,42,43,47,47,47,42,42,42,41,43,43,42,42,42,42,42,42,44,47,47,47,47,42,42,44,47,47,42,42,42,42,42,42,44,47,47,47,46,42,44,44,47,47,47,47,46,42,42,42,44,47,47,47,51,44,44,44,47,47,46,46,46,46,46,42,42,47,51,51,49,48,44,44,47,46,46,46,46,46,46,46,42,47,51,50,48,48,23,28,28,37,46,46,46,46,46,46,44,46,50,50,147,147,148,147,147,147,147,46,46,46,46,46,44,46,127,90,90,146,128,147,147,175,123,135,46,46,46,46,46,46,125,83,83,83,90,126,123,201,210,121,46,46,46,46,44,125,128,84,83,83,126,128,191,156,178,107,46,46,46,46,38,43,78,84,41,45,126,152,166,167,178,41,46,46,46,46,39,43,43,41,41,45,126,119,166,42,41,41,41,46,46,46,41,38,43,41,41,41,45,103,166,42,41,41,41,41,46,46,11,10,9,9,8,8,7,7,7,6,6,6,6,5,5,5,10,10,9,9,8,8,7,7,7,6,6,6,5,5,5,5,10,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,10,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,10,9,9,8,8,7,7,6,6,6,5,5,5,5,5,4,10,9,9,8,8,7,7,6,6,5,5,5,5,5,4,4,10,9,9,8,8,7,6,6,5,5,5,5,5,5,4,4,10,9,8,8,7,7,6,6,5,5,5,5,5,4,4,4,10,9,8,7,7,6,26,31,31,41,5,5,4,4,4,4,9,8,7,7,182,182,184,183,183,183,183,4,4,4,4,4,9,8,158,115,115,181,159,183,183,216,154,30,4,4,4,4,8,7,157,107,107,107,115,157,154,50,52,27,4,4,4,4,8,157,160,107,107,107,157,160,140,40,45,25,4,4,4,4,8,7,127,136,6,5,157,113,123,123,45,4,4,4,4,4,9,7,7,6,6,5,157,89,122,4,4,4,4,4,4,4,9,7,7,6,6,5,5,78,123,4,4,4,4,4,4,4,93,93,93,92,87,88,87,87,87,87,87,86,88,88,88,88,88,93,93,92,88,87,87,87,87,87,86,88,88,88,88,88,88,93,92,93,88,87,87,87,87,87,88,88,88,88,88,88,88,93,93,93,88,87,87,87,89,88,88,88,88,88,88,88,89,93,93,93,93,87,87,90,94,94,88,88,88,88,88,88,89,93,93,93,93,87,90,90,94,93,93,93,93,88,88,88,89,93,93,93,98,90,90,90,93,93,93,93,93,93,93,88,87,93,98,98,95,95,90,90,93,93,93,93,93,93,93,93,87,93,98,97,95,95,31,35,35,41,93,93,93,93,93,93,90,92,97,97,7,7,7,7,7,7,7,93,93,93,93,93,89,92,7,6,6,6,6,6,6,7,7,88,93,93,93,93,92,92,7,6,6,6,6,6,6,23,24,83,93,93,93,93,89,7,7,6,6,6,6,6,175,20,22,77,93,93,93,93,83,89,196,205,87,92,6,152,161,161,22,87,93,93,93,93,83,89,89,87,87,92,6,133,161,88,87,87,87,93,93,93,86,83,89,87,87,87,92,123,161,88,87,87,87,87,93,93] +[76,76,72,72,72,70,72,71,72,72,76,77,77,77,79,78,74,74,69,69,69,71,71,73,73,75,75,76,77,78,78,78,72,70,70,68,67,69,69,75,73,73,75,75,73,76,77,78,68,68,68,66,68,69,76,76,75,70,72,72,72,72,72,71,68,68,68,68,68,72,72,72,74,71,69,69,73,72,72,72,69,67,67,67,67,65,65,63,72,73,72,72,73,73,73,72,68,69,67,67,58,57,57,74,73,73,73,74,74,75,73,71,68,68,68,57,57,57,57,73,73,74,74,73,74,75,75,75,68,68,68,53,53,52,73,73,73,73,74,74,70,71,74,74,68,68,53,53,52,52,69,72,72,72,73,74,70,70,70,70,70,53,52,52,52,69,69,69,72,72,72,72,63,67,67,67,70,52,52,52,52,69,23,30,30,36,72,72,63,63,63,63,52,52,158,140,140,140,141,125,125,125,72,72,63,63,63,63,52,52,124,118,139,166,134,123,135,125,125,69,63,63,63,63,52,51,100,119,126,126,164,202,124,125,115,140,148,52,52,55,51,100,142,189,187,125,179,174,211,82,116,116,52,52,52,52,25,25,24,24,23,22,21,21,20,20,20,20,20,21,21,22,22,22,21,20,19,19,19,18,18,18,18,18,19,19,19,20,20,20,19,19,18,17,16,16,17,16,16,17,17,17,18,18,17,17,16,15,15,14,14,15,15,15,15,15,15,15,15,15,14,14,14,13,13,12,12,13,13,13,13,13,13,13,13,13,13,12,12,12,11,11,11,10,10,11,11,11,12,12,12,12,11,11,11,10,10,9,9,8,9,9,10,10,10,11,11,11,10,10,10,9,8,8,7,7,8,8,8,9,9,10,10,10,9,9,9,8,7,7,7,7,7,7,7,8,8,8,9,9,9,8,8,7,7,6,6,6,6,7,7,7,7,7,7,8,8,8,7,7,6,6,6,6,6,6,6,6,6,6,7,7,8,7,7,6,6,6,26,33,33,41,6,6,6,6,6,6,7,6,196,174,175,174,176,156,156,156,5,6,6,6,6,6,6,6,156,148,173,205,168,154,168,156,156,5,5,5,5,5,6,5,126,149,158,158,42,51,156,156,26,30,184,5,5,5,5,126,106,139,137,156,45,44,53,19,26,26,4,4,4,4,123,124,120,120,120,117,120,119,120,120,125,125,126,126,128,127,122,123,116,116,116,119,119,121,122,124,125,125,127,127,128,128,120,118,118,115,115,118,118,125,122,122,124,124,122,125,127,127,116,116,116,115,117,117,126,126,125,119,121,121,121,122,122,120,116,117,117,117,117,122,122,122,124,121,118,118,123,121,122,122,118,116,116,116,116,114,114,112,123,123,122,122,123,123,123,121,118,118,116,116,105,105,105,125,124,124,124,124,125,126,123,121,118,118,118,105,105,105,105,124,124,125,125,124,125,126,125,126,118,118,118,100,100,100,124,124,124,124,125,125,120,122,126,125,118,118,100,100,100,100,120,123,123,124,124,125,120,120,121,121,120,100,100,100,100,120,120,120,123,123,124,124,113,117,117,117,120,100,100,99,99,120,31,36,36,41,124,124,113,113,113,113,100,100,7,7,7,7,7,7,7,7,123,124,113,113,113,113,99,99,7,6,7,7,7,6,7,7,7,120,112,112,113,113,99,99,6,6,6,6,21,23,7,7,80,90,7,100,100,103,99,6,147,175,173,7,21,21,24,66,80,80,99,99,99,99] diff --git a/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_0.json b/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_0.json index 3a7a886a..b1d2781f 100644 --- a/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_0.json +++ b/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_0.json @@ -1 +1 @@ -[93,93,93,94,95,95,95,161,153,147,148,62,71,124,132,137,93,93,93,93,94,94,95,95,170,150,62,68,71,67,132,150,93,93,93,93,93,94,94,95,95,57,68,70,69,60,116,139,92,92,92,92,94,94,94,94,94,91,41,69,63,49,105,120,92,92,92,94,94,94,94,94,94,91,91,93,55,93,97,99,92,92,94,94,94,94,94,93,93,93,91,93,93,95,86,100,92,94,94,94,94,94,94,93,93,93,92,94,94,94,88,88,93,93,93,94,94,90,93,94,46,94,94,94,94,93,94,88,90,90,90,90,90,90,38,48,47,63,94,94,94,93,93,94,90,90,90,90,90,151,152,152,152,152,152,94,94,93,93,90,86,90,90,90,154,110,110,110,110,108,154,137,94,93,90,90,86,90,90,25,154,98,98,98,98,98,154,159,85,85,90,90,86,86,90,22,154,98,98,98,98,98,154,128,84,85,85,90,86,86,86,18,154,121,76,76,76,76,154,105,81,84,85,85,74,74,74,73,14,76,76,76,76,76,73,126,81,81,81,85,74,74,74,73,73,73,73,73,76,76,76,76,76,81,81,81,118,118,118,120,120,120,120,181,176,171,170,14,14,155,163,165,118,118,118,118,120,120,120,120,188,172,14,14,14,14,163,176,118,118,118,118,118,120,120,120,120,13,13,13,14,14,147,166,117,117,117,118,120,120,120,120,119,116,12,13,13,14,136,151,117,117,117,120,120,120,120,120,119,115,115,118,13,129,133,137,117,117,120,120,120,120,120,119,119,119,115,118,118,131,118,137,117,119,120,120,120,120,119,119,119,119,117,120,120,119,122,120,119,119,119,119,120,115,119,119,140,120,120,120,120,119,119,120,115,115,115,115,115,115,117,147,144,194,119,119,120,119,119,119,115,115,115,115,115,126,127,127,127,127,127,119,119,119,119,115,110,115,115,115,128,93,93,93,93,92,129,115,119,118,114,115,110,115,115,51,128,84,83,83,83,83,129,23,108,108,114,114,110,110,115,46,128,84,83,83,83,83,128,120,108,108,108,114,110,110,110,36,128,102,98,98,98,98,128,100,104,108,108,108,96,96,96,95,27,98,98,98,98,98,71,119,105,105,105,108,96,96,96,95,95,95,95,95,98,98,98,99,99,105,105,105,205,205,205,207,207,207,207,195,189,185,184,198,217,173,181,183,205,205,205,205,207,207,207,207,198,186,197,212,217,210,181,191,205,205,205,206,206,207,207,207,207,188,212,216,215,193,165,183,205,205,205,205,208,208,208,208,206,201,154,214,202,169,159,171,205,205,205,208,208,208,208,208,206,202,201,204,184,153,157,160,205,205,208,208,208,208,208,206,206,206,201,204,204,156,143,160,205,208,208,208,208,208,208,206,206,206,204,207,207,206,149,145,208,208,208,208,208,202,207,208,120,208,207,207,207,206,206,145,202,202,203,203,203,203,108,124,122,147,208,207,207,206,206,206,203,203,203,203,203,170,171,171,171,171,171,208,207,206,206,201,197,203,203,203,172,140,140,140,140,138,172,160,208,206,201,201,197,203,203,88,172,131,131,131,131,131,172,217,194,194,201,201,197,197,203,83,172,131,131,131,131,131,172,126,194,194,194,201,198,197,197,71,172,148,182,182,182,183,172,111,190,195,195,195,180,180,180,178,60,183,183,183,183,183,91,125,190,190,190,195,181,181,180,179,178,178,178,178,183,183,183,183,183,191,191,190] +[60,59,52,51,52,59,66,70,71,71,67,71,74,80,78,80,65,63,51,49,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,39,40,40,40,51,40,40,40,40,40,40,40,40,40,40,40,40,40,40,40,51,40,40,40,40,40,40,40,40,40,40,40,40,40,40,40,51,51,40,40,40,40,40,40,40,40,40,66,67,67,68,33,51,51,41,41,41,41,41,41,41,41,41,68,67,67,65,33,51,51,41,41,41,41,41,27,23,41,42,65,65,65,65,33,51,51,51,42,42,43,43,43,25,43,141,61,61,61,33,33,51,51,51,43,43,43,145,145,25,135,143,60,60,61,33,33,51,51,51,132,136,132,137,153,122,145,144,36,55,55,33,33,132,51,51,156,160,129,140,120,123,124,123,36,54,54,33,33,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,24,72,67,64,61,62,70,75,80,83,82,79,85,88,87,87,89,80,73,63,60,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,162,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,162,53,53,53,53,53,53,53,53,53,53,53,53,53,53,53,162,162,54,54,54,54,54,54,54,54,54,122,123,123,124,108,162,162,54,54,54,54,54,54,54,54,54,124,124,124,120,108,162,162,55,55,55,55,55,120,103,55,55,120,120,120,120,108,162,162,162,56,56,56,56,57,112,57,164,112,112,112,108,108,162,162,162,57,57,57,165,168,112,163,169,111,111,113,108,108,162,162,162,156,160,160,164,175,154,174,174,69,103,103,108,108,154,162,162,174,180,158,167,152,156,157,159,68,100,100,108,108,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,109,86,80,77,74,74,84,90,92,94,97,89,98,103,106,102,105,92,83,77,72,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,130,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,130,127,127,127,127,127,127,127,127,127,127,127,127,127,127,127,130,130,127,127,127,127,127,127,127,127,127,142,143,143,143,103,130,130,128,128,128,128,128,128,128,128,128,143,143,143,140,103,130,130,128,128,128,128,128,121,111,128,128,140,140,140,140,103,130,130,130,129,129,130,130,130,116,130,178,134,134,134,103,103,130,130,130,130,130,130,179,182,116,174,182,133,134,135,103,103,130,130,130,170,174,175,177,188,166,185,185,100,127,127,103,103,168,130,130,188,191,174,179,167,171,172,173,100,124,124,103,103,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114,114] diff --git a/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_1.json b/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_1.json index 8dddac6b..246284ff 100644 --- a/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_1.json +++ b/arenai_desktop/tests/resources/golden_images/golden_desktop_step30_tank_1.json @@ -1 +1 @@ -[56,56,56,54,53,53,53,53,52,52,52,50,50,49,49,49,55,54,54,53,53,53,53,52,50,50,50,50,49,49,49,48,54,54,53,53,53,53,52,50,50,50,50,49,49,48,48,48,54,55,53,53,52,94,100,77,165,50,49,48,50,50,50,50,55,55,53,119,117,117,118,155,146,194,50,50,50,50,50,49,55,100,149,118,117,117,117,131,124,114,50,50,49,49,49,49,53,81,97,118,82,82,82,82,107,103,176,49,49,49,49,48,53,67,99,149,82,82,82,82,80,160,149,49,49,52,50,50,52,57,54,87,149,82,82,82,82,134,127,208,52,52,52,50,52,52,52,21,19,148,82,82,82,82,82,110,52,52,51,51,52,52,52,18,17,24,148,82,82,82,82,86,137,51,51,51,52,52,48,48,48,74,108,148,51,51,51,51,51,51,51,51,52,48,48,48,48,48,51,51,51,51,51,51,51,51,51,51,51,50,50,50,50,53,49,49,48,48,48,48,48,48,48,48,50,50,50,50,53,53,49,48,48,48,48,48,48,48,48,48,50,50,50,53,53,53,53,48,48,48,48,48,48,48,48,48,12,11,10,10,9,9,8,8,7,7,7,7,6,6,6,6,11,10,10,9,9,8,8,7,7,7,7,6,6,6,6,6,10,10,9,9,8,8,7,7,7,6,6,6,6,6,5,5,10,9,9,8,8,119,127,99,122,6,6,6,6,5,5,5,9,9,8,149,147,147,147,115,108,142,6,5,5,5,5,5,9,127,185,149,147,147,147,98,93,86,5,5,5,5,5,5,8,132,123,149,105,105,105,105,81,79,44,5,5,5,4,4,8,111,160,185,105,105,105,105,102,41,38,5,5,4,4,4,8,95,90,142,185,105,105,105,105,35,33,52,4,4,4,4,7,7,7,143,133,184,105,105,105,105,105,25,4,4,4,4,7,7,6,121,116,170,184,105,105,105,105,20,30,4,4,4,7,7,6,6,5,7,9,184,4,4,4,4,4,4,4,3,7,6,6,5,5,5,4,4,4,4,4,4,4,3,3,3,6,6,5,5,5,4,4,4,4,4,4,3,3,3,3,3,6,5,5,5,4,4,4,4,4,4,4,3,3,3,3,3,5,5,5,4,4,4,4,4,4,4,3,3,3,3,3,3,103,104,103,101,101,100,100,100,99,99,99,97,97,96,96,96,101,101,101,100,100,100,100,99,97,97,97,97,96,96,96,96,101,101,100,100,100,100,100,97,97,97,97,96,96,95,95,95,101,103,100,100,100,6,5,5,160,97,96,95,97,97,97,97,103,103,100,6,6,6,6,154,149,177,97,97,97,97,97,97,102,6,6,6,6,6,6,140,136,130,97,97,97,97,97,97,100,201,6,6,6,5,5,5,126,124,21,97,97,97,97,96,100,181,229,6,6,6,6,6,5,20,19,97,97,100,97,97,99,165,161,211,7,6,6,6,6,18,18,23,100,100,100,97,99,99,99,28,27,7,6,6,6,6,6,78,100,100,100,100,99,99,99,26,25,31,7,6,6,6,6,68,89,99,99,99,99,99,95,95,95,90,113,7,99,99,99,99,99,99,99,99,99,95,95,95,95,95,99,99,99,99,99,99,99,99,99,99,98,98,97,97,97,101,96,96,96,96,96,96,96,96,96,96,98,97,97,97,101,101,96,96,96,96,96,96,96,96,96,96,97,97,97,101,101,101,101,96,96,96,96,96,96,96,96,96] +[136,140,137,146,136,132,124,110,118,133,151,208,155,228,144,195,136,136,156,160,152,137,128,118,134,149,145,155,234,212,198,171,142,132,125,133,119,112,116,123,142,143,143,145,158,156,180,166,101,141,136,125,123,123,123,134,142,141,140,155,152,144,147,161,83,89,92,104,135,121,122,124,127,135,140,146,59,77,77,148,75,80,82,83,85,90,89,126,127,56,58,58,61,67,68,68,170,71,76,77,71,70,70,56,55,58,58,61,67,67,67,55,69,74,93,93,78,63,55,59,58,61,61,61,67,55,56,54,72,73,73,63,61,54,59,57,60,60,60,55,55,55,53,53,61,68,68,68,59,59,59,59,59,63,63,55,53,53,53,51,69,69,56,63,63,59,63,63,63,63,55,53,51,51,59,59,63,63,62,62,63,63,23,30,30,36,51,51,59,59,59,59,66,66,66,66,66,142,141,139,140,140,140,58,58,58,58,58,60,60,60,60,140,140,155,140,141,137,156,156,58,58,58,58,66,66,66,66,156,155,102,119,102,102,156,157,54,67,67,67,66,66,66,103,155,155,102,102,102,102,156,156,204,67,67,67,154,157,155,166,158,152,146,133,140,152,169,219,172,233,165,206,153,153,170,177,170,161,157,144,156,167,164,173,240,219,207,185,158,149,146,156,148,142,144,146,161,162,162,164,175,172,191,182,87,158,153,149,151,151,146,156,161,160,159,173,170,163,165,177,51,61,61,81,158,150,150,151,150,154,157,164,12,11,12,165,30,40,43,52,56,63,66,154,154,13,12,11,10,10,10,10,201,89,21,22,22,23,23,13,11,11,10,10,10,10,10,9,16,18,24,25,20,13,12,10,10,10,9,9,9,9,8,8,13,13,13,12,12,10,10,9,9,9,9,8,8,8,7,7,12,11,11,11,10,9,9,9,8,8,8,8,7,7,7,6,10,10,9,9,8,8,8,8,8,7,7,7,6,6,6,6,7,7,8,8,8,7,26,34,34,40,6,6,6,5,5,5,6,6,6,6,6,176,175,173,174,175,175,5,5,5,5,5,5,5,5,5,174,174,192,175,175,170,193,193,5,5,5,5,4,4,5,5,193,192,129,150,129,129,194,194,4,4,4,4,4,4,4,167,193,193,129,129,129,129,193,193,149,4,4,4,168,172,169,177,171,163,160,147,154,167,181,226,186,237,179,214,169,169,182,186,180,173,171,159,170,181,178,187,246,223,212,195,173,165,165,172,160,152,157,160,176,176,177,178,188,185,200,194,139,173,169,165,163,163,160,170,176,175,174,187,184,177,179,190,126,131,134,143,172,164,164,164,164,169,173,178,106,128,128,181,122,125,127,126,127,131,129,169,168,103,106,106,110,117,117,117,225,93,125,125,119,117,117,103,102,106,106,109,117,117,117,102,117,122,92,92,127,111,102,107,106,109,109,109,116,102,103,101,122,123,123,111,109,101,107,105,109,109,109,103,103,103,101,101,109,117,117,117,107,107,108,108,108,113,113,103,101,101,101,98,119,119,103,112,112,107,112,112,112,112,103,101,98,98,108,108,113,113,112,112,112,112,31,36,36,41,98,98,107,108,108,108,117,117,117,117,117,7,7,7,7,7,7,107,107,107,107,107,109,109,109,109,6,6,7,6,6,6,7,7,107,107,107,107,117,117,117,117,7,7,6,6,6,6,7,7,102,118,118,118,117,117,117,236,7,7,6,6,6,6,7,7,183,118,118,118] diff --git a/arenai_desktop/tests/src/tests_e2e_agent_input.cpp b/arenai_desktop/tests/src/tests_e2e_agent_input.cpp index 19857f9f..d65a2523 100644 --- a/arenai_desktop/tests/src/tests_e2e_agent_input.cpp +++ b/arenai_desktop/tests/src/tests_e2e_agent_input.cpp @@ -21,7 +21,7 @@ using namespace arenai::desktop; // ======================================================================== // End-to-end: DesktopGameEnvironment -> AbstractAgent::act() -// Mirrors run_game()'s loop with the interceptor in the SAC agent's seat; +// Mirrors run_game()'s loop with the interceptor in the enemy agent's seat; // the player view goes to a headless no-op backend, the enemy visions are // rendered by the environment's real offscreen backend. The player tank is // part of the scene, hence desktop-specific golden images. @@ -54,7 +54,7 @@ namespace { const auto steps = env.step(FREQUENCY, actions); states.clear(); - for (const auto &[state, reward, done]: steps) states.push_back(state); + for (const auto &[state, reward, done, truncated]: steps) states.push_back(state); actions = agent.act(states, VISION_HEIGHT, VISION_WIDTH); } diff --git a/arenai_model/include/arenai_model/constants.h b/arenai_model/include/arenai_model/constants.h index fca7ecdf..17e4124f 100644 --- a/arenai_model/include/arenai_model/constants.h +++ b/arenai_model/include/arenai_model/constants.h @@ -18,6 +18,11 @@ namespace arenai::model { constexpr float ENEMY_TURRET_RADIAL_VELOCITY = std::numbers::pi * 1.f; + constexpr float CANON_AIM_DISTANCE = 100.f; + + constexpr float ZOOM_MAGNIFICATION = 2.f; + constexpr float ZOOM_TRANSITION_SECONDS = 0.25f; + }// namespace arenai::model #endif//ARENAI_MODEL_CONSTANTS_H diff --git a/arenai_model/include/arenai_model/engine.h b/arenai_model/include/arenai_model/engine.h index 40b8cbd9..67a5f0e2 100644 --- a/arenai_model/include/arenai_model/engine.h +++ b/arenai_model/include/arenai_model/engine.h @@ -6,8 +6,11 @@ #define ARENAI_ENGINE_H #include +#include #include +#include + #include "./item.h" namespace arenai::model { @@ -21,6 +24,9 @@ namespace arenai::model { virtual void step(float delta) = 0; + // first hit position along the segment [from, to], or nullopt when nothing is hit + virtual std::optional ray_cast(glm::vec3 from, glm::vec3 to) const = 0; + virtual std::vector> get_items() = 0; virtual void remove_bodies_and_constraints() = 0; diff --git a/arenai_model/include/arenai_model/tank.h b/arenai_model/include/arenai_model/tank.h index 6fea5749..a29ed7d3 100644 --- a/arenai_model/include/arenai_model/tank.h +++ b/arenai_model/include/arenai_model/tank.h @@ -65,6 +65,7 @@ namespace arenai::model { virtual bool is_first_frame_dead() const = 0; virtual bool is_suicide() const = 0; + virtual bool is_timeout() const = 0; virtual void on_death() = 0; }; diff --git a/arenai_model/include/arenai_model/tank_factory.h b/arenai_model/include/arenai_model/tank_factory.h index 9da73508..2636253f 100644 --- a/arenai_model/include/arenai_model/tank_factory.h +++ b/arenai_model/include/arenai_model/tank_factory.h @@ -22,7 +22,7 @@ namespace arenai::model { virtual std::unique_ptr make_enemy_tank( const std::shared_ptr &file_reader, - const std::string &tank_prefix_name, glm::vec3 chassis_pos) = 0; + const std::string &tank_prefix_name, glm::vec3 chassis_pos, bool apply_timeout) = 0; virtual std::unique_ptr make_player_tank( const std::shared_ptr &file_reader, diff --git a/arenai_model/src/jolt_engine.cpp b/arenai_model/src/jolt_engine.cpp index 0bf18374..4677c323 100644 --- a/arenai_model/src/jolt_engine.cpp +++ b/arenai_model/src/jolt_engine.cpp @@ -246,6 +246,12 @@ namespace arenai::model { remove_dead_items(); } + std::optional + JoltPhysicEngine::ray_cast(const glm::vec3 from, const glm::vec3 to) const { + if (const auto fraction = ray_test(from, to, {})) return from + *fraction * (to - from); + return std::nullopt; + } + std::optional JoltPhysicEngine::ray_test( const glm::vec3 from, const glm::vec3 to, const std::vector &excluded) const { diff --git a/arenai_model/src/jolt_engine.h b/arenai_model/src/jolt_engine.h index b6c7fa81..32e0cb6e 100644 --- a/arenai_model/src/jolt_engine.h +++ b/arenai_model/src/jolt_engine.h @@ -45,6 +45,8 @@ namespace arenai::model { void step(float delta) override; + std::optional ray_cast(glm::vec3 from, glm::vec3 to) const override; + std::vector> get_items() override; void remove_bodies_and_constraints() override; diff --git a/arenai_model/src/tank/jolt_enemy_tank.cpp b/arenai_model/src/tank/jolt_enemy_tank.cpp index 89c0167f..3e026d18 100644 --- a/arenai_model/src/tank/jolt_enemy_tank.cpp +++ b/arenai_model/src/tank/jolt_enemy_tank.cpp @@ -45,7 +45,7 @@ namespace arenai::model { JoltPhysicEngine &engine, const std::shared_ptr &file_reader, const std::string &tank_prefix_name, const glm::vec3 chassis_pos, - const float wanted_frame_frequency) + const float wanted_frame_frequency, const bool apply_timeout) : JoltTank( engine, file_reader, tank_prefix_name, chassis_pos, wanted_frame_frequency, [this](const ShellItem *shell, const ShellContactInfo &info, Item *item) { @@ -59,13 +59,17 @@ namespace arenai::model { }), max_frames_upside_down(static_cast(4.f / wanted_frame_frequency)), curr_frame_upside_down(0), miss_distance_scale(1.5f), miss_distance_exponent(1.f / 2.f), - hit_reward_scale(0.1f), hit_received_cost(0.15f), initial_nb_shells(10), + hit_reward_scale(0.1f), hit_received_cost(0.3f), initial_nb_shells(10), nb_shells(initial_nb_shells), max_shells(30), fire_cooldown_frames(static_cast(1.f / 6.f / wanted_frame_frequency)), curr_cooldown_frame(fire_cooldown_frames), shells_recharged_per_hit(5), nb_frames_per_shell_regen(static_cast(1.5f / wanted_frame_frequency)), - curr_frame_shell_regen(0), is_dead_already_triggered(false), has_touch(false), - has_kill(false), has_fired(false) {} + curr_frame_shell_regen(0), is_dead_already_triggered(false), apply_timeout(apply_timeout), + starved(false), max_frames_without_hit(static_cast(30.f / wanted_frame_frequency)), + remaining_frames(max_frames_without_hit), + nb_frames_added_when_hit(static_cast(3.f / wanted_frame_frequency)), + nb_frames_added_when_kill(static_cast(15.f / wanted_frame_frequency)), + has_hit(false), has_kill(false), has_fired(false) {} float JoltEnemyTank::compute_hit_reward( const glm::vec3 &fire_pos, const glm::vec3 &enemy_pos, const glm::vec3 &shell_pos) const { @@ -132,8 +136,9 @@ namespace arenai::model { float JoltEnemyTank::get_reward() const { RewardDetail detail; - // 1. dead / suicide penalty - detail.death = is_dead() ? -1.f : 0.f; + // 1. death penalty — starving out is a pure truncation: any penalty there is an + // unpredictable shock for the critic (the timer is not observable) + detail.death = is_dead() && !is_timeout() ? -1.f : 0.f; // 2. fired shells reward for (int i = static_cast(tracked_shells.size()) - 1; i >= 0; i--) { @@ -165,7 +170,7 @@ namespace arenai::model { // 4. total reward, kept split for the metrics last_reward_detail = detail; - return detail.aim + detail.hit + detail.received + detail.death; + return detail.hit + detail.received + detail.death; } RewardDetail JoltEnemyTank::get_last_reward_detail() const { return last_reward_detail; } @@ -212,6 +217,19 @@ namespace arenai::model { // 4. fire cooldown curr_cooldown_frame = std::min(fire_cooldown_frames, curr_cooldown_frame + 1); + + // 5. timeout + remaining_frames--; + + if (has_hit) remaining_frames += nb_frames_added_when_hit; + if (has_kill) remaining_frames += nb_frames_added_when_kill; + + // starving out (no hit for too long) ends the tank: reported as a truncation to + // the learner — only when the timer is what killed it + if (apply_timeout && remaining_frames <= 0) { + if (!is_dead()) starved = true; + kill_life_items(); + } } void JoltEnemyTank::on_shell_fired(const std::shared_ptr &shell) { @@ -246,11 +264,11 @@ namespace arenai::model { hit = true; killed = true; - has_touch = true; + has_hit = true; has_kill = true; } else if (!life_item->is_dead()) { hit = true; - has_touch = true; + has_hit = true; } } @@ -276,8 +294,8 @@ namespace arenai::model { } bool JoltEnemyTank::consume_has_hit() { - if (has_touch) { - has_touch = false; + if (has_hit) { + has_hit = false; return true; } return false; @@ -302,6 +320,8 @@ namespace arenai::model { return curr_frame_upside_down > max_frames_upside_down; } + bool JoltEnemyTank::is_timeout() const { return starved; } + void JoltEnemyTank::on_death() { if (is_dead() && !is_dead_already_triggered) { is_dead_already_triggered = true; diff --git a/arenai_model/src/tank/jolt_enemy_tank.h b/arenai_model/src/tank/jolt_enemy_tank.h index 3cadf86e..39d9674a 100644 --- a/arenai_model/src/tank/jolt_enemy_tank.h +++ b/arenai_model/src/tank/jolt_enemy_tank.h @@ -41,13 +41,14 @@ namespace arenai::model { JoltPhysicEngine &engine, const std::shared_ptr &file_reader, const std::string &tank_prefix_name, glm::vec3 chassis_pos, - float wanted_frame_frequency); + float wanted_frame_frequency, bool apply_timeout); float get_reward() const override; bool is_dead() const override; bool is_first_frame_dead() const override; bool is_suicide() const override; + bool is_timeout() const override; bool consume_has_hit() override; bool consume_has_kill() override; @@ -92,7 +93,16 @@ namespace arenai::model { bool is_dead_already_triggered; - bool has_touch; + bool apply_timeout; + // latched when the timer actually kills the tank: the timer keeps running after a + // combat death, so the live "remaining_frames <= 0" test would mislabel it later + bool starved; + int max_frames_without_hit; + int remaining_frames; + int nb_frames_added_when_hit; + int nb_frames_added_when_kill; + + bool has_hit; bool has_kill; bool has_fired; diff --git a/arenai_model/src/tank/jolt_tank_factory.cpp b/arenai_model/src/tank/jolt_tank_factory.cpp index 8c557376..54363ffe 100644 --- a/arenai_model/src/tank/jolt_tank_factory.cpp +++ b/arenai_model/src/tank/jolt_tank_factory.cpp @@ -17,9 +17,10 @@ namespace arenai::model { std::unique_ptr JoltTankFactory::make_enemy_tank( const std::shared_ptr &file_reader, - const std::string &tank_prefix_name, glm::vec3 chassis_pos) { + const std::string &tank_prefix_name, glm::vec3 chassis_pos, bool apply_timeout) { return std::make_unique( - engine, file_reader, tank_prefix_name, chassis_pos, wanted_frame_frequency); + engine, file_reader, tank_prefix_name, chassis_pos, wanted_frame_frequency, + apply_timeout); } std::unique_ptr JoltTankFactory::make_player_tank( diff --git a/arenai_model/src/tank/jolt_tank_factory.h b/arenai_model/src/tank/jolt_tank_factory.h index ce260feb..5081ad09 100644 --- a/arenai_model/src/tank/jolt_tank_factory.h +++ b/arenai_model/src/tank/jolt_tank_factory.h @@ -17,7 +17,8 @@ namespace arenai::model { std::unique_ptr make_enemy_tank( const std::shared_ptr &file_reader, - const std::string &tank_prefix_name, glm::vec3 chassis_pos) override; + const std::string &tank_prefix_name, glm::vec3 chassis_pos, + bool apply_timeout) override; std::unique_ptr make_player_tank( const std::shared_ptr &file_reader, diff --git a/arenai_model/src/tank/parts/canon.cpp b/arenai_model/src/tank/parts/canon.cpp index 71f501e0..31396e2d 100644 --- a/arenai_model/src/tank/parts/canon.cpp +++ b/arenai_model/src/tank/parts/canon.cpp @@ -5,14 +5,25 @@ #include "./canon.h" #include +#include #include +#include + using namespace arenai; using namespace arenai::model; namespace { + const float ZOOMED_FOV = + 2.f * std::atan(std::tan(arenai::view::DEFAULT_FOV / 2.f) / ZOOM_MAGNIFICATION); + + // Jolt's default position motor is a soft 2 Hz spring: the canon needs ~10 frames + // to reach a new target angle, a lag the aim loop cannot compensate + constexpr float MOTOR_FREQUENCY = 8.f; + constexpr float MOTOR_DAMPING = 1.f; + glm::mat4 to_glm(const JPH::RMat44 &m) { glm::mat4 result; for (int c = 0; c < 4; c++) { @@ -39,9 +50,9 @@ namespace arenai::model { std::make_shared( file_reader, std::filesystem::path("obj") / "anubis_canon.obj"), pos, scale, mass), - angle(0.f), file_reader(file_reader), will_fire(false), on_contact(on_contact), - on_shell_fired(on_shell_fired), can_fire(can_fire), - wanted_frame_frequency(wanted_frame_frequency) { + angle(0.f), zoom_engaged(false), current_fov(view::DEFAULT_FOV), file_reader(file_reader), + will_fire(false), on_contact(on_contact), on_shell_fired(on_shell_fired), + can_fire(can_fire), wanted_frame_frequency(wanted_frame_frequency) { JPH::HingeConstraintSettings settings; settings.mSpace = JPH::EConstraintSpace::LocalToBodyCOM; @@ -54,6 +65,8 @@ namespace arenai::model { settings.mHingeAxis2 = JPH::Vec3::sAxisX(); settings.mNormalAxis2 = JPH::Vec3::sAxisY(); + settings.mMotorSettings = JPH::MotorSettings(MOTOR_FREQUENCY, MOTOR_DAMPING); + auto *constraint = settings.Create(*turret, *ConvexItem::get_body()); // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) @@ -106,6 +119,8 @@ namespace arenai::model { hinge->SetTargetAngle(angle); if (input.fire_button.pressed && can_fire()) will_fire = true; + + zoom_engaged.store(input.zoom_button.pressed, std::memory_order_relaxed); } glm::vec3 CanonItem::pos() { @@ -117,9 +132,24 @@ namespace arenai::model { glm::vec3 CanonItem::look() { const glm::mat4 model_mat = to_glm(ConvexItem::get_body()->GetWorldTransform()); + return model_mat * glm::vec4(0, 0, CANON_AIM_DISTANCE, 1); + } + + glm::vec3 CanonItem::pivot() { + const glm::mat4 model_mat = to_glm(ConvexItem::get_body()->GetWorldTransform()); + return model_mat * glm::vec4(0, 0, 1, 1); } + float CanonItem::fov() { + const float target = + zoom_engaged.load(std::memory_order_relaxed) ? ZOOMED_FOV : view::DEFAULT_FOV; + const float step = + (view::DEFAULT_FOV - ZOOMED_FOV) * wanted_frame_frequency / ZOOM_TRANSITION_SECONDS; + current_fov += std::clamp(target - current_fov, -step, step); + return current_fov; + } + glm::vec3 CanonItem::up() { const glm::mat4 model_mat = to_glm(ConvexItem::get_body()->GetWorldTransform()); diff --git a/arenai_model/src/tank/parts/canon.h b/arenai_model/src/tank/parts/canon.h index d6a8b2d5..a28deb60 100644 --- a/arenai_model/src/tank/parts/canon.h +++ b/arenai_model/src/tank/parts/canon.h @@ -5,6 +5,7 @@ #ifndef ARENAI_CANON_H #define ARENAI_CANON_H +#include #include #include @@ -40,6 +41,8 @@ namespace arenai::model { glm::vec3 pos() override; glm::vec3 look() override; glm::vec3 up() override; + float fov() override; + glm::vec3 pivot() override; std::vector> get_constraints() override; @@ -48,6 +51,8 @@ namespace arenai::model { private: float angle; + std::atomic zoom_engaged; + float current_fov; JPH::Ref hinge; std::shared_ptr file_reader; bool will_fire; diff --git a/arenai_model/src/tank/parts/turret.cpp b/arenai_model/src/tank/parts/turret.cpp index 6c40b2b3..944cadef 100644 --- a/arenai_model/src/tank/parts/turret.cpp +++ b/arenai_model/src/tank/parts/turret.cpp @@ -10,6 +10,15 @@ using namespace arenai; using namespace arenai::model; +namespace { + + // Jolt's default position motor is a soft 2 Hz spring: the turret needs ~10 frames + // to reach a new target angle, a lag the aim loop cannot compensate + constexpr float MOTOR_FREQUENCY = 8.f; + constexpr float MOTOR_DAMPING = 1.f; + +}// namespace + namespace arenai::model { TurretItem::TurretItem( @@ -34,6 +43,8 @@ namespace arenai::model { settings.mHingeAxis2 = JPH::Vec3::sAxisY(); settings.mNormalAxis2 = JPH::Vec3::sAxisX(); + settings.mMotorSettings = JPH::MotorSettings(MOTOR_FREQUENCY, MOTOR_DAMPING); + auto *constraint = settings.Create(*chassis, *ConvexItem::get_body()); // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) diff --git a/arenai_model/tests/src/tests_enemy_tank/tests_enemy_tank.cpp b/arenai_model/tests/src/tests_enemy_tank/tests_enemy_tank.cpp index 74e7523e..5bbaea12 100644 --- a/arenai_model/tests/src/tests_enemy_tank/tests_enemy_tank.cpp +++ b/arenai_model/tests/src/tests_enemy_tank/tests_enemy_tank.cpp @@ -22,7 +22,7 @@ using namespace arenai::controller; TEST_F(EnemyTankTest, DeadWhenSingleWheelDestroyed) { add_ground(); - const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -53,7 +53,7 @@ TEST_F(EnemyTankTest, DeadWhenSingleWheelDestroyed) { TEST_F(EnemyTankTest, OnDeathMultipleCallsDoNotCrash) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -77,7 +77,7 @@ TEST_F(EnemyTankTest, OnDeathMultipleCallsDoNotCrash) { TEST_F(EnemyTankTest, OnDeathBeforeDeathDoesNothing) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -93,8 +93,8 @@ TEST_F(EnemyTankTest, OnDeathBeforeDeathDoesNothing) { TEST_F(EnemyTankTest, RewardWhenAllEnemiesDeadAndShellFired) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -129,7 +129,7 @@ TEST_F(EnemyTankTest, RewardWhenAllEnemiesDeadAndShellFired) { TEST_F(EnemyTankTest, RewardNoNaNWhenAloneInTankList) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -156,7 +156,7 @@ TEST_F(EnemyTankTest, RewardNoNaNWhenAloneInTankList) { TEST_F(EnemyTankTest, ShellHitsGroundNoRewardNoCrash) { add_ground(); // point the tank away from any enemy so the shell hits the ground - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -194,7 +194,7 @@ TEST_F(EnemyTankTest, ShellHitsGroundNoRewardNoCrash) { TEST_F(EnemyTankTest, SuicideDetectionWhenFlipped) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -214,8 +214,8 @@ TEST_F(EnemyTankTest, SuicideDetectionWhenFlipped) { TEST_F(EnemyTankTest, HasHitOtherTankResetsAfterCall) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -250,10 +250,11 @@ TEST_F(EnemyTankTest, ShellReserveRegeneratesOverTime) { constexpr int initial_shells = 10; constexpr float seconds_per_regen = 1.5f; + // proprioception ends with (reserve, cooldown) ratios constexpr int proprioception_reserve_index = ENEMY_PROPRIOCEPTION_SIZE - 2; add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); float last_reserve = tank->get_proprioception()[proprioception_reserve_index]; ASSERT_FLOAT_EQ( @@ -308,8 +309,11 @@ TEST_F(EnemyTankTest, ShellReserveRegenerationIsCappedAtMaximalReserve) { constexpr int max_shells = 30; constexpr float seconds_per_regen = 1.5f; + // proprioception ends with (reserve, cooldown) ratios + constexpr int proprioception_reserve_index = ENEMY_PROPRIOCEPTION_SIZE - 2; + add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < static_cast(nb_frames_one_second * seconds_per_regen * max_shells); i++) { @@ -317,7 +321,8 @@ TEST_F(EnemyTankTest, ShellReserveRegenerationIsCappedAtMaximalReserve) { tank->tick({}); } - ASSERT_FLOAT_EQ(tank->get_proprioception().back(), 1.f) << "the reserve should be full"; + ASSERT_FLOAT_EQ(tank->get_proprioception()[proprioception_reserve_index], 1.f) + << "the reserve should be full"; const std::shared_ptr shared_tank(tank.release()); const std::vector tanks{shared_tank}; @@ -325,6 +330,6 @@ TEST_F(EnemyTankTest, ShellReserveRegenerationIsCappedAtMaximalReserve) { // the reserve is already full: several periods must not push it above it for (int i = 0; i < 5 * nb_frames_one_second; i++) shared_tank->tick(tanks); - ASSERT_FLOAT_EQ(shared_tank->get_proprioception().back(), 1.f) + ASSERT_FLOAT_EQ(shared_tank->get_proprioception()[proprioception_reserve_index], 1.f) << "regeneration should never take the reserve above its max value"; } diff --git a/arenai_model/tests/src/tests_engine/tests_ray_cast.cpp b/arenai_model/tests/src/tests_engine/tests_ray_cast.cpp new file mode 100644 index 00000000..cd53348b --- /dev/null +++ b/arenai_model/tests/src/tests_engine/tests_ray_cast.cpp @@ -0,0 +1,38 @@ +// +// Created by claude on 21/09/2026. +// + +#include + +#include + +using namespace arenai; +using namespace arenai::model; + +// ======================================================================== +// ray_cast — hit position, miss cases +// ======================================================================== + +TEST_F(EngineTestFixture, RayCastHitsGroundWithoutStep) { + add_ground(); + + // no engine->step: spawn code casts rays right after adding bodies + const auto hit = engine->ray_cast(glm::vec3(10.f, 100.f, 10.f), glm::vec3(10.f, -100.f, 10.f)); + + ASSERT_TRUE(hit.has_value()); + EXPECT_NEAR(hit->x, 10.f, 1e-3f); + EXPECT_NEAR(hit->y, 0.f, 1e-2f); + EXPECT_NEAR(hit->z, 10.f, 1e-3f); +} + +TEST_F(EngineTestFixture, RayCastMissReturnsNullopt) { + // empty world + EXPECT_FALSE( + engine->ray_cast(glm::vec3(0.f, 100.f, 0.f), glm::vec3(0.f, -100.f, 0.f)).has_value()); + + add_ground(); + + // segment ending above the ground + EXPECT_FALSE( + engine->ray_cast(glm::vec3(0.f, 100.f, 0.f), glm::vec3(0.f, 50.f, 0.f)).has_value()); +} diff --git a/arenai_model/tests/src/tests_player_tank/tests_player_tank.cpp b/arenai_model/tests/src/tests_player_tank/tests_player_tank.cpp index 6ff0f73e..4ecbce7d 100644 --- a/arenai_model/tests/src/tests_player_tank/tests_player_tank.cpp +++ b/arenai_model/tests/src/tests_player_tank/tests_player_tank.cpp @@ -29,7 +29,7 @@ TEST_F(PlayerTankTest, ScoreZeroAtCreation) { TEST_F(PlayerTankTest, ScoreIncreasesOnHit) { add_ground(); const auto player = tank_factory->make_player_tank(file_reader, "player", {0.f, 5.f, 0.f}); - auto enemy = tank_factory->make_enemy_tank(file_reader, "enemy", {0.f, 5.f, 30.f}); + auto enemy = tank_factory->make_enemy_tank(file_reader, "enemy", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -47,7 +47,7 @@ TEST_F(PlayerTankTest, ScoreIncreasesOnHit) { TEST_F(PlayerTankTest, ScoreHigherOnKillThanHit) { add_ground(); const auto player = tank_factory->make_player_tank(file_reader, "player", {0.f, 5.f, 0.f}); - const auto enemy = tank_factory->make_enemy_tank(file_reader, "enemy", {0.f, 5.f, 30.f}); + const auto enemy = tank_factory->make_enemy_tank(file_reader, "enemy", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); diff --git a/arenai_model/tests/src/tests_proprioception/tests_proprioception.cpp b/arenai_model/tests/src/tests_proprioception/tests_proprioception.cpp index 3d6ae9d4..1ec2ac38 100644 --- a/arenai_model/tests/src/tests_proprioception/tests_proprioception.cpp +++ b/arenai_model/tests/src/tests_proprioception/tests_proprioception.cpp @@ -18,7 +18,7 @@ using namespace arenai::utils; TEST_F(ProprioceptionTest, ProprioceptionSizeCorrect) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -30,7 +30,7 @@ TEST_F(ProprioceptionTest, ProprioceptionSizeCorrect) { TEST_F(ProprioceptionTest, ProprioceptionNoNaN) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -45,7 +45,7 @@ TEST_F(ProprioceptionTest, ProprioceptionNoNaN) { TEST_F(ProprioceptionTest, ProprioceptionNoInfinity) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); for (int i = 0; i < 10; i++) engine->step(1.f / 60.f); @@ -59,7 +59,7 @@ TEST_F(ProprioceptionTest, ProprioceptionNoInfinity) { TEST_F(ProprioceptionTest, ProprioceptionContainsSubItemRelativePositions) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -89,7 +89,7 @@ TEST_F(ProprioceptionTest, ProprioceptionContainsSubItemRelativePositions) { TEST_F(ProprioceptionTest, ProprioceptionConsistentSize) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -106,7 +106,7 @@ TEST_F(ProprioceptionTest, ProprioceptionConsistentSize) { TEST_F(ProprioceptionTest, ProprioceptionForwardAndUpVectorsValid) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); diff --git a/arenai_model/tests/src/tests_reward/tests_reward.cpp b/arenai_model/tests/src/tests_reward/tests_reward.cpp index f2ac504d..d84326bc 100644 --- a/arenai_model/tests/src/tests_reward/tests_reward.cpp +++ b/arenai_model/tests/src/tests_reward/tests_reward.cpp @@ -21,8 +21,8 @@ using namespace arenai::controller; TEST_F(RewardTest, RewardZeroWhenAliveNoShot) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -32,16 +32,14 @@ TEST_F(RewardTest, RewardZeroWhenAliveNoShot) { const float reward_a = tanks[0]->get_reward(); const float reward_b = tanks[1]->get_reward(); - // the dense aim shaping leaves a negligible residue when the canon points ~90° - // away from the enemy, so the reward is near zero rather than exactly zero - ASSERT_NEAR(reward_a, 0.f, 1e-3f); - ASSERT_NEAR(reward_b, 0.f, 1e-3f); + ASSERT_FLOAT_EQ(reward_a, 0.f); + ASSERT_FLOAT_EQ(reward_b, 0.f); } TEST_F(RewardTest, RewardNegativeWhenDead) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -68,8 +66,8 @@ TEST_F(RewardTest, RewardNegativeWhenDead) { TEST_F(RewardTest, DeathPenaltyIsMinusOne) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 0.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {20.f, 0.f, 0.f}, false); engine->step(1.f / 60.f); @@ -87,8 +85,8 @@ TEST_F(RewardTest, DeathPenaltyIsMinusOne) { const float death_reward = tanks[1]->get_reward(); // death and suicide share the same penalty so early termination is never an escape; - // the fatal hit also counts as a received hit (-0.15) - ASSERT_FLOAT_EQ(death_reward, -1.15f); + // the fatal hit also counts as a received hit (-0.3) + ASSERT_FLOAT_EQ(death_reward, -1.3f); } // ======================================================================== @@ -98,8 +96,8 @@ TEST_F(RewardTest, DeathPenaltyIsMinusOne) { TEST_F(RewardTest, RewardPositiveOnHit) { add_ground(); // spawn tanks high enough so all parts start above ground and settle cleanly - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); // settle on ground (300 frames = 5s at 60fps) for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -136,8 +134,8 @@ TEST_F(RewardTest, RewardPositiveOnHit) { TEST_F(RewardTest, RewardUnderOneAfterHit) { add_ground(); // spawn tanks high enough so all parts start above ground and settle cleanly - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); // settle on ground (300 frames = 5s at 60fps) for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -169,7 +167,6 @@ TEST_F(RewardTest, RewardUnderOneAfterHit) { ASSERT_GE(max_reward_on_hit, 0.2f) << "reward should be greater than or equal to the hit bonus after hitting an enemy"; - // no fire, reward under the hit bonus constexpr user_input no_fire_input{ .left_joystick = {.x = 0.f, .y = 0.f}, .right_joystick = {.x = 0.f, .y = 0.f}, @@ -188,8 +185,42 @@ TEST_F(RewardTest, RewardUnderOneAfterHit) { ASSERT_FALSE(std::isnan(max_reward_on_no_hit)) << "reward should never be NaN"; ASSERT_FALSE(std::isinf(max_reward_on_no_hit)) << "reward should never be Inf"; - ASSERT_LE(max_reward_on_no_hit, 0.2f) - << "reward should stay under the hit bonus when no shell hit an enemy"; + ASSERT_FLOAT_EQ(max_reward_on_no_hit, 0.f) + << "reward should be exactly zero when no shell hit an enemy"; +} + +TEST_F(RewardTest, MissedShellPaysNothing) { + add_ground(); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {15.f, 5.f, 30.f}, false); + + for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); + + const std::shared_ptr shared_a(tank_a.release()); + const std::shared_ptr shared_b(tank_b.release()); + + constexpr user_input fire_input{ + .left_joystick = {.x = 0.f, .y = 0.f}, + .right_joystick = {.x = 0.f, .y = 0.f}, + .fire_button = {true}}; + for (const auto &ctrl: shared_a->get_controllers()) ctrl->apply_input(fire_input); + + const std::vector tanks{shared_a, shared_b}; + + float total_reward = 0.f; + int max_landed_shells = 0; + for (int i = 0; i < 180; i++) { + engine->step(1.f / 60.f); + shared_a->tick(tanks); + + total_reward += shared_a->get_reward(); + max_landed_shells = + std::max(shared_a->get_last_reward_detail().nb_landed_shells, max_landed_shells); + } + + ASSERT_FALSE(shared_a->consume_has_hit()) << "shell should have missed the enemy tank"; + ASSERT_GT(max_landed_shells, 0) << "the shell tracker should have sampled a landed shell"; + ASSERT_FLOAT_EQ(total_reward, 0.f) << "a shell that lands without hitting must pay nothing"; } // ======================================================================== @@ -198,12 +229,13 @@ TEST_F(RewardTest, RewardUnderOneAfterHit) { TEST_F(RewardTest, NoRewardWhenShootingAWreck) { + // proprioception ends with (reserve, cooldown) ratios constexpr int proprioception_reserve_index = ENEMY_PROPRIOCEPTION_SIZE - 2; add_ground(); // spawn tanks high enough so all parts start above ground and settle cleanly - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); // settle on ground (300 frames = 5s at 60fps) for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -249,8 +281,8 @@ TEST_F(RewardTest, NoRewardWhenShootingAWreck) { TEST_F(RewardTest, NoKillRewardWhenHittingAnotherPartOfAWreck) { add_ground(); // spawn tanks high enough so all parts start above ground and settle cleanly - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); // settle on ground (300 frames = 5s at 60fps) for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -300,7 +332,7 @@ TEST_F(RewardTest, NoKillRewardWhenHittingAnotherPartOfAWreck) { TEST_F(RewardTest, ZeroRewardWithEmptyTankList) { add_ground(); - auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); diff --git a/arenai_model/tests/src/tests_shell/tests_shell.cpp b/arenai_model/tests/src/tests_shell/tests_shell.cpp index 0819d6f8..5d61239f 100644 --- a/arenai_model/tests/src/tests_shell/tests_shell.cpp +++ b/arenai_model/tests/src/tests_shell/tests_shell.cpp @@ -19,7 +19,7 @@ using namespace arenai::controller; TEST_F(ShellTest, FireCreatesShellItem) { add_ground(); - const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -40,7 +40,7 @@ TEST_F(ShellTest, FireCreatesShellItem) { TEST_F(ShellTest, ShellDestroyedAfterLifetime) { add_ground(); - const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -68,8 +68,8 @@ TEST_F(ShellTest, ShellDestroyedAfterLifetime) { TEST_F(ShellTest, ShellHitsEnemyTank) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -89,8 +89,8 @@ TEST_F(ShellTest, ShellHitsEnemyTank) { TEST_F(ShellTest, ShellDestroyedOnContact) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -119,7 +119,7 @@ TEST_F(ShellTest, ShellDestroyedOnContact) { TEST_F(ShellTest, NoFireNoNewItems) { add_ground(); - const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); + const auto tank = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); engine->step(1.f / 60.f); @@ -140,8 +140,8 @@ TEST_F(ShellTest, NoFireNoNewItems) { TEST_F(ShellTest, ShellContactCallbackSetsReward) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -202,8 +202,8 @@ namespace { TEST_F(ShellTest, ShellImpactDealsExactlyOneDamage) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -226,8 +226,8 @@ TEST_F(ShellTest, ShellImpactDealsExactlyOneDamage) { TEST_F(ShellTest, TankSurvivesUntilAPartsHealthPointsAreSpent) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); @@ -265,8 +265,8 @@ TEST_F(ShellTest, TankSurvivesUntilAPartsHealthPointsAreSpent) { TEST_F(ShellTest, SpentShellDealsNoFurtherDamage) { add_ground(); - auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}); - auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}); + auto tank_a = tank_factory->make_enemy_tank(file_reader, "tank_a", {0.f, 5.f, 0.f}, false); + auto tank_b = tank_factory->make_enemy_tank(file_reader, "tank_b", {0.f, 5.f, 30.f}, false); for (int i = 0; i < 300; i++) engine->step(1.f / 60.f); diff --git a/arenai_view/include/arenai_view/camera.h b/arenai_view/include/arenai_view/camera.h index 8354ac7b..82f73fe2 100644 --- a/arenai_view/include/arenai_view/camera.h +++ b/arenai_view/include/arenai_view/camera.h @@ -7,21 +7,25 @@ #include #include +#include #include #include namespace arenai::view { + constexpr float DEFAULT_FOV = std::numbers::pi_v / 4.f; + class AbstractCamera { public: virtual ~AbstractCamera() = default; virtual glm::vec3 pos() = 0; - virtual glm::vec3 look() = 0; - virtual glm::vec3 up() = 0; + + virtual float fov(); + virtual glm::vec3 pivot(); }; class StaticCamera final : public AbstractCamera { @@ -29,9 +33,7 @@ namespace arenai::view { StaticCamera(glm::vec3 pos, glm::vec3 look, glm::vec3 up); glm::vec3 pos() override; - glm::vec3 look() override; - glm::vec3 up() override; private: @@ -40,16 +42,8 @@ namespace arenai::view { glm::vec3 up_vec; }; - // Fraction of the [from -> to] segment at which the world is first hit, in - // (0, 1]. std::nullopt when the path is free. using RaycastFunction = std::function(glm::vec3 from, glm::vec3 to)>; - /** - * Spring-arm decorator: keeps the wrapped camera's aim but pulls its position - * toward the look-at pivot when world geometry blocks the [pivot -> pos] - * segment, so the camera never goes behind walls or under the terrain. - * Retraction is instantaneous (no clipping), extension is smoothed. - */ class CollisionCamera final : public AbstractCamera { public: CollisionCamera( @@ -57,11 +51,12 @@ namespace arenai::view { float margin = 0.5f, float min_distance = 2.f, float extend_speed = 4.f); glm::vec3 pos() override; - glm::vec3 look() override; - glm::vec3 up() override; + float fov() override; + glm::vec3 pivot() override; + private: std::shared_ptr inner; RaycastFunction raycast; diff --git a/arenai_view/src/camera.cpp b/arenai_view/src/camera.cpp index 3f666e3d..4c29e576 100644 --- a/arenai_view/src/camera.cpp +++ b/arenai_view/src/camera.cpp @@ -10,6 +10,10 @@ namespace arenai::view { + float AbstractCamera::fov() { return DEFAULT_FOV; } + + glm::vec3 AbstractCamera::pivot() { return look(); } + StaticCamera::StaticCamera(const glm::vec3 pos, const glm::vec3 look, const glm::vec3 up) : pos_vec(pos), look_vec(look), up_vec(up) {} @@ -27,7 +31,7 @@ namespace arenai::view { current_distance(std::numeric_limits::max()) {} glm::vec3 CollisionCamera::pos() { - const glm::vec3 pivot = inner->look(); + const glm::vec3 pivot = inner->pivot(); const glm::vec3 desired = inner->pos(); const glm::vec3 offset = desired - pivot; @@ -52,4 +56,8 @@ namespace arenai::view { glm::vec3 CollisionCamera::up() { return inner->up(); } + float CollisionCamera::fov() { return inner->fov(); } + + glm::vec3 CollisionCamera::pivot() { return inner->pivot(); } + }// namespace arenai::view diff --git a/arenai_view/src/vulkan/scene/renderers/renderer.cpp b/arenai_view/src/vulkan/scene/renderers/renderer.cpp index 1af92d42..fa9ab16c 100644 --- a/arenai_view/src/vulkan/scene/renderers/renderer.cpp +++ b/arenai_view/src/vulkan/scene/renderers/renderer.cpp @@ -146,8 +146,7 @@ namespace arenai::view { // zero-to-one depth projection (Vulkan clip space); the y flip is // handled by the negative-height viewport, not by the matrix const glm::mat4 proj_matrix = glm::perspectiveRH_ZO( - static_cast(M_PI) / 4.f, - static_cast(get_width()) / static_cast(get_height()), 1.f, + camera_->fov(), static_cast(get_width()) / static_cast(get_height()), 1.f, 10000.f * std::sqrt(3.f)); last_view_proj_matrix_ = proj_matrix * view_matrix; diff --git a/resources/menu/ai.rml b/resources/menu/ai.rml new file mode 100644 index 00000000..410446af --- /dev/null +++ b/resources/menu/ai.rml @@ -0,0 +1,41 @@ + + + AI + + + +
+
+
←
+

Configure AI

+
+ +
+ Algorithm +
+
PPO
+
PPO Liquid
+
+
+ +

Trained AI model

+
{{current_dir}}
+
+
{{it}}
+
+
State dicts: {{agent_folder}}
+
Config: {{agent_config}}
+

{{agent_status}}

+

{{agent_status}}

+

No folder selected yet

+ +
+
Use this folder
+
+ +
+
Back
+
+
+ +
diff --git a/resources/menu/controls.rml b/resources/menu/controls.rml index a522b270..47305f3b 100644 --- a/resources/menu/controls.rml +++ b/resources/menu/controls.rml @@ -10,6 +10,14 @@

Controls

+
+ Input +
+
Mouse + Keyboard
+
Gamepad
+
+
+

Mouse + Keyboard

@@ -48,7 +56,7 @@
Reset to defaults
-
Back
+
Back
diff --git a/resources/menu/game_over.rml b/resources/menu/game_over.rml index 9e1c9bc7..f77ff967 100644 --- a/resources/menu/game_over.rml +++ b/resources/menu/game_over.rml @@ -7,7 +7,7 @@

Game Over

Score {{score}}
-
Retry
+
Retry
Main menu
Exit game
diff --git a/resources/menu/graphics.rml b/resources/menu/graphics.rml index 03dd6856..70086630 100644 --- a/resources/menu/graphics.rml +++ b/resources/menu/graphics.rml @@ -10,24 +10,32 @@

Graphics

+
+ Display +
+
Windowed
+
Fullscreen
+
+
+
Shadows
-
Off
-
Low
-
Medium
-
High
-
Ultra
+
Off
+
Low
+
Medium
+
High
+
Ultra
Anti-aliasing
-
Off
-
2x
-
4x
-
8x
+
Off
+
2x
+
4x
+
8x
@@ -45,7 +53,7 @@

Overridden by the ARENAI_VK_DEVICE environment variable.

-
Back
+
Back
diff --git a/resources/menu/main_menu.rml b/resources/menu/main_menu.rml index 17473d2f..941c8404 100644 --- a/resources/menu/main_menu.rml +++ b/resources/menu/main_menu.rml @@ -6,7 +6,7 @@

ArenAI

-
Play
+
Play
Parameters
Exit

Select a trained AI model folder in Parameters to play

diff --git a/resources/menu/menu.rcss b/resources/menu/menu.rcss index 4de2ae1f..4adc9cd1 100644 --- a/resources/menu/menu.rcss +++ b/resources/menu/menu.rcss @@ -4,7 +4,16 @@ Light tints derived from the accent (the palette has no light value): ink #EAF7FF, muted #8AC4E0, selection #7FD3FF, disabled #6FA3BF. Sharp corners (2–3dp), 3dp accent stripe on panels, uppercase display. - Display font Sora, technical labels IBM Plex Mono. */ + Display font Sora, technical labels IBM Plex Mono. + + One cursor mark for every control, mouse hover and gamepad focus alike: + the corner reticle (gui-registered corner-reticle decorator) — accent + ticks on dark chrome, night ticks on the accent-filled primary, ink + ticks on a selected surface. State cycle: normal → cursor (ticks + + #00a6fb17 wash) → pressed (#4DC0FC ticks + #00a6fb33 wash); selection + speaks through tint #00a6fb24 + accent text, never through the mark. + Sora controls answer with their surface only; IBM Plex Mono controls + (back arrow, keychips, list entries) also turn their text accent. */ body { font-family: "Sora"; @@ -109,8 +118,14 @@ h1 { .back-arrow:hover, body.gamepad-nav .back-arrow:focus { color: #00A6FB; - border-color: #00A6FB; background-color: #00a6fb17; + decorator: corner-reticle(#00A6FB); +} + +.back-arrow:active { + color: #00A6FB; + background-color: #00a6fb33; + decorator: corner-reticle(#4DC0FC); } h2 { @@ -147,41 +162,33 @@ h2 { documents with .gamepad-nav (mouse hover is cleared on that switch, and the first mouse move flips back). */ .button:hover, body.gamepad-nav .button:focus { - border-color: #00A6FB; background-color: #00a6fb17; + decorator: corner-reticle(#00A6FB); } -/* Primary buttons carry a detached cursor ring: RCSS cannot paint outside - an element, so the outer .button.primary is a 2dp ring + 2dp gap - (reserved in transparent — nothing shifts) around an inner .fill that - owns the accent background. The outer element keeps the focus, the - navigation and the click bindings. */ -.button.primary { - width: 302dp; - margin: 6dp auto; - padding: 2dp; - border: 2dp transparent; - /* concentric with the 2dp fill corners: 2dp + 2dp ring + 2dp gap */ - border-radius: 6dp; +.button:active { + background-color: #00a6fb33; + decorator: corner-reticle(#4DC0FC); } -.button.primary .fill { - display: block; - padding: 12dp 20dp; - border: 1dp #00A6FB; - border-radius: 2dp; +.button.primary { background-color: #00A6FB; + border-color: #00A6FB; color: #051923; } +/* night ticks: the mark keeps its contrast on the accent-filled surface */ .button.primary:hover, body.gamepad-nav .button.primary:focus { - border-color: #EAF7FF; - background-color: transparent; -} - -.button.primary:hover .fill, body.gamepad-nav .button.primary:focus .fill { background-color: #4DC0FC; border-color: #4DC0FC; + decorator: corner-reticle(#051923); +} + +.button.primary:active { + background-color: #0582CA; + border-color: #0582CA; + color: #EAF7FF; + decorator: corner-reticle(#EAF7FF); } .button.disabled { @@ -190,13 +197,14 @@ h2 { border-color: #006494; } -/* after the cursor rules, and restated under .gamepad-nav, so a disabled - Play keeps its greyed fill under either cursor (the ring still shows, - marking the position without promising a click) */ -.button.primary.disabled .fill, body.gamepad-nav .button.primary.disabled:focus .fill { +/* after the cursor rules, so a disabled Play keeps its greyed fill under + either cursor (the ticks still show, marking the position without + promising a click) */ +.button.disabled:hover, body.gamepad-nav .button.disabled:focus, .button.disabled:active { color: #6FA3BF; background-color: #003554; border-color: #006494; + decorator: corner-reticle(#00A6FB); } .hint { @@ -207,6 +215,15 @@ h2 { margin-top: 14dp; } +/* dry-run load verdict of the AI page */ +.hint.status-ok { + color: #35D0A5; +} + +.hint.status-error { + color: #FF6B6B; +} + .param-row { display: block; margin: 16dp 0; @@ -258,19 +275,18 @@ input.range sliderbar { border-radius: 2dp; } -/* Detached cursor ring, as on the primary button. The knob is a - widget-internal element that cannot be wrapped, so the ring + gap + fill - are drawn by the gui-registered cursor-ring decorator; the box grows to - 2dp ring + 2dp gap around the 16dp fill and recenters on the track. - Ring and fill reuse the base rule's 2dp border-radius (the decorator - follows the element's radius), so the ring stays square like the knob. */ +/* The knob is a widget-internal element that cannot host children, so the + decorator draws both the corner ticks and the 14dp fill: the box grows + to 24dp around the fill and recenters on the track. The fill reuses the + base rule's 2dp border-radius (the decorator follows the element's + radius), so it stays square like the idle knob. */ input.range sliderbar:hover, input.range sliderbar:active, body.gamepad-nav input.range:focus sliderbar { width: 24dp; height: 24dp; margin-top: -4dp; background-color: transparent; - decorator: cursor-ring(#EAF7FF #4DC0FC 2dp 2dp); + decorator: corner-reticle(#00A6FB #00A6FB 7dp 2dp 0dp 5dp); } input.range sliderarrowdec, input.range sliderarrowinc { @@ -282,15 +298,14 @@ input.range sliderarrowdec, input.range sliderarrowinc { display: block; } -/* same detached-ring structure as .button.primary: the outer .toggle is the - ring slot (2dp ring + 2dp gap, transparent when idle — its 4dp per side - also supply the old 8dp gap between neighbours), the inner .fill draws - the actual chip */ +/* the 4dp margin keeps the old 8dp gap between neighbours (previously + supplied by the transparent ring slot around each chip) */ .toggle { display: inline-block; - padding: 2dp; - border: 2dp transparent; - border-radius: 6dp; + margin: 4dp; + padding: 8dp 18dp; + border: 1dp #0582CA; + border-radius: 2dp; font-size: 13dp; font-weight: 600; text-transform: uppercase; @@ -301,61 +316,50 @@ input.range sliderarrowdec, input.range sliderarrowinc { nav: auto; } -.toggle .fill { - display: block; - padding: 8dp 18dp; - border: 1dp #0582CA; - border-radius: 2dp; +.toggle:hover, body.gamepad-nav .toggle:focus { + background-color: #00a6fb17; + decorator: corner-reticle(#00A6FB); } -/* unselected under the cursor: the plain-button feedback */ -.toggle:hover .fill, body.gamepad-nav .toggle:focus .fill { - background-color: #00a6fb17; - border-color: #00A6FB; +.toggle:active { + background-color: #00a6fb33; + decorator: corner-reticle(#4DC0FC); } -.toggle.selected .fill { +.toggle.selected { background-color: #00a6fb24; border-color: #00A6FB; color: #00A6FB; } -/* selected under the cursor: the detached ink ring, as on Play, and the +/* selected under the cursor: ink ticks over the accent-tinted chip, the selected tint restated so the cursor rules above cannot wash it out */ .toggle.selected:hover, body.gamepad-nav .toggle.selected:focus { - border-color: #EAF7FF; + background-color: #00a6fb24; + decorator: corner-reticle(#EAF7FF); } -.toggle.selected:hover .fill, body.gamepad-nav .toggle.selected:focus .fill { - background-color: #00a6fb24; - border-color: #00A6FB; +.toggle.selected:active { + background-color: #00a6fb33; + decorator: corner-reticle(#EAF7FF); } -/* Controls and Graphics pages — same design width as the parameters panel */ -#controls-panel, #graphics-panel { +/* Controls, Graphics and AI pages — same design width as the parameters panel */ +#controls-panel, #graphics-panel, #ai-panel { width: 560dp; } -/* the Controls row lays its toggles out as a flex line so Configure can - ride its auto margin to the right edge */ -.toggle-group.controls-row { - display: flex; - align-items: center; +/* the three sub-page buttons of the parameters panel: centered by the base + .button auto margins, and the same visible size as the Back button below + (220dp content + 40dp padding + 2dp border) — plain, only Back keeps the + primary highlight. The group stands clear of the sliders above it. */ +.configure-group { + display: block; + margin-top: 32dp; } -/* opens the controls page from the Controls row: chip-sized, pushed to the - right edge by its auto margin */ .button.configure { - display: inline-block; - width: auto; - margin: 0 0 0 auto; - padding: 9dp 18dp; - font-size: 13dp; -} - -/* the Graphics row's Configure sits flush left under its label */ -.button.configure.left { - margin-left: 0; + width: 220dp; } /* binding rows: a dark inset list, like the file explorer */ @@ -406,9 +410,15 @@ input.range sliderarrowdec, input.range sliderarrowinc { } .keychip:hover, body.gamepad-nav .keychip:focus { - border-color: #00A6FB; background-color: #00a6fb17; color: #00A6FB; + decorator: corner-reticle(#00A6FB); +} + +.keychip:active { + background-color: #00a6fb33; + color: #00A6FB; + decorator: corner-reticle(#4DC0FC); } @keyframes chip-pulse { @@ -441,6 +451,8 @@ input.range sliderarrowdec, input.range sliderarrowinc { cursor: default; tab-index: none; nav: none; + /* inert: the cursor mark of .keychip:hover must not land here */ + decorator: none; } .fixed-tag { @@ -476,9 +488,17 @@ input.range sliderarrowdec, input.range sliderarrowinc { nav: auto; } +/* dense list rows: smaller ticks (5dp), tighter inset (2dp) */ .device-entry:hover, body.gamepad-nav .device-entry:focus { background-color: #00a6fb17; color: #00A6FB; + decorator: corner-reticle(#00A6FB transparent 5dp 2dp 2dp); +} + +.device-entry:active { + background-color: #00a6fb33; + color: #00A6FB; + decorator: corner-reticle(#4DC0FC transparent 5dp 2dp 2dp); } .device-entry.selected { @@ -486,6 +506,11 @@ input.range sliderarrowdec, input.range sliderarrowinc { color: #00A6FB; } +.device-entry.selected:hover, body.gamepad-nav .device-entry.selected:focus { + background-color: #00a6fb24; + decorator: corner-reticle(#EAF7FF transparent 5dp 2dp 2dp); +} + .current-dir { display: block; padding: 8dp 10dp; @@ -523,6 +548,13 @@ input.range sliderarrowdec, input.range sliderarrowinc { .file-entry:hover, body.gamepad-nav .file-entry:focus { background-color: #00a6fb17; color: #00A6FB; + decorator: corner-reticle(#00A6FB transparent 5dp 2dp 2dp); +} + +.file-entry:active { + background-color: #00a6fb33; + color: #00A6FB; + decorator: corner-reticle(#4DC0FC transparent 5dp 2dp 2dp); } .selected-folder { @@ -547,12 +579,6 @@ input.range sliderarrowdec, input.range sliderarrowinc { margin: 0 8dp; } -/* the primary Back keeps the same visible size as its plain neighbours: - 220dp content + 40dp padding + 2dp border, carried by the ring wrapper */ -.row-buttons .button.primary { - width: 262dp; -} - /* wide gap above Back: visually detaches it from the explorer's "Use this folder" action right above */ .row-buttons.back-row { diff --git a/resources/menu/parameters.rml b/resources/menu/parameters.rml index 55cf472b..e55e5f4b 100644 --- a/resources/menu/parameters.rml +++ b/resources/menu/parameters.rml @@ -20,44 +20,15 @@
-
- Controls -
-
Mouse + Keyboard
-
Gamepad
-
Configure
-
-
- -
- Display -
-
Windowed
-
Fullscreen
-
-
- -
- Graphics -
Configure
-
- -

Trained AI model folder

-
{{current_dir}}
-
-
{{it}}
-
-
Selected: {{sac_folder}}
-

{{sac_status}}

-

{{sac_status}}

-

No folder selected yet

- -
-
Use this folder
+
+
Configure Graphics
+
Configure Controls
+
Configure AI
+

The AI model is not loaded yet — configure it before playing.

-
Back
+
Back
diff --git a/resources/menu/pause.rml b/resources/menu/pause.rml index 52f501cb..b0106989 100644 --- a/resources/menu/pause.rml +++ b/resources/menu/pause.rml @@ -7,7 +7,7 @@

Paused

Score {{score}}
-
Continue
+
Continue
Main menu
Exit game
diff --git a/resources/trained_models/ppo_liquid_409/config.json b/resources/trained_models/ppo_liquid_409/config.json new file mode 100644 index 00000000..87990476 --- /dev/null +++ b/resources/trained_models/ppo_liquid_409/config.json @@ -0,0 +1,96 @@ +{ + "agent": { + "actor_learning_rate": 9.999999747378752e-05, + "chunk_size": 30, + "clip_epsilon": 0.20000000298023224, + "continuous_target_entropy": [ + -0.20000000298023224, + -0.20000000298023224, + -0.8799999952316284, + -0.8799999952316284 + ], + "critic_learning_rate": 0.0003000000142492354, + "delta_t": 0.03333333507180214, + "discrete_target_entropy_factors": [ + 0.4000000059604645, + 0.4000000059604645 + ], + "epochs": 2, + "gae_lambda": 0.9900000095367432, + "gamma": 0.996999979019165, + "grad_norm_max": 0.5, + "group_norm_nums": [ + 2, + 3, + 4, + 6, + 8, + 12, + 16 + ], + "hidden_size_sensors": 128, + "initial_fire_proba": 0.4000000059604645, + "initial_sigma": 0.5, + "initial_zoom_proba": 0.25, + "metric_window_size": 256, + "minibatch_size": 1024, + "neuron_number": 128, + "rollout_size": 900, + "target_kl": 0.05000000074505806, + "unfolding_steps": 6, + "vision_channels": [ + [ + 3, + 16 + ], + [ + 16, + 24 + ], + [ + 24, + 32 + ], + [ + 32, + 48 + ], + [ + 48, + 64 + ], + [ + 64, + 96 + ], + [ + 96, + 128 + ] + ] + }, + "environment": { + "curriculum_boundary_proba": 0.20000000298023224, + "curriculum_delta": 50.0, + "curriculum_probe_window": 24, + "curriculum_ratio_high": 0.05999999865889549, + "curriculum_ratio_low": 0.029999999329447746, + "final_spawn_height": 1000.0, + "final_spawn_width": 1000.0, + "initial_spawn_height": 250.0, + "initial_spawn_width": 250.0, + "nb_tanks": 32, + "vision_height": 128, + "vision_num_threads": 16, + "vision_width": 256, + "wanted_frequency": 0.03333333507180214 + }, + "train": { + "cuda": true, + "max_episode_steps": 5400, + "nb_episodes": 6000, + "output_folder": "./outputs/train_409_ppo_liquid_separated-target-entropy", + "resources_folder": "/home/samuel/CLionProjects/ArenAI/resources", + "save_every": 108000 + } +} diff --git a/resources/trained_models/ppo_liquid_409/save_40/actor.pt b/resources/trained_models/ppo_liquid_409/save_40/actor.pt new file mode 100644 index 00000000..c412e1de Binary files /dev/null and b/resources/trained_models/ppo_liquid_409/save_40/actor.pt differ diff --git a/resources/trained_models/ppo_liquid_409/save_57/actor.pt b/resources/trained_models/ppo_liquid_409/save_57/actor.pt new file mode 100644 index 00000000..661f8d03 Binary files /dev/null and b/resources/trained_models/ppo_liquid_409/save_57/actor.pt differ diff --git a/resources/trained_models/ppo_run_386_save_30/actor.pt b/resources/trained_models/ppo_run_386_save_30/actor.pt deleted file mode 100644 index 721aa013..00000000 Binary files a/resources/trained_models/ppo_run_386_save_30/actor.pt and /dev/null differ diff --git a/resources/trained_models/ppo_run_386_save_66/actor.pt b/resources/trained_models/ppo_run_386_save_66/actor.pt deleted file mode 100644 index 1d29ce48..00000000 Binary files a/resources/trained_models/ppo_run_386_save_66/actor.pt and /dev/null differ diff --git a/version.txt b/version.txt index e8ea05db..c813fe11 100644 --- a/version.txt +++ b/version.txt @@ -1 +1 @@ -1.2.4 +1.2.5