From 7cb727095920e620623008ec4b31d79531db862c Mon Sep 17 00:00:00 2001 From: zhaoxi <535394140@qq.com> Date: Fri, 31 Jul 2026 01:32:17 +0800 Subject: [PATCH 1/2] =?UTF-8?q?docs:=20=E5=AE=8C=E5=96=84=20rdt=20?= =?UTF-8?q?=E6=96=87=E6=A1=A3=E6=8F=8F=E8=BF=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 增加 Windows 一键安装脚本的描述 2. 完善 lpss 命令 --- doc/tutorials/rdt/rdt_lpss.md | 265 ++++++++++++++++++++++++++- doc/tutorials/rdt/usage.md | 4 + modules/io/include/rmvl/io/async.hpp | 3 + modules/io/src/async.cpp | 14 +- 4 files changed, 276 insertions(+), 10 deletions(-) diff --git a/doc/tutorials/rdt/rdt_lpss.md b/doc/tutorials/rdt/rdt_lpss.md index 7d21ec95..7f2cf77e 100644 --- a/doc/tutorials/rdt/rdt_lpss.md +++ b/doc/tutorials/rdt/rdt_lpss.md @@ -2,8 +2,8 @@ LPSS CLI 工具 {#tutorial_rdt_lpss} ============ @author 赵曦 -@date 2026/06/06 -@version 1.0 +@date 2026/07/31 +@version 1.1 @brief LPSS 命令行工具的使用教程 @prev_tutorial{tutorial_rdt_rdt} @@ -35,6 +35,7 @@ LPSS 是一个轻量级的发布订阅通信框架,采用去中心化设计, @dl_item{create,创建一个依赖 lpss 的新项目} @dl_item{node,节点 CLI 工具} @dl_item{topic,话题 CLI 工具} +@dl_item{service,服务 CLI 工具} @dl_item{interface,内置消息接口查看工具} @dl_item{graph,节点图工具} @dl_item{viz,3D 可视化工具 LViz} @@ -87,9 +88,32 @@ LPSS 是一个轻量级的发布订阅通信框架,采用去中心化设计, @dl_begin{命令} @dl_item{help,显示此帮助信息} @dl_item{info,显示节点信息} -@dl_item{list,列出所有节点} +@dl_item{list,列出所有节点,可使用 `-c` 仅显示数量} @dl_end +#### list 子命令 + +列出当前发现的所有 LPSS 节点。 + +**用法** + +
+
lpss node list [-c]
+
+ +@dl_begin{选项} +@dl_item{\-c,仅输出节点数量} +@dl_end + +**示例** + +
+
# 列出节点名称
+
lpss node list
+
# 仅输出节点数量
+
lpss node list -c
+
+ #### info 子命令 查看指定节点的信息,输出形如以下的内容 @@ -102,6 +126,12 @@ Publish Topics: Subscribe Topics: xxx + +Server Services: + xxx + +Client Services: + xxx ``` **用法** @@ -126,13 +156,14 @@ Subscribe Topics: **用法**
-
lpss topic [help | info | list | echo | pub | type | hz | bw] [args...]
+
lpss topic [help | info | list | find | echo | pub | type | hz | bw] [args...]
@dl_begin{命令} @dl_item{help,显示此帮助信息} @dl_item{info,显示话题信息} -@dl_item{list,列出所有话题} +@dl_item{list,列出所有话题,可使用 `-c` 仅显示数量} +@dl_item{find,按消息类型查找话题,可使用 `-c` 仅显示数量} @dl_item{echo,显示话题内容} @dl_item{pub,发布话题} @dl_item{type,显示话题类型} @@ -140,6 +171,45 @@ Subscribe Topics: @dl_item{bw,测量话题带宽,单位为 MB/s、kB/s 或 B/s} @dl_end +#### list 子命令 + +列出当前发现的所有话题,话题名称按字典序输出。 + +**用法** + +
+
lpss topic list [-c]
+
+ +@dl_begin{选项} +@dl_item{\-c,仅输出话题数量} +@dl_end + +#### find 子命令 + +按完整消息类型查找话题,查询结果按话题名称的字典序输出。 + +**用法** + +
+
lpss topic find <msg_type> [-c]
+
+ +@param msg_type 消息类型,例如 `std/String` + +@dl_begin{选项} +@dl_item{\-c,仅输出匹配的话题数量} +@dl_end + +**示例** + +
+
# 查找所有使用 std/String 消息类型的话题
+
lpss topic find std/String
+
# 仅输出匹配话题的数量
+
lpss topic find std/String -c
+
+ #### info 子命令 查看指定话题的信息,输出形如以下的内容 @@ -259,6 +329,146 @@ Subscriber Node:
lpss topic bw /str
+### service 命令 + +服务工具,用于查看已发现的 LPSS 服务,并使用 JSON 请求调用内置服务。 + +**用法** + +
+
lpss service [help | info | list | type | find | call] [args...]
+
+ +@dl_begin{命令} +@dl_item{help,显示此帮助信息} +@dl_item{info,显示服务信息} +@dl_item{list,列出所有服务,可使用 `-c` 仅显示数量} +@dl_item{type,显示服务类型} +@dl_item{find,按服务类型查找服务,可使用 `-c` 仅显示数量} +@dl_item{call,使用 JSON 请求调用内置服务} +@dl_end + +#### list 子命令 + +列出当前发现的所有服务,服务名称按字典序输出。 + +**用法** + +
+
lpss service list [-c]
+
+ +@dl_begin{选项} +@dl_item{\-c,仅输出服务数量} +@dl_end + +#### info 子命令 + +查看指定服务的类型、服务端节点和客户端节点,输出形如以下内容。 + +``` +Type: std/SetBool + +Server Node: + server_node + +Client Node: + client_node +``` + +**用法** + +
+
lpss service info <service_name>
+
+ +@param service_name 服务名称 + +**示例** + +
+
lpss service info /set_enabled
+
+ +#### type 子命令 + +显示指定服务的服务类型。 + +**用法** + +
+
lpss service type <service_name>
+
+ +@param service_name 服务名称 + +**示例** + +
+
# 输出 std/SetBool
+
lpss service type /set_enabled
+
+ +#### find 子命令 + +按类型精确匹配服务。`service_type` 可以是完整服务类型,也可以是对应的请求或响应消息类型。查询结果按服务名称的字典序输出。 + +**用法** + +
+
lpss service find <service_type> [-c]
+
+ +@param service_type 服务类型,例如 `std/SetBool`;也可使用 `std/SetBool_Request` 或 `std/SetBool_Response` + +@dl_begin{选项} +@dl_item{\-c,仅输出匹配的服务数量} +@dl_end + +**示例** + +
+
# 列出所有 std/SetBool 服务
+
lpss service find std/SetBool
+
# 仅输出匹配服务的数量
+
lpss service find std/SetBool -c
+
+ +#### call 子命令 + +使用 JSON 对象作为请求调用指定服务,并将响应输出为 JSON。`json_request` 省略或为空时按 `{}` 处理,命令等待响应的超时时间为 3 秒。 + +**用法** + +
+
lpss service call <service_name> [json_request]
+
+ +@param service_name 服务名称 +@param json_request JSON 对象形式的请求,建议使用单引号包围,避免 Shell 改写引号等字符 + +目前 `call` 支持以下内置服务类型。 + +| 服务类型 | JSON 请求 | +| --- | --- | +| `std/Empty` | `{}`,可省略 | +| `std/Trigger` | `{}`,可省略 | +| `std/SetBool` | 必须包含布尔字段 `data` | +| `sensor/SetCameraInfo` | 必须包含对象字段 `camera_info`;其中 `D` 和 `K` 如果出现,必须分别包含 5 个和 9 个数字 | + +**示例** + +
+
# 调用无请求字段的 Trigger 服务
+
lpss service call /trigger
+
# 使用 JSON 请求调用 SetBool 服务
+
lpss service call /set_enabled '{"data":true}'
+
# 调用 SetCameraInfo 服务,未填写的 CameraInfo 字段使用默认值
+
lpss service call /set_camera_info '{"camera_info":{"height":1080,"width":1920}}'
+
+ +@note `call` 暂不支持自定义服务类型;发现到不支持的类型时会输出 `unsupported service type`。 + ### interface 命令 内置消息接口查看工具 @@ -346,10 +556,49 @@ geometry/Quaternion orientation ### graph 命令 -节点图工具 +节点图工具 LGraph,用于在浏览器中查看 LPSS 通信拓扑,包括节点、话题、服务及其发布/订阅、服务端/客户端关系。也可直接使用 `lgraph` 命令启动。 -@warning 未完成,敬请期待 +**用法** + +
+
lpss graph [name_subfix]
+
# 等价命令
+
lgraph [name_subfix]
+
+ +@param name_subfix 可选的实例名称,用于生成 LPSS 节点名 `lgraph_`;省略时自动生成 5 位随机标识 + +启动成功后,终端会输出本机和局域网访问地址。LGraph 默认监听 `17493` 端口,本机可访问: + +
+
http://localhost:17493
+
+ +**示例** + +
+
# 使用自动生成的实例名启动
+
lpss graph
+
# LPSS 节点名为 lgraph_debug
+
lgraph debug
+
+ +@note Web 服务使用固定端口 `17493`,同一主机上不能同时启动多个 LGraph 实例。按 `Ctrl+C` 可停止工具。 + +@note 如果提示 `%lpss graph 工具尚未安装`,请在 `rmvl-dev-tools` 中重新运行 `install.bash`。 ### viz 命令 -3D 可视化工具 LViz,也可直接使用 `lviz` 命令来启动。 +3D 可视化工具 LViz,也可直接使用 `lviz` 命令启动。 + +**用法** + +
+
lpss viz [name_subfix]
+
# 等价命令
+
lviz [name_subfix]
+
+ +@param name_subfix 可选的实例名称,用于生成 LPSS 节点名 `lviz_node_`;省略时自动生成 5 位随机标识 + +LViz 默认监听 `17492` 端口,启动后可通过 `http://localhost:17492` 访问,按 `Ctrl+C` 停止工具。 diff --git a/doc/tutorials/rdt/usage.md b/doc/tutorials/rdt/usage.md index 80b3d996..8a8f2a2d 100644 --- a/doc/tutorials/rdt/usage.md +++ b/doc/tutorials/rdt/usage.md @@ -7,6 +7,10 @@ RMVL Dev Tools(简称 `rdt`)是一个 CLI 命令行工具,能够极大幅
wget https://cv-rmvl.github.io/install -qO - | bash
+Windows 用户可以使用如下的一键安装命令,在开始菜单找到 Windows Powershell 后点击进入,或者按下 `Win+X` 唤起对话框,按下 `I` 进入 Windows Powershell,输入以下内容,根据提示操作即可完成 RMVL 以及 rdt 工具的安装: +
+
irm https://cv-rmvl.github.io/install-win | iex
+
`rdt` 的使用非常简单,安装完成后在终端输入以下命令即可查看帮助文档: diff --git a/modules/io/include/rmvl/io/async.hpp b/modules/io/include/rmvl/io/async.hpp index 913fbd7c..18bb31c8 100644 --- a/modules/io/include/rmvl/io/async.hpp +++ b/modules/io/include/rmvl/io/async.hpp @@ -447,6 +447,9 @@ class Timer { class TimerAwaiter : public AsyncIOAwaiter { public: TimerAwaiter(IOContext &ctx, FileDescriptor fd, double duration) : AsyncIOAwaiter(ctx, fd), _duration(duration) {} +#ifdef _WIN32 + ~TimerAwaiter(); +#endif //! @cond void await_suspend(std::coroutine_handle<> handle); diff --git a/modules/io/src/async.cpp b/modules/io/src/async.cpp index 5c26f64d..806f6e16 100644 --- a/modules/io/src/async.cpp +++ b/modules/io/src/async.cpp @@ -148,11 +148,20 @@ struct TimerContext { IocpOverlapped *ovl{}; }; -void CALLBACK timer_callback(PTP_CALLBACK_INSTANCE, PVOID context, PTP_TIMER timer) { +void CALLBACK timer_callback(PTP_CALLBACK_INSTANCE, PVOID context, PTP_TIMER) { auto timer_context = reinterpret_cast(context); // 手动投递完成包 PostQueuedCompletionStatus(timer_context->aioh, 0, 0, &timer_context->ovl->ov); +} + +Timer::TimerAwaiter::~TimerAwaiter() { + if (_fd == INVALID_FD) + return; + auto timer = reinterpret_cast(_fd); + SetThreadpoolTimer(timer, nullptr, 0, 0); + WaitForThreadpoolTimerCallbacks(timer, TRUE); CloseThreadpoolTimer(timer); + _fd = INVALID_FD; } void Timer::TimerAwaiter::await_suspend(std::coroutine_handle<> handle) { @@ -161,7 +170,8 @@ void Timer::TimerAwaiter::await_suspend(std::coroutine_handle<> handle) { auto timer_ctx = new (_ovl->info) TimerContext{_aioh, _ovl.get()}; _fd = CreateThreadpoolTimer(timer_callback, timer_ctx, nullptr); if (_fd == nullptr) { - handle.resume(); + _fd = INVALID_FD; + PostQueuedCompletionStatus(_aioh, 0, 0, &_ovl->ov); return; } From 38fbd7633f46d50401497df8e76b7cec68cd3aa3 Mon Sep 17 00:00:00 2001 From: zhaoxi <535394140@qq.com> Date: Sat, 1 Aug 2026 15:45:27 +0800 Subject: [PATCH 2/2] =?UTF-8?q?refactor(ml):=20=E5=8F=96=E6=B6=88=20OnnxNe?= =?UTF-8?q?t=20=E7=9A=84=E5=A4=9A=E6=80=81=E7=BB=93=E6=9E=84=EF=BC=8C?= =?UTF-8?q?=E6=94=B9=E4=B8=BA=E7=BA=AF=E7=BB=A7=E6=89=BF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../include/rmvl/detector/armor_detector.h | 4 +- .../include/rmvl/detector/gyro_detector.h | 2 +- extra/detector/src/armor_detector/find.cpp | 7 +- extra/detector/src/gyro_detector/find.cpp | 7 +- modules/ml/include/rmvl/ml/ort.h | 105 ++++++------------ modules/ml/src/ort/classification.cpp | 19 +++- modules/ml/src/ort/ort.cpp | 10 +- 7 files changed, 62 insertions(+), 92 deletions(-) diff --git a/extra/detector/include/rmvl/detector/armor_detector.h b/extra/detector/include/rmvl/detector/armor_detector.h index cf8b60eb..897c9f18 100755 --- a/extra/detector/include/rmvl/detector/armor_detector.h +++ b/extra/detector/include/rmvl/detector/armor_detector.h @@ -12,8 +12,8 @@ #pragma once #include "rmvl/combo/armor.h" -#include "rmvl/tracker/tracker.h" #include "rmvl/ml/ort.h" +#include "rmvl/tracker/tracker.h" namespace rm { @@ -39,7 +39,7 @@ struct RMVL_EXPORTS_W_AG ArmorDetectorInfo { class RMVL_EXPORTS_W ArmorDetector final { double _tick; //!< 每一帧对应的时间点 ImuData _imu_data; //!< 每一帧对应的 IMU 数据 - std::unique_ptr _ort; + std::unique_ptr _ort; std::unordered_map _robot_t; public: diff --git a/extra/detector/include/rmvl/detector/gyro_detector.h b/extra/detector/include/rmvl/detector/gyro_detector.h index 1790506c..4a703390 100644 --- a/extra/detector/include/rmvl/detector/gyro_detector.h +++ b/extra/detector/include/rmvl/detector/gyro_detector.h @@ -40,7 +40,7 @@ class RMVL_EXPORTS_W GyroDetector final { ImuData _imu_data; //!< 每一帧对应的 IMU 数据 int _armor_num; //!< 默认装甲板数目 - std::unique_ptr _ort; + std::unique_ptr _ort; std::unordered_map _robot_t; public: diff --git a/extra/detector/src/armor_detector/find.cpp b/extra/detector/src/armor_detector/find.cpp index 8e5d6f92..e5656d43 100755 --- a/extra/detector/src/armor_detector/find.cpp +++ b/extra/detector/src/armor_detector/find.cpp @@ -38,10 +38,9 @@ void ArmorDetector::find(cv::Mat &src, std::vector &features, std: rois.reserve(armors.size()); for (const auto &armor : armors) { cv::Mat roi = Armor::getNumberROI(src, armor); - PreprocessOptions preop; - preop.means = {para::armor_detector_param.MODEL_MEAN}; - preop.stds = {para::armor_detector_param.MODEL_STD}; - int idx = ClassificationNet::cast(_ort->inference({roi}, preop, {})).first; + int idx = _ort->inference({roi}, {para::armor_detector_param.MODEL_MEAN}, + {para::armor_detector_param.MODEL_STD}) + .first; armor->setType(_robot_t[idx]); rois.emplace_back(roi); } diff --git a/extra/detector/src/gyro_detector/find.cpp b/extra/detector/src/gyro_detector/find.cpp index d3dc49d9..5e236b61 100755 --- a/extra/detector/src/gyro_detector/find.cpp +++ b/extra/detector/src/gyro_detector/find.cpp @@ -38,10 +38,9 @@ void GyroDetector::find(cv::Mat &src, std::vector &features, std:: rois.reserve(armors.size()); for (const auto &armor : armors) { cv::Mat roi = Armor::getNumberROI(src, armor); - PreprocessOptions preop; - preop.means = {para::gyro_detector_param.MODEL_MEAN}; - preop.stds = {para::gyro_detector_param.MODEL_STD}; - int idx = ClassificationNet::cast(_ort->inference({roi}, preop, {})).first; + int idx = _ort->inference({roi}, {para::gyro_detector_param.MODEL_MEAN}, + {para::gyro_detector_param.MODEL_STD}) + .first; armor->setType(_robot_t[idx]); rois.emplace_back(roi); } diff --git a/modules/ml/include/rmvl/ml/ort.h b/modules/ml/include/rmvl/ml/ort.h index 934f3775..a99ca9ab 100755 --- a/modules/ml/include/rmvl/ml/ort.h +++ b/modules/ml/include/rmvl/ml/ort.h @@ -11,92 +11,60 @@ #pragma once -#include +#include +#include +#include +#include +#include + #include #include #include "rmvl/core/rmvldef.hpp" -namespace rm -{ +namespace rm { //! @addtogroup ml_ort //! @{ //! Ort 提供者 -enum class OrtProvider : uint8_t -{ +enum class OrtProvider : uint8_t { CPU, //!< 由 `CPU` 执行 CUDA, //!< 由 `CUDA` 执行 TensorRT, //!< 由 `TensorRT` 执行 OpenVINO //!< 由 `OpenVINO` 执行 }; -//! 预处理选项 -struct RMVL_EXPORTS_W_AG PreprocessOptions -{ - RMVL_W_RW std::vector means; //!< 均值 - RMVL_W_RW std::vector stds; //!< 标准差 -}; - -//! 后处理选项 -struct RMVL_EXPORTS_W_AG PostprocessOptions -{ - RMVL_W_RW uint8_t color{}; //!< 颜色通道 - RMVL_W_RW std::vector thresh{}; //!< 阈值向量 -}; - //! ONNX-Runtime (Ort) 部署库基类 \cite microsoft23ort -class RMVL_EXPORTS_W OnnxNet -{ +class RMVL_EXPORTS_W OnnxNet { public: - /** - * @brief 创建 OnnxNet 对象 - * - * @param[in] model_path 模型路径,如果该路径不存在,则程序将因错误而退出 - * @param[in] prov Ort 提供者 - */ - RMVL_W OnnxNet(std::string_view model_path, OrtProvider prov); - //! 打印环境信息 RMVL_W static void printEnvInfo() noexcept; //! 打印模型信息 RMVL_W void printModelInfo() noexcept; - /** - * @brief 推理 - * - * @param[in] images 所有输入图像 - * @param[in] preop 预处理选项,不同网络有不同的预处理选项 - * @param[in] postop 后处理选项,不同网络有不同的后处理选项 - * @return 使用 `std::any` 表示的所有推理结果,需要根据具体的网络进行解析 - * @note 可使用 `::cast` 函数对返回类型进行转换 - */ - RMVL_W std::any inference(const std::vector &images, const PreprocessOptions &preop, const PostprocessOptions &postop); - - virtual ~OnnxNet() = default; + //! @cond + ~OnnxNet() = default; + //! @endcond -private: +protected: /** - * @brief 预处理 + * @brief 创建 OnnxNet 对象 * - * @param[in] images 所有输入图像 - * @param[in] preop 预处理选项,不同网络有不同的预处理选项 - * @return 模型的输入 Tensors + * @param[in] model_path 模型路径,如果该路径不存在,则程序将因错误而退出 + * @param[in] prov Ort 提供者 */ - virtual std::vector preProcess(const std::vector &images, const PreprocessOptions &preop); + OnnxNet(std::string_view model_path, OrtProvider prov); /** - * @brief 后处理 + * @brief 执行 ONNX Runtime 推理 * - * @param[in] output_tensors 模型的输出 Tensors - * @param[in] postop 后处理选项,不同网络有不同的后处理选项 - * @return 使用 `std::any` 表示的所有推理结果,需要根据具体的网络进行解析 + * @param[in] input_tensors 模型的输入 Tensors + * @return 模型的输出 Tensors */ - virtual std::any postProcess(const std::vector &output_tensors, const PostprocessOptions &postop); + std::vector run(const std::vector &input_tensors); -protected: Ort::MemoryInfo _memory_info; //!< 内存分配信息 Ort::Env _env; //!< 环境配置 Ort::SessionOptions _session_options; //!< 会话选项 @@ -118,17 +86,8 @@ class RMVL_EXPORTS_W OnnxNet * @note * - 输出层为 `[1, n]`,其中 `n` 为类别数 */ -class RMVL_EXPORTS_W ClassificationNet : public OnnxNet -{ +class RMVL_EXPORTS_W ClassificationNet : public OnnxNet { public: - /** - * @brief 推理结果转换 - * - * @param[in] result 使用 `std::any` 表示的推理结果 - * @return 转换后的推理结果,为 `std::pair` 类型,表示分类结果及其置信度 - */ - RMVL_W static std::pair cast(const std::any &result) { return std::any_cast>(result); } - /** * @brief 创建分类网络对象 * @@ -137,24 +96,34 @@ class RMVL_EXPORTS_W ClassificationNet : public OnnxNet */ RMVL_W ClassificationNet(std::string_view model_path, OrtProvider prov = OrtProvider::CPU); + /** + * @brief 执行分类网络推理 + * + * @param[in] images 所有输入图像 + * @param[in] means 各通道的均值 + * @param[in] stds 各通道的标准差 + * @return 分类结果及其置信度 + */ + RMVL_W std::pair inference(const std::vector &images, const std::vector &means, const std::vector &stds); + private: /** * @brief 预处理 * * @param[in] images 所有输入图像 - * @param[in] options 预处理选项,包含各通道的均值和标准差 + * @param[in] means 各通道的均值 + * @param[in] stds 各通道的标准差 * @return 模型的输入 Tensors */ - std::vector preProcess(const std::vector &images, const PreprocessOptions &options) override; + std::vector preProcess(const std::vector &images, const std::vector &means, const std::vector &stds); /** * @brief 后处理 * * @param[in] output_tensors 模型的输出 Tensors - * @param[in] postop 无需后处理选项,传入 `{}` 即可 - * @return 用 `std::any` 表示的分类结果及其置信度,可使用 `rm::ClassificationNet::cast` 函数对返回类型进行转换 + * @return 分类结果及其置信度 */ - std::any postProcess(const std::vector &output_tensors, const PostprocessOptions &postop) override; + std::pair postProcess(const std::vector &output_tensors); std::vector> _iarrays; //!< 输入数组 }; diff --git a/modules/ml/src/ort/classification.cpp b/modules/ml/src/ort/classification.cpp index 81be5a58..7502f415 100644 --- a/modules/ml/src/ort/classification.cpp +++ b/modules/ml/src/ort/classification.cpp @@ -66,7 +66,14 @@ static void imageToVector(const cv::Mat &input_image, float mean, float std, std p_input_array[h * W + w] = (input_image.at(h, w) / 255.f - mean) / std; } -std::vector ClassificationNet::preProcess(const std::vector &images, const PreprocessOptions &options) +std::pair ClassificationNet::inference(const std::vector &images, const std::vector &means, + const std::vector &stds) +{ + return postProcess(run(preProcess(images, means, stds))); +} + +std::vector ClassificationNet::preProcess(const std::vector &images, const std::vector &means, + const std::vector &stds) { std::size_t input_count = _session->GetInputCount(); RMVL_Assert(input_count == 1 && images.size() == 1); @@ -82,11 +89,11 @@ std::vector ClassificationNet::preProcess(const std::vector RMVL_Assert(shape[1] == 3 || shape[1] == 1); shape[0] = 1; // img -> iarray - RMVL_Assert(!options.means.empty() && !options.stds.empty()); + RMVL_Assert(!means.empty() && !stds.empty()); if (shape[1] == 3) - imageToVector(img, options.means, options.stds, _iarrays.front()); + imageToVector(img, means, stds, _iarrays.front()); else - imageToVector(img, options.means.front(), options.stds.front(), _iarrays.front()); + imageToVector(img, means.front(), stds.front(), _iarrays.front()); // 更新每个输入层的数据 input_tensors.emplace_back(Ort::Value::CreateTensor( _memory_info, _iarrays.front().data(), _iarrays.front().size(), shape.data(), shape.size())); @@ -94,14 +101,14 @@ std::vector ClassificationNet::preProcess(const std::vector return input_tensors; } -std::any ClassificationNet::postProcess(const std::vector &output_tensors, const PostprocessOptions &) +std::pair ClassificationNet::postProcess(const std::vector &output_tensors) { RMVL_Assert(output_tensors.size() == 1); auto &output_tensor = output_tensors.front(); const float *output = output_tensor.GetTensorData(); std::size_t size{output_tensor.GetTensorTypeAndShapeInfo().GetElementCount()}; auto maxit = std::max_element(output, output + size); - return std::make_pair(static_cast(maxit - output), *maxit); + return std::make_pair(static_cast(maxit - output), *maxit); } } // namespace rm diff --git a/modules/ml/src/ort/ort.cpp b/modules/ml/src/ort/ort.cpp index a7dd5429..fc84d429 100755 --- a/modules/ml/src/ort/ort.cpp +++ b/modules/ml/src/ort/ort.cpp @@ -58,15 +58,11 @@ OnnxNet::OnnxNet(std::string_view model_path, OrtProvider prov) : _memory_info(O #endif } -std::vector OnnxNet::preProcess(const std::vector &, const PreprocessOptions &) { return {}; } -std::any OnnxNet::postProcess(const std::vector &, const PostprocessOptions &) { return {}; } - -std::any OnnxNet::inference(const std::vector &images, const PreprocessOptions &preop, const PostprocessOptions &postop) +std::vector OnnxNet::run(const std::vector &input_tensors) { RMVL_Assert(_session != nullptr); - auto itensors = preProcess(images, preop); #if ORT_API_VERSION < 12 - return postProcess(_session->Run(Ort::RunOptions{nullptr}, _inames.data(), itensors.data(), itensors.size(), _onames.data(), _onames.size()), postop); + return _session->Run(Ort::RunOptions{nullptr}, _inames.data(), input_tensors.data(), input_tensors.size(), _onames.data(), _onames.size()); #else std::vector input_names(_inames.size()); for (std::size_t i = 0; i < _inames.size(); i++) @@ -74,7 +70,7 @@ std::any OnnxNet::inference(const std::vector &images, const Preprocess std::vector output_names(_onames.size()); for (std::size_t i = 0; i < _onames.size(); i++) output_names[i] = _onames[i].get(); - return postProcess(_session->Run(Ort::RunOptions{nullptr}, input_names.data(), itensors.data(), itensors.size(), output_names.data(), output_names.size()), postop); + return _session->Run(Ort::RunOptions{nullptr}, input_names.data(), input_tensors.data(), input_tensors.size(), output_names.data(), output_names.size()); #endif }