hyperverse 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- hyperdrone/__init__.py +31 -0
- hyperdrone/_native/cmake/hyperdrone_jit.cmake +65 -0
- hyperdrone/_native/common/bindings.h +39 -0
- hyperdrone/_native/common/cuda_staging.h +75 -0
- hyperdrone/_native/common/jit_host.h +69 -0
- hyperdrone/_native/common/models.h +24 -0
- hyperdrone/_native/dynamics/CMakeLists.txt +61 -0
- hyperdrone/_native/dynamics/host.cpp +193 -0
- hyperdrone/_native/dynamics/iface.h +76 -0
- hyperdrone/_native/dynamics/impl.cpp +608 -0
- hyperdrone/_native/env/CMakeLists.txt +71 -0
- hyperdrone/_native/env/environment.cpp +124 -0
- hyperdrone/_native/env/environment.h +407 -0
- hyperdrone/_native/env/rotorcraft_bindings.cpp +51 -0
- hyperdrone/_native/render/CMakeLists.txt +94 -0
- hyperdrone/_native/render/renderer.cpp +962 -0
- hyperdrone/_native/render/rig_bindings.h +151 -0
- hyperdrone/_native/render/scene_bindings.cpp +466 -0
- hyperdrone/_vendor/rl-tools/CMakeLists.txt +205 -0
- hyperdrone/_vendor/rl-tools/LICENSE +21 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/all.cmake +34 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/git-diff.cmake +100 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/git-hash.cmake +24 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/summary.cmake +64 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier0-compiler.cmake +18 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier1-blas.cmake +51 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier2-json-hdf5-zlib.cmake +81 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier3-tensorboard.cmake +25 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier4-cuda.cmake +42 -0
- hyperdrone/_vendor/rl-tools/cmake/autodetect/tier5-cli11.cmake +24 -0
- hyperdrone/_vendor/rl-tools/cmake/dependencies/assimp.cmake +19 -0
- hyperdrone/_vendor/rl-tools/cmake/legacy_flags.cmake +33 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/assimp.cmake +30 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/googletest.cmake +20 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/metal.cmake +25 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/mujoco.cmake +89 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/optix.cmake +25 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/raytracing.cmake +105 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/raytracing_generic.cmake +2 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/vulkan.cmake +4 -0
- hyperdrone/_vendor/rl-tools/cmake/optional/webgpu.cmake +24 -0
- hyperdrone/_vendor/rl-tools/cmake/scripts/embed_text.cmake +7 -0
- hyperdrone/_vendor/rl-tools/cmake/scripts/generate-git-snapshot.cmake +138 -0
- hyperdrone/_vendor/rl-tools/include/conta/conta.h +1161 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/matrix.h +143 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_arm.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cpu.h +314 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cpu_accelerate.h +37 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cpu_blas.h +67 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cpu_mkl.h +38 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cpu_openblas.h +34 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_cuda.h +289 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_dummy.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_esp32.h +29 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_generic.h +944 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/operations_wasm32.h +13 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/matrix/persist_code.h +136 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_arm.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cpu.h +94 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cpu_accelerate.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cpu_blas.h +115 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cpu_mkl.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cpu_openblas.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_cuda.h +713 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/operations_generic.h +1356 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/persist_code.h +130 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/containers/tensor/tensor.h +490 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/arm.h +91 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu.h +209 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu_accelerate.h +23 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu_blas.h +21 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu_mkl.h +45 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu_openblas.h +56 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cpu_tensorboard.h +34 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/cuda.h +415 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/devices.h +97 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/dummy.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/esp32.h +85 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/generic.h +0 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/devices/wasm32.h +59 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/model.h +218 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/operations_generic.h +769 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/persist.h +278 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/policy_adapter.h +100 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/policy_adapter_persist.h +52 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/dyn/tensor_operations_generic.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/applications/l2f/c_backend.h +117 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/applications/l2f/c_interface.h +36 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/applications/l2f/l2f.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/applications/l2f/operations_dyn.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/applications/l2f/operations_generic.h +268 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/debugging_pool/c_backend.h +31 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/debugging_pool/c_interface.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/executor/c_backend.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/executor/c_interface.h +50 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/executor/executor.h +133 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/executor/helper.h +109 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/inference/executor/operations_generic.h +219 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_arduino.h +117 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_arm.h +56 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_cpu.h +93 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_cpu_tensorboard.h +151 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_cuda.h +65 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_dummy.h +45 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/logging/operations_wasm32.h +56 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_arm.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_cpu.h +138 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_cuda.h +380 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_dummy.h +93 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_esp32.h +92 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_generic.h +295 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/math/operations_wasm32.h +230 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/mode/mode.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/activation_functions.h +113 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/capability/capability.h +61 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/capability/persist_code.h +31 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/group.h +99 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/avg_pool2d/layer.h +122 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/avg_pool2d/operations_cuda.h +64 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/avg_pool2d/operations_generic.h +186 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/avg_pool2d/persist.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/layer.h +300 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/operations_cpu_mkl.h +590 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/operations_cuda.h +809 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/operations_generic.h +1219 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/persist.h +99 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/conv2d/persist_code.h +170 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/layer.h +145 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_arm/dsp.h +63 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_arm/opt.h +173 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cpu.h +13 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cpu_accelerate.h +66 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cpu_blas.h +282 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cpu_mkl.h +66 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cpu_openblas.h +65 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_cuda.h +519 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_dummy.h +8 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_esp32/dsp.h +58 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_esp32/opt.h +167 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/operations_generic.h +429 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/persist.h +62 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/persist_code.h +178 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dense/persist_common.h +59 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dynamic_conv2d/layer.h +188 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dynamic_conv2d/operations_cuda.h +356 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dynamic_conv2d/operations_generic.h +536 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/dynamic_conv2d/persist.h +32 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/embedding/layer.h +119 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/embedding/operations_generic.h +223 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/embedding/persist.h +46 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/flatten/layer.h +121 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/flatten/operations_generic.h +154 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/flatten/persist.h +44 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/flatten/persist_code.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/helper_operations_cuda.h +602 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/helper_operations_generic.h +251 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/layer.h +206 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/operations_generic.h +889 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/persist.h +75 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/gru/persist_code.h +199 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/layers.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/max_pool2d/layer.h +143 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/max_pool2d/operations_cuda.h +80 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/max_pool2d/operations_generic.h +216 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/max_pool2d/persist.h +39 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cpu.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cpu_accelerate.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cpu_blas.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cpu_mkl.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cpu_openblas.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_cuda.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_dummy.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/operations_generic.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/resnet_block/layer.h +223 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/resnet_block/operations_cuda.h +138 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/resnet_block/operations_generic.h +446 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/resnet_block/persist.h +70 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/sample_and_squash/layer.h +154 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/sample_and_squash/operations_cuda.h +119 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/sample_and_squash/operations_generic.h +502 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/sample_and_squash/persist.h +50 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/sample_and_squash/persist_code.h +173 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/standardize/layer.h +116 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/standardize/operations_cuda.h +156 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/standardize/operations_generic.h +295 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/standardize/persist.h +83 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/standardize/persist_code.h +163 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/td3_sampling/layer.h +123 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/td3_sampling/operations_generic.h +272 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/td3_sampling/persist.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/td3_sampling/persist_code.h +164 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/unflatten/layer.h +122 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/unflatten/operations_cuda.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/unflatten/operations_generic.h +160 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/unflatten/persist.h +47 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/unflatten/persist_code.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/upsample2d/layer.h +132 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/upsample2d/operations_cuda.h +141 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/layers/upsample2d/operations_generic.h +230 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/loss_functions/categorical_cross_entropy/operations_generic.h +102 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/loss_functions/mse/operations_cuda.h +97 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/loss_functions/mse/operations_generic.h +101 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/loss_functions/operations_generic.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/nn.h +14 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu_accelerate.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu_blas.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu_mkl.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu_mux.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cpu_openblas.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_cuda.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_dummy.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/operations_generic.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/adam.h +121 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/instance/operations_cuda.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/instance/operations_generic.h +150 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/instance/persist.h +39 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/instance/persist_code.h +80 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/operations_cuda.h +85 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/operations_generic.h +109 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/adam/persist.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/lamb/instance/operations_generic.h +107 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/lamb/lamb.h +40 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/lamb/operations_generic.h +68 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/optimizers.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/instance/operations_cuda.h +85 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/instance/operations_generic.h +140 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/instance/persist.h +37 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/instance/persist_code.h +72 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/operations_cuda.h +49 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/operations_generic.h +50 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/sgd/sgd.h +92 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/optimizers/update.h +54 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/parameters/operations_cuda.h +0 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/parameters/operations_generic.h +98 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/parameters/parameters.h +79 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/parameters/persist.h +34 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/parameters/persist_code.h +136 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/persist.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn/persist_code.h +27 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/behaviors/skip_parameter_gradients.h +65 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/group/README.md +107 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/group/model.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/group/persist_code.h +77 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/network.h +174 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/operations_cuda.h +10 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/operations_dummy.h +8 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/operations_generic.h +407 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/persist.h +56 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp/persist_code.h +137 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp_unconditional_stddev/network.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp_unconditional_stddev/operations_generic.h +86 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp_unconditional_stddev/persist.h +43 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/mlp_unconditional_stddev/persist_code.h +142 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/models.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/multi_agent_wrapper/model.h +148 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/multi_agent_wrapper/operations_generic.h +288 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/multi_agent_wrapper/persist.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/multi_agent_wrapper/persist_code.h +71 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/operations_cpu.h +0 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/operations_cuda.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/operations_dummy.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/operations_generic.h +4 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/parallel/model.h +290 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/parallel/operations_cuda.h +174 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/parallel/operations_generic.h +623 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/parallel/persist.h +80 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/parallel/persist_code.h +200 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/persist.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/persist_code.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/random_uniform/model.h +43 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/random_uniform/operations_generic.h +73 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/resnet/resnet.h +97 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/sequential/model.h +270 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/sequential/operations_generic.h +625 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/sequential/persist.h +99 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/nn_models/sequential/persist_code.h +115 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/numeric_types/bf16.h +28 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/numeric_types/categories.h +22 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/numeric_types/persist_code.h +54 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/numeric_types/policy.h +48 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/arm/group_1.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/arm/group_2.h +14 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/arm/group_3.h +14 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/arm.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu/group_1.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu/group_3.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_accelerate/group_1.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_accelerate/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_accelerate/group_3.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_accelerate.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_mkl/group_1.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_mkl/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_mkl/group_3.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_mkl.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_mux.h +148 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_openblas/group_1.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_openblas/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_openblas/group_3.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_openblas.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_tensorboard/group_1.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_tensorboard/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_tensorboard/group_3.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cpu_tensorboard.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cuda/group_1.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cuda/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cuda/group_3.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/cuda.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/dummy/group_1.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/dummy/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/dummy/group_3.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/dummy.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/esp32/group_1.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/esp32/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/esp32/group_3.h +13 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/esp32.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/wasm32/group_1.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/wasm32/group_2.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/wasm32/group_3.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/operations/wasm32.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/hdf5/hdf5.h +54 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/hdf5/operations_cpu.h +419 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/tar/io.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/tar/operations_cpu.h +58 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/tar/operations_generic.h +620 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/tar/operations_posix.h +63 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/backends/tar/tar.h +74 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/persist/code.h +38 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_arm.h +39 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_cpu.h +81 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_cuda.h +169 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_dummy.h +72 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_esp32.h +45 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_generic.h +133 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_generic_array.h +73 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/operations_wasm32.h +57 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/random/persist.h +77 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/camera.h +105 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/annotations/cache.h +26 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/annotations/free_space.h +60 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/annotations/operations_cpu.h +408 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/glb/operations_cpu.h +1313 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/operations_cpu.h +96 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/procthor/conversion.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/procthor/operations_cpu.h +141 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/datasets/procthor/procthor.h +40 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/generic/operations_cpu.h +733 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/generic/operations_generic.h +1299 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/metal/context.h +198 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/metal/device_source.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/metal/operations_cpu.h +1140 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/optix/device.h +233 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/optix/device_impl.h +1084 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/optix/operations_cuda.h +1808 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/optix/overlay_accel.h +64 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/specialization.h +38 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/vulkan/context.h +248 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/vulkan/device_source.h +48 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/vulkan/operations_cpu.h +2056 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/webgpu/bvh_sah.h +182 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/webgpu/context.h +303 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/webgpu/device_source.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/backends/webgpu/operations_cpu.h +1455 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/operations_cpu_common.h +916 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/operations_cpu_mux.h +31 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/operations_cpu_post.h +85 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/renderer.h +478 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/save_cpu.h +315 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/scene.h +44 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/transforms_generic.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/raytracing/types.h +45 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/rig/operations_cpu.h +119 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/rig/operations_generic.h +121 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/rig/rig.h +70 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/scene.h +134 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/segmentation_cpu.h +145 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/transforms.h +204 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rendering/types.h +55 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/algorithms.h +6 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/operations_cpu.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/operations_cpu_mkl.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/operations_generic.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/loop/core/config.h +243 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/loop/core/operations_cuda.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/loop/core/operations_generic.h +205 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/loop/core/persist.h +109 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/loop/core/state.h +51 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/mixed_imitation.h +63 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/operations_cuda.h +190 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/operations_generic.h +332 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/operations_generic_mixed_imitation.h +81 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/persist.h +28 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/ppo/ppo.h +105 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/loop/core/approximators_gru.h +84 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/loop/core/approximators_mlp.h +61 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/loop/core/config.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/loop/core/operations_generic.h +193 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/loop/core/state.h +52 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_cpu.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_cpu_accelerate.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_cpu_mkl.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_cpu_mux.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_cuda.h +131 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/operations_generic.h +580 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/sac/sac.h +151 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/loop/core/approximators_gru.h +72 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/loop/core/approximators_mlp.h +54 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/loop/core/config.h +107 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/loop/core/operations_generic.h +174 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/loop/core/state.h +46 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_cpu.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_cpu_accelerate.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_cpu_mkl.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_cpu_mux.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_cuda.h +117 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/operations_generic.h +359 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/algorithms/td3/td3.h +141 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/components.h +4 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/off_policy_runner.h +225 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_cpu.h +95 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_cpu_accelerate.h +21 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_cpu_mkl.h +21 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_cuda.h +237 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_generic.h +498 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/off_policy_runner/operations_generic_per_env.h +112 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/collection.h +57 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/on_policy_runner.h +122 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_cpu.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_cpu_accelerate.h +8 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_cpu_mkl.h +8 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_cpu_mux.h +12 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_cuda.h +183 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_generic.h +366 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/operations_generic_common.h +74 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/on_policy_runner/persist.h +47 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/operations_cpu.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/operations_cpu_accelerate.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/operations_cpu_mkl.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/operations_generic.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/replay_buffer/operations_cpu.h +9 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/replay_buffer/operations_generic.h +123 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/replay_buffer/persist.h +57 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/replay_buffer/replay_buffer.h +93 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/running_normalizer/operations_generic.h +76 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/running_normalizer/persist.h +31 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/components/running_normalizer/running_normalizer.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environment_wrappers/operations_generic.h +58 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environment_wrappers/scale_observations/operations_generic.h +28 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environment_wrappers/scale_observations/wrapper.h +26 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environment_wrappers/wrappers.h +29 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/acrobot/acrobot.h +88 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/acrobot/operations_cpu.h +152 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/acrobot/operations_generic.h +202 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/batch/environment.h +40 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/batch/operations_cuda.h +204 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/batch/operations_generic.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/batch/operations_generic_common.h +26 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/batch/persist.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/car.h +119 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/operations_cpu.h +110 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/operations_generic.h +206 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/operations_json.h +59 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/track.h +23 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/car/ui.h +189 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/environments.h +28 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/flag/environment.h +87 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/flag/operations_cpu.h +212 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/flag/operations_generic.h +215 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/observation.h +75 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/operations_cpu.h +938 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/operations_cuda.h +380 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/pose.h +228 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/presets.h +148 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/rig/operations_cpu.h +189 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/rig/operations_generic.h +32 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/rig/rig.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/moving_gate/moving_gate.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/moving_gate/operations_cpu.h +262 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/moving_gate/operations_cuda.h +254 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/target_frame/operations_cpu.h +224 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/target_frame/operations_cuda.h +218 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/target_frame/operations_generic_position_hold.h +128 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/target_frame/position_hold.h +22 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/target_frame/target_frame.h +77 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/autopilot.h +61 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/baseline.h +46 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/calibration.h +62 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/metrics.h +101 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/operations_cpu.h +389 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/tasks/visual_inertial_localization/visual_inertial_localization.h +123 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/hyperdrone/world.h +318 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/inertial_velocity.h +27 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/metrics.h +266 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/multirotor.h +1138 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_cpu.h +2246 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/05_state_is_nan.h +148 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/10_sample_initial_parameters.h +228 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/20_initial_state.h +208 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/30_sample_initial_state.h +300 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/35_get_desired_state.h +63 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/40_observe.h +637 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/50_state_algebra.h +57 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/60_dynamics.h +134 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/70_post_integration.h +368 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic/80_abs_diff.h +312 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic.h +367 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_generic_inertial_velocity.h +44 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_helper_generic.h +43 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_multitask_generic.h +70 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/operations_multitask_generic_forward.h +24 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/default.h +209 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/arpl.h +112 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/crazyflie.h +127 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/crazyflie_openmv.h +119 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/flightmare.h +111 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/fs.h +113 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/mrs.h +110 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/soft.h +117 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/soft_rigid.h +118 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/x500.h +88 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/x500_real.h +88 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/dynamics/x500_sim.h +114 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/init/default.h +56 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/registry.h +96 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/reward_functions/default.h +64 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/reward_functions/reward_functions.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/reward_functions/squared/operations_generic.h +202 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/reward_functions/squared/squared.h +52 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/termination/default.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/trajectories/lissajous.h +72 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/parameters/trajectories/trajectory.h +78 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/persist.h +21 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/persist_code.h +53 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/quaternion_helper.h +103 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/l2f/ui.h +180 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/memory/environment.h +69 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/memory/operations_cpu.h +46 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/memory/operations_generic.h +78 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/README.MD +5 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/ant.h +77 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/model.h +415 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/operations_cpu.h +162 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/persist.h +86 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/ant/ui.h +141 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/mujoco/mujoco.h +1 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/multi_agent/bottleneck/bottleneck.h +113 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/multi_agent/bottleneck/operations_cpu.h +289 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/multi_agent/bottleneck/operations_generic.h +495 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/multi_agent/environments.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/observation.h +37 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/operations_cpu.h +2 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/operations_generic.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/pendulum/operations_cpu.h +146 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/pendulum/operations_generic.h +262 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/pendulum/pendulum.h +197 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/pendulum/ui.h +97 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/pendulum/ui_xeus.h +124 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/reacher/operations_cpu.h +426 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/reacher/operations_generic.h +424 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/environments/reacher/reacher.h +137 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/loop.h +15 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/checkpoint/config.h +28 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/checkpoint/operations_cpu.h +233 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/checkpoint/persist.h +30 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/checkpoint/state.h +21 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/curriculum/config.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/curriculum/operations_generic.h +47 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/curriculum/persist.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/curriculum/state.h +17 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/evaluation/config.h +47 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/evaluation/operations_generic.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/evaluation/persist.h +40 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/evaluation/state.h +36 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/extrack/config.h +22 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/extrack/operations_cpu.h +44 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/extrack/persist.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/extrack/state.h +26 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/nn_analytics/config.h +37 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/nn_analytics/operations_cpu.h +113 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/nn_analytics/persist.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/nn_analytics/state.h +24 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/save_trajectories/config.h +43 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/save_trajectories/operations_cpu.h +190 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/save_trajectories/persist.h +30 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/save_trajectories/state.h +37 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/steps.h +4 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/timing/config.h +27 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/timing/operations_cpu.h +52 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/timing/persist.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/loop/steps/timing/state.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/operations_generic.h +3 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/rl.h +5 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/utils/evaluation/evaluation.h +106 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/utils/evaluation/operations_cpu.h +48 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/utils/evaluation/operations_generic.h +275 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/utils/validation.h +269 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl/utils/validation_analysis.h +85 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/rl_tools.h +80 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/ui_server/client/client.h +50 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/ui_server/client/operations_boost.h +92 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/ui_server/client/operations_cpu.h +174 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/ui_server/client/operations_websocket.h +208 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/ui_server/server.h +448 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/universe.h +4 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/declarations_cpu.h +13 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_arm.h +29 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_cpu.h +25 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_cuda.h +33 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_dummy.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_esp32.h +20 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_generic.h +19 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/assert/operations_wasm32.h +18 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/env/operations_cpu.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/env/operations_generic.h +14 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/extrack/extrack.h +129 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/extrack/operations_cpu.h +588 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/integrators.h +103 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/memcpy.h +16 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/tuple/operations_generic.h +48 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/tuple/tuple.h +147 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/typing.h +122 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/generic/vector_operations.h +172 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/polyak/operations_cuda.h +53 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/polyak/operations_generic.h +93 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/string/operations_generic.h +201 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/utils/zlib/operations_cpu.h +45 -0
- hyperdrone/_vendor/rl-tools/include/rl_tools/version.h +28 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/CMakeLists.txt +15 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/README.md +193 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/CMakeLists.txt +11 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/generic/CMakeLists.txt +17 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/generic/freestanding_check.cpp +135 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/metal/CMakeLists.txt +31 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/metal/device.metal +953 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/metal/metal_cpp_impl.cpp +6 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/optix/CMakeLists.txt +51 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/optix/device.cu +1 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/optix/overlay_accel.cu +311 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/vulkan/CMakeLists.txt +37 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/vulkan/device.comp +1013 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/vulkan/device_spirv.cpp +44 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/webgpu/CMakeLists.txt +31 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/backends/webgpu/device.wgsl +1448 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/CMakeLists.txt +110 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/README.md +55 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/benchmark.cpp +180 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/benchmark_dynamic.cpp +491 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/rig.cpp +119 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/simulator_matrix.cpp +1130 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/benchmark/simulator_matrix_config.h.in +13 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/camera_orbit.h +51 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/CMakeLists.txt +199 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/antialiasing.cpp +214 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/drone.cpp +320 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/drone_device.cu +359 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/environment/environment.h +96 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/environment/operations_cpu.h +263 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/example.cpp +273 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/fixed_pose.cpp +221 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/flow.cpp +120 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/interactive.cpp +1015 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/make_drone_glb.py +172 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/minimal.cpp +82 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/minimal_overlay.cpp +137 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/motion_blur.cpp +232 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/render_pose_trace.cpp +1829 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/render_pose_trace_motion_blur.cpp +647 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/web/assimp/CMakeLists.txt +8 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/web/build.sh +80 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/example/web/drone_web.cpp +234 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/CMakeLists.txt +147 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/README.md +15 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/job.sbatch +28 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model.h +260 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model_config.h +24 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model_forward_cuda.h +397 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model_operations.h +418 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model_persist.h +45 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/model_student.h +70 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/scene.cpp +222 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/scene.h +68 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/.gitignore +4 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/CMakeLists.txt +4 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Info.plist +24 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Info_ios.plist +36 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Sources/CameraManager.swift +120 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Sources/ContentView.swift +703 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Sources/DatasetRecorder.swift +285 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/Sources/YawPredictorApp.swift +25 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/YawPredictor.entitlements +8 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/build.sh +66 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/build_ios.sh +133 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/generate.sh +6 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/h5_to_tar.cpp +82 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/inference.cpp +166 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/inference.h +26 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/project.yml +30 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/swift_app/test_tar_vs_hdf5.cpp +173 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_cuda.cu +699 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_dataset_overlay.cpp +475 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_distill_cuda.cu +995 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_eval.cu +457 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_real_world_eval.cpp +660 -0
- hyperdrone/_vendor/rl-tools/src/rendering/raytracing/yaw_prediction/yaw_prediction_viewer.cu +443 -0
- hyperdrone/cuda/__init__.py +78 -0
- hyperdrone/dynamics/__init__.py +22 -0
- hyperdrone/dynamics/_component.py +35 -0
- hyperdrone/dynamics/_config.py +48 -0
- hyperdrone/dynamics/_sim.py +195 -0
- hyperdrone/env/__init__.py +17 -0
- hyperdrone/env/_multi_environment.py +407 -0
- hyperdrone/env/_rotorcraft.py +31 -0
- hyperdrone/examples/.gitignore +2 -0
- hyperdrone/examples/Benchmark.ipynb +102 -0
- hyperdrone/examples/HyperDrone.ipynb +514 -0
- hyperdrone/examples/Orbiter.ipynb +462 -0
- hyperdrone/examples/Renderer.ipynb +173 -0
- hyperdrone/examples/__init__.py +16 -0
- hyperdrone/examples/articulation.py +55 -0
- hyperdrone/examples/benchmark.py +236 -0
- hyperdrone/examples/benchmark_minimal.py +65 -0
- hyperdrone/examples/data.py +70 -0
- hyperdrone/examples/drone_flythrough.py +169 -0
- hyperdrone/examples/hyperdrone_videos/my_world_preview.h +17 -0
- hyperdrone/examples/my_world.h +17 -0
- hyperdrone/examples/render_procthor.py +41 -0
- hyperdrone/examples/visual_inertial_localization.py +159 -0
- hyperdrone/gym/__init__.py +63 -0
- hyperdrone/jit/__init__.py +129 -0
- hyperdrone/jit/_build.py +113 -0
- hyperdrone/jit/_lock.py +23 -0
- hyperdrone/jit/_patches.py +99 -0
- hyperdrone/jit/_toolchain.py +17 -0
- hyperdrone/jit/_workspace.py +68 -0
- hyperdrone/render/__init__.py +53 -0
- hyperdrone/render/_component.py +51 -0
- hyperdrone/render/_config.py +92 -0
- hyperdrone/render/_renderer.py +447 -0
- hyperdrone/render/_rig.py +136 -0
- hyperdrone/render/_sampling.py +69 -0
- hyperdrone/render/_scene.py +36 -0
- hyperverse-0.0.1.dist-info/METADATA +370 -0
- hyperverse-0.0.1.dist-info/RECORD +746 -0
- hyperverse-0.0.1.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,151 @@
|
|
|
1
|
+
#include "../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_SAC_SAC_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_SAC_SAC_H
|
|
5
|
+
//#include "../../../nn_models/output_view/model.h"
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
9
|
+
namespace rl_tools::rl::algorithms::sac {
|
|
10
|
+
template<typename TYPE_POLICY, typename TI, TI ACTION_DIM=1>
|
|
11
|
+
struct DefaultParameters {
|
|
12
|
+
using T = typename TYPE_POLICY::DEFAULT;
|
|
13
|
+
static constexpr T GAMMA = 0.99;
|
|
14
|
+
static constexpr TI ACTOR_BATCH_SIZE = 32;
|
|
15
|
+
static constexpr TI CRITIC_BATCH_SIZE = 32;
|
|
16
|
+
static constexpr TI CRITIC_TRAINING_INTERVAL = 1;
|
|
17
|
+
static constexpr TI ACTOR_TRAINING_INTERVAL = 1;
|
|
18
|
+
static constexpr TI CRITIC_TARGET_UPDATE_INTERVAL = 1;
|
|
19
|
+
static constexpr T ACTOR_POLYAK = 1.0 - 0.005;
|
|
20
|
+
static constexpr T CRITIC_POLYAK = 1.0 - 0.005;
|
|
21
|
+
static constexpr bool IGNORE_TERMINATION = false; // ignoring the termination flag is useful for training on environments with negative rewards, where the agent would try to terminate the episode as soon as possible otherwise
|
|
22
|
+
static constexpr TI SEQUENCE_LENGTH = 1; // note that this implementation does only show next_observation sequences to the target actor and critic. Hence they have one step (the initial one in the sequence) less information. This makes the sequence length deterministic (otherwise it would depend on the number of resets in the batch). For most environments and for larger sequences the information gain should be negligible but for some (mostly artifiical) environments the first state matters (e.g. the FlagMemory environment). A possible mitigations is repeating the initial observation in the environment
|
|
23
|
+
static constexpr bool ENTROPY_BONUS = true;
|
|
24
|
+
static constexpr bool ENTROPY_BONUS_NEXT_STEP = true;
|
|
25
|
+
static constexpr bool MASK_NON_TERMINAL = true;
|
|
26
|
+
|
|
27
|
+
static constexpr T TARGET_ENTROPY = -((T)ACTION_DIM);
|
|
28
|
+
static constexpr T ALPHA = 0.5;
|
|
29
|
+
static constexpr bool ADAPTIVE_ALPHA = true;
|
|
30
|
+
static constexpr T LOG_STD_LOWER_BOUND = -20;
|
|
31
|
+
static constexpr T LOG_STD_UPPER_BOUND = 2;
|
|
32
|
+
static constexpr T LOG_PROBABILITY_EPSILON = 1e-6;
|
|
33
|
+
};
|
|
34
|
+
|
|
35
|
+
template<
|
|
36
|
+
typename T_TYPE_POLICY,
|
|
37
|
+
typename T_TI,
|
|
38
|
+
typename T_ENVIRONMENT,
|
|
39
|
+
typename T_ACTOR_NETWORK_TYPE,
|
|
40
|
+
typename T_CRITIC_NETWORK_TYPE,
|
|
41
|
+
typename T_CRITIC_TARGET_NETWORK_TYPE,
|
|
42
|
+
typename T_ALPHA_PARAMETER_TYPE,
|
|
43
|
+
typename T_ACTOR_OPTIMIZER,
|
|
44
|
+
typename T_CRITIC_OPTIMIZER,
|
|
45
|
+
typename T_ALPHA_OPTIMIZER,
|
|
46
|
+
typename T_PARAMETERS,
|
|
47
|
+
bool T_INCLUDE_FIRST_STEP_IN_TARGETS
|
|
48
|
+
>
|
|
49
|
+
struct Specification{
|
|
50
|
+
using TYPE_POLICY = T_TYPE_POLICY;
|
|
51
|
+
using TI = T_TI;
|
|
52
|
+
using ENVIRONMENT = T_ENVIRONMENT;
|
|
53
|
+
using ACTOR_NETWORK_TYPE = T_ACTOR_NETWORK_TYPE;
|
|
54
|
+
using CRITIC_NETWORK_TYPE = T_CRITIC_NETWORK_TYPE;
|
|
55
|
+
using CRITIC_TARGET_NETWORK_TYPE = T_CRITIC_TARGET_NETWORK_TYPE;
|
|
56
|
+
using ALPHA_PARAMETER_TYPE = T_ALPHA_PARAMETER_TYPE;
|
|
57
|
+
using ACTOR_OPTIMIZER = T_ACTOR_OPTIMIZER;
|
|
58
|
+
using CRITIC_OPTIMIZER = T_CRITIC_OPTIMIZER;
|
|
59
|
+
using ALPHA_OPTIMIZER = T_ALPHA_OPTIMIZER;
|
|
60
|
+
using PARAMETERS = T_PARAMETERS;
|
|
61
|
+
static constexpr bool INCLUDE_FIRST_STEP_IN_TARGETS = T_INCLUDE_FIRST_STEP_IN_TARGETS;
|
|
62
|
+
};
|
|
63
|
+
|
|
64
|
+
template <typename T_SPEC, bool T_DYNAMIC_ALLOCATION>
|
|
65
|
+
struct ActorTrainingBuffersSpecification{
|
|
66
|
+
using SPEC = T_SPEC;
|
|
67
|
+
static constexpr bool DYNAMIC_ALLOCATION = T_DYNAMIC_ALLOCATION;
|
|
68
|
+
};
|
|
69
|
+
|
|
70
|
+
template<typename T_SPEC>
|
|
71
|
+
struct ActorTrainingBuffers{
|
|
72
|
+
using SPEC = typename T_SPEC::SPEC;
|
|
73
|
+
using TYPE_POLICY = typename SPEC::TYPE_POLICY;
|
|
74
|
+
using T = typename TYPE_POLICY::template GET<numeric_types::categories::Buffer>;
|
|
75
|
+
using TI = typename SPEC::TI;
|
|
76
|
+
static constexpr bool DYNAMIC_ALLOCATION = T_SPEC::DYNAMIC_ALLOCATION;
|
|
77
|
+
static constexpr TI SEQUENCE_LENGTH = SPEC::PARAMETERS::SEQUENCE_LENGTH;
|
|
78
|
+
static constexpr TI BATCH_SIZE = SPEC::PARAMETERS::ACTOR_BATCH_SIZE;
|
|
79
|
+
static constexpr TI ACTOR_INPUT_DIM = get_last(typename SPEC::ACTOR_NETWORK_TYPE::INPUT_SHAPE{});
|
|
80
|
+
static constexpr TI ACTION_DIM = SPEC::ENVIRONMENT::ACTION_DIM;
|
|
81
|
+
static constexpr TI CRITIC_OBSERVATION_DIM = get_last(typename SPEC::CRITIC_NETWORK_TYPE::INPUT_SHAPE{}) - SPEC::ENVIRONMENT::ACTION_DIM;
|
|
82
|
+
|
|
83
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, CRITIC_OBSERVATION_DIM + ACTION_DIM>, DYNAMIC_ALLOCATION>> state_action_value_input;
|
|
84
|
+
template<typename SPEC::TI DIM>
|
|
85
|
+
using STATE_ACTION_VALUE_VIEW = typename decltype(state_action_value_input)::template VIEW_RANGE<tensor::ViewSpec<2, DIM>>;
|
|
86
|
+
STATE_ACTION_VALUE_VIEW<CRITIC_OBSERVATION_DIM> observations;
|
|
87
|
+
STATE_ACTION_VALUE_VIEW<ACTION_DIM> actions;
|
|
88
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> d_output;
|
|
89
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, CRITIC_OBSERVATION_DIM + ACTION_DIM>, DYNAMIC_ALLOCATION>> d_critic_1_input, d_critic_2_input;
|
|
90
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTION_DIM>, DYNAMIC_ALLOCATION>> d_critic_action_input;
|
|
91
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTION_DIM>, DYNAMIC_ALLOCATION>> action_sample, action_noise;
|
|
92
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTION_DIM>, DYNAMIC_ALLOCATION>> d_actor_output_squashing;
|
|
93
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTION_DIM * 2>, DYNAMIC_ALLOCATION>> d_squashing_input;
|
|
94
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTION_DIM * 2>, DYNAMIC_ALLOCATION>> d_actor_output;
|
|
95
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, ACTOR_INPUT_DIM>, DYNAMIC_ALLOCATION>> d_actor_input;
|
|
96
|
+
Tensor<tensor::Specification<T, TI, tensor::Shape<TI, 1>, DYNAMIC_ALLOCATION>> loss_weight;
|
|
97
|
+
};
|
|
98
|
+
template <typename T_SPEC, bool T_DYNAMIC_ALLOCATION>
|
|
99
|
+
struct CriticTrainingBuffersSpecification{
|
|
100
|
+
using SPEC = T_SPEC;
|
|
101
|
+
static constexpr bool DYNAMIC_ALLOCATION = T_DYNAMIC_ALLOCATION;
|
|
102
|
+
};
|
|
103
|
+
template<typename T_SPEC>
|
|
104
|
+
struct CriticTrainingBuffers{
|
|
105
|
+
using SPEC = typename T_SPEC::SPEC;
|
|
106
|
+
using TYPE_POLICY = typename SPEC::TYPE_POLICY;
|
|
107
|
+
using T_BUFFER = typename TYPE_POLICY::template GET<numeric_types::categories::Buffer>;
|
|
108
|
+
using TI = typename SPEC::TI;
|
|
109
|
+
static constexpr bool DYNAMIC_ALLOCATION = T_SPEC::DYNAMIC_ALLOCATION;
|
|
110
|
+
static constexpr TI SEQUENCE_LENGTH = SPEC::PARAMETERS::SEQUENCE_LENGTH;
|
|
111
|
+
static constexpr TI NEXT_SEQUENCE_LENGTH = SPEC::INCLUDE_FIRST_STEP_IN_TARGETS ? SEQUENCE_LENGTH + 1 : SEQUENCE_LENGTH;
|
|
112
|
+
static constexpr TI BATCH_SIZE = SPEC::PARAMETERS::CRITIC_BATCH_SIZE;
|
|
113
|
+
static constexpr TI ACTION_DIM = SPEC::ENVIRONMENT::ACTION_DIM;
|
|
114
|
+
static constexpr TI CRITIC_OBSERVATION_DIM = get_last(typename SPEC::CRITIC_NETWORK_TYPE::INPUT_SHAPE{}) - SPEC::ENVIRONMENT::ACTION_DIM;
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, NEXT_SEQUENCE_LENGTH, BATCH_SIZE, CRITIC_OBSERVATION_DIM + ACTION_DIM>, DYNAMIC_ALLOCATION>> next_state_action_value_input;
|
|
118
|
+
template<typename SPEC::TI DIM>
|
|
119
|
+
using NEXT_STATE_ACTION_VALUE_VIEW = typename decltype(next_state_action_value_input)::template VIEW_RANGE<tensor::ViewSpec<2, DIM>>;
|
|
120
|
+
NEXT_STATE_ACTION_VALUE_VIEW<CRITIC_OBSERVATION_DIM> next_observations;
|
|
121
|
+
NEXT_STATE_ACTION_VALUE_VIEW<ACTION_DIM> next_actions;
|
|
122
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> action_value;
|
|
123
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> target_action_value;
|
|
124
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, NEXT_SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> next_state_action_value_critic_1;
|
|
125
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, NEXT_SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> next_state_action_value_critic_2;
|
|
126
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, CRITIC_OBSERVATION_DIM + ACTION_DIM>, DYNAMIC_ALLOCATION>> d_input;
|
|
127
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, SEQUENCE_LENGTH, BATCH_SIZE, 1>, DYNAMIC_ALLOCATION>> d_output;
|
|
128
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, 1>, DYNAMIC_ALLOCATION>> loss_weight;
|
|
129
|
+
Tensor<tensor::Specification<T_BUFFER, TI, tensor::Shape<TI, NEXT_SEQUENCE_LENGTH * BATCH_SIZE>, DYNAMIC_ALLOCATION>> next_action_log_probs;
|
|
130
|
+
};
|
|
131
|
+
|
|
132
|
+
template<typename T_SPEC>
|
|
133
|
+
struct ActorCritic {
|
|
134
|
+
using SPEC = T_SPEC;
|
|
135
|
+
// using T = typename SPEC::T;
|
|
136
|
+
using TI = typename SPEC::TI;
|
|
137
|
+
|
|
138
|
+
typename SPEC::ACTOR_NETWORK_TYPE actor;
|
|
139
|
+
typename SPEC::CRITIC_NETWORK_TYPE critics[2];
|
|
140
|
+
typename SPEC::CRITIC_TARGET_NETWORK_TYPE critics_target[2];
|
|
141
|
+
|
|
142
|
+
typename SPEC::ACTOR_OPTIMIZER actor_optimizer;
|
|
143
|
+
typename SPEC::CRITIC_OPTIMIZER critic_optimizers[2];
|
|
144
|
+
typename SPEC::ALPHA_OPTIMIZER alpha_optimizer;
|
|
145
|
+
};
|
|
146
|
+
}
|
|
147
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
#endif
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
#include "../../../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_APPROXIMATORS_GRU_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_APPROXIMATORS_GRU_H
|
|
5
|
+
|
|
6
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
7
|
+
namespace rl_tools::rl::algorithms::td3::loop::core{
|
|
8
|
+
template<typename T, typename TI, typename ENVIRONMENT, typename PARAMETERS, bool DYNAMIC_ALLOCATION>
|
|
9
|
+
struct ConfigApproximatorsGRU{
|
|
10
|
+
// static constexpr bool USE_GRU = true;
|
|
11
|
+
using TD3_PARAMETERS = typename PARAMETERS::TD3_PARAMETERS;
|
|
12
|
+
template <typename CAPABILITY>
|
|
13
|
+
struct Actor{
|
|
14
|
+
using INPUT_SHAPE = tensor::Shape<TI, TD3_PARAMETERS::SEQUENCE_LENGTH, TD3_PARAMETERS::ACTOR_BATCH_SIZE, ENVIRONMENT::Observation::DIM>;
|
|
15
|
+
using INPUT_LAYER_CONFIG = nn::layers::dense::Configuration<T, TI, PARAMETERS::ACTOR_HIDDEN_DIM, PARAMETERS::CRITIC_ACTIVATION_FUNCTION, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Input>;
|
|
16
|
+
using INPUT_LAYER = nn::layers::dense::BindConfiguration<INPUT_LAYER_CONFIG>;
|
|
17
|
+
using GRU_CONFIG = nn::layers::gru::Configuration<T, TI, PARAMETERS::ACTOR_HIDDEN_DIM, nn::parameters::groups::Normal, true>;
|
|
18
|
+
using GRU = nn::layers::gru::BindConfiguration<GRU_CONFIG>;
|
|
19
|
+
using GRU2_CONFIG = nn::layers::gru::Configuration<T, TI, PARAMETERS::ACTOR_HIDDEN_DIM, nn::parameters::groups::Normal, true>;
|
|
20
|
+
using GRU2 = nn::layers::gru::BindConfiguration<GRU2_CONFIG>;
|
|
21
|
+
using DENSE_LAYER_CONFIG = nn::layers::dense::Configuration<T, TI, PARAMETERS::ACTOR_HIDDEN_DIM, PARAMETERS::ACTOR_ACTIVATION_FUNCTION, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Normal>;
|
|
22
|
+
using DENSE_LAYER = nn::layers::dense::BindConfiguration<DENSE_LAYER_CONFIG>;
|
|
23
|
+
using OUTPUT_CONFIG = nn::layers::dense::Configuration<T, TI, ENVIRONMENT::ACTION_DIM, nn::activation_functions::ActivationFunction::IDENTITY, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Output>;
|
|
24
|
+
using OUTPUT = nn::layers::dense::BindConfiguration<OUTPUT_CONFIG>;
|
|
25
|
+
|
|
26
|
+
using MODULE_GRU = nn_models::sequential::Module<GRU, nn_models::sequential::Module<OUTPUT>>;
|
|
27
|
+
using MODULE_GRU_TWO_LAYER = nn_models::sequential::Module<GRU, nn_models::sequential::Module<GRU2, nn_models::sequential::Module<OUTPUT>>>;
|
|
28
|
+
using MODULE_GRU_THREE_LAYER = nn_models::sequential::Module<GRU, nn_models::sequential::Module<GRU2, nn_models::sequential::Module<DENSE_LAYER, nn_models::sequential::Module<OUTPUT>>>>;
|
|
29
|
+
|
|
30
|
+
using SELECTED_MODULE = rl_tools::utils::typing::conditional_t<PARAMETERS::CRITIC_NUM_LAYERS == 3, MODULE_GRU, rl_tools::utils::typing::conditional_t<PARAMETERS::CRITIC_NUM_LAYERS == 4, MODULE_GRU_TWO_LAYER, MODULE_GRU_THREE_LAYER>>;
|
|
31
|
+
static_assert(PARAMETERS::CRITIC_NUM_LAYERS == 3 || PARAMETERS::CRITIC_NUM_LAYERS == 4 || PARAMETERS::CRITIC_NUM_LAYERS == 5, "Only 3/4/5 layers (1 input + 1/2 GRU + 1/2 Output) are supported right now");
|
|
32
|
+
using MODEL = nn_models::sequential::Build<CAPABILITY, SELECTED_MODULE, INPUT_SHAPE>;
|
|
33
|
+
};
|
|
34
|
+
template <typename CAPABILITY>
|
|
35
|
+
struct Critic{
|
|
36
|
+
using INPUT_SHAPE = tensor::Shape<TI, TD3_PARAMETERS::SEQUENCE_LENGTH, TD3_PARAMETERS::CRITIC_BATCH_SIZE, ENVIRONMENT::ObservationPrivileged::DIM + ENVIRONMENT::ACTION_DIM>;
|
|
37
|
+
using INPUT_LAYER_CONFIG = nn::layers::dense::Configuration<T, TI, PARAMETERS::CRITIC_HIDDEN_DIM, PARAMETERS::CRITIC_ACTIVATION_FUNCTION, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Input>;
|
|
38
|
+
using INPUT_LAYER = nn::layers::dense::BindConfiguration<INPUT_LAYER_CONFIG>;
|
|
39
|
+
using GRU_CONFIG = nn::layers::gru::Configuration<T, TI, PARAMETERS::CRITIC_HIDDEN_DIM, nn::parameters::groups::Normal, true>;
|
|
40
|
+
using GRU = nn::layers::gru::BindConfiguration<GRU_CONFIG>;
|
|
41
|
+
using GRU2_CONFIG = nn::layers::gru::Configuration<T, TI, PARAMETERS::CRITIC_HIDDEN_DIM, nn::parameters::groups::Normal, true>;
|
|
42
|
+
using GRU2 = nn::layers::gru::BindConfiguration<GRU2_CONFIG>;
|
|
43
|
+
using DENSE_LAYER_CONFIG = nn::layers::dense::Configuration<T, TI, PARAMETERS::CRITIC_HIDDEN_DIM, PARAMETERS::CRITIC_ACTIVATION_FUNCTION, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Normal>;
|
|
44
|
+
using DENSE_LAYER = nn::layers::dense::BindConfiguration<DENSE_LAYER_CONFIG>;
|
|
45
|
+
using OUTPUT_CONFIG = nn::layers::dense::Configuration<T, TI, 1, nn::activation_functions::ActivationFunction::IDENTITY, nn::layers::dense::DefaultInitializer<T, TI>, nn::parameters::groups::Output>;
|
|
46
|
+
using OUTPUT = nn::layers::dense::BindConfiguration<OUTPUT_CONFIG>;
|
|
47
|
+
static constexpr TI INPUT_DIM = ENVIRONMENT::ObservationPrivileged::DIM+ENVIRONMENT::ACTION_DIM;
|
|
48
|
+
|
|
49
|
+
using MODULE_GRU = nn_models::sequential::Module<GRU, nn_models::sequential::Module<OUTPUT>>;
|
|
50
|
+
using MODULE_GRU_TWO_LAYER = nn_models::sequential::Module<GRU, nn_models::sequential::Module<GRU2, nn_models::sequential::Module<OUTPUT>>>;
|
|
51
|
+
using MODULE_GRU_THREE_LAYER = nn_models::sequential::Module<GRU, nn_models::sequential::Module<GRU2, nn_models::sequential::Module<DENSE_LAYER, nn_models::sequential::Module<OUTPUT>>>>;
|
|
52
|
+
|
|
53
|
+
using SELECTED_MODULE = rl_tools::utils::typing::conditional_t<PARAMETERS::CRITIC_NUM_LAYERS == 3, MODULE_GRU, rl_tools::utils::typing::conditional_t<PARAMETERS::CRITIC_NUM_LAYERS == 4, MODULE_GRU_TWO_LAYER, MODULE_GRU_THREE_LAYER>>;
|
|
54
|
+
static_assert(PARAMETERS::CRITIC_NUM_LAYERS == 3 || PARAMETERS::CRITIC_NUM_LAYERS == 4 || PARAMETERS::CRITIC_NUM_LAYERS == 5, "Only 3/4/5 layers (1 input + 1/2 GRU + 1/2 Output) are supported right now");
|
|
55
|
+
using MODEL = nn_models::sequential::Build<CAPABILITY, SELECTED_MODULE, INPUT_SHAPE>;
|
|
56
|
+
};
|
|
57
|
+
|
|
58
|
+
using CAPABILITY_ACTOR = nn::capability::Gradient<nn::parameters::Adam, DYNAMIC_ALLOCATION>;
|
|
59
|
+
using CAPABILITY_CRITIC = nn::capability::Gradient<nn::parameters::Adam, DYNAMIC_ALLOCATION>;
|
|
60
|
+
using ACTOR_TYPE = typename Actor<CAPABILITY_ACTOR>::MODEL;
|
|
61
|
+
using CRITIC_TYPE = typename Critic<CAPABILITY_CRITIC>::MODEL;
|
|
62
|
+
using CRITIC_TARGET_TYPE = typename Critic<nn::capability::Forward<>>::MODEL;
|
|
63
|
+
using ACTOR_TARGET_TYPE = typename Actor<nn::capability::Forward<>>::MODEL;
|
|
64
|
+
using OPTIMIZER_SPEC = nn::optimizers::adam::Specification<T, TI, typename PARAMETERS::OPTIMIZER_PARAMETERS, DYNAMIC_ALLOCATION>;
|
|
65
|
+
using OPTIMIZER = nn::optimizers::Adam<OPTIMIZER_SPEC>;
|
|
66
|
+
|
|
67
|
+
};
|
|
68
|
+
}
|
|
69
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
70
|
+
#endif
|
|
71
|
+
|
|
72
|
+
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
#include "../../../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_APPROXIMATORS_MLP_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_APPROXIMATORS_MLP_H
|
|
5
|
+
|
|
6
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
7
|
+
namespace rl_tools::rl::algorithms::td3::loop::core{
|
|
8
|
+
// The approximator config sets up any types that support the usual rl_tools::forward and rl_tools::backward operations (can be custom as well)
|
|
9
|
+
// We provide approximators based on the sequential and mlp models. The latter (mlp) allows for a variable number of layers, but is restricted to a uniform hidden layer size while the former allows for arbitrary layers to be combined in a sequential manner. Both support compile-time autodiff
|
|
10
|
+
template<typename TYPE_POLICY, typename TI, typename ENVIRONMENT, typename PARAMETERS, bool DYNAMIC_ALLOCATION>
|
|
11
|
+
struct ConfigApproximatorsMLP{
|
|
12
|
+
using TD3_PARAMETERS = typename PARAMETERS::TD3_PARAMETERS;
|
|
13
|
+
template <typename CAPABILITY>
|
|
14
|
+
struct ACTOR{
|
|
15
|
+
using INPUT_SHAPE = tensor::Shape<TI, TD3_PARAMETERS::SEQUENCE_LENGTH, TD3_PARAMETERS::ACTOR_BATCH_SIZE, ENVIRONMENT::Observation::DIM>;
|
|
16
|
+
using MLP_CONFIG = nn_models::mlp::Configuration<TYPE_POLICY, TI, ENVIRONMENT::ACTION_DIM, PARAMETERS::ACTOR_NUM_LAYERS, PARAMETERS::ACTOR_HIDDEN_DIM, PARAMETERS::ACTOR_ACTIVATION_FUNCTION, nn::activation_functions::ActivationFunction::TANH>;
|
|
17
|
+
using MLP = nn_models::mlp::BindConfiguration<MLP_CONFIG>;
|
|
18
|
+
struct SAMPLING_PARAMETERS: nn::layers::td3_sampling::DefaultParameters<TYPE_POLICY>{
|
|
19
|
+
static constexpr typename TYPE_POLICY::DEFAULT STD = PARAMETERS::EXPLORATION_NOISE;
|
|
20
|
+
};
|
|
21
|
+
using SAMPLING_CONFIG = nn::layers::td3_sampling::Configuration<TYPE_POLICY, TI, SAMPLING_PARAMETERS>;
|
|
22
|
+
using SAMPLING = nn::layers::td3_sampling::BindConfiguration<SAMPLING_CONFIG>;
|
|
23
|
+
|
|
24
|
+
using MODULE_CHAIN = nn_models::sequential::Module<MLP, nn_models::sequential::Module<SAMPLING>>;
|
|
25
|
+
using MODEL = nn_models::sequential::Build<CAPABILITY, MODULE_CHAIN, INPUT_SHAPE>;
|
|
26
|
+
};
|
|
27
|
+
|
|
28
|
+
template <typename CAPABILITY>
|
|
29
|
+
struct CRITIC{
|
|
30
|
+
static constexpr TI HIDDEN_DIM = PARAMETERS::CRITIC_HIDDEN_DIM;
|
|
31
|
+
static constexpr auto ACTIVATION_FUNCTION = PARAMETERS::CRITIC_ACTIVATION_FUNCTION;
|
|
32
|
+
|
|
33
|
+
using INPUT_SHAPE = tensor::Shape<TI, TD3_PARAMETERS::SEQUENCE_LENGTH, TD3_PARAMETERS::CRITIC_BATCH_SIZE, ENVIRONMENT::ObservationPrivileged::DIM + ENVIRONMENT::ACTION_DIM>;
|
|
34
|
+
using MLP_CONFIG = nn_models::mlp::Configuration<TYPE_POLICY, TI, 1, PARAMETERS::ACTOR_NUM_LAYERS, PARAMETERS::ACTOR_HIDDEN_DIM, PARAMETERS::ACTOR_ACTIVATION_FUNCTION, nn::activation_functions::ActivationFunction::IDENTITY>;
|
|
35
|
+
using MLP = nn_models::mlp::BindConfiguration<MLP_CONFIG>;
|
|
36
|
+
|
|
37
|
+
using MODULE_CHAIN = nn_models::sequential::Module<MLP>;
|
|
38
|
+
using MODEL = nn_models::sequential::Build<CAPABILITY, MODULE_CHAIN, INPUT_SHAPE>;
|
|
39
|
+
};
|
|
40
|
+
|
|
41
|
+
using OPTIMIZER_SPEC = nn::optimizers::adam::Specification<TYPE_POLICY, TI, typename PARAMETERS::OPTIMIZER_PARAMETERS, DYNAMIC_ALLOCATION>;
|
|
42
|
+
|
|
43
|
+
using OPTIMIZER = nn::optimizers::Adam<OPTIMIZER_SPEC>;
|
|
44
|
+
|
|
45
|
+
using ACTOR_TYPE = typename ACTOR<nn::capability::Gradient<nn::parameters::Adam, DYNAMIC_ALLOCATION>>::MODEL;
|
|
46
|
+
using ACTOR_TARGET_TYPE = typename ACTOR<nn::capability::Forward<DYNAMIC_ALLOCATION>>::MODEL;
|
|
47
|
+
using CRITIC_TYPE = typename CRITIC<nn::capability::Gradient<nn::parameters::Adam, DYNAMIC_ALLOCATION>>::MODEL;
|
|
48
|
+
using CRITIC_TARGET_TYPE = typename CRITIC<nn::capability::Forward<DYNAMIC_ALLOCATION>>::MODEL;
|
|
49
|
+
};
|
|
50
|
+
}
|
|
51
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
#endif
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
#include "../../../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_CONFIG_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_CONFIG_H
|
|
5
|
+
|
|
6
|
+
#include "../../../../../nn/layers/td3_sampling/layer.h"
|
|
7
|
+
#include "../../../../../nn_models/mlp/network.h"
|
|
8
|
+
#include "../../../../../nn_models/random_uniform/model.h"
|
|
9
|
+
#include "../../../../../nn_models/sequential/model.h"
|
|
10
|
+
#include "../../../../../rl/algorithms/td3/td3.h"
|
|
11
|
+
#include "../../../../../nn/optimizers/adam/adam.h"
|
|
12
|
+
#include "state.h"
|
|
13
|
+
#include "approximators_mlp.h"
|
|
14
|
+
#include "approximators_gru.h"
|
|
15
|
+
|
|
16
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
17
|
+
namespace rl_tools::rl::algorithms::td3::loop::core{
|
|
18
|
+
// Config State (Init/Step)
|
|
19
|
+
|
|
20
|
+
template<typename TYPE_POLICY, typename TI, typename ENVIRONMENT>
|
|
21
|
+
struct DefaultParameters{
|
|
22
|
+
using T = typename TYPE_POLICY::DEFAULT;
|
|
23
|
+
using TD3_PARAMETERS = rl::algorithms::td3::DefaultParameters<TYPE_POLICY, TI>;
|
|
24
|
+
static constexpr TI N_ENVIRONMENTS = 1;
|
|
25
|
+
static constexpr TI N_WARMUP_STEPS = 100; // Exploration executed with a uniform random policy for N_WARMUP_STEPS steps
|
|
26
|
+
static constexpr TI N_WARMUP_STEPS_CRITIC = 100; // Number of steps before critic training starts
|
|
27
|
+
static constexpr TI N_WARMUP_STEPS_ACTOR = 100; // Number of steps before actor training starts
|
|
28
|
+
static constexpr TI STEP_LIMIT = 10000;
|
|
29
|
+
static constexpr TI REPLAY_BUFFER_CAP = STEP_LIMIT; // Note: when inheriting from this class for overwriting the default STEP_LIMIT you need to set the REPLAY_BUFFER_CAP as well otherwise it will be the default step limit
|
|
30
|
+
static constexpr TI EPISODE_STEP_LIMIT = ENVIRONMENT::EPISODE_STEP_LIMIT;
|
|
31
|
+
|
|
32
|
+
static constexpr TI ACTOR_HIDDEN_DIM = 64;
|
|
33
|
+
static constexpr TI ACTOR_NUM_LAYERS = 3;
|
|
34
|
+
static constexpr auto ACTOR_ACTIVATION_FUNCTION = nn::activation_functions::ActivationFunction::RELU;
|
|
35
|
+
static constexpr TI CRITIC_HIDDEN_DIM = 64;
|
|
36
|
+
static constexpr TI CRITIC_NUM_LAYERS = 3;
|
|
37
|
+
static constexpr auto CRITIC_ACTIVATION_FUNCTION = nn::activation_functions::ActivationFunction::RELU;
|
|
38
|
+
|
|
39
|
+
static constexpr bool COLLECT_EPISODE_STATS = true;
|
|
40
|
+
static constexpr TI EPISODE_STATS_BUFFER_SIZE = 1000;
|
|
41
|
+
|
|
42
|
+
static constexpr T EXPLORATION_NOISE = 0.1;
|
|
43
|
+
static constexpr bool SHARED_BATCH = true;
|
|
44
|
+
|
|
45
|
+
using BATCH_SAMPLING_PARAMETERS = rl::components::off_policy_runner::SequentialBatchParameters<TYPE_POLICY, TI, TD3_PARAMETERS::SEQUENCE_LENGTH, rl::components::off_policy_runner::SequentialBatchParametersDefault>;
|
|
46
|
+
|
|
47
|
+
using OPTIMIZER_PARAMETERS = nn::optimizers::adam::DEFAULT_PARAMETERS_TENSORFLOW<TYPE_POLICY>;
|
|
48
|
+
};
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
template<typename T_TYPE_POLICY, typename T_TI, typename T_RNG, typename T_ENVIRONMENT, typename T_PARAMETERS = DefaultParameters<T_TYPE_POLICY, T_TI, T_ENVIRONMENT>, template<typename, typename, typename, typename, bool> class APPROXIMATOR_CONFIG=ConfigApproximatorsMLP, bool T_DYNAMIC_ALLOCATION=true>
|
|
52
|
+
struct Config{
|
|
53
|
+
using TYPE_POLICY = T_TYPE_POLICY;
|
|
54
|
+
using T_CONFIG = typename TYPE_POLICY::DEFAULT;
|
|
55
|
+
using TI = T_TI;
|
|
56
|
+
using RNG = T_RNG;
|
|
57
|
+
using ENVIRONMENT = T_ENVIRONMENT;
|
|
58
|
+
using ENVIRONMENT_EVALUATION = T_ENVIRONMENT;
|
|
59
|
+
static constexpr bool DYNAMIC_ALLOCATION = T_DYNAMIC_ALLOCATION;
|
|
60
|
+
|
|
61
|
+
using NN = APPROXIMATOR_CONFIG<TYPE_POLICY, TI, T_ENVIRONMENT, T_PARAMETERS, DYNAMIC_ALLOCATION>;
|
|
62
|
+
// using NN = ConfigApproximatorsMLP<T, TI, T_ENVIRONMENT, T_PARAMETERS>;
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
using CORE_PARAMETERS = T_PARAMETERS;
|
|
66
|
+
|
|
67
|
+
static constexpr TI ENVIRONMENT_STEPS_PER_LOOP_STEP = CORE_PARAMETERS::N_ENVIRONMENTS;
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
using EXPLORATION_POLICY_SPEC = nn_models::random_uniform::Specification<TYPE_POLICY, TI, ENVIRONMENT::Observation::DIM, ENVIRONMENT::ACTION_DIM, nn_models::random_uniform::Range::MINUS_ONE_TO_ONE>;
|
|
71
|
+
using EXPLORATION_POLICY = nn_models::RandomUniform<EXPLORATION_POLICY_SPEC>;
|
|
72
|
+
|
|
73
|
+
using ACTOR_CRITIC_SPEC = rl::algorithms::td3::Specification<TYPE_POLICY, TI, ENVIRONMENT, typename NN::ACTOR_TYPE, typename NN::ACTOR_TARGET_TYPE, typename NN::CRITIC_TYPE, typename NN::CRITIC_TARGET_TYPE, typename NN::OPTIMIZER, typename CORE_PARAMETERS::TD3_PARAMETERS, CORE_PARAMETERS::BATCH_SAMPLING_PARAMETERS::INCLUDE_FIRST_STEP_IN_TARGETS>;
|
|
74
|
+
using ACTOR_CRITIC_TYPE = rl::algorithms::td3::ActorCritic<ACTOR_CRITIC_SPEC>;
|
|
75
|
+
static constexpr TI NUM_NNS = 3;
|
|
76
|
+
using EVAL_ACTOR_TYPE = typename NN::ACTOR_TYPE::template CHANGE_BATCH_SIZE<TI, CORE_PARAMETERS::N_ENVIRONMENTS>;
|
|
77
|
+
|
|
78
|
+
struct OFF_POLICY_RUNNER_PARAMETERS{
|
|
79
|
+
static constexpr TI N_ENVIRONMENTS = CORE_PARAMETERS::N_ENVIRONMENTS;
|
|
80
|
+
static constexpr bool ASYMMETRIC_OBSERVATIONS = !rl_tools::utils::typing::is_same_v<typename ENVIRONMENT::Observation, typename ENVIRONMENT::ObservationPrivileged>;
|
|
81
|
+
static constexpr TI REPLAY_BUFFER_CAPACITY = CORE_PARAMETERS::REPLAY_BUFFER_CAP;
|
|
82
|
+
static constexpr TI EPISODE_STEP_LIMIT = CORE_PARAMETERS::EPISODE_STEP_LIMIT;
|
|
83
|
+
static constexpr bool COLLECT_EPISODE_STATS = CORE_PARAMETERS::COLLECT_EPISODE_STATS;
|
|
84
|
+
static constexpr TI EPISODE_STATS_BUFFER_SIZE = CORE_PARAMETERS::EPISODE_STATS_BUFFER_SIZE;
|
|
85
|
+
static constexpr T_CONFIG EXPLORATION_NOISE = CORE_PARAMETERS::EXPLORATION_NOISE;
|
|
86
|
+
static constexpr bool SAMPLE_PARAMETERS = true;
|
|
87
|
+
};
|
|
88
|
+
using POLICIES = rl_tools::utils::Tuple<TI, EXPLORATION_POLICY, EVAL_ACTOR_TYPE>;
|
|
89
|
+
|
|
90
|
+
using OFF_POLICY_RUNNER_SPEC = rl::components::off_policy_runner::Specification<TYPE_POLICY, TI, ENVIRONMENT, POLICIES, OFF_POLICY_RUNNER_PARAMETERS, T_DYNAMIC_ALLOCATION>;
|
|
91
|
+
static_assert(ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::ACTOR_BATCH_SIZE == ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::CRITIC_BATCH_SIZE);
|
|
92
|
+
static constexpr TI TARGET_SEQUENCE_LENGTH = CORE_PARAMETERS::TD3_PARAMETERS::SEQUENCE_LENGTH + (CORE_PARAMETERS::BATCH_SAMPLING_PARAMETERS::INCLUDE_FIRST_STEP_IN_TARGETS ? 1 : 0);
|
|
93
|
+
|
|
94
|
+
static constexpr TI DEFAULT_SEQUENCE_LENGTH = DefaultParameters<TYPE_POLICY, TI, ENVIRONMENT>::TD3_PARAMETERS::SEQUENCE_LENGTH;
|
|
95
|
+
static constexpr bool USING_DEFAULT_BATCH_SAMPLING_PARAMETERS = rl_tools::utils::typing::is_base_of_v<components::off_policy_runner::SequentialBatchParametersDefault, typename CORE_PARAMETERS::BATCH_SAMPLING_PARAMETERS>;
|
|
96
|
+
static_assert(CORE_PARAMETERS::TD3_PARAMETERS::SEQUENCE_LENGTH == DEFAULT_SEQUENCE_LENGTH || !USING_DEFAULT_BATCH_SAMPLING_PARAMETERS, "When setting a custom sequence length, please study and set custom batch sampling parameters.");
|
|
97
|
+
|
|
98
|
+
using CRITIC_BATCH_SPEC = rl::components::off_policy_runner::SequentialBatchSpecification<OFF_POLICY_RUNNER_SPEC, ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::SEQUENCE_LENGTH, ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::CRITIC_BATCH_SIZE, typename CORE_PARAMETERS::BATCH_SAMPLING_PARAMETERS, DYNAMIC_ALLOCATION>;
|
|
99
|
+
using ACTOR_BATCH_SPEC = rl::components::off_policy_runner::SequentialBatchSpecification<OFF_POLICY_RUNNER_SPEC, ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::SEQUENCE_LENGTH, ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::ACTOR_BATCH_SIZE, typename CORE_PARAMETERS::BATCH_SAMPLING_PARAMETERS, DYNAMIC_ALLOCATION>;
|
|
100
|
+
template <typename CONFIG>
|
|
101
|
+
using State = State<CONFIG>;
|
|
102
|
+
};
|
|
103
|
+
}
|
|
104
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
105
|
+
|
|
106
|
+
#endif
|
|
107
|
+
|
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
#include "../../../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_OPERATIONS_GENERIC_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_OPERATIONS_GENERIC_H
|
|
5
|
+
|
|
6
|
+
#include "../../../../../nn/optimizers/adam/instance/operations_generic.h"
|
|
7
|
+
#include "../../../../../nn/layers/td3_sampling/operations_generic.h"
|
|
8
|
+
#include "../../../../../nn_models/operations_generic.h"
|
|
9
|
+
#include "../../../../../nn_models/sequential/operations_generic.h"
|
|
10
|
+
#include "../../../../../rl/algorithms/td3/operations_generic.h"
|
|
11
|
+
#include "../../../../../nn_models/random_uniform/operations_generic.h"
|
|
12
|
+
#include "../../../../../rl/components/off_policy_runner/operations_generic.h"
|
|
13
|
+
#include "../../../../../random/operations_generic_array.h"
|
|
14
|
+
|
|
15
|
+
#include "config.h"
|
|
16
|
+
|
|
17
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
18
|
+
namespace rl_tools{
|
|
19
|
+
template <typename DEVICE, typename T_CONFIG>
|
|
20
|
+
RL_TOOLS_FUNCTION_PLACEMENT void malloc(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
21
|
+
malloc(device, ts.rng);
|
|
22
|
+
malloc(device, ts.actor_critic);
|
|
23
|
+
malloc(device, ts.off_policy_runner);
|
|
24
|
+
malloc(device, ts.critic_batch);
|
|
25
|
+
malloc(device, ts.critic_training_buffers);
|
|
26
|
+
malloc(device, ts.critic_buffers[0]);
|
|
27
|
+
malloc(device, ts.critic_buffers[1]);
|
|
28
|
+
malloc(device, ts.critic_target_buffers[0]);
|
|
29
|
+
malloc(device, ts.critic_target_buffers[1]);
|
|
30
|
+
malloc(device, ts.actor_batch);
|
|
31
|
+
malloc(device, ts.actor_training_buffers);
|
|
32
|
+
malloc(device, ts.actor_buffers_eval);
|
|
33
|
+
malloc(device, ts.actor_buffers[0]);
|
|
34
|
+
malloc(device, ts.actor_buffers[1]);
|
|
35
|
+
malloc(device, ts.actor_target_buffers[0]);
|
|
36
|
+
malloc(device, ts.actor_target_buffers[1]);
|
|
37
|
+
for(auto& env: ts.envs){
|
|
38
|
+
rl_tools::malloc(device, env);
|
|
39
|
+
}
|
|
40
|
+
}
|
|
41
|
+
template <typename DEVICE, typename T_CONFIG>
|
|
42
|
+
RL_TOOLS_FUNCTION_PLACEMENT void free(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
43
|
+
free(device, ts.rng);
|
|
44
|
+
free(device, ts.actor_critic);
|
|
45
|
+
free(device, ts.off_policy_runner);
|
|
46
|
+
free(device, ts.critic_batch);
|
|
47
|
+
free(device, ts.critic_training_buffers);
|
|
48
|
+
free(device, ts.critic_buffers[0]);
|
|
49
|
+
free(device, ts.critic_buffers[1]);
|
|
50
|
+
free(device, ts.critic_target_buffers[0]);
|
|
51
|
+
free(device, ts.critic_target_buffers[1]);
|
|
52
|
+
free(device, ts.actor_batch);
|
|
53
|
+
free(device, ts.actor_training_buffers);
|
|
54
|
+
free(device, ts.actor_buffers_eval);
|
|
55
|
+
free(device, ts.actor_buffers[0]);
|
|
56
|
+
free(device, ts.actor_buffers[1]);
|
|
57
|
+
free(device, ts.actor_target_buffers[0]);
|
|
58
|
+
free(device, ts.actor_target_buffers[1]);
|
|
59
|
+
for(auto& env: ts.envs){
|
|
60
|
+
rl_tools::free(device, env);
|
|
61
|
+
}
|
|
62
|
+
}
|
|
63
|
+
template <typename DEVICE, typename T_CONFIG>
|
|
64
|
+
RL_TOOLS_FUNCTION_PLACEMENT void init(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts, typename T_CONFIG::TI seed = 0){
|
|
65
|
+
using CONFIG = T_CONFIG;
|
|
66
|
+
using TYPE_POLICY = typename CONFIG::TYPE_POLICY;
|
|
67
|
+
using TI = typename DEVICE::index_t;
|
|
68
|
+
|
|
69
|
+
init(device, ts.rng, seed);
|
|
70
|
+
|
|
71
|
+
init(device, ts.actor_critic, ts.rng);
|
|
72
|
+
|
|
73
|
+
for(TI env_i = 0; env_i < CONFIG::CORE_PARAMETERS::N_ENVIRONMENTS; env_i ++){
|
|
74
|
+
rl_tools::init(device, ts.envs[env_i]);
|
|
75
|
+
}
|
|
76
|
+
init(device, ts.off_policy_runner);
|
|
77
|
+
|
|
78
|
+
ts.step = 0;
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
template <typename DEVICE, typename T_CONFIG>
|
|
83
|
+
RL_TOOLS_FUNCTION_PLACEMENT bool step(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
84
|
+
using CONFIG = T_CONFIG;
|
|
85
|
+
set_step(device, device.logger, ts.step);
|
|
86
|
+
bool finished = false;
|
|
87
|
+
if(ts.step >= CONFIG::CORE_PARAMETERS::N_WARMUP_STEPS){
|
|
88
|
+
step<1>(device, ts.off_policy_runner, get_actor(ts), ts.actor_buffers_eval, ts.rng);
|
|
89
|
+
}
|
|
90
|
+
else{
|
|
91
|
+
typename CONFIG::EXPLORATION_POLICY exploration_policy;
|
|
92
|
+
typename CONFIG::EXPLORATION_POLICY::template Buffer<> exploration_policy_buffer;
|
|
93
|
+
step<0>(device, ts.off_policy_runner, exploration_policy, exploration_policy_buffer, ts.rng);
|
|
94
|
+
}
|
|
95
|
+
|
|
96
|
+
bool train_critic_flag = ts.step >= CONFIG::CORE_PARAMETERS::N_WARMUP_STEPS_CRITIC && ts.step % CONFIG::CORE_PARAMETERS::TD3_PARAMETERS::CRITIC_TRAINING_INTERVAL == 0;
|
|
97
|
+
bool update_critic_targets_flag = ts.step >= CONFIG::CORE_PARAMETERS::N_WARMUP_STEPS_CRITIC && ts.step % CONFIG::CORE_PARAMETERS::TD3_PARAMETERS::CRITIC_TARGET_UPDATE_INTERVAL == 0;
|
|
98
|
+
bool train_actor_flag = ts.step >= CONFIG::CORE_PARAMETERS::N_WARMUP_STEPS_ACTOR && ts.step % CONFIG::CORE_PARAMETERS::TD3_PARAMETERS::ACTOR_TRAINING_INTERVAL == 0;
|
|
99
|
+
bool update_actor_targets_flag = ts.step >= CONFIG::CORE_PARAMETERS::N_WARMUP_STEPS_ACTOR && ts.step % CONFIG::CORE_PARAMETERS::TD3_PARAMETERS::ACTOR_TARGET_UPDATE_INTERVAL == 0;
|
|
100
|
+
|
|
101
|
+
if(CONFIG::CORE_PARAMETERS::SHARED_BATCH && (train_critic_flag || train_actor_flag)){
|
|
102
|
+
gather_batch(device, ts.off_policy_runner, ts.critic_batch, ts.rng);
|
|
103
|
+
auto action_noise_matrix_view = matrix_view(device, ts.critic_training_buffers.target_next_action_noise);
|
|
104
|
+
target_action_noise(device, ts.actor_critic, action_noise_matrix_view, ts.rng);
|
|
105
|
+
}
|
|
106
|
+
if(train_critic_flag){
|
|
107
|
+
for(int critic_i = 0; critic_i < 2; critic_i++){
|
|
108
|
+
if constexpr(!CONFIG::CORE_PARAMETERS::SHARED_BATCH) {
|
|
109
|
+
gather_batch(device, ts.off_policy_runner, ts.critic_batch, ts.rng);
|
|
110
|
+
auto action_noise_matrix_view = matrix_view(device, ts.critic_training_buffers.target_next_action_noise);
|
|
111
|
+
target_action_noise(device, ts.actor_critic, action_noise_matrix_view, ts.rng);
|
|
112
|
+
}
|
|
113
|
+
train_critic(device, ts.actor_critic, ts.actor_critic.critics[critic_i], ts.critic_batch, ts.actor_critic.critic_optimizers[critic_i], ts.actor_buffers[critic_i], ts.actor_target_buffers[critic_i], ts.critic_buffers[critic_i], ts.critic_target_buffers[critic_i], ts.critic_training_buffers, ts.rng);
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
if(update_critic_targets_flag){
|
|
117
|
+
update_critic_targets(device, ts.actor_critic);
|
|
118
|
+
}
|
|
119
|
+
if(train_actor_flag){
|
|
120
|
+
if constexpr(CONFIG::CORE_PARAMETERS::SHARED_BATCH) {
|
|
121
|
+
train_actor(device, ts.actor_critic, ts.critic_batch, ts.actor_critic.actor_optimizer, ts.actor_buffers[0], ts.critic_buffers[0], ts.actor_training_buffers, ts.rng);
|
|
122
|
+
}
|
|
123
|
+
else{
|
|
124
|
+
gather_batch(device, ts.off_policy_runner, ts.actor_batch, ts.rng);
|
|
125
|
+
train_actor(device, ts.actor_critic, ts.actor_batch, ts.actor_critic.actor_optimizer, ts.actor_buffers[0], ts.critic_buffers[0], ts.actor_training_buffers, ts.rng);
|
|
126
|
+
}
|
|
127
|
+
}
|
|
128
|
+
if(update_actor_targets_flag){
|
|
129
|
+
update_actor_target(device, ts.actor_critic);
|
|
130
|
+
}
|
|
131
|
+
ts.step++;
|
|
132
|
+
if(ts.step > CONFIG::CORE_PARAMETERS::STEP_LIMIT){
|
|
133
|
+
return true;
|
|
134
|
+
}
|
|
135
|
+
else{
|
|
136
|
+
return finished;
|
|
137
|
+
}
|
|
138
|
+
}
|
|
139
|
+
// the following operations are for nn_analytics iterating the neural networks
|
|
140
|
+
template <auto INDEX, typename DEVICE, typename T_CONFIG>
|
|
141
|
+
RL_TOOLS_FUNCTION_PLACEMENT constexpr auto& get_nn(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
142
|
+
static_assert(INDEX < T_CONFIG::NUM_NNS, "Index out of bounds, there are only 3 neural networks in the TD3");
|
|
143
|
+
if constexpr(INDEX == 0){
|
|
144
|
+
return ts.actor_critic.actor;
|
|
145
|
+
}
|
|
146
|
+
else{
|
|
147
|
+
if constexpr(INDEX == 1){
|
|
148
|
+
return ts.actor_critic.critics[0];
|
|
149
|
+
}
|
|
150
|
+
else{
|
|
151
|
+
return ts.actor_critic.critics[1];
|
|
152
|
+
}
|
|
153
|
+
}
|
|
154
|
+
}
|
|
155
|
+
template <auto INDEX, typename DEVICE, typename T_CONFIG>
|
|
156
|
+
RL_TOOLS_FUNCTION_PLACEMENT constexpr auto& get_nn_name(DEVICE& device, rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
157
|
+
static_assert(INDEX < T_CONFIG::NUM_NNS, "Index out of bounds, there are only 3 neural networks in the TD3");
|
|
158
|
+
if constexpr(INDEX == 0){
|
|
159
|
+
return "actor";
|
|
160
|
+
}
|
|
161
|
+
else{
|
|
162
|
+
if constexpr(INDEX == 1){
|
|
163
|
+
return "critic[0]";
|
|
164
|
+
}
|
|
165
|
+
else{
|
|
166
|
+
return "critic[1]";
|
|
167
|
+
}
|
|
168
|
+
}
|
|
169
|
+
}
|
|
170
|
+
}
|
|
171
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
#endif
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
#include "../../../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_STATE_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_LOOP_CORE_STATE_H
|
|
5
|
+
|
|
6
|
+
#include "../../../../../rl/algorithms/td3/td3.h"
|
|
7
|
+
#include "../../../../../rl/components/off_policy_runner/off_policy_runner.h"
|
|
8
|
+
|
|
9
|
+
RL_TOOLS_NAMESPACE_WRAPPER_START
|
|
10
|
+
namespace rl_tools{
|
|
11
|
+
namespace rl::algorithms::td3::loop::core{
|
|
12
|
+
// Config State (Init/Step)
|
|
13
|
+
template<typename T_CONFIG>
|
|
14
|
+
struct State{
|
|
15
|
+
using CONFIG = T_CONFIG;
|
|
16
|
+
using TI = typename CONFIG::TI;
|
|
17
|
+
typename CONFIG::RNG rng;
|
|
18
|
+
rl::components::OffPolicyRunner<typename CONFIG::OFF_POLICY_RUNNER_SPEC> off_policy_runner;
|
|
19
|
+
typename CONFIG::ENVIRONMENT envs[decltype(off_policy_runner)::N_ENVIRONMENTS];
|
|
20
|
+
typename CONFIG::ENVIRONMENT::Parameters env_parameters[decltype(off_policy_runner)::N_ENVIRONMENTS];
|
|
21
|
+
typename CONFIG::ACTOR_CRITIC_TYPE actor_critic;
|
|
22
|
+
// typename CONFIG::NN::ACTOR_TYPE::template Buffer<CONFIG::DYNAMIC_ALLOCATION> actor_deterministic_evaluation_buffers;
|
|
23
|
+
// rl::components::off_policy_runner::Batch<rl::components::off_policy_runner::BatchSpecification<typename decltype(off_policy_runner)::SPEC, CONFIG::ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::CRITIC_BATCH_SIZE>> critic_batch;
|
|
24
|
+
rl::components::off_policy_runner::SequentialBatch<typename CONFIG::CRITIC_BATCH_SPEC> critic_batch;
|
|
25
|
+
rl::algorithms::td3::CriticTrainingBuffers<rl::algorithms::td3::CriticTrainingBuffersSpecification<typename CONFIG::ACTOR_CRITIC_SPEC, CONFIG::DYNAMIC_ALLOCATION>> critic_training_buffers;
|
|
26
|
+
typename CONFIG::NN::CRITIC_TYPE::template Buffer<CONFIG::DYNAMIC_ALLOCATION> critic_buffers[2];
|
|
27
|
+
using TARGET_CRITIC = typename rl_tools::utils::typing::remove_reference_t<decltype(actor_critic.critics_target[0])>::template CHANGE_SEQUENCE_LENGTH<TI, CONFIG::TARGET_SEQUENCE_LENGTH>;
|
|
28
|
+
typename TARGET_CRITIC::template Buffer<CONFIG::DYNAMIC_ALLOCATION> critic_target_buffers[2];
|
|
29
|
+
rl::components::off_policy_runner::SequentialBatch<typename CONFIG::ACTOR_BATCH_SPEC> actor_batch;
|
|
30
|
+
// rl::components::off_policy_runner::Batch<rl::components::off_policy_runner::BatchSpecification<typename decltype(off_policy_runner)::SPEC, CONFIG::ACTOR_CRITIC_TYPE::SPEC::PARAMETERS::ACTOR_BATCH_SIZE>> actor_batch;
|
|
31
|
+
rl::algorithms::td3::ActorTrainingBuffers<rl::algorithms::td3::ActorTrainingBuffersSpecification<typename CONFIG::ACTOR_CRITIC_TYPE::SPEC, CONFIG::DYNAMIC_ALLOCATION>> actor_training_buffers;
|
|
32
|
+
typename CONFIG::NN::ACTOR_TYPE::template Buffer<CONFIG::DYNAMIC_ALLOCATION> actor_buffers[2];
|
|
33
|
+
using TARGET_ACTOR = typename decltype(actor_critic.actor)::template CHANGE_SEQUENCE_LENGTH<TI, CONFIG::TARGET_SEQUENCE_LENGTH>;
|
|
34
|
+
typename TARGET_ACTOR::template Buffer<CONFIG::DYNAMIC_ALLOCATION> actor_target_buffers[2];
|
|
35
|
+
typename CONFIG::EVAL_ACTOR_TYPE::template Buffer<CONFIG::DYNAMIC_ALLOCATION> actor_buffers_eval;
|
|
36
|
+
TI step;
|
|
37
|
+
};
|
|
38
|
+
}
|
|
39
|
+
template <typename T_CONFIG>
|
|
40
|
+
constexpr auto& get_actor(rl::algorithms::td3::loop::core::State<T_CONFIG>& ts){
|
|
41
|
+
return ts.actor_critic.actor;
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
}
|
|
45
|
+
RL_TOOLS_NAMESPACE_WRAPPER_END
|
|
46
|
+
#endif
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
#include "../../../version.h"
|
|
2
|
+
#if (defined(RL_TOOLS_DISABLE_INCLUDE_GUARDS) || !defined(RL_TOOLS_RL_ALGORITHMS_TD3_OPERATIONS_CPU_H)) && (RL_TOOLS_USE_THIS_VERSION == 1)
|
|
3
|
+
#pragma once
|
|
4
|
+
#define RL_TOOLS_RL_ALGORITHMS_TD3_OPERATIONS_CPU_H
|
|
5
|
+
|
|
6
|
+
#include "operations_generic.h"
|
|
7
|
+
#include "../../../rl/components/operations_cpu.h"
|
|
8
|
+
|
|
9
|
+
#endif
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
#if defined(RL_TOOLS_BACKEND_ENABLE_MKL) && !defined(RL_TOOLS_BACKEND_DISABLE_BLAS)
|
|
2
|
+
#include "../../../rl/algorithms/td3/operations_cpu_mkl.h"
|
|
3
|
+
#else
|
|
4
|
+
#if defined(RL_TOOLS_BACKEND_ENABLE_ACCELERATE) && !defined(RL_TOOLS_BACKEND_DISABLE_BLAS)
|
|
5
|
+
#include "../../../rl/algorithms/td3/operations_cpu_accelerate.h"
|
|
6
|
+
#else
|
|
7
|
+
#include "../../../rl/algorithms/td3/operations_cpu.h"
|
|
8
|
+
#endif
|
|
9
|
+
#endif
|