| 188 | return kernel_code |
| 189 | |
| 190 | def gen_body_core_code(bm, by): |
| 191 | length = 4 |
| 192 | all_code = "" |
| 193 | for i in range(length): |
| 194 | core_code = "\n\ |
| 195 | uint8x16_t vec_a_{0} = vld1q_u8(a + i * KK / 2 + k * 32 * 2 + {0} * 16);\n\ |
| 196 | uint8x16_t vec_a{0}_top = vshrq_n_u8(vec_a_{0}, 4);\n\ |
| 197 | uint8x16_t vec_a{0}_bot = vandq_u8(vec_a_{0}, vec_mask);\n\ |
| 198 | int8x16_t vec_v_{0}_left_tmp0 = vqtbl1q_s8(vec_lut[{1} * k + {2}], vec_a{0}_top);\n\ |
| 199 | int8x16_t vec_v_{0}_left_tmp1 = vqtbl1q_s8(vec_lut[{1} * k + {3}], vec_a{0}_top);\n\ |
| 200 | int8x16_t vec_v_{0}_right_tmp0 = vqtbl1q_s8(vec_lut[{1} * k + {4}], vec_a{0}_bot);\n\ |
| 201 | int8x16_t vec_v_{0}_right_tmp1 = vqtbl1q_s8(vec_lut[{1} * k + {5}], vec_a{0}_bot);\n\ |
| 202 | int8x16x2_t vec_v_left_{0} = vzipq_s8(vec_v_{0}_left_tmp1, vec_v_{0}_left_tmp0);\n\ |
| 203 | int8x16x2_t vec_v_right_{0} = vzipq_s8(vec_v_{0}_right_tmp1, vec_v_{0}_right_tmp0);\n\ |
| 204 | vec_c[{6}] += vec_v_left_{0}.val[0];\n\ |
| 205 | vec_c[{6}] += vec_v_right_{0}.val[0];\n\ |
| 206 | vec_c[{7}] += vec_v_left_{0}.val[1];\n\ |
| 207 | vec_c[{7}] += vec_v_right_{0}.val[1];\n\ |
| 208 | ".format(i, 2 * by // 2, (4 * i) % (2 * by // 2), (4 * i + 1) % (2 * by // 2), (4 * i + 2) % (2 * by // 2), (4 * i + 3) % (2 * by // 2), (i * 2) // (by // 2) * 2 + 0, (i * 2) // (by // 2) * 2 + 1) |
| 209 | |
| 210 | all_code = "".join([all_code, core_code]) |
| 211 | |
| 212 | all_code = "".join([all_code, "\n }\n\n"]) |
| 213 | |
| 214 | for i in range(bm // 8): |
| 215 | core_code = "\ |
| 216 | int32x4_t vec_v_bot_low_low_{0} = vmovl_s16(vget_low_s16(vec_c[{0}]));\n\ |
| 217 | int32x4_t vec_v_bot_low_high_{0} = vmovl_high_s16(vec_c[{0}]);\n\ |
| 218 | vst1q_s32(c + i + {1}, vld1q_s32(c + i + {1}) + vec_v_bot_low_low_{0});\n\ |
| 219 | vst1q_s32(c + i + {2}, vld1q_s32(c + i + {2}) + vec_v_bot_low_high_{0});\n".format(i, i * 8, i * 8 + 4) |
| 220 | all_code = "".join([all_code, core_code]) |
| 221 | |
| 222 | return all_code |
| 223 | |
| 224 | def gen_tbl_impl(pre, BM, BK, bm, k): |
| 225 | |