MCPcopy Create free account
hub / github.com/MegEngine/MegCC / GetKernelBody

Method GetKernelBody

compiler/lib/KernelGen/BareMetal/IndexingOneHot.cpp:26–69  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24}
25
26std::string IndexingOneHotKernel::GetKernelBody(TContext* context) const {
27 std::stringstream axis_init_ss;
28 std::stringstream writer;
29 writer << "#include <string.h>\n";
30 writer << GenCommonRet() << " ";
31 writer << GetKernelSignature(context) << "{\n";
32 int axis = context->getAttrInt("axis");
33 // clang-format off
34 writer << StringTemplate::StringTemplateArgs(context).add("axis", axis).render(
35 R"(
36 float* src = (float*)inputs[0]->ptr;
37 int* idx = (int*)inputs[1]->ptr;
38 float* dst = (float*)outputs[0]->ptr;
39
40 const Tensor* src_tensor = inputs[0];
41 const Tensor* dst_tensor = outputs[0];
42 const Layout src_layout = src_tensor->layout;
43 const Layout dst_layout = dst_tensor->layout;
44
45 int axis = ${axis};
46 int batch = 1;
47 int elems = 1;
48 for (int i = 0; i < axis; ++i){
49 batch *= src_layout.dims[i];
50 }
51 for (int i = axis + 1; i < src_layout.nr_dim; ++i){
52 elems *= src_layout.dims[i];
53 }
54 int batch_stride = src_layout.dims[axis] * elems;
55 for (int bid = 0; bid < batch; ++bid){
56 float* src_ptr = src + bid * batch_stride;
57 for(int id = 0; id < elems; ++id){
58 *dst = src_ptr[*idx * elems + id];
59 ++dst;
60 ++idx;
61 }
62 }
63
64 return TinyNN_SUCCESS;
65 })"
66 );
67 // clang-format on
68 return writer.str();
69}
70
71} // namespace BareMetal
72} // namespace KernelGen

Callers

nothing calls this directly

Calls 4

GenCommonRetFunction · 0.85
StringTemplateArgsClass · 0.85
renderMethod · 0.45
addMethod · 0.45

Tested by

no test coverage detected