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 节点。
+
+**用法**
+
+
+
+@dl_begin{选项}
+@dl_item{\-c,仅输出节点数量}
+@dl_end
+
+**示例**
+
+
+
+
lpss node list
+
+
lpss node list
+
+
#### info 子命令
查看指定节点的信息,输出形如以下的内容
@@ -102,6 +126,12 @@ Publish Topics:
Subscribe Topics:
xxx
+
+Server Services:
+ xxx
+
+Client Services:
+ xxx
```
**用法**
@@ -126,13 +156,14 @@ Subscribe Topics:
**用法**
-
lpss topic [args...]
+
lpss topic [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 子命令
+
+列出当前发现的所有话题,话题名称按字典序输出。
+
+**用法**
+
+
+
+@dl_begin{选项}
+@dl_item{\-c,仅输出话题数量}
+@dl_end
+
+#### find 子命令
+
+按完整消息类型查找话题,查询结果按话题名称的字典序输出。
+
+**用法**
+
+
+
+@param msg_type 消息类型,例如 `std/String`
+
+@dl_begin{选项}
+@dl_item{\-c,仅输出匹配的话题数量}
+@dl_end
+
+**示例**
+
+
+
+
lpss topic find std/String
+
+
lpss topic find std/String
+
+
#### info 子命令
查看指定话题的信息,输出形如以下的内容
@@ -259,6 +329,146 @@ Subscriber Node:
lpss topic bw /str
+### service 命令
+
+服务工具,用于查看已发现的 LPSS 服务,并使用 JSON 请求调用内置服务。
+
+**用法**
+
+
+
lpss service [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 子命令
+
+列出当前发现的所有服务,服务名称按字典序输出。
+
+**用法**
+
+
+
+@dl_begin{选项}
+@dl_item{\-c,仅输出服务数量}
+@dl_end
+
+#### info 子命令
+
+查看指定服务的类型、服务端节点和客户端节点,输出形如以下内容。
+
+```
+Type: std/SetBool
+
+Server Node:
+ server_node
+
+Client Node:
+ client_node
+```
+
+**用法**
+
+
+
+@param service_name 服务名称
+
+**示例**
+
+
+
lpss service info /set_enabled
+
+
+#### type 子命令
+
+显示指定服务的服务类型。
+
+**用法**
+
+
+
+@param service_name 服务名称
+
+**示例**
+
+
+
+
lpss service type /set_enabled
+
+
+#### find 子命令
+
+按类型精确匹配服务。`service_type` 可以是完整服务类型,也可以是对应的请求或响应消息类型。查询结果按服务名称的字典序输出。
+
+**用法**
+
+
+
+@param service_type 服务类型,例如 `std/SetBool`;也可使用 `std/SetBool_Request` 或 `std/SetBool_Response`
+
+@dl_begin{选项}
+@dl_item{\-c,仅输出匹配的服务数量}
+@dl_end
+
+**示例**
+
+
+
+
lpss service find std/SetBool
+
+
lpss service find std/SetBool
+
+
+#### call 子命令
+
+使用 JSON 对象作为请求调用指定服务,并将响应输出为 JSON。`json_request` 省略或为空时按 `{}` 处理,命令等待响应的超时时间为 3 秒。
+
+**用法**
+
+
+
+@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 个数字 |
+
+**示例**
+
+
+
+
lpss service call /trigger
+
+
lpss service call /set_enabled '{"data":true}'
+
+
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 []
+
+
lgraph []
+
+
+@param name_subfix 可选的实例名称,用于生成 LPSS 节点名 `lgraph_`;省略时自动生成 5 位随机标识
+
+启动成功后,终端会输出本机和局域网访问地址。LGraph 默认监听 `17493` 端口,本机可访问:
+
+
+
http://localhost:17493
+
+
+**示例**
+
+
+
+
lpss graph
+
+
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 []
+
+
lviz []
+
+
+@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 - | 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
}