| 75 | } |
| 76 | |
| 77 | std::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 | |
| 117 | std::string get_format(TContext* ctx) { |
| 118 | auto format_str = ctx->getAttrStr("format"); |