首先将固定的转换rknn代码迁移到自己的代码中,mean_values根据训练代码是否归一化,进行自我配置,去Twinlitenet代码里面没有看到归一化的代码。配置如下:

class RKNNModel:
def __init__(self, onnx_file, rknn_file):
    # 创建 RKNN 对象
    self.rknn = RKNN(verbose=True)
    self.input_width = 640  # 模型输入宽度
    self.input_height = 360  # 模型输入高度
    self.rknn_file = rknn_file
    print('Configuring RKNN...')
    self.rknn.config(mean_values=[[0,0,0]], std_values=[[1,1,1]],quantized_algorithm='normal',quantized_method='channel',target_platform='rk3588'
                     )

    # 加载 ONNX 模型
    print('Loading ONNX model...')
    ret = self.rknn.load_onnx(model=onnx_file)
    if ret != 0:
        print('Failed to load ONNX model!')
        exit()

    # 编译 RKNN 模型
    print('Compiling RKNN model...')
    ret = self.rknn.build(do_quantization=QUANTIZE_ON, dataset='./image_list.txt', rknn_batch_size=1)
    if ret != 0:
        print('Failed to compile RKNN model!')
        exit()

    # 导出 RKNN 模型
    self.rknn.export_rknn(rknn_file)
    print(f'RKNN model saved to {rknn_file}')

    # 初始化运行时环境
    print('-->Initializing runtime environment...')
    ret = self.rknn.init_runtime()
    if ret != 0:
        print('Init runtime environment failed')
        exit()
    print('done')

对代码的图片后处理,代码如下,需要修改的部分将之前onnx获得输出的模式改为rknn获得输出的模式,输出两个分割区域可行驶区域da和车道线lanes

    def forward(self, img):
        start_time = time.time()
        # 复制原图像,防止修改原图
        img_ = img.copy()
        img=cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
        # 调整图像大小为模型输入大小
        img = cv2.resize(img, (self.input_width, self.input_height))
        img_rs = img.copy()  # 复制调整后的图像,后续用于绘制结果
        # 图像归一化,将像素缩放到 [0, 1] 之间
        img = img.astype(np.float32) / 255.0
        # 为了符合模型输入要求,调整维度并转置
        #img = np.transpose(np.float32(img[:, :, :, np.newaxis]), (3, 2, 0, 1))  # 新增维度并调整顺序
        img=np.expand_dims(img,axis=0)
        img = np.ascontiguousarray(img)  # 确保内存连续,避免内存分配问题
        outputs = self.rknn.inference(inputs=[img])
        da = outputs[0]
        lanes = outputs[1]
        da = np.argmax(da, 1)  # 提取驾驶区域
        lanes = np.argmax(lanes, 1)  # 提取车道线

        # 2. 转换为 uint8 类型并进行缩放
        da = da.astype('uint8')  # 将类别索引转为 uint8 格式
        da = da[0] * 255  # 将类别映射到 0-255 范围

        lanes = lanes.astype('uint8')
        lanes = lanes[0] * 255  # 同样将车道线类别映射到 0-255 范围

        # 3. 绘制掩码区域
        img_rs[da > 100] = [255, 0, 0]  # 红色区域为驾驶区域
        img_rs[lanes > 100] = [0, 255, 0]  # 绿色区域为车道线部分
        elapsed_time = time.time() - start_time  # 计算推理所用的时间
        # 显示推理时间
        cv2.putText(
            img_rs,
            f"Elapsed Time: {elapsed_time * 1000:.1f} ms",
            (10, 30),
            cv2.FONT_HERSHEY_SIMPLEX,
            0.8,
            (0, 255, 0),  # 绿色字体
            2,
            cv2.LINE_AA,  # 平滑线条
        )
        return img_rs  # 返回处理后的图像

完整视频推理代码如下,指定自己的onnx路径以及生成的rknn路径,就可以在PC端进行RKNN视频推理

def parse_opt(known=False):
    parser = argparse.ArgumentParser()
    parser.add_argument('--onnx', type=str, required=True, help='ONNX model file path')
    parser.add_argument('--output', type=str, required=True, help='Output RKNN model file path')
    parser.add_argument('--video', type=str, required=True, help='Image file for inference')
    parser.add_argument('--dataset', type=str, default='./image_list.txt', help='Dataset file for quantization (optional)')
    parser.add_argument('--quantize', action='store_true', help='Enable quantization')
    opt = parser.parse_known_args()[0] if known else parser.parse_args()
    return opt

def main(opt):
    # 转换 ONNX 为 RKNN
    onnx_file = opt.onnx
    rknn_file = opt.output
    video_path = opt.video
    model = RKNNModel(onnx_file, rknn_file)
    #img = cv2.imread(video_path)
    video_capture=cv2.VideoCapture(video_path)
    while True:
        ret,frame=video_capture.read()
        if not ret:
            break
        result = model.forward(frame)

        # 显示并保存结果
        cv2.imshow("Segmentation Result", result)
        key=cv2.waitKey(1)
        if key==27:
            break
    
    video_capture.release()
    cv2.destroyAllWindows()

if name == ‘main’:
opt = parse_opt()
main(opt)

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐