FT_DSP.gitlink.net/master/html/_sources/functionlib/dsplib/select.rst.txt

189 lines
11 KiB
ReStructuredText
Raw 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.

Select
=================
根据条件张量逐元素选择输入值。对于每个输出位置如果条件为真True则选择 `input0` 的值;否则选择 `input1` 的值。该算子支持广播机制。
.. math::
\text{output}_i = \begin{cases}
\text{input0}[idx2], & \text{if } \text{condition}[idx1] = \text{True} \\
\text{input1}[idx3], & \text{if } \text{condition}[idx1] = \text{False}
\end{cases}
其中,当不需要广播时(`is_broadcast = 0``idx1 = idx2 = idx3 = i`;当需要广播时(`is_broadcast = 1`),使用索引映射 `index_list1``index_list2``index_list3` 来确定各个输入张量的索引。
输入:
- **input0** - 第一个输入数据地址。当条件为真时选择此值。
- **input1** - 第二个输入数据地址。当条件为假时选择此值。
- **condition** - 条件数据地址bool类型。决定选择哪个输入的值。
- **output_dims** - 输出张量的维度信息数组。
- **output_dims_num** - 输出张量的维度数。
- **index_list1** - 条件张量的索引映射数组,用于广播场景。大小为输出总元素数。
- **index_list2** - input0 的索引映射数组,用于广播场景。大小为输出总元素数。
- **index_list3** - input1 的索引映射数组,用于广播场景。大小为输出总元素数。
- **is_broadcast** - 是否需要广播的标志。0 表示不需要广播1 表示需要广播。
- **core_mask** - 核掩码(仅共享存储版本需要)。
输出:
- **output** - 输出数据地址,其形状由 `output_dims``output_dims_num` 确定。
支持平台:
``FT78NE``
``MT7004``
.. note::
- FT78NE 支持fp32, int8, int16, int32, fp64, cplx64, cplx128
- MT7004 支持fp16, fp32, int16, int32, cplx64
**共享存储版本:**
.. c:function:: void i8_select_s(int8_t* input0, int8_t* input1, bool* condition, int8_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void i16_select_s(int16_t* input0, int16_t* input1, bool* condition, int16_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void i32_select_s(int32_t* input0, int32_t* input1, bool* condition, int32_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void hp_select_s(half* input0, half* input1, bool* condition, half* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void fp_select_s(float* input0, float* input1, bool* condition, float* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void dp_select_s(double* input0, double* input1, bool* condition, double* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void c64_select_s(float* input0, float* input1, bool* condition, float* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
.. c:function:: void c128_select_s(double* input0, double* input1, bool* condition, double* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast, int core_mask)
**C调用示例无广播**
.. code-block:: c
:linenos:
:emphasize-lines: 34-35
//FT78NE示例
#include <stdio.h>
#include <select.h>
int main(int argc, char* argv[]) {
// 假设在DDR空间
float *input0 = (float *)0xA0000000;
float *input1 = (float *)0xA1000000;
bool *condition = (bool *)0xA2000000;
float *output = (float *)0xB0000000;
// 输出形状 [2, 3, 4]
unsigned long long output_dims[] = {2, 3, 4};
unsigned long long output_dims_num = 3;
// 计算总元素数
unsigned long long total_elements = 2 * 3 * 4; // 24
// 索引映射数组无广播时可以为NULL或与输出索引相同
unsigned long long *index_list1 = (unsigned long long *)0xC0000000;
unsigned long long *index_list2 = (unsigned long long *)0xC0100000;
unsigned long long *index_list3 = (unsigned long long *)0xC0200000;
// 初始化索引映射(无广播时直接使用顺序索引)
for (unsigned long long i = 0; i < total_elements; i++) {
index_list1[i] = i;
index_list2[i] = i;
index_list3[i] = i;
}
long long is_broadcast = 0; // 不需要广播
int core_mask = 0xff;
fp_select_s(input0, input1, condition, output, output_dims, output_dims_num,
index_list1, index_list2, index_list3, is_broadcast, core_mask);
return 0;
}
**C调用示例有广播**
.. code-block:: c
:linenos:
:emphasize-lines: 35-36
//FT78NE示例
#include <stdio.h>
#include <select.h>
int main(int argc, char* argv[]) {
// 假设在DDR空间
float *input0 = (float *)0xA0000000; // 形状 [3, 4]
float *input1 = (float *)0xA1000000; // 标量或形状 [1]
bool *condition = (bool *)0xA2000000; // 形状 [2, 3, 4]
float *output = (float *)0xB0000000; // 形状 [2, 3, 4]
// 输出形状 [2, 3, 4]
unsigned long long output_dims[] = {2, 3, 4};
unsigned long long output_dims_num = 3;
// 计算总元素数
unsigned long long total_elements = 2 * 3 * 4; // 24
// 索引映射数组(需要根据广播规则预先计算)
unsigned long long *index_list1 = (unsigned long long *)0xC0000000;
unsigned long long *index_list2 = (unsigned long long *)0xC0100000;
unsigned long long *index_list3 = (unsigned long long *)0xC0200000;
// 初始化索引映射(示例:需要根据实际广播规则计算)
// 这里假设 condition 和 output 形状相同input0 需要广播
for (unsigned long long i = 0; i < total_elements; i++) {
index_list1[i] = i; // condition 索引
index_list2[i] = i % 12; // input0 索引(假设需要广播)
index_list3[i] = 0; // input1 索引(标量)
}
long long is_broadcast = 1; // 需要广播
int core_mask = 0xff;
fp_select_s(input0, input1, condition, output, output_dims, output_dims_num,
index_list1, index_list2, index_list3, is_broadcast, core_mask);
return 0;
}
**私有存储版本:**
.. c:function:: void i8_select_p(int8_t* input0, int8_t* input1, bool* condition, int8_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void i16_select_p(int16_t* input0, int16_t* input1, bool* condition, int16_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void i32_select_p(int32_t* input0, int32_t* input1, bool* condition, int32_t* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void hp_select_p(half* input0, half* input1, bool* condition, half* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void fp_select_p(float* input0, float* input1, bool* condition, float* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void dp_select_p(double* input0, double* input1, bool* condition, double* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void c64_select_p(float* input0, float* input1, bool* condition, float* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
.. c:function:: void c128_select_p(double* input0, double* input1, bool* condition, double* output, unsigned long long* output_dims, unsigned long long output_dims_num, unsigned long long* index_list1, unsigned long long* index_list2, unsigned long long* index_list3, long long is_broadcast)
**C调用示例私有存储版本**
.. code-block:: c
:linenos:
:emphasize-lines: 30-31
//FT78NE示例
#include <stdio.h>
#include <select.h>
int main(int argc, char* argv[]) {
// 假设在L2空间
float *input0 = (float *)0x10000000;
float *input1 = (float *)0x10001000;
bool *condition = (bool *)0x10002000;
float *output = (float *)0x10003000;
// 输出形状 [2, 3, 4]
unsigned long long output_dims[] = {2, 3, 4};
unsigned long long output_dims_num = 3;
unsigned long long total_elements = 2 * 3 * 4;
unsigned long long *index_list1 = (unsigned long long *)0x10004000;
unsigned long long *index_list2 = (unsigned long long *)0x10005000;
unsigned long long *index_list3 = (unsigned long long *)0x10006000;
// 初始化索引映射(无广播)
for (unsigned long long i = 0; i < total_elements; i++) {
index_list1[i] = i;
index_list2[i] = i;
index_list3[i] = i;
}
long long is_broadcast = 0;
fp_select_p(input0, input1, condition, output, output_dims, output_dims_num,
index_list1, index_list2, index_list3, is_broadcast);
return 0;
}