diff --git a/.clang-format b/.clang-format index 7dd2922..2e4daf6 100644 --- a/.clang-format +++ b/.clang-format @@ -1,42 +1,54 @@ Language: Cpp +BasedOnStyle: LLVM -BasedOnStyle: WebKit - -# 使用空格而不是 Tab +# RMCS uses spaces and four-space indentation. UseTab: Never - -# 缩进 4 字符 IndentWidth: 4 TabWidth: 4 - -# 访问修饰符(public)等靠左对齐 +ContinuationIndentWidth: 4 +ConstructorInitializerIndentWidth: 4 AccessModifierOffset: -4 -# 注释不对齐 -AlignTrailingComments: false - -# 大括号的换行规则 -BreakBeforeBraces: Attach - -# 模板定义的换行规则 -BreakTemplateDeclarations: "Yes" - -# 最大列宽度 +Standard: c++23 ColumnLimit: 100 +MaxEmptyLinesToKeep: 1 -# 是否允许短的函数,语句块等单独一行 -AllowShortIfStatementsOnASingleLine: AllIfsAndElse -AllowShortBlocksOnASingleLine: Empty -AllowShortFunctionsOnASingleLine: All - +BreakBeforeBraces: Attach +BreakConstructorInitializers: BeforeComma +BreakInheritanceList: BeforeComma BreakBeforeBinaryOperators: NonAssignment -PenaltyBreakAssignment: 2 -PenaltyBreakString: 2 +AlignAfterOpenBracket: AlwaysBreak +AlignOperands: AlignAfterOperator +AlignTrailingComments: + Kind: Always + OverEmptyLines: 64 AlignConsecutiveAssignments: + Enabled: false +AlignConsecutiveBitFields: + Enabled: true + AcrossEmptyLines: false + AcrossComments: false +AlignConsecutiveDeclarations: + Enabled: false +AlignConsecutiveMacros: Enabled: true AcrossEmptyLines: false AcrossComments: false - AlignCompound: false -SpaceInEmptyBlock: true +PointerAlignment: Left +IndentPPDirectives: AfterHash +PPIndentWidth: 1 +IndentWrappedFunctionNames: true + +AllowAllArgumentsOnNextLine: true +AllowAllParametersOfDeclarationOnNextLine: true +AllowShortBlocksOnASingleLine: Empty +AllowShortCaseLabelsOnASingleLine: true + +FixNamespaceComments: true +IncludeBlocks: Preserve + +AlwaysBreakTemplateDeclarations: Yes +IndentRequiresClause: false +RequiresClausePosition: SingleLine diff --git a/CMakeLists.txt b/CMakeLists.txt index 0cb45b8..91bd886 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -57,11 +57,6 @@ ament_auto_add_library(rmcs_rl_bridge SHARED ) target_link_libraries(rmcs_rl_bridge ${cpp_typesupport_target}) -ament_auto_add_library(rmcs_rl_legacy SHARED - src/rl_controller.cpp -) -target_link_libraries(rmcs_rl_legacy ${ONNXRUNTIME_ROOT_DIR}/lib/libonnxruntime.so) - ament_auto_add_executable(policy_server src/policy_server.cpp ) diff --git a/README.md b/README.md index 03d2a16..70d50a5 100644 --- a/README.md +++ b/README.md @@ -7,10 +7,10 @@ RMCS(RoboMaster Control System)的 **RL 策略桥**:在「传统 RMCS 结 详细文档: -- [架构与组件](doc/architecture.md) -- [桥式重构设计(v2,定稿方案)](doc/bridge-design.md) -- [构建、配置与部署](doc/deployment.md) -- [策略模型合同(张量形状 + 元数据要求)](doc/model-contract.md) +- [架构与组件](planning/docs/architecture.md) +- [桥式重构设计(v2,定稿方案)](planning/docs/bridge-design.md) +- [构建、配置与部署](planning/docs/deployment.md) +- [策略模型合同(张量形状 + 元数据要求)](planning/docs/model-contract.md) ## 两条硬边界 @@ -26,15 +26,11 @@ topic 存在的**唯一**原因是「策略进程独立」。想消除这 3 条 | 名称 | 位置 | 依赖 ONNX Runtime | 职责 | |---|---|---|---| -| `rmcs::rl::RlBridge` | 库 `rmcs_rl_bridge` | **否** | 观测组装、定频发布、动作回写、`valid/healthy/action_age` 事实位、合同指纹 | +| `rmcs_rl::RlBridge` | 库 `rmcs_rl_bridge` | **否** | 观测组装、定频发布、动作回写、`valid/healthy/action_age` 事实位、合同指纹 | +| `rmcs_rl::PolicyServerLauncher` | 库 `rmcs_rl_bridge` | **否** | 随 executor 生命周期 fork/exec 独立的 `policy_server` 子进程,可退避重启、防孤儿 | | `policy_server` | 可执行文件(`ros2 run rmcs_rl policy_server`) | **是**(唯一链接 ORT 的可执行文件) | 收一帧 obs → 归一化 → ONNX → 回一帧 action(纯反应式,无定时器) | -| `rmcs::rl::RlController` | 库 `rmcs_rl_legacy` | 是 | **遗留控制器**(旧配置格式,P1 切换完成后删除) | | 消息 `rmcs_rl/msg/*` | 本包 `msg/`(rosidl 生成,类型全名如 `rmcs_rl/msg/Observation`) | 否 | `Observation` / `Action` / `PolicyStatus` | -> ⚠️ `RlController` 是遗留路径:它自带 ONNX 推理、PD、FSM,配置键(`rl_inference_frequency`、 -> `position_pd_joints`、`action_terms: joint=... mode=... kp=...`)与本文档描述的桥式配置**完全不同**, -> 不要混用、也不要与 `RlBridge` 同时挂载。桥式链路取代它之后即删除。 - ## 数据流 ``` @@ -46,7 +42,7 @@ RMCS 侧 output 接口 ──(桥:按词条拼 obs)──> obs 向量 ──to ## 部署到真机 本包**不含台架夹具**:链路必须挂到真实 RMCS 组件上跑。简要步骤(完整流程见 -[doc/deployment.md](doc/deployment.md)): +[planning/docs/deployment.md](planning/docs/deployment.md)): 1. 改 `config/executor.yaml`:观测/动作词条与 `joint_*` 指向真机的接口路径, `policy_server.rl_model_path` 指向已盖章的模型; @@ -59,7 +55,7 @@ ros2 run rmcs_rl policy_server --ros-args --params-file ``` `layout_hash` 两侧必须一致;`valid=0` 的原因会直接打在桥的日志里(见 -[doc/deployment.md](doc/deployment.md) 的排障表)。注意 P1 的 RMCS 侧 consumer 尚未实现, +[planning/docs/deployment.md](planning/docs/deployment.md) 的排障表)。注意 P1 的 RMCS 侧 consumer 尚未实现, 真机上桥恒 `valid=0`、不会输出权威动作(见「现状与边界」)。 ## 加一台新车型 / 加一条观测词条 @@ -105,13 +101,11 @@ bash src/rmcs_rl/tool/test_layout_contract.sh ## 现状与边界 - **P0 已完成**:桥、策略进程、消息定义、合同指纹(v2)、工具链。 -- **P1 未实现**:RMCS 侧 consumer(FSM / PREPARE / kp,kd / 限位 / NaN 让位)与 `RlController` → 桥的切换。 - 目前没有组件写 `/wheel_leg/rl/enable`,也没有组件消费 `/wheel_leg/rl/action/*`,因此实车上 - `enable_default: false` → **桥恒 `valid=0`**,不会输出权威动作(这是刻意的:没人宣告权威就不许输出)。 +- **P1 进行中**:RMCS 侧 consumer(FSM / PREPARE / kp,kd / 限位 / NaN 让位)。 + deformable 已有消费侧落地(`DeformableRlSuspension` / 仲裁);wheel-leg 尚未迁到桥式, + 因此轮腿上桥若挂载且 `enable_default: false` → **桥恒 `valid=0`**,不会输出权威动作。 - 桥**从不写电机 `control_*` 接口**:executor 禁止同名 output,电机控制权始终在 RMCS 侧消费组件手上, - 所以「RL 接管 / 退回传统控制」不需要仲裁组件(P1 的 consumer 用写 NaN 让位)。 -- `RlController`(库 `rmcs_rl_legacy`)**仍然是唯一的整机可用路径**,P1 完成后删除; - 它带 ORT 进控制进程,配置格式与桥式完全不同,不要混配。 + 所以「RL 接管 / 退回传统控制」不需要仲裁组件(consumer 用写 NaN 让位)。 - 观测接口运行期热加不支持(配对只在启动做一次);`contract_ok` 一旦因合同不符锁存为 false, 当前实现**只能靠重启 executor 进程恢复**。 @@ -122,26 +116,29 @@ rmcs_rl/ ├── README.md # 本页(入口) ├── config/ │ └── executor.yaml # 实车配置模板(轮腿;复制到 rmcs_bringup/config/.yaml) -├── doc/ -│ ├── architecture.md # 进程/组件、数据流、接口清单、valid 状态机 -│ ├── bridge-design.md # 重构定稿方案(权威设计) -│ ├── deployment.md # 构建 → 模型 → 配置 → 运行 → 验证 → 交接 -│ └── model-contract.md # 模型张量合同与 metadata 清单 +├── planning/ +│ └── docs/ +│ ├── architecture.md # 进程/组件、数据流、接口清单、valid 状态机 +│ ├── bridge-design.md # 重构定稿方案(权威设计) +│ ├── deployment.md # 构建 → 模型 → 配置 → 运行 → 验证 → 交接 +│ ├── model-contract.md # 模型张量合同与 metadata 清单 +│ └── deformable-rl-pipeline.md # deformable 消费侧落地说明 ├── models/ # 策略 ONNX(安装到 share/rmcs_rl/models/) ├── msg/ # Observation / Action / PolicyStatus(rosidl 生成,类型名 rmcs_rl/msg/*) ├── src/ │ ├── rl_bridge.cpp # 桥 │ ├── rl_layout.hpp # FNV-1a64 / layout_hash / model_id(与 tool/rl_layout.py 同构) -│ ├── policy_server.cpp # 策略进程 -│ ├── onnxruntime_inference.hpp -│ └── rl_controller.cpp # 遗留控制器(P1 删除) +│ ├── policy_server.cpp # 策略进程(独立可执行文件,非 executor 组件) +│ ├── policy_server_launcher.cpp # PolicyServerLauncher 组件(拉起/重启策略进程) +│ └── onnxruntime_inference.hpp ├── tool/ # 见上表 -├── plugins.xml # rmcs_rl_bridge / rmcs_rl_legacy 的 pluginlib 导出 +├── plugins.xml # rmcs_rl_bridge 的 pluginlib 导出 └── CMakeLists.txt ``` ## 相关 -- 集成示例:`rmcs_bringup/config/wheel-leg-infantry-rl.yaml`(当前是遗留 `RlController` 配置) +- 集成示例:`rmcs_bringup/config/deformable-infantry-omni-rl.yaml`(桥式 + launcher); + `wheel-leg-infantry-rl.yaml` 尚未挂 RL 桥(待迁到桥式) - 训练侧:任何能导出 `obs[1,N] → actions[1,M]` 且带 `rmcs_obs_layout` / `rmcs_actions_layout` metadata 的 ONNX 仓库(Isaac Lab / legged_gym / rsl_rl …)都可以接 diff --git a/config/executor.yaml b/config/executor.yaml index 10462a4..02b0244 100644 --- a/config/executor.yaml +++ b/config/executor.yaml @@ -5,8 +5,8 @@ rmcs_executor: components: - rmcs_core::hardware::WheelLegInfantryRL -> wheel_leg_infantry_rl - rmcs_core::controller::chassis::WheelLegChassisController -> wheel_leg_chassis_controller - - rmcs::rl::RlBridge -> rl_bridge - - rmcs::rl::PolicyServer -> policy_server + - rmcs_rl::RlBridge -> rl_bridge + - rmcs_rl::PolicyServerLauncher -> policy_server_launcher wheel_leg_infantry_rl: ros__parameters: @@ -48,7 +48,7 @@ rl_bridge: joint_velocity_suffix: "/velocity" joint_torque_suffix: "/torque" default_joint_pos: - left_hip_joint: -0.5v + left_hip_joint: -0.5 left_knee_joint: -0.35 left_wheel: 0.0 right_hip_joint: 0.5 @@ -86,3 +86,14 @@ policy_server: normalization_from_metadata: true publish_status: false status_rate: 2.0 + +# 由 executor 进程内组件随 rmcs 一起拉起 policy_server(不改 bringup/launch/服务)。 +# params_file 相对 share/rmcs_bringup/config 解析,也可写绝对路径; +# 复制本模板到 rmcs_bringup/config/.yaml 后,改成那个文件名。 +policy_server_launcher: + ros__parameters: + autostart: true + params_file: "executor.yaml" + respawn: true + respawn_delay: 1.0 + poll_interval: 0.5 diff --git a/msg/Observation.msg b/msg/Observation.msg index dbf1861..2cee0f3 100644 --- a/msg/Observation.msg +++ b/msg/Observation.msg @@ -1,4 +1,4 @@ std_msgs/Header header # stamp = 采样时刻(桥填 ROS 时间);frame_id 留空 uint64 obs_seq # 单调递增的观测序号,桥生成 -uint64 layout_hash # 桥按 YAML 词条算出的合同指纹(FNV-1a64,见 doc/bridge-design.md §6.2) +uint64 layout_hash # 桥按 YAML 词条算出的合同指纹(FNV-1a64,见 planning/docs/bridge-design.md §6.2) float64[] obs # 长度 == obs_size diff --git a/plugins.xml b/plugins.xml index 9f556d8..4025b71 100644 --- a/plugins.xml +++ b/plugins.xml @@ -1,7 +1,4 @@ - - - - - + + diff --git a/src/onnxruntime_inference.hpp b/src/onnxruntime_inference.hpp index 95614c6..fd7dd3d 100644 --- a/src/onnxruntime_inference.hpp +++ b/src/onnxruntime_inference.hpp @@ -13,21 +13,21 @@ #include #include -namespace rmcs::rl { +namespace rmcs_rl { class OnnxRuntimeInference { public: struct Config { std::string model_path; - std::string input_name = "obs"; + std::string input_name = "obs"; std::string output_name = "actions"; - std::size_t input_size = 0; + std::size_t input_size = 0; std::size_t output_size = 0; }; OnnxRuntimeInference() = default; - OnnxRuntimeInference(const OnnxRuntimeInference&) = delete; + OnnxRuntimeInference(const OnnxRuntimeInference&) = delete; OnnxRuntimeInference& operator=(const OnnxRuntimeInference&) = delete; bool load(const Config& config) { @@ -48,27 +48,27 @@ class OnnxRuntimeInference { if (session_->GetInputCount() != 1 || session_->GetOutputCount() != 1) { error = "model must have exactly 1 input and 1 output, got " - + std::to_string(session_->GetInputCount()) + " input(s) / " - + std::to_string(session_->GetOutputCount()) + " output(s)"; + + std::to_string(session_->GetInputCount()) + " input(s) / " + + std::to_string(session_->GetOutputCount()) + " output(s)"; session_.reset(); return false; } - const auto actual_input = session_->GetInputNameAllocated(0, allocator_); + const auto actual_input = session_->GetInputNameAllocated(0, allocator_); const auto actual_output = session_->GetOutputNameAllocated(0, allocator_); if (actual_input.get() != config_.input_name || actual_output.get() != config_.output_name) { error = "tensor names must be '" + config_.input_name + "' / '" - + config_.output_name + "', got '" + actual_input.get() + "' / '" - + actual_output.get() + "'"; + + config_.output_name + "', got '" + actual_input.get() + "' / '" + + actual_output.get() + "'"; session_.reset(); return false; } - const auto input_type_info = session_->GetInputTypeInfo(0); + const auto input_type_info = session_->GetInputTypeInfo(0); const auto output_type_info = session_->GetOutputTypeInfo(0); - const auto input_info = input_type_info.GetTensorTypeAndShapeInfo(); - const auto output_info = output_type_info.GetTensorTypeAndShapeInfo(); + const auto input_info = input_type_info.GetTensorTypeAndShapeInfo(); + const auto output_info = output_type_info.GetTensorTypeAndShapeInfo(); if (input_info.GetElementType() != ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT || output_info.GetElementType() != ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT) { error = "tensor element type must be float32"; @@ -76,7 +76,7 @@ class OnnxRuntimeInference { return false; } - const auto input_shape = input_info.GetShape(); + const auto input_shape = input_info.GetShape(); const auto output_shape = output_info.GetShape(); if (input_shape.size() != 2 || output_shape.size() != 2) { error = "tensor rank must be 2 ([1, N])"; @@ -89,22 +89,22 @@ class OnnxRuntimeInference { return false; } - const auto model_input_size = static_cast(input_shape[1]); + const auto model_input_size = static_cast(input_shape[1]); const auto model_output_size = static_cast(output_shape[1]); if (config_.input_size != 0 && config_.input_size != model_input_size) { error = "configured rl_obs_size=" + std::to_string(config_.input_size) - + " but model input shape is [1," + std::to_string(model_input_size) + "]"; + + " but model input shape is [1," + std::to_string(model_input_size) + "]"; session_.reset(); return false; } if (config_.output_size != 0 && config_.output_size != model_output_size) { error = "configured rl_action_size=" + std::to_string(config_.output_size) - + " but model output shape is [1," + std::to_string(model_output_size) + "]"; + + " but model output shape is [1," + std::to_string(model_output_size) + "]"; session_.reset(); return false; } - config_.input_size = model_input_size; + config_.input_size = model_input_size; config_.output_size = model_output_size; input_shape_.assign(input_shape.begin(), input_shape.end()); output_shape_.assign(output_shape.begin(), output_shape.end()); @@ -113,7 +113,7 @@ class OnnxRuntimeInference { memory_info_ = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault); return true; } catch (const Ort::Exception& exception) { - error = std::string { "ONNX Runtime: " } + exception.what(); + error = std::string{"ONNX Runtime: "} + exception.what(); session_.reset(); return false; } @@ -126,12 +126,13 @@ class OnnxRuntimeInference { [[nodiscard]] std::size_t output_size() const { return config_.output_size; } [[nodiscard]] std::optional metadata(const std::string& key) const { - if (!session_) return std::nullopt; + if (!session_) + return std::nullopt; try { auto model_metadata = session_->GetModelMetadata(); - auto value = - model_metadata.LookupCustomMetadataMapAllocated(key.c_str(), allocator_); - if (!value) return std::nullopt; + auto value = model_metadata.LookupCustomMetadataMapAllocated(key.c_str(), allocator_); + if (!value) + return std::nullopt; return std::string(value.get()); } catch (const Ort::Exception&) { return std::nullopt; @@ -142,21 +143,25 @@ class OnnxRuntimeInference { if (!session_ || input.size() < config_.input_size || output.size() < config_.output_size) return false; try { - std::copy(input.begin(), input.begin() + static_cast(config_.input_size), + std::copy( + input.begin(), input.begin() + static_cast(config_.input_size), input_buffer_.begin()); - Ort::Value input_tensor = Ort::Value::CreateTensor(memory_info_, - input_buffer_.data(), config_.input_size, input_shape_.data(), input_shape_.size()); - Ort::Value output_tensor = Ort::Value::CreateTensor(memory_info_, - output_buffer_.data(), config_.output_size, output_shape_.data(), + Ort::Value input_tensor = Ort::Value::CreateTensor( + memory_info_, input_buffer_.data(), config_.input_size, input_shape_.data(), + input_shape_.size()); + Ort::Value output_tensor = Ort::Value::CreateTensor( + memory_info_, output_buffer_.data(), config_.output_size, output_shape_.data(), output_shape_.size()); - const char* input_names[] = {config_.input_name.c_str()}; + const char* input_names[] = {config_.input_name.c_str()}; const char* output_names[] = {config_.output_name.c_str()}; - auto outputs = session_->Run( - Ort::RunOptions { nullptr }, input_names, &input_tensor, 1, output_names, 1); - if (outputs.size() != 1 || !outputs[0].IsTensor()) return false; + auto outputs = session_->Run( + Ort::RunOptions{nullptr}, input_names, &input_tensor, 1, output_names, 1); + if (outputs.size() != 1 || !outputs[0].IsTensor()) + return false; const float* data = outputs[0].GetTensorData(); - std::copy(data, data + static_cast(config_.output_size), output.begin()); + std::copy( + data, data + static_cast(config_.output_size), output.begin()); return true; } catch (const Ort::Exception&) { return false; @@ -164,10 +169,10 @@ class OnnxRuntimeInference { } private: - Ort::Env env_ { ORT_LOGGING_LEVEL_WARNING, "rmcs_rl" }; + Ort::Env env_{ORT_LOGGING_LEVEL_WARNING, "rmcs_rl"}; Ort::SessionOptions session_options_; Ort::AllocatorWithDefaultOptions allocator_; - Ort::MemoryInfo memory_info_ { nullptr }; + Ort::MemoryInfo memory_info_{nullptr}; std::unique_ptr session_; std::vector input_shape_; std::vector output_shape_; @@ -176,4 +181,4 @@ class OnnxRuntimeInference { Config config_; }; -} // namespace rmcs::rl +} // namespace rmcs_rl diff --git a/src/policy_server.cpp b/src/policy_server.cpp index 9e2aa91..c457e1c 100644 --- a/src/policy_server.cpp +++ b/src/policy_server.cpp @@ -28,64 +28,68 @@ #include "onnxruntime_inference.hpp" #include "rl_layout.hpp" -namespace rmcs::rl { +namespace rmcs_rl { namespace { - std::vector parse_float_list(const std::string& text, const char* what) { - std::vector values; - std::size_t begin = 0; - while (begin <= text.size()) { - const auto end = text.find_first_of(", \t", begin); - const auto piece = - text.substr(begin, end == std::string::npos ? std::string::npos : end - begin); - if (!piece.empty()) { - try { - std::size_t consumed = 0; - const double value = std::stod(piece, &consumed); - if (consumed != piece.size()) throw std::invalid_argument("trailing"); - values.push_back(value); - } catch (const std::exception&) { - throw std::invalid_argument( - std::string { "policy_server: metadata " } + what + " has a non-number '" - + piece + "'"); - } +std::vector parse_float_list(const std::string& text, const char* what) { + std::vector values; + std::size_t begin = 0; + while (begin <= text.size()) { + const auto end = text.find_first_of(", \t", begin); + const auto piece = + text.substr(begin, end == std::string::npos ? std::string::npos : end - begin); + if (!piece.empty()) { + try { + std::size_t consumed = 0; + const double value = std::stod(piece, &consumed); + if (consumed != piece.size()) + throw std::invalid_argument("trailing"); + values.push_back(value); + } catch (const std::exception&) { + throw std::invalid_argument( + std::string{"policy_server: metadata "} + what + " has a non-number '" + piece + + "'"); } - if (end == std::string::npos) break; - begin = end + 1; } - return values; + if (end == std::string::npos) + break; + begin = end + 1; } + return values; +} - std::optional parse_float(const std::string& text) { - try { - std::size_t consumed = 0; - const double value = std::stod(text, &consumed); - if (consumed != text.size()) return std::nullopt; - return value; - } catch (const std::exception&) { +std::optional parse_float(const std::string& text) { + try { + std::size_t consumed = 0; + const double value = std::stod(text, &consumed); + if (consumed != text.size()) return std::nullopt; - } + return value; + } catch (const std::exception&) { + return std::nullopt; } +} - double parse_clip_metadata(const std::string& text, const char* key) { - const auto value = parse_float(text); - if (!value.has_value()) - throw std::runtime_error( - std::string { "policy_server: metadata " } + key + "=" + text + " is not a float"); - if (!std::isfinite(*value) || *value <= 0.0) - throw std::runtime_error(std::string { "policy_server: metadata " } + key - + " must be finite and > 0, got " + text); - return *value; - } +double parse_clip_metadata(const std::string& text, const char* key) { + const auto value = parse_float(text); + if (!value.has_value()) + throw std::runtime_error( + std::string{"policy_server: metadata "} + key + "=" + text + " is not a float"); + if (!std::isfinite(*value) || *value <= 0.0) + throw std::runtime_error( + std::string{"policy_server: metadata "} + key + " must be finite and > 0, got " + text); + return *value; +} } // namespace class PolicyServer final : public rclcpp::Node { public: PolicyServer() - : Node("policy_server", - rclcpp::NodeOptions { }.automatically_declare_parameters_from_overrides(true)) { + : Node( + "policy_server", + rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { rl_base_ = string_or_("rl_base", "/rl"); @@ -96,14 +100,14 @@ class PolicyServer final : public rclcpp::Node { std::string load_error; OnnxRuntimeInference::Config inference_config; - inference_config.model_path = resolved_model_path_; - inference_config.input_name = string_or_("input_name", "obs"); + inference_config.model_path = resolved_model_path_; + inference_config.input_name = string_or_("input_name", "obs"); inference_config.output_name = string_or_("output_name", "actions"); if (!inference_.load(inference_config, load_error)) throw std::runtime_error( "policy_server: cannot load '" + resolved_model_path_ + "': " + load_error); - obs_size_ = inference_.input_size(); + obs_size_ = inference_.input_size(); action_size_ = inference_.output_size(); if (obs_size_ == 0 || action_size_ == 0) throw std::runtime_error("policy_server: model has zero-sized obs/action tensor"); @@ -115,13 +119,16 @@ class PolicyServer final : public rclcpp::Node { const auto obs_layout = inference_.metadata("rmcs_obs_layout"); const auto actions_layout = inference_.metadata("rmcs_actions_layout"); if (!obs_layout || obs_layout->empty() || !actions_layout || actions_layout->empty()) - throw std::runtime_error("policy_server: model '" + resolved_model_path_ + throw std::runtime_error( + "policy_server: model '" + resolved_model_path_ + "' is missing metadata 'rmcs_obs_layout' / 'rmcs_actions_layout'; stamp it with " - "tool/stamp_layout_metadata.py before deploying (see doc/bridge-design.md §6)"); - obs_signature_ = *obs_layout; + "tool/stamp_layout_metadata.py before deploying (see " + "planning/docs/bridge-design.md §6)"); + obs_signature_ = *obs_layout; actions_signature_ = *actions_layout; - policy_version_ = inference_.metadata("policy_version").value_or(""); - layout_hash_ = rmcs::rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); + policy_version_ = inference_.metadata("policy_version").value_or(""); + layout_hash_ = + rmcs_rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); if (const auto declared = inference_.metadata("policy_layout_hash"); declared && !declared->empty()) { @@ -129,22 +136,24 @@ class PolicyServer final : public rclcpp::Node { try { declared_hash = std::stoull(*declared, nullptr, 16); } catch (const std::exception&) { - throw std::runtime_error("policy_server: metadata 'policy_layout_hash' is not a " - "hex string: '" + throw std::runtime_error( + "policy_server: metadata 'policy_layout_hash' is not a " + "hex string: '" + *declared + "'"); } if (declared_hash != layout_hash_) - throw std::runtime_error("policy_server: metadata 'policy_layout_hash'=" - + *declared + " does not match the hash computed from the layout metadata (" + throw std::runtime_error( + "policy_server: metadata 'policy_layout_hash'=" + *declared + + " does not match the hash computed from the layout metadata (" + hex16(layout_hash_) + "); re-stamp the model"); } load_normalization_(); action_publisher_ = create_publisher( - rl_base_ + "/action", rclcpp::QoS { rclcpp::KeepLast(1) }.best_effort()); + rl_base_ + "/action", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort()); observation_subscription_ = create_subscription( - rl_base_ + "/obs", rclcpp::QoS { rclcpp::KeepLast(1) }.best_effort(), + rl_base_ + "/obs", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort(), [this](rmcs_rl::msg::Observation::UniquePtr message) { on_observation_(std::move(message)); }); @@ -152,23 +161,26 @@ class PolicyServer final : public rclcpp::Node { if (bool_or_("publish_status", false)) { status_publisher_ = create_publisher( rl_base_ + "/policy_status", - rclcpp::QoS { rclcpp::KeepLast(1) }.transient_local().best_effort()); + rclcpp::QoS{rclcpp::KeepLast(1)}.transient_local().best_effort()); const double rate = std::max(number_or_("status_rate", 2.0), 0.1); - status_timer_ = create_wall_timer( + status_timer_ = create_wall_timer( std::chrono::duration(1.0 / rate), [this]() { publish_status_(); }); } RCLCPP_INFO(get_logger(), "policy loaded: %s", resolved_model_path_.c_str()); RCLCPP_INFO(get_logger(), " model_id : %s", hex16(model_id_).c_str()); - RCLCPP_INFO(get_logger(), " policy_version: %s", + RCLCPP_INFO( + get_logger(), " policy_version: %s", policy_version_.empty() ? "(none)" : policy_version_.c_str()); RCLCPP_INFO(get_logger(), " obs_size=%zu action_size=%zu", obs_size_, action_size_); RCLCPP_INFO(get_logger(), " obs signature : %s", obs_signature_.c_str()); RCLCPP_INFO(get_logger(), " action signature : %s", actions_signature_.c_str()); - RCLCPP_INFO(get_logger(), " layout_hash : %s (bridge must log the same value)", + RCLCPP_INFO( + get_logger(), " layout_hash : %s (bridge must log the same value)", hex16(layout_hash_).c_str()); if (!obs_mean_.empty()) - RCLCPP_INFO(get_logger(), " normalization : mean/std from metadata, obs_clip=%s", + RCLCPP_INFO( + get_logger(), " normalization : mean/std from metadata, obs_clip=%s", obs_clip_.has_value() ? std::to_string(*obs_clip_).c_str() : "(none)"); RCLCPP_INFO(get_logger(), "waiting for obs on %s/obs", rl_base_.c_str()); } @@ -177,18 +189,21 @@ class PolicyServer final : public rclcpp::Node { std::string string_or_(const std::string& name, const std::string& fallback) const { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return fallback; + if (!get_parameter(name, parameter)) + return fallback; } catch (const std::exception&) { return fallback; } - if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_STRING) return fallback; + if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_STRING) + return fallback; return parameter.as_string(); } double number_or_(const std::string& name, double fallback) const { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return fallback; + if (!get_parameter(name, parameter)) + return fallback; } catch (const std::exception&) { return fallback; } @@ -202,16 +217,19 @@ class PolicyServer final : public rclcpp::Node { bool bool_or_(const std::string& name, bool fallback) const { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return fallback; + if (!get_parameter(name, parameter)) + return fallback; } catch (const std::exception&) { return fallback; } - if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_BOOL) return fallback; + if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_BOOL) + return fallback; return parameter.as_bool(); } static std::string resolve_model_path_(const std::string& path) { - if (path.empty() || path.front() == '/') return path; + if (path.empty() || path.front() == '/') + return path; try { return ament_index_cpp::get_package_share_directory("rmcs_rl") + "/" + path; } catch (const std::exception&) { @@ -227,18 +245,19 @@ class PolicyServer final : public rclcpp::Node { if (const auto value = inference_.metadata("rmcs_obs_std"); value && !value->empty()) obs_std_ = parse_float_list(*value, "rmcs_obs_std"); if (!obs_mean_.empty() && obs_mean_.size() != obs_size_) - throw std::runtime_error("policy_server: metadata rmcs_obs_mean has " - + std::to_string(obs_mean_.size()) + " values, expected " - + std::to_string(obs_size_)); + throw std::runtime_error( + "policy_server: metadata rmcs_obs_mean has " + std::to_string(obs_mean_.size()) + + " values, expected " + std::to_string(obs_size_)); if (!obs_std_.empty() && obs_std_.size() != obs_size_) - throw std::runtime_error("policy_server: metadata rmcs_obs_std has " - + std::to_string(obs_std_.size()) + " values, expected " - + std::to_string(obs_size_)); + throw std::runtime_error( + "policy_server: metadata rmcs_obs_std has " + std::to_string(obs_std_.size()) + + " values, expected " + std::to_string(obs_size_)); if (!obs_mean_.empty() && !obs_std_.empty()) { for (const double sigma : obs_std_) if (!(sigma > 0.0) || !std::isfinite(sigma)) throw std::runtime_error( - "policy_server: metadata rmcs_obs_std must be finite and > 0"); + "policy_server: metadata rmcs_obs_std must be " + "finite and > 0"); } else { obs_mean_.clear(); obs_std_.clear(); @@ -251,9 +270,11 @@ class PolicyServer final : public rclcpp::Node { } const double obs_clip_param = number_or_("obs_clip", -1.0); - if (obs_clip_param >= 0.0) obs_clip_ = obs_clip_param; + if (obs_clip_param >= 0.0) + obs_clip_ = obs_clip_param; const double action_clip_param = number_or_("action_clip", -1.0); - if (action_clip_param >= 0.0) action_clip_ = action_clip_param; + if (action_clip_param >= 0.0) + action_clip_ = action_clip_param; obs_buffer_.assign(obs_size_, 0.0F); action_buffer_.assign(action_size_, 0.0F); @@ -263,7 +284,8 @@ class PolicyServer final : public rclcpp::Node { if (message->layout_hash != layout_hash_) { if (!layout_mismatch_logged_) { layout_mismatch_logged_ = true; - RCLCPP_FATAL(get_logger(), + RCLCPP_FATAL( + get_logger(), "layout_hash mismatch: bridge sent %s but this model expects %s " "(model '%s'); refusing to answer (bridge will report valid=0). " "Check that the deployment YAML and the model metadata describe the same " @@ -275,8 +297,9 @@ class PolicyServer final : public rclcpp::Node { return; } if (message->obs.size() != obs_size_) { - RCLCPP_ERROR_THROTTLE(get_logger(), *get_clock(), 1000, - "obs has %zu values, model expects %zu; ignoring", message->obs.size(), obs_size_); + RCLCPP_ERROR_THROTTLE( + get_logger(), *get_clock(), 1000, "obs has %zu values, model expects %zu; ignoring", + message->obs.size(), obs_size_); ++rejected_count_; return; } @@ -285,11 +308,13 @@ class PolicyServer final : public rclcpp::Node { for (std::size_t i = 0; i < obs_size_; ++i) { double value = message->obs[i]; - if (!obs_mean_.empty()) value = (value - obs_mean_[i]) / obs_std_[i]; - if (obs_clip_.has_value()) value = std::clamp(value, -*obs_clip_, *obs_clip_); + if (!obs_mean_.empty()) + value = (value - obs_mean_[i]) / obs_std_[i]; + if (obs_clip_.has_value()) + value = std::clamp(value, -*obs_clip_, *obs_clip_); if (!std::isfinite(value)) { - RCLCPP_WARN_THROTTLE(get_logger(), *get_clock(), 1000, - "obs[%zu] is not finite; ignoring frame", i); + RCLCPP_WARN_THROTTLE( + get_logger(), *get_clock(), 1000, "obs[%zu] is not finite; ignoring frame", i); ++rejected_count_; return; } @@ -297,15 +322,17 @@ class PolicyServer final : public rclcpp::Node { } if (!inference_.run(obs_buffer_, action_buffer_)) { - RCLCPP_ERROR_THROTTLE(get_logger(), *get_clock(), 1000, - "ONNX Runtime rejected the obs frame; ignoring"); + RCLCPP_ERROR_THROTTLE( + get_logger(), *get_clock(), 1000, "ONNX Runtime rejected the obs frame; ignoring"); ++rejected_count_; return; } - if (!std::all_of(action_buffer_.begin(), action_buffer_.end(), - [](float value) { return std::isfinite(value); })) { - RCLCPP_WARN_THROTTLE(get_logger(), *get_clock(), 1000, + if (!std::all_of(action_buffer_.begin(), action_buffer_.end(), [](float value) { + return std::isfinite(value); + })) { + RCLCPP_WARN_THROTTLE( + get_logger(), *get_clock(), 1000, "policy output contains non-finite values; ignoring frame " "(bridge will see a stale action and drop authority)"); ++rejected_count_; @@ -314,19 +341,20 @@ class PolicyServer final : public rclcpp::Node { rmcs_rl::msg::Action action; action.header.stamp = get_clock()->now(); - action.obs_seq = message->obs_seq; - action.layout_hash = layout_hash_; - action.model_id = model_id_; + action.obs_seq = message->obs_seq; + action.layout_hash = layout_hash_; + action.model_id = model_id_; action.action.resize(action_size_); for (std::size_t i = 0; i < action_size_; ++i) { double value = static_cast(action_buffer_[i]); - if (action_clip_.has_value()) value = std::clamp(value, -*action_clip_, *action_clip_); + if (action_clip_.has_value()) + value = std::clamp(value, -*action_clip_, *action_clip_); action.action[i] = value; } action_publisher_->publish(action); - const auto elapsed = std::chrono::duration( - std::chrono::steady_clock::now() - started); + const auto elapsed = + std::chrono::duration(std::chrono::steady_clock::now() - started); record_inference_time_(elapsed.count()); ++served_count_; } @@ -334,31 +362,33 @@ class PolicyServer final : public rclcpp::Node { void record_inference_time_(double microseconds) { inference_window_[inference_window_cursor_] = microseconds; inference_window_cursor_ = (inference_window_cursor_ + 1) % inference_window_.size(); - inference_window_count_ = std::min(inference_window_count_ + 1, inference_window_.size()); + inference_window_count_ = std::min(inference_window_count_ + 1, inference_window_.size()); } double percentile_(double fraction) const { - if (inference_window_count_ == 0) return 0.0; + if (inference_window_count_ == 0) + return 0.0; std::vector samples( inference_window_.begin(), inference_window_.begin() + inference_window_count_); std::sort(samples.begin(), samples.end()); - const auto index = static_cast( - fraction * static_cast(samples.size() - 1) + 0.5); + const auto index = + static_cast(fraction * static_cast(samples.size() - 1) + 0.5); return samples[std::min(index, samples.size() - 1)]; } void publish_status_() { - if (!status_publisher_) return; + if (!status_publisher_) + return; rmcs_rl::msg::PolicyStatus status; - status.header.stamp = get_clock()->now(); - status.model_name = resolved_model_path_; - status.model_id = model_id_; - status.policy_version = policy_version_; - status.layout_hash = layout_hash_; - status.obs_size = static_cast(obs_size_); - status.action_size = static_cast(action_size_); - status.inference_p50_us = percentile_(0.50); - status.inference_p99_us = percentile_(0.99); + status.header.stamp = get_clock()->now(); + status.model_name = resolved_model_path_; + status.model_id = model_id_; + status.policy_version = policy_version_; + status.layout_hash = layout_hash_; + status.obs_size = static_cast(obs_size_); + status.action_size = static_cast(action_size_); + status.inference_p50_us = percentile_(0.50); + status.inference_p99_us = percentile_(0.99); status_publisher_->publish(status); } @@ -367,10 +397,10 @@ class PolicyServer final : public rclcpp::Node { std::string obs_signature_; std::string actions_signature_; std::string policy_version_; - std::size_t obs_size_ = 0; + std::size_t obs_size_ = 0; std::size_t action_size_ = 0; std::uint64_t layout_hash_ = 0; - std::uint64_t model_id_ = 0; + std::uint64_t model_id_ = 0; std::vector obs_mean_; std::vector obs_std_; @@ -380,13 +410,13 @@ class PolicyServer final : public rclcpp::Node { std::vector obs_buffer_; std::vector action_buffer_; - std::uint64_t served_count_ = 0; + std::uint64_t served_count_ = 0; std::uint64_t rejected_count_ = 0; - bool layout_mismatch_logged_ = false; + bool layout_mismatch_logged_ = false; - std::array inference_window_ { }; + std::array inference_window_{}; std::size_t inference_window_cursor_ = 0; - std::size_t inference_window_count_ = 0; + std::size_t inference_window_count_ = 0; OnnxRuntimeInference inference_; rclcpp::Publisher::SharedPtr action_publisher_; @@ -395,12 +425,12 @@ class PolicyServer final : public rclcpp::Node { rclcpp::TimerBase::SharedPtr status_timer_; }; -} // namespace rmcs::rl +} // namespace rmcs_rl int main(int argc, char** argv) { rclcpp::init(argc, argv); try { - auto node = std::make_shared(); + auto node = std::make_shared(); rclcpp::spin(node); } catch (const std::exception& error) { fprintf(stderr, "[Fatal] policy_server startup failed: %s\n", error.what()); diff --git a/src/policy_server_launcher.cpp b/src/policy_server_launcher.cpp index 6f18f1b..a22561d 100644 --- a/src/policy_server_launcher.cpp +++ b/src/policy_server_launcher.cpp @@ -19,7 +19,7 @@ #include #include -namespace rmcs::rl { +namespace rmcs_rl { // 随 executor 生命周期拉起独立的 policy_server 子进程: // - 组件只存在于需要 RL 的配置里,非 RL 车不受影响; @@ -104,7 +104,8 @@ class PolicyServerLauncher } if (executable_.empty() || ::access(executable_.c_str(), X_OK) != 0) { - RCLCPP_ERROR(get_logger(), + RCLCPP_ERROR( + get_logger(), "policy_server executable not found (looked for '%s'); autostart disabled", executable_.c_str()); autostart_ = false; @@ -115,18 +116,22 @@ class PolicyServerLauncher params_path_ = params_file_; } else { try { - const auto bringup_share = - ament_index_cpp::get_package_share_directory("rmcs_bringup"); + const auto bringup_share = ament_index_cpp::get_package_share_directory( + "rmcs_" + "bringu" + "p"); params_path_ = bringup_share + "/config/" + params_file_; } catch (const std::exception& error) { - RCLCPP_ERROR(get_logger(), "cannot locate rmcs_bringup share directory: %s", + RCLCPP_ERROR( + get_logger(), "cannot locate rmcs_bringup share directory: %s", error.what()); } } } if (params_path_.empty() || ::access(params_path_.c_str(), R_OK) != 0) { - RCLCPP_ERROR(get_logger(), + RCLCPP_ERROR( + get_logger(), "policy_server params file not found (params_file='%s'); autostart disabled", params_file_.c_str()); autostart_ = false; @@ -173,8 +178,9 @@ class PolicyServerLauncher child_pid_ = pid; next_start_time_ = SteadyClock::now() + respawn_period_; - RCLCPP_INFO(get_logger(), "started policy_server (pid=%d) with params '%s'", - static_cast(pid), params_path_.c_str()); + RCLCPP_INFO( + get_logger(), "started policy_server (pid=%d) with params '%s'", static_cast(pid), + params_path_.c_str()); } void reap_child_(SteadyClock::time_point now) { @@ -187,7 +193,8 @@ class PolicyServerLauncher return; if (result < 0 && errno != ECHILD) { - RCLCPP_WARN(get_logger(), "waitpid(%d) failed: %s", static_cast(child_pid_), + RCLCPP_WARN( + get_logger(), "waitpid(%d) failed: %s", static_cast(child_pid_), std::strerror(errno)); } @@ -196,11 +203,13 @@ class PolicyServerLauncher next_start_time_ = now + respawn_period_; if (result == pid && WIFEXITED(status)) - RCLCPP_WARN(get_logger(), "policy_server (pid=%d) exited with code %d", - static_cast(pid), WEXITSTATUS(status)); + RCLCPP_WARN( + get_logger(), "policy_server (pid=%d) exited with code %d", static_cast(pid), + WEXITSTATUS(status)); else if (result == pid && WIFSIGNALED(status)) - RCLCPP_WARN(get_logger(), "policy_server (pid=%d) killed by signal %d", - static_cast(pid), WTERMSIG(status)); + RCLCPP_WARN( + get_logger(), "policy_server (pid=%d) killed by signal %d", static_cast(pid), + WTERMSIG(status)); else RCLCPP_WARN(get_logger(), "policy_server (pid=%d) is gone", static_cast(pid)); } @@ -221,7 +230,8 @@ class PolicyServerLauncher ::kill(pid, SIGKILL); ::waitpid(pid, nullptr, 0); - RCLCPP_WARN(get_logger(), "policy_server (pid=%d) did not stop, killed", static_cast(pid)); + RCLCPP_WARN( + get_logger(), "policy_server (pid=%d) did not stop, killed", static_cast(pid)); } bool autostart_ = true; @@ -235,13 +245,13 @@ class PolicyServerLauncher SteadyClock::duration poll_period_ = std::chrono::milliseconds(500); SteadyClock::duration respawn_period_ = std::chrono::seconds(1); - SteadyClock::time_point last_poll_time_ { }; - SteadyClock::time_point next_start_time_ { }; + SteadyClock::time_point last_poll_time_{}; + SteadyClock::time_point next_start_time_{}; pid_t child_pid_ = -1; }; -} // namespace rmcs::rl +} // namespace rmcs_rl #include -PLUGINLIB_EXPORT_CLASS(rmcs::rl::PolicyServerLauncher, rmcs_executor::Component) +PLUGINLIB_EXPORT_CLASS(rmcs_rl::PolicyServerLauncher, rmcs_executor::Component) diff --git a/src/rl_bridge.cpp b/src/rl_bridge.cpp index 444429e..76841a2 100644 --- a/src/rl_bridge.cpp +++ b/src/rl_bridge.cpp @@ -33,94 +33,106 @@ #include "rl_layout.hpp" -namespace rmcs::rl { +namespace rmcs_rl { namespace { - constexpr std::size_t kNoSlot = std::numeric_limits::max(); - - std::vector split_by(const std::string& text, char delimiter) { - std::vector parts; - std::string current; - std::istringstream stream { text }; - while (std::getline(stream, current, delimiter)) - parts.push_back(current); - return parts; - } - - std::vector split_whitespace(const std::string& text) { - std::vector tokens; - std::istringstream stream { text }; - std::string token; - while (stream >> token) - tokens.push_back(token); - return tokens; +constexpr std::size_t kNoSlot = std::numeric_limits::max(); + +std::vector split_by(const std::string& text, char delimiter) { + std::vector parts; + std::string current; + std::istringstream stream{text}; + while (std::getline(stream, current, delimiter)) + parts.push_back(current); + return parts; +} + +std::vector split_whitespace(const std::string& text) { + std::vector tokens; + std::istringstream stream{text}; + std::string token; + while (stream >> token) + tokens.push_back(token); + return tokens; +} + +std::string format_number(double value) { + char buffer[32]; + std::snprintf(buffer, sizeof(buffer), "%.6g", value); + return buffer; +} + +std::string join(const std::vector& parts, const char* separator) { + std::string text; + for (std::size_t i = 0; i < parts.size(); ++i) { + if (i != 0) + text += separator; + text += parts[i]; } - - std::string format_number(double value) { - char buffer[32]; - std::snprintf(buffer, sizeof(buffer), "%.6g", value); - return buffer; - } - - std::string join(const std::vector& parts, const char* separator) { - std::string text; - for (std::size_t i = 0; i < parts.size(); ++i) { - if (i != 0) text += separator; - text += parts[i]; - } - return text; - } - - double parse_double(const std::string& text, const std::string& context) { - try { - std::size_t consumed = 0; - const double value = std::stod(text, &consumed); - if (consumed != text.size()) throw std::invalid_argument("trailing characters"); - return value; - } catch (const std::exception&) { - throw std::invalid_argument( - "RlBridge: '" + text + "' is not a number (" + context + ")"); - } - } - - std::int64_t parse_integer(const std::string& text, const std::string& context) { - try { - std::size_t consumed = 0; - const auto value = std::stoll(text, &consumed); - if (consumed != text.size()) throw std::invalid_argument("trailing characters"); - return value; - } catch (const std::exception&) { - throw std::invalid_argument( - "RlBridge: '" + text + "' is not an integer (" + context + ")"); - } + return text; +} + +double parse_double(const std::string& text, const std::string& context) { + try { + std::size_t consumed = 0; + const double value = std::stod(text, &consumed); + if (consumed != text.size()) + throw std::invalid_argument("trailing characters"); + return value; + } catch (const std::exception&) { + throw std::invalid_argument("RlBridge: '" + text + "' is not a number (" + context + ")"); } - - bool parse_boolean(const std::string& text, const std::string& context) { - if (text == "true" || text == "1") return true; - if (text == "false" || text == "0") return false; - throw std::invalid_argument( - "RlBridge: '" + text + "' is not a boolean (true/false) (" + context + ")"); +} + +std::int64_t parse_integer(const std::string& text, const std::string& context) { + try { + std::size_t consumed = 0; + const auto value = std::stoll(text, &consumed); + if (consumed != text.size()) + throw std::invalid_argument("trailing characters"); + return value; + } catch (const std::exception&) { + throw std::invalid_argument("RlBridge: '" + text + "' is not an integer (" + context + ")"); } +} - bool is_finite(double value) { return std::isfinite(value); } - - std::string pretty_type(const std::type_info& type) { - if (type == typeid(double)) return "double"; - if (type == typeid(float)) return "float"; - if (type == typeid(bool)) return "bool"; - if (type == typeid(int)) return "int"; - if (type == typeid(std::size_t)) return "std::size_t"; - if (type == typeid(Eigen::Vector3d)) return "Eigen::Vector3d"; - if (type == typeid(rmcs_description::BaseLink::DirectionVector)) - return "rmcs_description::BaseLink::DirectionVector"; - if (type == typeid(Eigen::Quaterniond)) return "Eigen::Quaterniond"; - return std::string { type.name() } + " (unknown)"; - } +bool parse_boolean(const std::string& text, const std::string& context) { + if (text == "true" || text == "1") + return true; + if (text == "false" || text == "0") + return false; + throw std::invalid_argument( + "RlBridge: '" + text + "' is not a boolean (true/false) (" + context + ")"); +} + +bool is_finite(double value) { return std::isfinite(value); } + +std::string pretty_type(const std::type_info& type) { + if (type == typeid(double)) + return "double"; + if (type == typeid(float)) + return "float"; + if (type == typeid(bool)) + return "bool"; + if (type == typeid(int)) + return "int"; + if (type == typeid(std::size_t)) + return "std::size_t"; + if (type == typeid(Eigen::Vector3d)) + return "Eigen::Vector3d"; + if (type == typeid(rmcs_description::BaseLink::DirectionVector)) + return "rmcs_description::BaseLink::DirectionVector"; + if (type == typeid(Eigen::Quaterniond)) + return "Eigen::Quaterniond"; + return std::string{type.name()} + " (unknown)"; +} } // namespace -class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { +class RlBridge final + : public rmcs_executor::Component + , public rclcpp::Node { enum class TermKind { kPath, kJointPos, kJointVel, kJointTorque, kLastAction, kConstant }; @@ -138,24 +150,24 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { struct ObsTerm { TermKind kind = TermKind::kPath; - Take take = Take::kScalar; + Take take = Take::kScalar; std::string path; std::string id; - std::size_t dim = 1; + std::size_t dim = 1; - std::size_t index = 0; - bool has_index = false; - int component = 0; - bool has_binding = false; - Binding binding = Binding::kDouble; + std::size_t index = 0; + bool has_index = false; + int component = 0; + bool has_binding = false; + Binding binding = Binding::kDouble; std::vector joints; std::vector joint_names; std::vector joint_slots; std::vector joint_defaults; bool relative = false; - bool zero = false; + bool zero = false; std::vector action_indices; std::vector constants; @@ -183,7 +195,7 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { struct Slot { std::string path; Binding binding = Binding::kDouble; - bool required = true; + bool required = true; std::unique_ptr> double_value; std::unique_ptr> bool_value; std::unique_ptr> int_value; @@ -195,17 +207,18 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { }; struct ActionSnapshot { - std::uint64_t obs_seq = 0; + std::uint64_t obs_seq = 0; std::uint64_t layout_hash = 0; - std::uint64_t model_id = 0; + std::uint64_t model_id = 0; std::vector action; - std::chrono::steady_clock::time_point received { }; + std::chrono::steady_clock::time_point received{}; }; public: explicit RlBridge() - : Node(get_component_name(), - rclcpp::NodeOptions { }.automatically_declare_parameters_from_overrides(true)) { + : Node( + get_component_name(), + rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { rl_base_ = string_or_("rl_base", "/rl"); @@ -219,26 +232,29 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (const auto& spec : action_specs) { auto term = parse_action_term_(spec); if (!output_paths.emplace(term.output).second) - throw std::invalid_argument("RlBridge: two action terms map to the same output " - "interface '" + throw std::invalid_argument( + "RlBridge: two action terms map to the same output " + "interface '" + term.output + "'"); if (!action_indices.emplace(term.index).second) - throw std::invalid_argument("RlBridge: duplicate action index " - + std::to_string(term.index) + " (each slot must appear exactly once)"); + throw std::invalid_argument( + "RlBridge: duplicate action index " + std::to_string(term.index) + + " (each slot must appear exactly once)"); action_terms_.push_back(std::move(term)); } for (std::size_t i = 0; i < action_terms_.size(); ++i) if (action_terms_[i].index != i) - throw std::invalid_argument("RlBridge: action indices must form the complete " - "permutation 0.." + throw std::invalid_argument( + "RlBridge: action indices must form the complete " + "permutation 0.." + std::to_string(action_terms_.size() - 1) + "; got index=" - + std::to_string(action_terms_[i].index) + " at position " - + std::to_string(i)); + + std::to_string(action_terms_[i].index) + " at position " + std::to_string(i)); action_size_ = action_terms_.size(); - if (const auto declared = integer_or_("rl_action_size"); declared.has_value() - && static_cast(*declared) != action_size_) - throw std::invalid_argument("RlBridge: rl_action_size=" + std::to_string(*declared) - + " but action_terms has " + std::to_string(action_size_) + " entries"); + if (const auto declared = integer_or_("rl_action_size"); + declared.has_value() && static_cast(*declared) != action_size_) + throw std::invalid_argument( + "RlBridge: rl_action_size=" + std::to_string(*declared) + " but action_terms has " + + std::to_string(action_size_) + " entries"); last_actions_.assign(action_size_, 0.0); written_.assign(action_size_, 0.0); @@ -251,24 +267,28 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { const auto observation_specs = string_array_or_("observation_terms"); if (observation_specs.empty()) throw std::invalid_argument( - "RlBridge: required parameter 'observation_terms' is missing"); + "RlBridge: required parameter 'observation_terms' is " + "missing"); for (const auto& spec : observation_specs) obs_terms_.push_back(parse_obs_term_(spec)); std::size_t cursor = 0; for (auto& term : obs_terms_) { if (term.has_index && term.index != cursor) - throw std::invalid_argument("RlBridge: observation term '" + term.id + "' declares " - "index=" + throw std::invalid_argument( + "RlBridge: observation term '" + term.id + + "' declares " + "index=" + std::to_string(term.index) + " but the running offset is " + std::to_string(cursor) + " (index= is an assertion, not a reorder)"); term.index = cursor; cursor += term.dim; } obs_size_ = cursor; - if (const auto declared = integer_or_("rl_obs_size"); declared.has_value() - && static_cast(*declared) != obs_size_) - throw std::invalid_argument("RlBridge: rl_obs_size=" + std::to_string(*declared) + if (const auto declared = integer_or_("rl_obs_size"); + declared.has_value() && static_cast(*declared) != obs_size_) + throw std::invalid_argument( + "RlBridge: rl_obs_size=" + std::to_string(*declared) + " but observation_terms sum to " + std::to_string(obs_size_)); policy_rate_ = number_or_("policy_rate", 50.0); @@ -282,33 +302,37 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { expected_model_id_ = parse_u64_(string_or_("expected_model_id", "0"), "expected_model_id"); const std::string invalid = string_or_("invalid_value", "nan"); - if (invalid == "nan") invalid_mode_ = InvalidMode::kNaN; - else if (invalid == "zero") invalid_mode_ = InvalidMode::kZero; - else if (invalid == "hold") invalid_mode_ = InvalidMode::kHold; + if (invalid == "nan") + invalid_mode_ = InvalidMode::kNaN; + else if (invalid == "zero") + invalid_mode_ = InvalidMode::kZero; + else if (invalid == "hold") + invalid_mode_ = InvalidMode::kHold; else - throw std::invalid_argument("RlBridge: invalid_value must be nan|zero|hold " - "(quote it in YAML: invalid_value: \"nan\" — an unquoted " - "nan is parsed as a float), got '" + throw std::invalid_argument( + "RlBridge: invalid_value must be nan|zero|hold " + "(quote it in YAML: invalid_value: \"nan\" — an unquoted " + "nan is parsed as a float), got '" + string_or_("invalid_value", "") + "'"); - enable_path_ = string_or_("enable_interface", rl_base_ + "/enable"); + enable_path_ = string_or_("enable_interface", rl_base_ + "/enable"); enable_default_ = bool_or_("enable_default", false); - reset_path_ = string_or_("reset_interface", ""); + reset_path_ = string_or_("reset_interface", ""); register_output(rl_base_ + "/valid", valid_output_, 0.0); register_output(rl_base_ + "/healthy", healthy_output_, 0.0); - register_output(rl_base_ + "/action_age", action_age_output_, - std::numeric_limits::quiet_NaN()); - register_output(rl_base_ + "/obs_seq", obs_seq_output_, std::size_t { 0 }); + register_output( + rl_base_ + "/action_age", action_age_output_, std::numeric_limits::quiet_NaN()); + register_output(rl_base_ + "/obs_seq", obs_seq_output_, std::size_t{0}); own_output_paths_.insert(rl_base_ + "/valid"); own_output_paths_.insert(rl_base_ + "/healthy"); own_output_paths_.insert(rl_base_ + "/action_age"); own_output_paths_.insert(rl_base_ + "/obs_seq"); obs_publisher_ = create_publisher( - rl_base_ + "/obs", rclcpp::QoS { rclcpp::KeepLast(1) }.best_effort()); + rl_base_ + "/obs", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort()); action_subscription_ = create_subscription( - rl_base_ + "/action", rclcpp::QoS { rclcpp::KeepLast(1) }.best_effort(), + rl_base_ + "/action", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort(), [this](rmcs_rl::msg::Action::UniquePtr message) { on_action_(std::move(message)); }); incoming_.action.assign(action_size_, 0.0); @@ -319,19 +343,18 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (auto& term : obs_terms_) { switch (term.kind) { case TermKind::kConstant: - case TermKind::kLastAction: - break; + case TermKind::kLastAction: break; case TermKind::kJointPos: case TermKind::kJointVel: case TermKind::kJointTorque: { for (std::size_t j = 0; j < term.joints.size(); ++j) { const auto& joint = joint_names_[term.joints[j]]; - const char field = term.kind == TermKind::kJointPos ? 'a' - : term.kind == TermKind::kJointVel ? 'v' - : 't'; - const auto path = joint_path_(joint, field); - term.joint_slots.push_back(acquire_slot_( - path, { Binding::kDouble }, true, output_map, "joint term")); + const char field = term.kind == TermKind::kJointPos ? 'a' + : term.kind == TermKind::kJointVel ? 'v' + : 't'; + const auto path = joint_path_(joint, field); + term.joint_slots.push_back( + acquire_slot_(path, {Binding::kDouble}, true, output_map, "joint term")); } break; } @@ -339,16 +362,13 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::vector candidates; switch (term.take) { case Take::kScalar: - candidates = { Binding::kDouble, Binding::kBool, Binding::kInt, - Binding::kSize }; + candidates = {Binding::kDouble, Binding::kBool, Binding::kInt, Binding::kSize}; break; case Take::kComponent: case Take::kVector: - candidates = { Binding::kVector3, Binding::kDirectionVector }; - break; - case Take::kGravity: - candidates = { Binding::kQuaternion }; + candidates = {Binding::kVector3, Binding::kDirectionVector}; break; + case Take::kGravity: candidates = {Binding::kQuaternion}; break; } term.slot = acquire_slot_( term.path, candidates, !term.has_default, output_map, "observation term"); @@ -358,31 +378,37 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { } if (!enable_path_.empty()) - enable_slot_ = acquire_slot_(enable_path_, { Binding::kBool, Binding::kDouble }, false, - output_map, "enable interface"); + enable_slot_ = acquire_slot_( + enable_path_, {Binding::kBool, Binding::kDouble}, false, output_map, + "enable interface"); if (!reset_path_.empty()) reset_slot_ = acquire_slot_( - reset_path_, { Binding::kSize, Binding::kInt, Binding::kDouble }, false, output_map, + reset_path_, {Binding::kSize, Binding::kInt, Binding::kDouble}, false, output_map, "reset interface"); for (const auto& slot : slots_) if (own_output_paths_.count(slot->path) != 0) - throw std::runtime_error("RlBridge: interface \"" + slot->path + throw std::runtime_error( + "RlBridge: interface \"" + slot->path + "\" is produced by RlBridge itself (self reference)"); - obs_signature_ = obs_layout_signature_(); + obs_signature_ = obs_layout_signature_(); actions_signature_ = actions_layout_signature_(); - layout_hash_ = rmcs::rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); + layout_hash_ = + rmcs_rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); log_layout_(); - RCLCPP_INFO(get_logger(), "contract: obs_size=%zu action_size=%zu policy_rate=%.3f Hz", - obs_size_, action_size_, policy_rate_); - RCLCPP_INFO(get_logger(), "layout_hash=%s (obs/action signatures below)", + RCLCPP_INFO( + get_logger(), "contract: obs_size=%zu action_size=%zu policy_rate=%.3f Hz", obs_size_, + action_size_, policy_rate_); + RCLCPP_INFO( + get_logger(), "layout_hash=%s (obs/action signatures below)", hex16(layout_hash_).c_str()); RCLCPP_INFO(get_logger(), " obs signature : %s", obs_signature_.c_str()); RCLCPP_INFO(get_logger(), " action signature : %s", actions_signature_.c_str()); - RCLCPP_INFO(get_logger(), "topics: %s/obs -> %s/action (best_effort, keep_last=1)", + RCLCPP_INFO( + get_logger(), "topics: %s/obs -> %s/action (best_effort, keep_last=1)", rl_base_.c_str(), rl_base_.c_str()); } @@ -394,7 +420,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { if (read_unsigned_(reset_slot_, reset_count) && reset_count != last_reset_count_) { last_reset_count_ = reset_count; reset_runtime_(); - RCLCPP_INFO(get_logger(), "RMCS reset: last_action cleared (reset_count=%llu)", + RCLCPP_INFO( + get_logger(), "RMCS reset: last_action cleared (reset_count=%llu)", static_cast(reset_count)); } } @@ -406,7 +433,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { ++pub_ok_count_; } else { ++obs_invalid_count_; - RCLCPP_WARN_THROTTLE(get_logger(), *get_clock(), 1000, + RCLCPP_WARN_THROTTLE( + get_logger(), *get_clock(), 1000, "observation sources not ready or non-finite; no obs published (%llu frames " "skipped)", static_cast(obs_invalid_count_)); @@ -414,20 +442,22 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { } ActionSnapshot& snapshot = read_snapshot_; - const bool has_snapshot = try_read_action_(snapshot); + const bool has_snapshot = try_read_action_(snapshot); double age = std::numeric_limits::quiet_NaN(); bool seq_ok = false; bool finite = true; if (has_snapshot) { - age = obs_age_of_(snapshot.obs_seq, now); - seq_ok = (snapshot.obs_seq == pub_seq_) || (snapshot.obs_seq + 1 == pub_seq_); - finite = std::all_of(snapshot.action.begin(), snapshot.action.end(), - [](double value) { return std::isfinite(value); }); + age = obs_age_of_(snapshot.obs_seq, now); + seq_ok = (snapshot.obs_seq == pub_seq_) || (snapshot.obs_seq + 1 == pub_seq_); + finite = std::all_of(snapshot.action.begin(), snapshot.action.end(), [](double value) { + return std::isfinite(value); + }); if (snapshot.layout_hash != layout_hash_) { if (contract_ok_) { contract_ok_ = false; - RCLCPP_FATAL(get_logger(), + RCLCPP_FATAL( + get_logger(), "contract mismatch: action layout_hash=%s != bridge layout_hash=%s " "(policy process and bridge disagree; refusing to output actions)", hex16(snapshot.layout_hash).c_str(), hex16(layout_hash_).c_str()); @@ -435,33 +465,36 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { } else if (expected_model_id_ != 0 && snapshot.model_id != expected_model_id_) { if (contract_ok_) { contract_ok_ = false; - RCLCPP_FATAL(get_logger(), - "model mismatch: action model_id=%s != expected_model_id=%s", + RCLCPP_FATAL( + get_logger(), "model mismatch: action model_id=%s != expected_model_id=%s", hex16(snapshot.model_id).c_str(), hex16(expected_model_id_).c_str()); } } } const bool enabled = read_enable_(); - const bool fresh = has_snapshot && is_finite(age) && age <= max_action_age_; - const bool valid = enabled && contract_ok_ && has_snapshot && fresh && seq_ok && finite; + const bool fresh = has_snapshot && is_finite(age) && age <= max_action_age_; + const bool valid = enabled && contract_ok_ && has_snapshot && fresh && seq_ok && finite; write_actions_(valid, snapshot); - if (valid) last_actions_ = snapshot.action; + if (valid) + last_actions_ = snapshot.action; - *valid_output_ = valid ? 1.0 : 0.0; - *healthy_output_ = (contract_ok_ && fresh) ? 1.0 : 0.0; + *valid_output_ = valid ? 1.0 : 0.0; + *healthy_output_ = (contract_ok_ && fresh) ? 1.0 : 0.0; *action_age_output_ = age; - *obs_seq_output_ = pub_seq_; + *obs_seq_output_ = pub_seq_; if (valid != last_valid_) { if (valid) { - RCLCPP_INFO(get_logger(), "valid=1 (obs_seq=%llu model_id=%s action_age=%.4f s)", + RCLCPP_INFO( + get_logger(), "valid=1 (obs_seq=%llu model_id=%s action_age=%.4f s)", static_cast(snapshot.obs_seq), hex16(snapshot.model_id).c_str(), age); } else { - RCLCPP_WARN(get_logger(), "valid=0: %s", + RCLCPP_WARN( + get_logger(), "valid=0: %s", invalid_reason_(enabled, contract_ok_, has_snapshot, fresh, seq_ok, finite) .c_str()); } @@ -475,25 +508,24 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::optional number_(const std::string& name) { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return std::nullopt; + if (!get_parameter(name, parameter)) + return std::nullopt; } catch (const std::exception&) { return std::nullopt; } switch (parameter.get_type()) { - case rclcpp::ParameterType::PARAMETER_DOUBLE: - return parameter.as_double(); + case rclcpp::ParameterType::PARAMETER_DOUBLE: return parameter.as_double(); case rclcpp::ParameterType::PARAMETER_INTEGER: return static_cast(parameter.as_int()); - case rclcpp::ParameterType::PARAMETER_NOT_SET: - return std::nullopt; - default: - throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a number"); + case rclcpp::ParameterType::PARAMETER_NOT_SET: return std::nullopt; + default: throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a number"); } } double number_or_(const std::string& name, double fallback) { const auto value = number_(name); - if (!value.has_value()) return fallback; + if (!value.has_value()) + return fallback; if (!is_finite(*value)) throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); return *value; @@ -502,20 +534,22 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { static std::uint64_t parse_u64_(const std::string& text, const std::string& name) { try { std::size_t consumed = 0; - const bool hex = text.rfind("0x", 0) == 0 || text.rfind("0X", 0) == 0; - const auto value = std::stoull(hex ? text.substr(2) : text, &consumed, hex ? 16 : 10); + const bool hex = text.rfind("0x", 0) == 0 || text.rfind("0X", 0) == 0; + const auto value = std::stoull(hex ? text.substr(2) : text, &consumed, hex ? 16 : 10); if (consumed != (hex ? text.size() - 2 : text.size())) throw std::invalid_argument("trailing characters"); return static_cast(value); } catch (const std::exception&) { - throw std::invalid_argument("RlBridge: parameter '" + name + throw std::invalid_argument( + "RlBridge: parameter '" + name + "' must be a decimal or 0x-prefixed 64-bit id, got '" + text + "'"); } } std::optional integer_or_(const std::string& name) { const auto value = number_(name); - if (!value.has_value()) return std::nullopt; + if (!value.has_value()) + return std::nullopt; const auto rounded = std::llround(*value); if (std::abs(*value - static_cast(rounded)) > 1e-9) throw std::invalid_argument("RlBridge: parameter '" + name + "' must be an integer"); @@ -525,11 +559,13 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { bool bool_or_(const std::string& name, bool fallback) { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return fallback; + if (!get_parameter(name, parameter)) + return fallback; } catch (const std::exception&) { return fallback; } - if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) return fallback; + if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) + return fallback; if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_BOOL) throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a boolean"); return parameter.as_bool(); @@ -538,11 +574,13 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::string string_or_(const std::string& name, const std::string& fallback) { rclcpp::Parameter parameter; try { - if (!get_parameter(name, parameter)) return fallback; + if (!get_parameter(name, parameter)) + return fallback; } catch (const std::exception&) { return fallback; } - if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) return fallback; + if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) + return fallback; if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_STRING) throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a string"); return parameter.as_string(); @@ -551,7 +589,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::vector string_array_or_(const std::string& name) { std::vector value; try { - if (!get_parameter(name, value)) return { }; + if (!get_parameter(name, value)) + return {}; } catch (const std::exception&) { throw std::invalid_argument( "RlBridge: parameter '" + name + "' must be a list of strings"); @@ -564,7 +603,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { joint_base_path_ = string_or_("joint_base_path", ""); if (!joint_names_.empty() && joint_base_path_.empty()) throw std::invalid_argument( - "RlBridge: joint_base_path is required when joint_names is set"); + "RlBridge: joint_base_path is required when joint_names is " + "set"); for (std::size_t i = 0; i < joint_names_.size(); ++i) { if (joint_names_[i].empty()) throw std::invalid_argument("RlBridge: joint_names contains an empty name"); @@ -573,22 +613,24 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { "RlBridge: duplicate joint name '" + joint_names_[i] + "'"); } - const std::string default_angle = string_or_("joint_angle_suffix", "/angle"); + const std::string default_angle = string_or_("joint_angle_suffix", "/angle"); const std::string default_velocity = string_or_("joint_velocity_suffix", "/velocity"); - const std::string default_torque = string_or_("joint_torque_suffix", "/torque"); + const std::string default_torque = string_or_("joint_torque_suffix", "/torque"); for (const auto& joint : joint_names_) { - angle_suffix_[joint] = default_angle; + angle_suffix_[joint] = default_angle; velocity_suffix_[joint] = default_velocity; - torque_suffix_[joint] = default_torque; + torque_suffix_[joint] = default_torque; } constexpr const char* prefix = "joint_interface_overrides."; - for (const auto& name : list_parameters({ "joint_interface_overrides" }, 10).names) { - if (name.rfind(prefix, 0) != 0) continue; - const std::string remainder = name.substr(std::string { prefix }.size()); - const auto separator = remainder.find_last_of('.'); + for (const auto& name : list_parameters({"joint_interface_overrides"}, 10).names) { + if (name.rfind(prefix, 0) != 0) + continue; + const std::string remainder = name.substr(std::string{prefix}.size()); + const auto separator = remainder.find_last_of('.'); if (separator == std::string::npos) - throw std::invalid_argument("RlBridge: joint_interface_overrides entry '" + name + throw std::invalid_argument( + "RlBridge: joint_interface_overrides entry '" + name + "' must be ."); const std::string joint = remainder.substr(0, separator); const std::string field = remainder.substr(separator + 1); @@ -596,26 +638,31 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { throw std::invalid_argument( "RlBridge: joint_interface_overrides references unknown joint '" + joint + "'"); const std::string value = string_or_(name, ""); - if (field == "angle") angle_suffix_[joint] = value; - else if (field == "velocity") velocity_suffix_[joint] = value; - else if (field == "torque") torque_suffix_[joint] = value; + if (field == "angle") + angle_suffix_[joint] = value; + else if (field == "velocity") + velocity_suffix_[joint] = value; + else if (field == "torque") + torque_suffix_[joint] = value; else - throw std::invalid_argument("RlBridge: joint_interface_overrides field '" + field + throw std::invalid_argument( + "RlBridge: joint_interface_overrides field '" + field + "' is not angle/velocity/torque"); } } std::string joint_path_(const std::string& joint, char field) const { const auto suffix = field == 'a' ? angle_suffix_.at(joint) - : field == 'v' ? velocity_suffix_.at(joint) + : field == 'v' ? velocity_suffix_.at(joint) : torque_suffix_.at(joint); return joint_base_path_ + "/" + joint + suffix; } std::optional default_joint_pos_(std::size_t joint) { - const auto name = "default_joint_pos." + joint_names_[joint]; + const auto name = "default_joint_pos." + joint_names_[joint]; const auto value = number_(name); - if (!value.has_value()) return std::nullopt; + if (!value.has_value()) + return std::nullopt; if (!is_finite(*value)) throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); return value; @@ -623,55 +670,46 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { static bool binding_matches_type_(Binding binding, const std::type_info& type) { switch (binding) { - case Binding::kDouble: - return type == typeid(double); - case Binding::kBool: - return type == typeid(bool); - case Binding::kInt: - return type == typeid(int); - case Binding::kSize: - return type == typeid(std::size_t); - case Binding::kVector3: - return type == typeid(Eigen::Vector3d); + case Binding::kDouble: return type == typeid(double); + case Binding::kBool: return type == typeid(bool); + case Binding::kInt: return type == typeid(int); + case Binding::kSize: return type == typeid(std::size_t); + case Binding::kVector3: return type == typeid(Eigen::Vector3d); case Binding::kDirectionVector: return type == typeid(rmcs_description::BaseLink::DirectionVector); - case Binding::kQuaternion: - return type == typeid(Eigen::Quaterniond); + case Binding::kQuaternion: return type == typeid(Eigen::Quaterniond); } return false; } static const char* binding_name_(Binding binding) { switch (binding) { - case Binding::kDouble: - return "double"; - case Binding::kBool: - return "bool"; - case Binding::kInt: - return "int"; - case Binding::kSize: - return "std::size_t"; - case Binding::kVector3: - return "Eigen::Vector3d"; - case Binding::kDirectionVector: - return "BaseLink::DirectionVector"; - case Binding::kQuaternion: - return "Eigen::Quaterniond"; + case Binding::kDouble: return "double"; + case Binding::kBool: return "bool"; + case Binding::kInt: return "int"; + case Binding::kSize: return "std::size_t"; + case Binding::kVector3: return "Eigen::Vector3d"; + case Binding::kDirectionVector: return "BaseLink::DirectionVector"; + case Binding::kQuaternion: return "Eigen::Quaterniond"; } return "unknown"; } - std::size_t acquire_slot_(const std::string& path, const std::vector& candidates, - bool required, const OutputInfoMap& output_map, const char* context) { + std::size_t acquire_slot_( + const std::string& path, const std::vector& candidates, bool required, + const OutputInfoMap& output_map, const char* context) { const auto output = output_map.find(path); if (output == output_map.end()) { - if (!required) return kNoSlot; - throw std::runtime_error("RlBridge: required input interface \"" + path + if (!required) + return kNoSlot; + throw std::runtime_error( + "RlBridge: required input interface \"" + path + "\" was not produced by any component (" + context + "); check the interface path or add default= to make it optional"); } if (output->second.kind != rmcs_executor::InterfaceKind::Normal) - throw std::runtime_error("RlBridge: input interface \"" + path + throw std::runtime_error( + "RlBridge: input interface \"" + path + "\" exists but is an Event interface; Normal required (" + context + ")"); const std::type_info& producer_type = output->second.type.get(); @@ -680,26 +718,26 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { return i; Binding selected = Binding::kDouble; - bool found = false; + bool found = false; for (const auto binding : candidates) if (binding_matches_type_(binding, producer_type)) { selected = binding; - found = true; + found = true; break; } if (!found) { std::vector expected; for (const auto binding : candidates) expected.push_back(binding_name_(binding)); - throw std::runtime_error("RlBridge: cannot bind observation interface \"" + path - + "\": producer declares \"" + pretty_type(producer_type) - + "\" but the term accepts { " + join(expected, ", ") + throw std::runtime_error( + "RlBridge: cannot bind observation interface \"" + path + "\": producer declares \"" + + pretty_type(producer_type) + "\" but the term accepts { " + join(expected, ", ") + " }. Either fix the term (take=/transform=) or pin type= explicitly."); } - auto slot = std::make_unique(); - slot->path = path; - slot->binding = selected; + auto slot = std::make_unique(); + slot->path = path; + slot->binding = selected; slot->required = required; switch (selected) { case Binding::kDouble: @@ -741,9 +779,9 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (const auto& token : split_whitespace(spec)) { const auto equals = token.find('='); if (equals == std::string::npos || equals == 0) - throw std::invalid_argument("RlBridge: term token '" + token - + "' must be key=value (term '" + spec + "')"); - const std::string key = token.substr(0, equals); + throw std::invalid_argument( + "RlBridge: term token '" + token + "' must be key=value (term '" + spec + "')"); + const std::string key = token.substr(0, equals); const std::string value = token.substr(equals + 1); if (!tokens.emplace(key, value).second) throw std::invalid_argument( @@ -752,8 +790,9 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { return tokens; } - static void validate_tokens_(const std::map& tokens, - const std::vector& allowed, const std::string& spec) { + static void validate_tokens_( + const std::map& tokens, const std::vector& allowed, + const std::string& spec) { for (const auto& [key, ignored] : tokens) { (void)ignored; if (std::find(allowed.begin(), allowed.end(), key) == allowed.end()) @@ -762,10 +801,12 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { } } - static double double_token_(const std::map& tokens, - const std::string& key, double fallback, const std::string& spec) { + static double double_token_( + const std::map& tokens, const std::string& key, double fallback, + const std::string& spec) { const auto token = tokens.find(key); - if (token == tokens.end()) return fallback; + if (token == tokens.end()) + return fallback; const double value = parse_double(token->second, key + " in term '" + spec + "'"); if (!is_finite(value)) throw std::invalid_argument( @@ -773,24 +814,30 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { return value; } - static bool bool_token_(const std::map& tokens, - const std::string& key, bool fallback, const std::string& spec) { + static bool bool_token_( + const std::map& tokens, const std::string& key, bool fallback, + const std::string& spec) { const auto token = tokens.find(key); - if (token == tokens.end()) return fallback; + if (token == tokens.end()) + return fallback; return parse_boolean(token->second, key + " in term '" + spec + "'"); } - static std::string string_token_(const std::map& tokens, - const std::string& key, const std::string& fallback) { + static std::string string_token_( + const std::map& tokens, const std::string& key, + const std::string& fallback) { const auto token = tokens.find(key); - if (token == tokens.end()) return fallback; + if (token == tokens.end()) + return fallback; return token->second; } - static void parse_clip_(const std::map& tokens, - const std::string& spec, bool& has_clip, double& clip_min, double& clip_max) { + static void parse_clip_( + const std::map& tokens, const std::string& spec, bool& has_clip, + double& clip_min, double& clip_max) { const auto clip = tokens.find("clip"); - if (clip == tokens.end()) return; + if (clip == tokens.end()) + return; const auto separator = clip->second.find(':'); if (separator == std::string::npos) { const double symmetric = parse_double(clip->second, "clip"); @@ -808,38 +855,40 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { has_clip = true; } - static void parse_index_(const std::map& tokens, - const std::string& spec, bool& has_index, std::size_t& index) { + static void parse_index_( + const std::map& tokens, const std::string& spec, bool& has_index, + std::size_t& index) { const auto token = tokens.find("index"); - if (token == tokens.end()) return; + if (token == tokens.end()) + return; const auto value = parse_integer(token->second, "index in term '" + spec + "'"); if (value < 0) throw std::invalid_argument("RlBridge: index must be >= 0 (term '" + spec + "')"); has_index = true; - index = static_cast(value); + index = static_cast(value); } - static std::vector parse_index_list_( - const std::string& text, std::size_t limit, const std::string& what) { + static std::vector + parse_index_list_(const std::string& text, std::size_t limit, const std::string& what) { std::vector indices; for (const auto& piece : split_by(text, ',')) { if (piece.empty()) - throw std::invalid_argument( - "RlBridge: empty entry in " + what + " '" + text + "'"); + throw std::invalid_argument("RlBridge: empty entry in " + what + " '" + text + "'"); const auto range = piece.find(".."); if (range == std::string::npos) { const auto index = parse_integer(piece, what); if (index < 0 || static_cast(index) >= limit) - throw std::invalid_argument("RlBridge: " + what + " index " - + std::to_string(index) + " out of range [0, " + std::to_string(limit) - + ")"); + throw std::invalid_argument( + "RlBridge: " + what + " index " + std::to_string(index) + + " out of range [0, " + std::to_string(limit) + ")"); indices.push_back(static_cast(index)); } else { const auto first = parse_integer(piece.substr(0, range), what); - const auto last = parse_integer(piece.substr(range + 2), what); + const auto last = parse_integer(piece.substr(range + 2), what); if (first < 0 || last < first || static_cast(last) >= limit) - throw std::invalid_argument("RlBridge: invalid " + what + " range '" + piece - + "' (limit " + std::to_string(limit) + ")"); + throw std::invalid_argument( + "RlBridge: invalid " + what + " range '" + piece + "' (limit " + + std::to_string(limit) + ")"); for (auto index = first; index <= last; ++index) indices.push_back(static_cast(index)); } @@ -857,7 +906,7 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { if (type == "joint_pos" || type == "joint_vel" || type == "joint_torque") { term.kind = type == "joint_pos" ? TermKind::kJointPos - : type == "joint_vel" ? TermKind::kJointVel + : type == "joint_vel" ? TermKind::kJointVel : TermKind::kJointTorque; const auto joints = tokens.find("joints"); if (joints == tokens.end() || joints->second.empty()) @@ -875,7 +924,7 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { term.joint_names.push_back(name); } term.relative = bool_token_(tokens, "relative", false, spec); - term.zero = bool_token_(tokens, "zero", false, spec); + term.zero = bool_token_(tokens, "zero", false, spec); if (term.kind != TermKind::kJointPos && (term.relative || term.zero)) throw std::invalid_argument( "RlBridge: relative/zero are only valid for joint_pos (term '" + spec + "')"); @@ -886,7 +935,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { if (term.kind == TermKind::kJointPos && term.relative) { const auto base = default_joint_pos_(joint); if (!base.has_value()) - throw std::invalid_argument("RlBridge: joint_pos relative term '" + spec + throw std::invalid_argument( + "RlBridge: joint_pos relative term '" + spec + "' requires default_joint_pos." + joint_names_[joint]); term.joint_defaults.push_back(*base); } else { @@ -896,17 +946,19 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { term.dim = term.joints.size(); const std::string flag = term.zero ? "zero" : (term.relative ? "rel" : "abs"); const std::string prefix = type == "joint_pos" ? "joint_pos" - : type == "joint_vel" ? "joint_vel" + : type == "joint_vel" ? "joint_vel" : "joint_torque"; term.id = prefix + ":" + flag + ":" + join(term.joint_names, "+"); - validate_tokens_(tokens, - { "type", "joints", "relative", "zero", "index", "scale", "clip", "name" }, spec); + validate_tokens_( + tokens, {"type", "joints", "relative", "zero", "index", "scale", "clip", "name"}, + spec); return finish_common_(term, tokens, spec); } if (type == "last_action") { term.kind = TermKind::kLastAction; if (const auto indices = tokens.find("indices"); indices != tokens.end()) - term.action_indices = parse_index_list_(indices->second, action_size_, "last_action"); + term.action_indices = + parse_index_list_(indices->second, action_size_, "last_action"); else for (std::size_t i = 0; i < action_size_; ++i) term.action_indices.push_back(i); @@ -915,11 +967,11 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (const auto index : term.action_indices) names.push_back(std::to_string(index)); term.id = "last_action:" + join(names, "+"); - validate_tokens_(tokens, { "type", "indices", "index", "scale", "clip", "name" }, spec); + validate_tokens_(tokens, {"type", "indices", "index", "scale", "clip", "name"}, spec); return finish_common_(term, tokens, spec); } if (type == "constant") { - term.kind = TermKind::kConstant; + term.kind = TermKind::kConstant; const auto value = tokens.find("value"); if (value == tokens.end() || value->second.empty()) throw std::invalid_argument( @@ -931,62 +983,67 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (const double item : term.constants) names.push_back(format_number(item)); term.id = "constant:" + join(names, "+"); - validate_tokens_(tokens, { "type", "value", "index", "scale", "clip", "name" }, spec); + validate_tokens_(tokens, {"type", "value", "index", "scale", "clip", "name"}, spec); return finish_common_(term, tokens, spec); } const auto path = tokens.find("path"); if (path == tokens.end() || path->second.empty()) - throw std::invalid_argument("RlBridge: term requires path=... (term '" + spec + throw std::invalid_argument( + "RlBridge: term requires path=... (term '" + spec + "') unless it is joint_*/last_action/constant"); - term.kind = TermKind::kPath; - term.path = path->second; + term.kind = TermKind::kPath; + term.path = path->second; const std::string take = string_token_(tokens, "take", ""); const std::string transform = string_token_(tokens, "transform", ""); if (transform == "projected_gravity") { if (!take.empty() && take != "gravity") throw std::invalid_argument( - "RlBridge: transform=projected_gravity conflicts with take=" + take - + " (term '" + spec + "')"); + "RlBridge: transform=projected_gravity conflicts with " + "take=" + + take + " (term '" + spec + "')"); term.take = Take::kGravity; - term.id = "gravity:" + term.path; - term.dim = 3; + term.id = "gravity:" + term.path; + term.dim = 3; if (!type.empty()) { if (type != "quaternion") - throw std::invalid_argument("RlBridge: projected_gravity requires a " - "quaternion interface; type=" + throw std::invalid_argument( + "RlBridge: projected_gravity requires a " + "quaternion interface; type=" + type + " is not quaternion (term '" + spec + "')"); term.has_binding = true; - term.binding = Binding::kQuaternion; + term.binding = Binding::kQuaternion; } } else if (transform == "gravity") { throw std::invalid_argument( - "RlBridge: transform=gravity is not supported; use transform=projected_gravity " + "RlBridge: transform=gravity is not supported; use " + "transform=projected_gravity " "(term '" + spec + "')"); } else if (!transform.empty()) { - throw std::invalid_argument("RlBridge: unknown transform='" + transform + "' (term '" - + spec + "')"); + throw std::invalid_argument( + "RlBridge: unknown transform='" + transform + "' (term '" + spec + "')"); } else if (take.empty()) { term.take = Take::kScalar; - term.id = term.path; - term.dim = 1; + term.id = term.path; + term.dim = 1; } else if (take == "x" || take == "y" || take == "z") { - term.take = Take::kComponent; - term.component = std::string { "xyz" }.find(take); - term.id = "vec3c:" + term.path + ":" + take; - term.dim = 1; + term.take = Take::kComponent; + term.component = std::string{"xyz"}.find(take); + term.id = "vec3c:" + term.path + ":" + take; + term.dim = 1; } else if (take == "vec3" || take == "vec" || take == "all" || take == "vector") { term.take = Take::kVector; - term.id = "vec3:" + term.path; - term.dim = 3; + term.id = "vec3:" + term.path; + term.dim = 3; } else if (take == "gravity") { term.take = Take::kGravity; - term.id = "gravity:" + term.path; - term.dim = 3; + term.id = "gravity:" + term.path; + term.dim = 3; } else { - throw std::invalid_argument("RlBridge: unknown take='" + take + throw std::invalid_argument( + "RlBridge: unknown take='" + take + "' (use x|y|z|vec3|gravity, or omit for scalar) (term '" + spec + "')"); } @@ -1007,59 +1064,67 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { } else if (type == "quaternion") { term.binding = Binding::kQuaternion; } else { - throw std::invalid_argument("RlBridge: unknown type='" + type + "' (term '" + spec - + "')"); + throw std::invalid_argument( + "RlBridge: unknown type='" + type + "' (term '" + spec + "')"); } - const bool scalar_binding = term.binding == Binding::kDouble - || term.binding == Binding::kBool || term.binding == Binding::kInt - || term.binding == Binding::kSize; - const bool vector_binding = term.binding == Binding::kVector3 - || term.binding == Binding::kDirectionVector; + const bool scalar_binding = + term.binding == Binding::kDouble || term.binding == Binding::kBool + || term.binding == Binding::kInt || term.binding == Binding::kSize; + const bool vector_binding = + term.binding == Binding::kVector3 || term.binding == Binding::kDirectionVector; if (term.take == Take::kScalar && !scalar_binding) - throw std::invalid_argument("RlBridge: type=" + type + throw std::invalid_argument( + "RlBridge: type=" + type + " needs take=vec3 or take=x|y|z (a whole vector is not a scalar) (term '" + spec + "')"); if ((term.take == Take::kComponent || term.take == Take::kVector) && !vector_binding) - throw std::invalid_argument("RlBridge: take=" + take + " needs a vector interface " - "(type=vector3|direction_vector), got type=" + throw std::invalid_argument( + "RlBridge: take=" + take + + " needs a vector interface " + "(type=vector3|direction_vector), got type=" + type + " (term '" + spec + "')"); if (term.take == Take::kGravity && term.binding != Binding::kQuaternion) throw std::invalid_argument( - "RlBridge: transform=projected_gravity needs type=quaternion (term '" + spec - + "')"); + "RlBridge: transform=projected_gravity needs " + "type=quaternion (term '" + + spec + "')"); } - validate_tokens_(tokens, - { "path", "take", "transform", "type", "index", "scale", "clip", "default", "name" }, + validate_tokens_( + tokens, + {"path", "take", "transform", "type", "index", "scale", "clip", "default", "name"}, spec); return finish_common_(term, tokens, spec); } - ObsTerm finish_common_(ObsTerm term, const std::map& tokens, - const std::string& spec) { + ObsTerm finish_common_( + ObsTerm term, const std::map& tokens, const std::string& spec) { parse_index_(tokens, spec, term.has_index, term.index); term.scale = double_token_(tokens, "scale", 1.0, spec); if (term.scale == 0.0) - throw std::invalid_argument( - "RlBridge: scale=0 is not allowed (term '" + spec + "')"); + throw std::invalid_argument("RlBridge: scale=0 is not allowed (term '" + spec + "')"); parse_clip_(tokens, spec, term.has_clip, term.clip_min, term.clip_max); const auto default_value = tokens.find("default"); if (default_value != tokens.end()) { if (!(term.kind == TermKind::kPath && term.take == Take::kScalar)) - throw std::invalid_argument("RlBridge: default= is only valid for scalar path terms " - "(term '" + throw std::invalid_argument( + "RlBridge: default= is only valid for scalar path " + "terms " + "(term '" + spec + "')"); term.default_value = parse_double(default_value->second, "default"); if (!is_finite(term.default_value)) - throw std::invalid_argument("RlBridge: default must be finite (term '" + spec + "')"); + throw std::invalid_argument( + "RlBridge: default must be finite (term '" + spec + "')"); term.has_default = true; } const auto name = tokens.find("name"); if (name != tokens.end()) { if (name->second.empty()) - throw std::invalid_argument("RlBridge: name= must not be empty (term '" + spec + "')"); + throw std::invalid_argument( + "RlBridge: name= must not be empty (term '" + spec + "')"); term.id = name->second; } return term; @@ -1068,13 +1133,14 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { ActionTerm parse_action_term_(const std::string& spec) { ActionTerm term; const auto tokens = parse_tokens_(spec); - const auto index = tokens.find("index"); + const auto index = tokens.find("index"); if (index == tokens.end()) throw std::invalid_argument( "RlBridge: action term requires index= (term '" + spec + "')"); const auto index_value = parse_integer(index->second, "action index"); if (index_value < 0) - throw std::invalid_argument("RlBridge: action index must be >= 0 (term '" + spec + "')"); + throw std::invalid_argument( + "RlBridge: action index must be >= 0 (term '" + spec + "')"); term.index = static_cast(index_value); const auto output = tokens.find("output"); @@ -1082,16 +1148,16 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { throw std::invalid_argument( "RlBridge: action term requires output= (term '" + spec + "')"); term.output = output->second; - term.id = string_token_(tokens, "name", term.output); + term.id = string_token_(tokens, "name", term.output); if (term.id.empty()) - throw std::invalid_argument("RlBridge: action name= must not be empty (term '" + spec - + "')"); + throw std::invalid_argument( + "RlBridge: action name= must not be empty (term '" + spec + "')"); term.scale = double_token_(tokens, "scale", 1.0, spec); if (term.scale == 0.0) - throw std::invalid_argument("RlBridge: action scale=0 is not allowed (term '" + spec - + "')"); + throw std::invalid_argument( + "RlBridge: action scale=0 is not allowed (term '" + spec + "')"); parse_clip_(tokens, spec, term.has_clip, term.clip_min, term.clip_max); - validate_tokens_(tokens, { "index", "output", "name", "scale", "clip" }, spec); + validate_tokens_(tokens, {"index", "output", "name", "scale", "clip"}, spec); return term; } @@ -1099,7 +1165,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::string signature = "v2"; for (const auto& term : obs_terms_) { signature += "|" + term.id; - if (term.scale != 1.0) signature += "*" + format_number(term.scale); + if (term.scale != 1.0) + signature += "*" + format_number(term.scale); signature += "@" + std::to_string(term.dim); } return signature; @@ -1109,7 +1176,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::string signature = "v2"; for (const auto& term : action_terms_) { signature += "|#" + std::to_string(term.index) + ":" + term.id; - if (term.scale != 1.0) signature += "*" + format_number(term.scale); + if (term.scale != 1.0) + signature += "*" + format_number(term.scale); } return signature; } @@ -1120,15 +1188,17 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::ostringstream line; line << " [" << term.index << ":" << term.index + term.dim << ") " << term.id << " scale=" << format_number(term.scale); - if (term.has_default) line << " default=" << format_number(term.default_value); + if (term.has_default) + line << " default=" << format_number(term.default_value); if (term.has_clip) line << " clip=" << format_number(term.clip_min) << ":" << format_number(term.clip_max); if (!term.path.empty()) { line << " <- " << term.path; - if (term.slot != kNoSlot) line << " (" << binding_name_(slots_[term.slot]->binding) - << ")"; - else line << " (default)"; + if (term.slot != kNoSlot) + line << " (" << binding_name_(slots_[term.slot]->binding) << ")"; + else + line << " (default)"; } RCLCPP_INFO(get_logger(), "%s", line.str().c_str()); } @@ -1136,7 +1206,8 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (const auto& term : action_terms_) { std::ostringstream line; line << " #" << term.index << " " << term.id << " -> " << term.output; - if (term.scale != 1.0) line << " scale=" << format_number(term.scale); + if (term.scale != 1.0) + line << " scale=" << format_number(term.scale); RCLCPP_INFO(get_logger(), "%s", line.str().c_str()); } } @@ -1145,48 +1216,58 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { const auto& entry = *slots_[slot]; switch (entry.binding) { case Binding::kDouble: - if (!entry.double_value->ready()) return false; + if (!entry.double_value->ready()) + return false; value = **entry.double_value; return true; case Binding::kBool: - if (!entry.bool_value->ready()) return false; + if (!entry.bool_value->ready()) + return false; value = **entry.bool_value ? 1.0 : 0.0; return true; case Binding::kInt: - if (!entry.int_value->ready()) return false; + if (!entry.int_value->ready()) + return false; value = static_cast(**entry.int_value); return true; case Binding::kSize: - if (!entry.size_value->ready()) return false; + if (!entry.size_value->ready()) + return false; value = static_cast(**entry.size_value); return true; - default: - return false; + default: return false; } } bool read_unsigned_(std::size_t slot, std::uint64_t& value) const { double raw = 0.0; - if (!read_double_(slot, raw)) return false; - if (!is_finite(raw) || raw < 0.0) return false; + if (!read_double_(slot, raw)) + return false; + if (!is_finite(raw) || raw < 0.0) + return false; value = static_cast(raw); return true; } bool read_enable_() const { - if (enable_slot_ == kNoSlot) return enable_default_; + if (enable_slot_ == kNoSlot) + return enable_default_; double raw = 0.0; - if (!read_double_(enable_slot_, raw)) return enable_default_; + if (!read_double_(enable_slot_, raw)) + return enable_default_; return raw != 0.0; } bool build_observation_(std::vector& obs) const { obs.assign(obs_size_, 0.0); - const auto push = [&obs](std::size_t index, double value, double scale, bool has_clip, + const auto push = [&obs]( + std::size_t index, double value, double scale, bool has_clip, double clip_min, double clip_max, bool& ok) { double result = value * scale; - if (has_clip) result = std::clamp(result, clip_min, clip_max); - if (!is_finite(result)) ok = false; + if (has_clip) + result = std::clamp(result, clip_min, clip_max); + if (!is_finite(result)) + ok = false; obs[index] = result; }; @@ -1195,47 +1276,59 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { switch (term.kind) { case TermKind::kPath: { if (term.slot == kNoSlot) { - push(term.index, term.default_value, term.scale, term.has_clip, term.clip_min, + push( + term.index, term.default_value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); break; } const auto& entry = *slots_[term.slot]; if (term.take == Take::kScalar) { double value = 0.0; - if (!read_double_(term.slot, value)) return false; - push(term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, + if (!read_double_(term.slot, value)) + return false; + push( + term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } else if (term.take == Take::kComponent) { double value = 0.0; if (entry.binding == Binding::kVector3) { - if (!entry.vector3_value->ready()) return false; + if (!entry.vector3_value->ready()) + return false; value = (**entry.vector3_value)[term.component]; } else { - if (!entry.direction_vector_value->ready()) return false; + if (!entry.direction_vector_value->ready()) + return false; value = (**entry.direction_vector_value).vector[term.component]; } - push(term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, + push( + term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } else if (term.take == Take::kVector) { if (entry.binding == Binding::kVector3) { - if (!entry.vector3_value->ready()) return false; + if (!entry.vector3_value->ready()) + return false; const auto& value = **entry.vector3_value; for (std::size_t i = 0; i < 3; ++i) - push(term.index + i, value[static_cast(i)], term.scale, + push( + term.index + i, value[static_cast(i)], term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } else { - if (!entry.direction_vector_value->ready()) return false; + if (!entry.direction_vector_value->ready()) + return false; const auto& value = (**entry.direction_vector_value).vector; for (std::size_t i = 0; i < 3; ++i) - push(term.index + i, value[static_cast(i)], term.scale, + push( + term.index + i, value[static_cast(i)], term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } } else { - if (!entry.quaternion_value->ready()) return false; + if (!entry.quaternion_value->ready()) + return false; const Eigen::Vector3d gravity = **entry.quaternion_value * Eigen::Vector3d(0.0, 0.0, -1.0); for (std::size_t i = 0; i < 3; ++i) - push(term.index + i, gravity[static_cast(i)], term.scale, + push( + term.index + i, gravity[static_cast(i)], term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } break; @@ -1244,10 +1337,13 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { for (std::size_t j = 0; j < term.joints.size(); ++j) { double value = 0.0; if (!term.zero) { - if (!read_double_(term.joint_slots[j], value)) return false; - if (term.relative) value -= term.joint_defaults[j]; + if (!read_double_(term.joint_slots[j], value)) + return false; + if (term.relative) + value -= term.joint_defaults[j]; } - push(term.index + j, value, term.scale, term.has_clip, term.clip_min, + push( + term.index + j, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } break; @@ -1256,52 +1352,57 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { case TermKind::kJointTorque: { for (std::size_t j = 0; j < term.joints.size(); ++j) { double value = 0.0; - if (!read_double_(term.joint_slots[j], value)) return false; - push(term.index + j, value, term.scale, term.has_clip, term.clip_min, + if (!read_double_(term.joint_slots[j], value)) + return false; + push( + term.index + j, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); } break; } case TermKind::kLastAction: { for (std::size_t j = 0; j < term.action_indices.size(); ++j) - push(term.index + j, last_actions_[term.action_indices[j]], term.scale, + push( + term.index + j, last_actions_[term.action_indices[j]], term.scale, term.has_clip, term.clip_min, term.clip_max, ok); break; } case TermKind::kConstant: { for (std::size_t j = 0; j < term.constants.size(); ++j) - push(term.index + j, term.constants[j], term.scale, term.has_clip, term.clip_min, + push( + term.index + j, term.constants[j], term.scale, term.has_clip, term.clip_min, term.clip_max, ok); break; } } - if (!ok) return false; + if (!ok) + return false; } return true; } - void publish_observation_(const std::vector& obs, - std::chrono::steady_clock::time_point now) { + void publish_observation_( + const std::vector& obs, std::chrono::steady_clock::time_point now) { rmcs_rl::msg::Observation message; - message.header.stamp = get_clock()->now(); + message.header.stamp = get_clock()->now(); message.header.frame_id = ""; - message.obs_seq = ++pub_seq_; - message.layout_hash = layout_hash_; - message.obs = obs; + message.obs_seq = ++pub_seq_; + message.layout_hash = layout_hash_; + message.obs = obs; obs_publisher_->publish(message); - prev_pub_seq_ = pub_seq_ - 1; + prev_pub_seq_ = pub_seq_ - 1; prev_pub_time_ = last_pub_time_; last_pub_time_ = now; - pub_started_ = true; + pub_started_ = true; } void on_action_(rmcs_rl::msg::Action::UniquePtr message) { const auto received = std::chrono::steady_clock::now(); if (message->action.size() != action_size_) { - RCLCPP_ERROR_THROTTLE(get_logger(), *get_clock(), 1000, - "ignoring action with %zu values (expected %zu)", message->action.size(), - action_size_); + RCLCPP_ERROR_THROTTLE( + get_logger(), *get_clock(), 1000, "ignoring action with %zu values (expected %zu)", + message->action.size(), action_size_); return; } @@ -1309,10 +1410,10 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { action_sequence_.store(sequence + 1, std::memory_order_relaxed); std::atomic_thread_fence(std::memory_order_release); - incoming_.obs_seq = message->obs_seq; + incoming_.obs_seq = message->obs_seq; incoming_.layout_hash = message->layout_hash; - incoming_.model_id = message->model_id; - incoming_.received = received; + incoming_.model_id = message->model_id; + incoming_.received = received; std::copy(message->action.begin(), message->action.end(), incoming_.action.begin()); std::atomic_thread_fence(std::memory_order_release); @@ -1322,17 +1423,20 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { bool try_read_action_(ActionSnapshot& snapshot) { for (int attempt = 0; attempt < 4; ++attempt) { const std::uint64_t before = action_sequence_.load(std::memory_order_relaxed); - if (before == 0 || (before & 1U) != 0) return false; + if (before == 0 || (before & 1U) != 0) + return false; std::atomic_thread_fence(std::memory_order_acquire); snapshot = incoming_; std::atomic_thread_fence(std::memory_order_acquire); - if (action_sequence_.load(std::memory_order_relaxed) == before) return true; + if (action_sequence_.load(std::memory_order_relaxed) == before) + return true; } return false; } double obs_age_of_(std::uint64_t obs_seq, std::chrono::steady_clock::time_point now) const { - if (!pub_started_) return std::numeric_limits::quiet_NaN(); + if (!pub_started_) + return std::numeric_limits::quiet_NaN(); if (obs_seq == pub_seq_) return std::chrono::duration(now - last_pub_time_).count(); if (obs_seq + 1 == pub_seq_) @@ -1343,22 +1447,17 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { void write_actions_(bool valid, const ActionSnapshot& snapshot) { for (std::size_t i = 0; i < action_terms_.size(); ++i) { const auto& term = action_terms_[i]; - double value = 0.0; + double value = 0.0; if (valid) { value = snapshot.action[i] * term.scale; - if (term.has_clip) value = std::clamp(value, term.clip_min, term.clip_max); + if (term.has_clip) + value = std::clamp(value, term.clip_min, term.clip_max); written_[i] = value; } else { switch (invalid_mode_) { - case InvalidMode::kNaN: - value = std::numeric_limits::quiet_NaN(); - break; - case InvalidMode::kZero: - value = 0.0; - break; - case InvalidMode::kHold: - value = written_[i]; - break; + case InvalidMode::kNaN: value = std::numeric_limits::quiet_NaN(); break; + case InvalidMode::kZero: value = 0.0; break; + case InvalidMode::kHold: value = written_[i]; break; } } **action_outputs_[i] = value; @@ -1368,18 +1467,25 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { void reset_runtime_() { std::fill(last_actions_.begin(), last_actions_.end(), 0.0); std::fill(written_.begin(), written_.end(), 0.0); - pub_started_ = false; + pub_started_ = false; prev_pub_seq_ = 0; } - std::string invalid_reason_(bool enabled, bool contract_ok, bool has_snapshot, bool fresh, - bool seq_ok, bool finite) const { - if (!contract_ok) return "contract/model mismatch (latched; restart or re-enable)"; - if (!enabled) return "disabled by enable interface"; - if (!has_snapshot) return "no action received yet"; - if (!finite) return "action contains non-finite values"; - if (!seq_ok) return "action answers an obs frame older than the previous one"; - if (!fresh) return "action is stale (age > max_action_age)"; + std::string invalid_reason_( + bool enabled, bool contract_ok, bool has_snapshot, bool fresh, bool seq_ok, + bool finite) const { + if (!contract_ok) + return "contract/model mismatch (latched; restart or re-enable)"; + if (!enabled) + return "disabled by enable interface"; + if (!has_snapshot) + return "no action received yet"; + if (!finite) + return "action contains non-finite values"; + if (!seq_ok) + return "action answers an obs frame older than the previous one"; + if (!fresh) + return "action is stale (age > max_action_age)"; return "unknown"; } @@ -1396,19 +1502,19 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::vector>> action_outputs_; std::unordered_set own_output_paths_; - std::size_t obs_size_ = 0; + std::size_t obs_size_ = 0; std::size_t action_size_ = 0; std::string rl_base_; - double policy_rate_ = 50.0; - double max_action_age_ = 0.04; + double policy_rate_ = 50.0; + double max_action_age_ = 0.04; std::uint64_t expected_model_id_ = 0; InvalidMode invalid_mode_ = InvalidMode::kNaN; std::string enable_path_; bool enable_default_ = false; std::string reset_path_; std::size_t enable_slot_ = kNoSlot; - std::size_t reset_slot_ = kNoSlot; + std::size_t reset_slot_ = kNoSlot; std::string obs_signature_; std::string actions_signature_; @@ -1417,33 +1523,33 @@ class RlBridge final : public rmcs_executor::Component, public rclcpp::Node { std::vector last_actions_; std::vector written_; - std::chrono::nanoseconds pub_period_ { std::chrono::milliseconds(20) }; - std::chrono::steady_clock::time_point last_pub_time_ { }; - std::chrono::steady_clock::time_point prev_pub_time_ { }; - bool pub_started_ = false; - std::uint64_t pub_seq_ = 0; + std::chrono::nanoseconds pub_period_{std::chrono::milliseconds(20)}; + std::chrono::steady_clock::time_point last_pub_time_{}; + std::chrono::steady_clock::time_point prev_pub_time_{}; + bool pub_started_ = false; + std::uint64_t pub_seq_ = 0; std::uint64_t prev_pub_seq_ = 0; std::uint64_t last_reset_count_ = 0; std::uint64_t obs_invalid_count_ = 0; std::uint64_t pub_ok_count_ = 0; bool contract_ok_ = true; - bool last_valid_ = false; + bool last_valid_ = false; ActionSnapshot incoming_; ActionSnapshot read_snapshot_; - alignas(64) std::atomic action_sequence_ { 0 }; + alignas(64) std::atomic action_sequence_{0}; - OutputInterface valid_output_ { }; - OutputInterface healthy_output_ { }; - OutputInterface action_age_output_ { }; - OutputInterface obs_seq_output_ { }; + OutputInterface valid_output_{}; + OutputInterface healthy_output_{}; + OutputInterface action_age_output_{}; + OutputInterface obs_seq_output_{}; rclcpp::Publisher::SharedPtr obs_publisher_; rclcpp::Subscription::SharedPtr action_subscription_; }; -} // namespace rmcs::rl +} // namespace rmcs_rl #include -PLUGINLIB_EXPORT_CLASS(rmcs::rl::RlBridge, rmcs_executor::Component) +PLUGINLIB_EXPORT_CLASS(rmcs_rl::RlBridge, rmcs_executor::Component) diff --git a/src/rl_controller.cpp b/src/rl_controller.cpp deleted file mode 100644 index 63e94b3..0000000 --- a/src/rl_controller.cpp +++ /dev/null @@ -1,652 +0,0 @@ -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include - -#include "onnxruntime_inference.hpp" - -namespace rmcs::rl { - -class RlController - : public rmcs_executor::Component - , public rclcpp::Node { - - enum class State : std::uint8_t { kInit = 0, kIdle = 1, kPrepare = 2, kRl = 3 }; - -public: - explicit RlController() - : Node( - get_component_name(), - rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { - - joint_names_ = param_or("joint_names", std::vector{}); - joint_base_path_ = param_or("joint_base_path", ""); - position_pd_joints_ = param_or( - "position_pd_joints", std::vector{}); - velocity_pd_joints_ = param_or( - "velocity_pd_joints", std::vector{}); - - if (joint_base_path_.empty()) - throw std::invalid_argument( - "joint_base_path must be configured for this robot (e.g. /chassis, /wheel_leg)"); - - position_group_angle_suffix_ = - param_or("position_group_angle_suffix", "/angle"); - position_group_velocity_suffix_ = - param_or("position_group_velocity_suffix", "/velocity"); - velocity_group_angle_suffix_ = - param_or("velocity_group_angle_suffix", "/angle"); - velocity_group_velocity_suffix_ = - param_or("velocity_group_velocity_suffix", "/velocity"); - - const std::size_t dof = joint_names_.size(); - if (dof == 0 || dof > 32) - throw std::invalid_argument( - "joint_names must be configured (1..32 joints, training DOF order)"); - for (const auto idx : position_pd_joints_) - if (idx < 0 || static_cast(idx) >= dof) - throw std::invalid_argument("position_pd_joints out of range"); - for (const auto idx : velocity_pd_joints_) - if (idx < 0 || static_cast(idx) >= dof) - throw std::invalid_argument("velocity_pd_joints out of range"); - if (position_pd_joints_.empty() && velocity_pd_joints_.empty()) - throw std::invalid_argument( - "position_pd_joints/velocity_pd_joints: at least one PD group must be configured"); - - dof_ = dof; - - const auto configured_obs = require_param("rl_obs_size"); - const auto configured_act = require_param("rl_action_size"); - if (configured_obs <= 0 || configured_act <= 0) - throw std::invalid_argument("rl_obs_size / rl_action_size must be positive"); - rl_obs_size_ = static_cast(configured_obs); - rl_action_size_ = static_cast(configured_act); - std::int64_t max_grouped_joint = -1; - for (const auto idx : position_pd_joints_) - max_grouped_joint = std::max(max_grouped_joint, idx); - for (const auto idx : velocity_pd_joints_) - max_grouped_joint = std::max(max_grouped_joint, idx); - if (configured_act < max_grouped_joint + 1) - throw std::invalid_argument( - "rl_action_size=" + std::to_string(configured_act) - + " too small: PD group joints need action slots up to index " - + std::to_string(max_grouped_joint)); - const std::size_t layout_obs_size = 10 + 2 * dof + rl_action_size_; - if (rl_obs_size_ != layout_obs_size) - throw std::invalid_argument( - "rl_obs_size=" + std::to_string(rl_obs_size_) - + " inconsistent with this controller's observation layout length " - "(10 + 2*" + std::to_string(dof) + " + rl_action_size=" - + std::to_string(layout_obs_size) + "); see doc/model-contract.md"); - default_dof_pos_ = param_or("default_dof_pos", std::vector(dof, 0.0)); - dof_pos_limits_lower_ = param_or("dof_pos_limits_lower", std::vector(dof, -100.0)); - dof_pos_limits_upper_ = param_or("dof_pos_limits_upper", std::vector(dof, 100.0)); - position_action_scale_ = param_or("position_action_scale", 1.0); - velocity_action_scale_ = param_or("velocity_action_scale", 10.0); - max_velocity_ = param_or("max_velocity", 100.0); - position_kp_ = param_or("position_kp", 200.0); - position_kd_ = param_or("position_kd", 4.0); - velocity_kp_ = param_or("velocity_kp", 20.0); - velocity_kd_ = param_or("velocity_kd", 0.5); - position_torque_max_ = param_or("position_torque_max", 20.0); - velocity_torque_max_ = param_or("velocity_torque_max", 6.0); - - obs_height_scale_ = param_or("obs_height_scale", 5.0); - obs_ang_vel_scale_ = param_or("obs_ang_vel_scale", 0.5); - obs_gravity_scale_ = param_or("obs_gravity_scale", 1.0); - obs_dof_pos_scale_ = param_or("obs_dof_pos_scale", 1.0); - obs_dof_vel_scale_ = param_or("obs_dof_vel_scale", 0.1); - clip_observations_ = param_or("clip_observations", 100.0); - clip_actions_ = param_or("clip_actions", 100.0); - - motion_linear_x_min_ = param_or("motion_linear_x_min", 0.0); - motion_linear_x_max_ = param_or("motion_linear_x_max", 0.0); - motion_angular_z_min_ = param_or("motion_angular_z_min", 0.0); - motion_angular_z_max_ = param_or("motion_angular_z_max", 0.0); - command_height_min_ = param_or("command_height_min", 0.0); - command_height_max_ = param_or("command_height_max", 10.0); - default_command_height_ = param_or("default_command_height", 0.0); - - auto_enter_rl_ = param_or("auto_enter_rl", false); - prepare_dof_pos_ = param_or( - "prepare_dof_pos", - std::vector(position_pd_joints_.size(), 0.0)); - if (!position_pd_joints_.empty() - && prepare_dof_pos_.size() != position_pd_joints_.size()) - throw std::invalid_argument( - "prepare_dof_pos size must match position_pd_joints size"); - prepare_kp_ = param_or("prepare_kp", 80.0); - prepare_kd_ = param_or("prepare_kd", 2.0); - prepare_max_velocity_ = param_or("prepare_max_velocity", 1.0); - prepare_reach_threshold_ = param_or("prepare_reach_threshold", 0.02); - - rl_model_path_ = param_or("rl_model_path", ""); - rl_inference_frequency_ = param_or("rl_inference_frequency", 100.0); - rl_publish_network_io_ = param_or("rl_publish_network_io", false); - - joint_angle_input_ = std::make_unique[]>(dof); - joint_velocity_input_ = std::make_unique[]>(dof); - joint_control_torque_output_ = std::make_unique[]>(dof); - for (std::size_t i = 0; i < dof; ++i) { - const std::string base = joint_base_path_ + "/" + joint_names_[i]; - const bool is_position_joint = std::find( - position_pd_joints_.begin(), position_pd_joints_.end(), - static_cast(i)) - != position_pd_joints_.end(); - register_input( - base - + (is_position_joint ? position_group_angle_suffix_ - : velocity_group_angle_suffix_), - joint_angle_input_[i]); - register_input( - base - + (is_position_joint ? position_group_velocity_suffix_ - : velocity_group_velocity_suffix_), - joint_velocity_input_[i]); - register_output(base + "/control_torque", joint_control_torque_output_[i], 0.0); - } - - register_input(joint_base_path_ + "/imu/quaternion", imu_quaternion_); - register_input(joint_base_path_ + "/imu/angular_velocity", imu_angular_velocity_); - - register_input(joint_base_path_ + "/command/vx", command_vx_, false); - register_input(joint_base_path_ + "/command/yaw_rate", command_yaw_rate_, false); - register_input(joint_base_path_ + "/command/height", command_height_, false); - register_input(joint_base_path_ + "/command/state", command_state_, false); - register_input(joint_base_path_ + "/reset_count", reset_count_, false); - - if (rl_publish_network_io_) { - register_output( - joint_base_path_ + "/rl/observation", rl_observation_output_, std::vector{}); - register_output( - joint_base_path_ + "/rl/action", rl_action_output_, std::vector{}); - observation_publisher_ = create_publisher( - joint_base_path_ + "/rl/observation", 1); - action_publisher_ = create_publisher( - joint_base_path_ + "/rl/action", 1); - state_publisher_ = create_publisher( - joint_base_path_ + "/rl/state", 1); - } - register_output(joint_base_path_ + "/rl/state", rl_state_output_, 0); - - action_.assign(rl_action_size_, 0.0); - last_actions_.assign(rl_action_size_, 0.0); - prepare_pos_.assign(position_pd_joints_.size(), 0.0); - - inference_ready_ = false; - const std::string resolved_model_path = resolve_model_path_(rl_model_path_); - if (!resolved_model_path.empty()) { - inference_ready_ = inference_.load(OnnxRuntimeInference::Config{ - .model_path = resolved_model_path, - .input_name = "obs", - .output_name = "actions", - .input_size = rl_obs_size_, - .output_size = rl_action_size_, - }); - if (inference_ready_) { - RCLCPP_INFO( - get_logger(), - "RL policy loaded: %s ([1,%zu] -> [1,%zu], %.1f Hz, %zu joints)", - resolved_model_path.c_str(), rl_obs_size_, rl_action_size_, - rl_inference_frequency_, dof); - } else { - RCLCPP_ERROR( - get_logger(), "Failed to load RL policy '%s'; RL state unavailable", - resolved_model_path.c_str()); - } - } else { - RCLCPP_ERROR(get_logger(), "rl_model_path not set; RL state unavailable"); - } - - height_ = default_command_height_; - } - - void update() override { - const auto now = std::chrono::steady_clock::now(); - const double dt = std::clamp( - std::chrono::duration(now - last_update_time_).count(), 0.0, 0.1); - last_update_time_ = now; - - if (reset_count_.ready() && *reset_count_ != last_reset_count_) { - last_reset_count_ = *reset_count_; - reset_runtime_(); - return; - } - - read_commands_(); - update_state_machine_(); - - switch (state_) { - case State::kInit: - case State::kIdle: - write_zero_outputs_(); - break; - case State::kPrepare: - prepare_step_(dt); - break; - case State::kRl: - rl_step_(); - break; - } - - *rl_state_output_ = static_cast(state_); - - if (rl_publish_network_io_ && should_publish_io_()) { - std_msgs::msg::Int32 state_msg; - state_msg.data = static_cast(state_); - state_publisher_->publish(state_msg); - } - - if (dof_ > 0) { - std::ostringstream oss; - oss << "joint_q_deg=["; - for (std::size_t i = 0; i < dof_; ++i) { - if (i > 0) - oss << ' '; - oss << read_joint_angle_(i) * (180.0 / std::numbers::pi); - } - oss << "] (state=" << static_cast(state_) << ")"; - RCLCPP_WARN_THROTTLE( - get_logger(), *get_clock(), 1000, "%s", oss.str().c_str()); - } - } - -private: - template - T param_or(const std::string& name, const T& default_value) { - T value; - try { - if (get_parameter(name, value)) - return value; - } catch (const rclcpp::exceptions::InvalidParameterValueException&) { - } - RCLCPP_WARN(get_logger(), "Parameter '%s' not set, using default", name.c_str()); - return default_value; - } - - template - T require_param(const std::string& name) { - T value{}; - try { - if (get_parameter(name, value)) - return value; - } catch (const std::exception& error) { - throw std::invalid_argument( - "required parameter '" + name + "' is invalid: " + error.what()); - } - throw std::invalid_argument( - "missing required parameter '" + name + "' (no default; configure per robot)"); - } - - static std::string resolve_model_path_(const std::string& path) { - if (path.empty() || path.front() == '/') - return path; - try { - return ament_index_cpp::get_package_share_directory("rmcs_rl") + "/" + path; - } catch (const std::exception&) { - return path; - } - } - - void read_commands_() { - const double vx = command_vx_.ready() ? *command_vx_ : 0.0; - const double yaw = command_yaw_rate_.ready() ? *command_yaw_rate_ : 0.0; - const double height = command_height_.ready() ? *command_height_ : default_command_height_; - vx_ = std::clamp(vx, motion_linear_x_min_, motion_linear_x_max_); - yaw_rate_ = std::clamp(yaw, motion_angular_z_min_, motion_angular_z_max_); - height_ = std::clamp(height, command_height_min_, command_height_max_); - } - - bool resolve_state_command_(int raw, State& target) const { - switch (raw) { - case 0: target = State::kInit; return true; - case 1: target = State::kIdle; return true; - case 2: target = State::kPrepare; return true; - case 3: - target = inference_ready_ ? State::kRl : State::kIdle; - return inference_ready_; - default: return false; - } - } - - void update_state_machine_() { - if (command_state_.ready()) { - const int raw = *command_state_; - State target; - if (resolve_state_command_(raw, target)) { - if (target == State::kRl && state_ != State::kRl - && !(state_ == State::kPrepare && prepare_reached_)) { - RCLCPP_WARN_THROTTLE( - get_logger(), *get_clock(), 1000, - "Refusing RL: state=%d prepare_reached=%d joint_q=[%.3f %.3f %.3f %.3f] " - "(send 2 first, wait PREPARE done, then 3)", - static_cast(state_), prepare_reached_ ? 1 : 0, - read_joint_angle_(0), read_joint_angle_(1), read_joint_angle_(2), - read_joint_angle_(3)); - target = state_; - } - if (target != state_) - enter_state_(target); - } else { - RCLCPP_WARN_THROTTLE( - get_logger(), *get_clock(), 1000, "Invalid state command %d (use 0..3)", raw); - } - } - if (auto_enter_rl_ && inference_ready_ && state_ == State::kPrepare && prepare_reached_) { - enter_state_(State::kRl); - auto_enter_rl_ = false; - } - } - - void enter_state_(State target) { - state_ = target; - reset_policy_runtime_(); - if (target == State::kPrepare) { - for (std::size_t k = 0; k < position_pd_joints_.size(); ++k) - prepare_pos_[k] = read_joint_angle_(static_cast(position_pd_joints_[k])); - prepare_reached_ = false; - } - RCLCPP_INFO(get_logger(), "Entering state %d", static_cast(target)); - } - - bool is_velocity_pd_joint_(std::size_t joint_index) const { - return std::find( - velocity_pd_joints_.begin(), velocity_pd_joints_.end(), - static_cast(joint_index)) - != velocity_pd_joints_.end(); - } - - bool build_observation_(std::vector& obs) { - obs.assign(rl_obs_size_, 0.0); - std::size_t k = 0; - obs[k++] = vx_; - obs[k++] = 0.0; - obs[k++] = yaw_rate_; - obs[k++] = height_ * obs_height_scale_; - const Eigen::Vector3d ang_vel = *imu_angular_velocity_; - for (std::size_t i = 0; i < 3; ++i) - obs[k++] = ang_vel[i] * obs_ang_vel_scale_; - const Eigen::Vector3d gravity = *imu_quaternion_ * Eigen::Vector3d(0.0, 0.0, -1.0); - for (std::size_t i = 0; i < 3; ++i) - obs[k++] = gravity[i] * obs_gravity_scale_; - for (std::size_t i = 0; i < dof_; ++i) { - obs[k++] = is_velocity_pd_joint_(i) - ? 0.0 - : (read_joint_angle_(i) - default_dof_pos_[i]) * obs_dof_pos_scale_; - } - for (std::size_t i = 0; i < dof_; ++i) - obs[k++] = read_joint_velocity_(i) * obs_dof_vel_scale_; - for (std::size_t i = 0; i < rl_action_size_; ++i) - obs[k++] = last_actions_[i]; - - if (k != rl_obs_size_) - return false; - - for (double& value : obs) - value = std::clamp(value, -clip_observations_, clip_observations_); - return std::all_of(obs.begin(), obs.end(), [](double v) { return std::isfinite(v); }); - } - - bool should_infer_() { - const auto now = std::chrono::steady_clock::now(); - if (!last_inference_time_initialized_) { - last_inference_time_ = now; - last_inference_time_initialized_ = true; - return true; - } - const double period = 1.0 / std::max(rl_inference_frequency_, 1.0); - if (std::chrono::duration(now - last_inference_time_).count() < period) - return false; - last_inference_time_ = now; - return true; - } - - bool should_publish_io_() { - const auto now = std::chrono::steady_clock::now(); - if (std::chrono::duration(now - last_io_publish_time_).count() < 10.0) - return false; - last_io_publish_time_ = now; - return true; - } - - void rl_step_() { - if (!inference_ready_) { - fail_safe_("policy session is not ready"); - return; - } - std::vector obs; - if (!build_observation_(obs)) { - fail_safe_("policy observation size/validity mismatch"); - return; - } - if (rl_publish_network_io_) - (*rl_observation_output_) = obs; - - if (should_infer_()) { - std::vector obs_f(obs.begin(), obs.end()); - std::vector act_f(rl_action_size_, 0.0F); - if (!inference_.run(obs_f, act_f)) { - fail_safe_("ONNX Runtime rejected the policy input"); - return; - } - for (std::size_t i = 0; i < rl_action_size_; ++i) - action_[i] = std::clamp(static_cast(act_f[i]), -clip_actions_, clip_actions_); - last_actions_ = action_; - if (rl_publish_network_io_) { - (*rl_action_output_) = action_; - std_msgs::msg::Float64MultiArray obs_msg; - obs_msg.data = obs; - observation_publisher_->publish(obs_msg); - std_msgs::msg::Float64MultiArray act_msg; - act_msg.data = action_; - action_publisher_->publish(act_msg); - } - } - - apply_pd_(); - } - - void apply_pd_() { - std::vector torques(dof_, 0.0); - - for (const auto idx : position_pd_joints_) { - const std::size_t i = static_cast(idx); - const double target = std::clamp( - position_action_scale_ * action_[i] + default_dof_pos_[i], - dof_pos_limits_lower_[i], dof_pos_limits_upper_[i]); - const double tau = position_kp_ * (target - read_joint_angle_(i)) - - position_kd_ * read_joint_velocity_(i); - torques[i] = std::clamp(tau, -position_torque_max_, position_torque_max_); - } - for (const auto idx : velocity_pd_joints_) { - const std::size_t i = static_cast(idx); - const double vel_target = std::clamp( - velocity_action_scale_ * action_[i], -max_velocity_, max_velocity_); - const double tau = velocity_kp_ * (vel_target - read_joint_velocity_(i)) - - velocity_kd_ * read_joint_velocity_(i); - torques[i] = std::clamp(tau, -velocity_torque_max_, velocity_torque_max_); - } - - write_outputs_(torques); - } - - void prepare_step_(double dt) { - std::vector torques(dof_, 0.0); - bool reached = true; - for (std::size_t k = 0; k < position_pd_joints_.size(); ++k) { - const std::size_t i = static_cast(position_pd_joints_[k]); - const double target = prepare_dof_pos_[k]; - const double step = prepare_max_velocity_ * dt; - const double diff = target - prepare_pos_[k]; - if (std::abs(diff) > step) { - prepare_pos_[k] += std::copysign(step, diff); - reached = false; - } else { - prepare_pos_[k] = target; - } - const double tau = prepare_kp_ * (prepare_pos_[k] - read_joint_angle_(i)) - - prepare_kd_ * read_joint_velocity_(i); - torques[i] = std::clamp(tau, -position_torque_max_, position_torque_max_); - if (std::abs(prepare_pos_[k] - read_joint_angle_(i)) > prepare_reach_threshold_) - reached = false; - } - prepare_reached_ = reached; - write_outputs_(torques); - } - - void write_outputs_(const std::vector& torques) { - for (std::size_t i = 0; i < dof_; ++i) { - const double value = std::isfinite(torques[i]) ? torques[i] : 0.0; - *joint_control_torque_output_[i] = value; - } - } - - void write_zero_outputs_() { write_outputs_(std::vector(dof_, 0.0)); } - - void fail_safe_(const std::string& reason) { - RCLCPP_ERROR_THROTTLE(get_logger(), *get_clock(), 1000, "RL safety stop: %s", reason.c_str()); - write_zero_outputs_(); - reset_policy_runtime_(); - if (state_ != State::kIdle) { - state_ = State::kIdle; - RCLCPP_INFO(get_logger(), "Fail-safe: forced to IDLE"); - } - } - - void reset_policy_runtime_() { - std::fill(last_actions_.begin(), last_actions_.end(), 0.0); - std::fill(action_.begin(), action_.end(), 0.0); - std::fill(prepare_pos_.begin(), prepare_pos_.end(), 0.0); - prepare_reached_ = false; - last_inference_time_initialized_ = false; - } - - void reset_runtime_() { - write_zero_outputs_(); - reset_policy_runtime_(); - state_ = State::kIdle; - vx_ = 0.0; - yaw_rate_ = 0.0; - height_ = default_command_height_; - RCLCPP_INFO(get_logger(), "RMCS reset: outputs stopped, state cleared"); - } - - double read_joint_angle_(std::size_t index) const { - return joint_angle_input_[index].ready() ? *joint_angle_input_[index] : 0.0; - } - double read_joint_velocity_(std::size_t index) const { - return joint_velocity_input_[index].ready() ? *joint_velocity_input_[index] : 0.0; - } - - std::vector joint_names_; - std::string joint_base_path_; - std::vector position_pd_joints_; - std::vector velocity_pd_joints_; - std::string position_group_angle_suffix_ = "/angle"; - std::string position_group_velocity_suffix_ = "/velocity"; - std::string velocity_group_angle_suffix_ = "/angle"; - std::string velocity_group_velocity_suffix_ = "/velocity"; - std::size_t dof_ = 0; - - std::vector default_dof_pos_; - std::vector dof_pos_limits_lower_; - std::vector dof_pos_limits_upper_; - double position_action_scale_ = 1.0; - double velocity_action_scale_ = 10.0; - double max_velocity_ = 100.0; - double position_kp_ = 200.0; - double position_kd_ = 4.0; - double velocity_kp_ = 20.0; - double velocity_kd_ = 0.5; - double position_torque_max_ = 20.0; - double velocity_torque_max_ = 6.0; - - std::size_t rl_obs_size_ = 0; - std::size_t rl_action_size_ = 0; - double obs_height_scale_ = 5.0; - double obs_ang_vel_scale_ = 0.5; - double obs_gravity_scale_ = 1.0; - double obs_dof_pos_scale_ = 1.0; - double obs_dof_vel_scale_ = 0.1; - double clip_observations_ = 100.0; - double clip_actions_ = 100.0; - - double motion_linear_x_min_ = 0.0; - double motion_linear_x_max_ = 0.0; - double motion_angular_z_min_ = 0.0; - double motion_angular_z_max_ = 0.0; - double command_height_min_ = 0.0; - double command_height_max_ = 10.0; - double default_command_height_ = 0.0; - - bool auto_enter_rl_ = false; - std::vector prepare_dof_pos_; - double prepare_kp_ = 80.0; - double prepare_kd_ = 2.0; - double prepare_max_velocity_ = 1.0; - double prepare_reach_threshold_ = 0.02; - - std::string rl_model_path_; - double rl_inference_frequency_ = 100.0; - bool rl_publish_network_io_ = false; - - State state_ = State::kInit; - double vx_ = 0.0; - double yaw_rate_ = 0.0; - double height_ = 0.0; - std::vector action_; - std::vector last_actions_; - std::vector prepare_pos_; - bool prepare_reached_ = false; - bool inference_ready_ = false; - - std::chrono::steady_clock::time_point last_update_time_{}; - std::chrono::steady_clock::time_point last_inference_time_{}; - std::chrono::steady_clock::time_point last_io_publish_time_{}; - bool last_inference_time_initialized_ = false; - std::size_t last_reset_count_ = 0; - - OnnxRuntimeInference inference_; - - std::unique_ptr[]> joint_angle_input_; - std::unique_ptr[]> joint_velocity_input_; - std::unique_ptr[]> joint_control_torque_output_; - - rmcs_executor::Component::InputInterface imu_quaternion_; - rmcs_executor::Component::InputInterface imu_angular_velocity_; - - rmcs_executor::Component::InputInterface command_vx_; - rmcs_executor::Component::InputInterface command_yaw_rate_; - rmcs_executor::Component::InputInterface command_height_; - rmcs_executor::Component::InputInterface command_state_; - rmcs_executor::Component::InputInterface reset_count_; - - rmcs_executor::Component::OutputInterface> rl_observation_output_; - rmcs_executor::Component::OutputInterface> rl_action_output_; - rmcs_executor::Component::OutputInterface rl_state_output_; - - rclcpp::Publisher::SharedPtr observation_publisher_; - rclcpp::Publisher::SharedPtr action_publisher_; - rclcpp::Publisher::SharedPtr state_publisher_; -}; - -} // namespace rmcs::rl - -#include - -PLUGINLIB_EXPORT_CLASS(rmcs::rl::RlController, rmcs_executor::Component) diff --git a/src/rl_layout.hpp b/src/rl_layout.hpp index 7885788..c184053 100644 --- a/src/rl_layout.hpp +++ b/src/rl_layout.hpp @@ -1,6 +1,5 @@ #pragma once - #include #include #include @@ -9,10 +8,10 @@ #include #include -namespace rmcs::rl { +namespace rmcs_rl { inline constexpr std::uint64_t kFnv1a64Offset = 0xCBF29CE484222325ULL; -inline constexpr std::uint64_t kFnv1a64Prime = 0x100000001B3ULL; +inline constexpr std::uint64_t kFnv1a64Prime = 0x100000001B3ULL; inline std::uint64_t fnv1a64(std::string_view bytes) { std::uint64_t hash = kFnv1a64Offset; @@ -29,11 +28,11 @@ inline std::string hex16(std::uint64_t value) { return buffer; } -inline std::uint64_t layout_hash(std::string_view obs_signature, - std::string_view actions_signature, std::size_t obs_size, std::size_t actions_size) { +inline std::uint64_t layout_hash( + std::string_view obs_signature, std::string_view actions_signature, std::size_t obs_size, + std::size_t actions_size) { std::string canonical; - canonical.reserve( - obs_signature.size() + actions_signature.size() + 48); + canonical.reserve(obs_signature.size() + actions_signature.size() + 48); canonical.append(obs_signature); canonical.append("||"); canonical.append(actions_signature); @@ -45,7 +44,7 @@ inline std::uint64_t layout_hash(std::string_view obs_signature, } inline bool model_id_of_file(const std::string& path, std::uint64_t& out, std::string& error) { - std::ifstream stream { path, std::ios::binary }; + std::ifstream stream{path, std::ios::binary}; if (!stream) { error = "cannot open model file '" + path + "'"; return false; @@ -55,10 +54,11 @@ inline bool model_id_of_file(const std::string& path, std::uint64_t& out, std::s while (stream) { stream.read(buffer.data(), static_cast(buffer.size())); const auto read = static_cast(stream.gcount()); - if (read == 0) break; + if (read == 0) + break; for (std::size_t i = 0; i < read; ++i) hash = (hash ^ static_cast(static_cast(buffer[i]))) - * kFnv1a64Prime; + * kFnv1a64Prime; } if (stream.bad()) { error = "failed while reading model file '" + path + "'"; @@ -68,4 +68,4 @@ inline bool model_id_of_file(const std::string& path, std::uint64_t& out, std::s return true; } -} // namespace rmcs::rl +} // namespace rmcs_rl diff --git a/tool/check_policy_contract.py b/tool/check_policy_contract.py index de25373..c862d64 100644 --- a/tool/check_policy_contract.py +++ b/tool/check_policy_contract.py @@ -6,7 +6,7 @@ * self-check (no --config; used by CI, deployment YAML lives in the RMCS repo): model loads; one input "obs" / one output "actions"; float32; rank 2; batch 1; concrete (non-dynamic) shapes. Layout metadata is OPTIONAL: - - missing / v1 -> SKIP (legacy RlController models are not forced to be stamped) + - missing / v1 -> SKIP (stamped layout metadata is optional for model-only self-check) - present v2 -> signatures must be internally consistent with the tensor sizes, and policy_layout_hash must match those signatures (catches a bad stamp without any YAML) diff --git a/tool/rl_layout.py b/tool/rl_layout.py index 5a65283..3cd93ff 100644 --- a/tool/rl_layout.py +++ b/tool/rl_layout.py @@ -1,7 +1,7 @@ #!/usr/bin/env python3 """rmcs_rl 布局单一真源:词条语法 → 规范串(v2) → layout_hash / model_id。 -C++ 桥(rl_bridge)与策略进程各自实现同一份规范(doc/bridge-design.md §6.2/§6.3), +C++ 桥(rl_bridge)与策略进程各自实现同一份规范(planning/docs/bridge-design.md §6.2/§6.3), Python 侧的真源就是本模块:stamp_layout_metadata.py / check_policy_contract.py / gen_synthetic_policy.py 全部从这里取语法与 hash,避免同一份契约靠人写两遍(drift)。 diff --git a/tool/stamp_layout_metadata.py b/tool/stamp_layout_metadata.py index 5293b4e..c01d9b3 100644 --- a/tool/stamp_layout_metadata.py +++ b/tool/stamp_layout_metadata.py @@ -3,7 +3,7 @@ 桥(rl_bridge)在启动时看不到模型文件,因此两侧靠 layout_hash 运行期握手: 本工具把 YAML 推导出的规范串(v2)与 policy_layout_hash 写进模型 metadata, -策略进程据此与桥的 Observation.layout_hash 对账(doc/bridge-design.md §6.2)。 +策略进程据此与桥的 Observation.layout_hash 对账(planning/docs/bridge-design.md §6.2)。 **不需要重训**:换布局/换版本重跑本工具即可。 写入的键: