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

Function gen_inline_addr

compiler/lib/KernelGen/BareMetal/ConvKernel.cpp:77–115  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

75}
76
77std::string gen_inline_addr(std::string format_str, std::string sparse) {
78 std::stringstream ss;
79 if (format_str == "NCHW_NCHW44") {
80 ss << GenFormatIter::gen_inline_format_iter_body("NCHW");
81 ss << GenFormatIter::gen_inline_format_iter_body("NCHW44");
82 } else {
83 ss << GenFormatIter::gen_inline_format_iter_body(format_str);
84 }
85 ss << R"(static inline size_t get_filter_addr_)" << format_str << "_" << sparse;
86 ss << R"((const int group, const int ocpg,
87 const int icpg, const int fh, const int fw,
88 const int* stride) {)";
89 if (format_str == "NCHW") {
90 ss << R"(return (size_t)group * stride[0] + ocpg * stride[1] + icpg * stride[2] +
91 fh * stride[3] + fw * stride[4];)";
92 } else if (format_str == "NCHW_NCHW44") {
93 CC_ASSERT(sparse == "DENSE");
94 ss << R"(return ocpg / 4 * stride[0] + fh * stride[1] + fw * stride[2] + icpg * stride[3] + (ocpg % 4) * stride[4];)";
95 } else if (format_str == "NCHW44") {
96 if (sparse == "DENSE") {
97 ss << R"(return (size_t)group * stride[0] + ocpg / 4 * stride[1] + icpg / 4 * stride[2] +
98 fh * stride[3] + fw * stride[4] + (icpg % 4) * stride[5] + (ocpg % 4) * stride[6];)";
99 } else {
100 CC_ASSERT(sparse == "GROUP") << "spare must be GOURP or DENSE\n";
101 ss << R"(return (size_t)group / 4 * stride[0] + fh * stride[1] + fw * stride[2] + (group % 4) * stride[3];)";
102 }
103 } else {
104 CC_ASSERT(format_str == "NCHW88") << "format not support\n";
105 if (sparse == "DENSE") {
106 ss << R"(return (size_t)group * stride[0] + ocpg / 8 * stride[1] + icpg / 8 * stride[2] +
107 fh * stride[3] + fw * stride[4] + (icpg % 8) * stride[5] + (ocpg % 8) * stride[6];)";
108 } else {
109 CC_ASSERT(sparse == "GROUP") << "spare must be GOURP or DENSE\n";
110 ss << R"(return (size_t)group / 8 * stride[0] + fh * stride[1] + fw * stride[2] + (group % 8) * stride[3];)";
111 }
112 }
113 ss << "}\n";
114 return ss.str();
115}
116
117std::string get_format(TContext* ctx) {
118 auto format_str = ctx->getAttrStr("format");

Callers 1

GetKernelBodyMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected