| 30 | return ss.str(); |
| 31 | } |
| 32 | std::string gen_inline_addr(std::string format_str, std::string sparse) { |
| 33 | std::stringstream ss; |
| 34 | ss << GenFormatIter::gen_inline_format_iter_body(format_str); |
| 35 | ss << R"(static inline size_t get_filter_addr_)" << format_str << "_" << sparse; |
| 36 | ss << R"((const int group, const int ocpg, |
| 37 | const int icpg, const int fh, const int fw, |
| 38 | const int* stride) {)"; |
| 39 | if (format_str == "NCHW") { |
| 40 | ss << R"(return (size_t)group * stride[0] + ocpg * stride[1] + icpg * stride[2] + |
| 41 | fh * stride[3] + fw * stride[4];)"; |
| 42 | } else { |
| 43 | CC_ASSERT(format_str == "NCHW44") << "format not support\n"; |
| 44 | if (sparse == "DENSE") { |
| 45 | ss << R"(return (size_t)group * stride[0] + ocpg / 4 * stride[1] + icpg / 4 * stride[2] + |
| 46 | fh * stride[3] + fw * stride[4] + (icpg % 4) * stride[5] + (ocpg % 4) * stride[6];)"; |
| 47 | } else { |
| 48 | CC_ASSERT(sparse == "GROUP") << "spare must be GOURP or DENSE\n"; |
| 49 | ss << R"(return (size_t)group / 4 * stride[0] + fh * stride[1] + fw * stride[2] + (group % 4) * stride[3];)"; |
| 50 | } |
| 51 | } |
| 52 | ss << "}\n"; |
| 53 | return ss.str(); |
| 54 | } |
| 55 | |
| 56 | std::string gen_dep() { |
| 57 | return R"( |