270 lines
12 KiB
ReStructuredText
270 lines
12 KiB
ReStructuredText
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;
|
||
}
|
||
|