pnna2/include/preprocess.h

247 lines
7.7 KiB
C
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/**
* @file preprocess.h
* @author @xun (xunyingya@hngwg.com)
* @brief 对输入张量执行统一的归一化与量化前处理
* @version 2.3.0
* @date 2025-08-28 16:42
*
* Copyright (c) 2025 by CCYH, All Rights Reserved.
*
*/
#ifndef _PREPROCESS_H_
#define _PREPROCESS_H_
#ifdef __cplusplus
extern "C" {
#endif
#include "quantize.h"
/**
* @brief 归一化参数
*
*/
typedef struct _normalization_params
{
int mean[3];
float scale;
} normalization_params_t;
/**
* @brief 设置归一化参数
*
* @param norm_params [OUT] 将参数赋值给该结构体
* @param mean [IN] 数组,每个通道的平均值
* @param scale [IN] 缩放因子
*/
void set_normalize_params(normalization_params_t *norm_params, \
int *mean, float scale);
/**
* @brief 对单个 uint8 输入执行归一化
*
* @param input [IN] 输入像素值
* @param mean [IN] 均值
* @param scale [IN] 缩放因子
* @return float 归一化后的浮点值
*/
float normalize_u8(uint8_t input, int mean, float scale);
/**
* @brief 归一化和量化接口,包含多输入输出,将输入数据归一化后量化
*
* @param output [OUT] 量化后的输出 buffer_t 数组
* @param input [IN] 原始输入 buffer_t 数组
* @param norm_params [IN] 归一化参数数组
* @param quant_params [IN] 量化参数数组
* @param counts [IN] 输入/输出 buffer_t 数组的数量
* @return int 处理后的状态码0表示成功其他表示失败
*/
int normalize_quantize(buffer_t *output, buffer_t *input,
normalization_params_t *norm_params,
pnna_buffer_param_t *quant_params,
uint32_t counts);
/**
* @brief 量化接口,包含多输入输出,将输入数据量化为指定格式
*
* @param output [OUT] 量化后的输出 buffer_t 数组
* @param input [IN] 原始输入 buffer_t 数组
* @param quant_params [IN] 量化参数数组
* @param counts [IN] 网络模型输入tensor的数量
* @return int 处理后的状态码0表示成功其他表示失败
*/
int quantize(buffer_t *output, buffer_t *input,
pnna_buffer_param_t *quant_params,
uint32_t counts);
/**
* @brief 根据量化参数数组创建对应的 buffer_t 数组
*
* 遍历每个 `pnna_buffer_param_t` 参数,计算其数据大小,
* 并为每个 buffer_t 中的 data 分配内存
*
* 该函数适用于根据模型的输入/输出量化信息自动构建 buffer_t 数组
*
* @param qp [IN] 量化参数数组,包含维度、格式等信息
* @param count [IN] buffer_t 数组数量
*
* @return buffer_t*
* - 成功返回指向 buffer_t 数组的指针
* - 失败返回 NULL
*/
buffer_t *create_buffer_from_qparams(pnna_buffer_param_t *qp, uint32_t count);
/**
* @brief 释放由 create_buffer_from_qparams 或其他形式创建的 buffer_t 数组
*
* @param buff_array [IN] 指向 buffer_t 数组的指针
* @param count [IN] 数组元素个数,必须与创建时保持一致
*/
void destroy_buffer_array(buffer_t *buff_array, uint32_t count);
#if defined(__linux__) || defined(PLATFORM_ARM_RTT)
/**
* @brief 从指定文件名加载二进制数据到缓冲区结构体
*
* @param name [IN] 要读取的二进制文件名(含路径)
* @return buffer_t
* - 成功时: 返回包含文件内容的结构体指针
* - 失败时: 返回NULL可能原因参数为空、内存分配失败、文件加载失败
*
* @note 内存管理规则:
* - 返回的 buffer_t 及其内部 data 指针必须通过 free_buffer() 统一释放
* - 若文件加载失败,函数内部会自动清理已分配内存,避免内存泄漏
*
* @warning 调用者必须检查返回值有效性:
* @code
* buffer_t *buf = load_binary_to_buffer("network_binary.nb");
* if (!buf) {
* // 错误处理
* }
* @endcode
*
* @see free_buffer(), load_binary_file()
*/
buffer_t load_binary_to_buffer(const char *name);
/**
* @brief 安全释放由 load_binary_to_buffer() 等函数创建的 buffer_t 缓冲区资源
*
* @param buffer [IN] 要释放的缓冲区指针
*
*/
void free_buffer(buffer_t *buffer);
/**
* @brief 从文件路径加载图像数据并封装为内存缓冲区
*
* @param filename [IN] 输入图像文件路径(支持常见格式如 JPEG/PNG/BMP
* @return buffer_t
* - 成功时: 返回包含图像数据的缓冲区
* - 失败时: 返回 { .data = NULL, .size = 0 }
*
* @note 数据存储规则:
* - 图像数据存储在连续内存中,布局为 HWC高度、宽度、通道格式
* - 张量维度固定为 4 维,形状为 [1, height, width, channels]
* - 数据类型固定为 8 位无符号整型,像素值范围 [0, 255]
*
* @par 张量结构示例:
* | 维度索引 | 0 | 1 (h) | 2 (w) | 3 (c) |
* |----------|-----|-------|-------|-------|
* | 值 | 1 | 480 | 640 | 3 |
*
* @warning 注意事项:
* - 调用者必须使用 `free_image_buffer()` 释放返回的内存
* - 图像通道数由文件决定RGB 图像为 3 通道RGBA 为 4 通道)
* - 底层依赖 stb_image 库,需确保链接正确
*
*/
buffer_t load_image(const char *filename);
/**
* @brief 从内存缓冲区加载图像数据并封装为数据缓冲区
*
* @param buffer [IN] 图像文件的二进制数据指针(必须为完整的图像文件内容)
* @param len [IN] 缓冲区长度(字节数)
*
* @return buffer_t
* - 成功时: 返回包含解码后图像数据的缓冲区(.data != NULL, .size > 0
* - 失败时: 返回 { .data = NULL, .size = 0 }(可能原因:参数无效、内存不足、解码失败)
*
* @note 数据存储规则:
* - 图像数据存储在连续内存中,布局为 HWC高度 × 宽度 × 通道)格式
* - 通道数由输入图像决定(如 RGB=3RGBA=4
* - 数据类型固定为 8 位无符号整型uint8_t像素值范围 [0, 255]
*
* @warning 注意事项:
* - 返回的缓冲区中的 data 内存由 stb_image 分配,调用者必须使用 `free_image_buffer()` 释放
* - 本函数不会对图像进行缩放或归一化,仅解码原始像素
* - 输入数据必须是完整的图像文件内容,否则解码将失败
*
* @code
* // 示例用法
* buffer_t img = load_image_from_memory(bin_data, bin_len);
* if (img.data) {
* printf("Image buffer size: %u bytes\n", img.size);
* free_image_buffer(img); // 必须释放
* }
* @endcode
*/
buffer_t load_image_from_memory(const uint8_t *buffer, int len);
/**
* @brief 释放图像缓冲区内存
*
* @param buffer [IN] 待释放的图像缓冲区(由 stbi_load/stbi_load_from_memory 分配)
*
* @note
* - 仅释放 buffer.data 指向的内存,不负责释放 buffer 本身
* - 调用后 buffer.data 被置为 NULL局部生效外部调用者需注意
*/
void free_image_buffer(buffer_t *buffer);
/**
* @brief 将图像数据从 HWC 格式转换为 CHW 格式
*
* @param dst [OUT] 输出缓冲区,必须预先分配足够内存
* @param src [IN] 输入缓冲区HWC 格式的图像数据
* @param h [IN] 图像高度
* @param w [IN] 图像宽度
* @param c [IN] 图像通道数通常为3或4
*/
void hwc2chw_neon(uint8_t *dst, const uint8_t *src, int h, int w, int c);
/**
* @brief 将图像数据从 HWC 格式转换为 CHW 格式
*
* @param filename [IN] 输入图像文件路径
* @return buffer_t
* - 成功时: 返回包含图像数据的缓冲区
* - 失败时: 返回 { .data = NULL, .size = 0 }
*/
buffer_t load_image_chw(const char *filename);
#endif
#ifdef __cplusplus
}
#endif
#endif // !_PREPROCESS_H_