| 42 | } |
| 43 | |
| 44 | std::string IndexingMultiAxisKernel::GetKernelBody(TContext* context) const { |
| 45 | std::stringstream axis_init_ss; |
| 46 | int nr_operand = context->getAttrInt("nr_operands"); |
| 47 | for (int i = 0; i < nr_operand - 2; ++i) { |
| 48 | axis_init_ss << "axis_vec[" << i |
| 49 | << "] = " << context->getAttrInt("axis:" + std::to_string(i)) |
| 50 | << ";\n"; |
| 51 | } |
| 52 | std::string dtype_specifier = Utils::cvt_dtype_specifier( |
| 53 | SymbolHelper::gen_valid_dtype(Utils::get_last_operand(context).dtype)); |
| 54 | std::stringstream writer; |
| 55 | writer << R"( |
| 56 | #include "tensor_util.h" |
| 57 | )"; |
| 58 | writer << "#include <string.h>\n"; |
| 59 | if (dtype_specifier == "gi_float16_t") |
| 60 | writer << gen_fp16_define(); |
| 61 | writer << GenCommonRet() << " "; |
| 62 | writer << GetKernelSignature(context) << "{\n"; |
| 63 | // clang-format off |
| 64 | writer << StringTemplate::StringTemplateArgs(context).add("axis_init_str",axis_init_ss.str()).add("dtype_specifier", dtype_specifier).render( |
| 65 | R"( |
| 66 | ${dtype_specifier}* src = (${dtype_specifier}*)inputs[0]->ptr; |
| 67 | ${dtype_specifier}* dst = (${dtype_specifier}*)outputs[0]->ptr; |
| 68 | |
| 69 | const Tensor* src_tensor = inputs[0]; |
| 70 | const Tensor* dst_tensor = outputs[0]; |
| 71 | const Layout src_layout = src_tensor->layout; |
| 72 | const Layout dst_layout = dst_tensor->layout; |
| 73 | int nr_index = nr_input - 1; |
| 74 | const Tensor* idx_tensors[7]; |
| 75 | int axis_vec[7]; |
| 76 | ${axis_init_str} |
| 77 | |
| 78 | for (int i = 0; i < nr_index; ++i) { |
| 79 | idx_tensors[i] = inputs[i + 1]; |
| 80 | } |
| 81 | // compute idx_axis start |
| 82 | size_t idx_axis = 0; |
| 83 | { |
| 84 | int contig_idx = 1; |
| 85 | for (size_t i = 1; i < nr_index; ++i) { |
| 86 | if (axis_vec[i] != axis_vec[i - 1] + 1) { |
| 87 | contig_idx = 0; |
| 88 | break; |
| 89 | } |
| 90 | } |
| 91 | if (contig_idx) { |
| 92 | idx_axis = axis_vec[0]; |
| 93 | } |
| 94 | } |
| 95 | // compute idx_axis end |
| 96 | |
| 97 | // compute nonidx axes start |
| 98 | size_t nonidx_axes[7], nr_nonidx_axes = 0; |
| 99 | { |
| 100 | size_t idx = 0; |
| 101 | for (size_t i = 0; i < src_layout.nr_dim; ++i) { |
nothing calls this directly
no test coverage detected