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

270 lines
12 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类型。决定选择哪个输入的值。
- **params** - 其他参数打包成数组。
- **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, long long *params, int core_mask)
.. c:function:: void i16_select_s(int16_t* input0, int16_t* input1, bool* condition, int16_t* output, long long *params, int core_mask)
.. c:function:: void i32_select_s(int32_t* input0, int32_t* input1, bool* condition, int32_t* output, long long *params, int core_mask)
.. c:function:: void hp_select_s(half* input0, half* input1, bool* condition, half* output, long long *params, int core_mask)
.. c:function:: void fp_select_s(float* input0, float* input1, bool* condition, float* output, long long *params, int core_mask)
.. c:function:: void dp_select_s(double* input0, double* input1, bool* condition, double* output, long long *params, int core_mask)
.. c:function:: void c64_select_s(float* input0, float* input1, bool* condition, float* output, long long *params, int core_mask)
.. c:function:: void c128_select_s(double* input0, double* input1, bool* condition, double* output, long long *params, 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[]) {
float *input0 = (float *)0x81000000;
float *input1 = (float *)0x82000000;
bool *condition = (bool *)0x83000000;
float *output = (float *)0x84000000;
float *checkoutput = (float *)0x85000000;
unsigned long long *index_list1 = (unsigned long long *)0x86000000;
unsigned long long *index_list2 = (unsigned long long *)0x87000000;
unsigned long long *index_list3 = (unsigned long long *)0x87800000;
long long is_broadcast = 0;
unsigned long long *input0_dims = global_input0_dims;
unsigned long long *input1_dims = global_input1_dims;
unsigned long long *cond_dims= global_cond_dims;
unsigned long long *output_dims = global_output_dims;
unsigned long long input0_dims_num = global_input0_dims_num;
unsigned long long input1_dims_num = global_input1_dims_num;
unsigned long long cond_dims_num = global_cond_dims_num;
unsigned long long output_dims_num = global_output_dims_num;
unsigned long long params[10];
params[0] = (unsigned long long)output_dims;
params[2] = (unsigned long long)index_list1;
params[3] = (unsigned long long)index_list2;
params[4] = (unsigned long long)index_list3;
//先计算is_broadcast
unsigned long long input0_num = get_total_elements(input0_dims_num, input0_dims);
unsigned long long input1_num = get_total_elements(input1_dims_num, input1_dims);
unsigned long long cond_num = get_total_elements(cond_dims_num, cond_dims);
unsigned long long output_num = get_total_elements(output_dims_num, output_dims);
if((input0_num == output_num && input1_num == output_num) && (cond_num == output_num)) {
is_broadcast = 0;
} else {
is_broadcast = 1;
}
params[1] = (unsigned long long)output_dims_num;
params[5] = (unsigned long long)is_broadcast;
srand(seed++);
int i;
//初始化input0, input1, condition
for (i = 0; i < input0_num; ++i) {
input0[i] = (float)(rand() % 100) / 10.0f;
}
for (i = 0; i < input1_num; ++i) {
input1[i] = (float)(rand() % 100) / 10.0f;
}
for (i = 0; i < cond_num; ++i) {
condition[i] = (bool)(rand() % 2);
}
int core_mask = 0x0f;
if(is_broadcast) {
GetBroadCastIndex(cond_dims, cond_dims_num, output_dims, output_dims_num, index_list1);
GetBroadCastIndex(input0_dims, input0_dims_num, output_dims, output_dims_num, index_list2);
GetBroadCastIndex(input1_dims, input1_dims_num, output_dims, output_dims_num, index_list3);
fp_select_s(input0, input1, condition, output, params, core_mask);
} else {
fp_select_s(input0, input1, condition, output, params, core_mask);
}
return 0;
}
**私有存储版本:**
.. c:function:: void i8_select_p(int8_t* input0, int8_t* input1, bool* condition, int8_t* output, long long *params)
.. c:function:: void i16_select_p(int16_t* input0, int16_t* input1, bool* condition, int16_t* output, long long *params)
.. c:function:: void i32_select_p(int32_t* input0, int32_t* input1, bool* condition, int32_t* output, long long *params)
.. c:function:: void hp_select_p(half* input0, half* input1, bool* condition, half* output, long long *params)
.. c:function:: void fp_select_p(float* input0, float* input1, bool* condition, float* output, long long *params)
.. c:function:: void dp_select_p(double* input0, double* input1, bool* condition, double* output, long long *params)
.. c:function:: void c64_select_p(float* input0, float* input1, bool* condition, float* output, long long *params)
.. c:function:: void c128_select_p(double* input0, double* input1, bool* condition, double* output, long long *params)
**C调用示例私有存储版本**
.. code-block:: c
:linenos:
:emphasize-lines: 30-31
//FT78NE示例
#include <stdio.h>
#include <select.h>
int main(int argc, char* argv[]) {
float *input0 = (float *)0x10010000;
float *input1 = (float *)0x10020000;
bool *condition = (bool *)0x10030000;
float *output = (float *)0x10040000;
float *checkoutput = (float *)0x10050000;
unsigned long long *index_list1 = (unsigned long long *)0x10060000;
unsigned long long *index_list2 = (unsigned long long *)0x10070000;
unsigned long long *index_list3 = (unsigned long long *)0x10078000;
long long is_broadcast = 0;
unsigned long long *input0_dims = global_input0_dims;
unsigned long long *input1_dims = global_input1_dims;
unsigned long long *cond_dims= global_cond_dims;
unsigned long long *output_dims = global_output_dims;
unsigned long long input0_dims_num = global_input0_dims_num;
unsigned long long input1_dims_num = global_input1_dims_num;
unsigned long long cond_dims_num = global_cond_dims_num;
unsigned long long output_dims_num = global_output_dims_num;
unsigned long long params[10];
params[0] = (unsigned long long)output_dims;
params[2] = (unsigned long long)index_list1;
params[3] = (unsigned long long)index_list2;
params[4] = (unsigned long long)index_list3;
//先计算is_broadcast
unsigned long long input0_num = get_total_elements(input0_dims_num, input0_dims);
unsigned long long input1_num = get_total_elements(input1_dims_num, input1_dims);
unsigned long long cond_num = get_total_elements(cond_dims_num, cond_dims);
unsigned long long output_num = get_total_elements(output_dims_num, output_dims);
if((input0_num == output_num && input1_num == output_num) && (cond_num == output_num)) {
is_broadcast = 0;
} else {
is_broadcast = 1;
}
params[1] = (unsigned long long)output_dims_num;
params[5] = (unsigned long long)is_broadcast;
srand(seed++);
int i;
//初始化input0, input1, condition
for (i = 0; i < input0_num; ++i) {
input0[i] = (float)(rand() % 100) / 10.0f;
}
for (i = 0; i < input1_num; ++i) {
input1[i] = (float)(rand() % 100) / 10.0f;
}
for (i = 0; i < cond_num; ++i) {
condition[i] = (bool)(rand() % 2);
}
if(is_broadcast) {
GetBroadCastIndex(cond_dims, cond_dims_num, output_dims, output_dims_num, index_list1);
GetBroadCastIndex(input0_dims, input0_dims_num, output_dims, output_dims_num, index_list2);
GetBroadCastIndex(input1_dims, input1_dims_num, output_dims, output_dims_num, index_list3);
fp_select_p(input0, input1, condition, output, params);
} else {
fp_select_p(input0, input1, condition, output, params);
}
return 0;
}