已合并
loaddatav2 support fp4 #2018
吴洋创建于 5月8日
loaddatav2 support fp4 #2018
已合并
共 1 个文件变更+54-50
| @@ -110,9 +110,10 @@ __aicore__ inline void LoadData2DL12L0BCal(__cb__ T* dst, __cbuf__ T* src, const | |||
| 110 | template <typename T> | 110 | template <typename T> |
| 111 | __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const Load2DBitModeParam &loadDataParam) | 111 | __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const Load2DBitModeParam &loadDataParam) |
| 112 | { | 112 | { |
| 113 | - static_assert(SupportType<T, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, bfloat16_t, | 113 | + static_assert( |
| 114 | - float, int32_t, uint32_t>(), | 114 | + SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, |
| 115 | - "LoadData 2dv2 only support uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | 115 | + bfloat16_t, float, int32_t, uint32_t>(), |
| 116 | + "LoadData 2dv2 only support fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | ||
| 116 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); | 117 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); |
| 117 | if ASCEND_IS_AIC { | 118 | if ASCEND_IS_AIC { |
| 118 | 119 | ||
| @@ -160,9 +161,10 @@ __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const | |||
| 160 | template <typename T> | 161 | template <typename T> |
| 161 | __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const Load2DBitModeParam &loadDataParam) | 162 | __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const Load2DBitModeParam &loadDataParam) |
| 162 | { | 163 | { |
| 163 | - static_assert(SupportType<T, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, bfloat16_t, | 164 | + static_assert( |
| 164 | - float, int32_t, uint32_t>(), | 165 | + SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, |
| 165 | - "LoadData 2dv2 only support uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | 166 | + bfloat16_t, float, int32_t, uint32_t>(), |
| 167 | + "LoadData 2dv2 only support fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | ||
| 166 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); | 168 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); |
| 167 | if ASCEND_IS_AIC { | 169 | if ASCEND_IS_AIC { |
| 168 | 170 | ||
| @@ -210,31 +212,32 @@ __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const | |||
| 210 | template <typename T> | 212 | template <typename T> |
| 211 | __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const LoadData2DParamsV2 &loadDataParam) | 213 | __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const LoadData2DParamsV2 &loadDataParam) |
| 212 | { | 214 | { |
| 213 | - static_assert(SupportType<T, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, bfloat16_t, | 215 | + static_assert( |
| 214 | - float, int32_t, uint32_t>(), | 216 | + SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, |
| 215 | - "LoadData 2dv2 only support uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | 217 | + bfloat16_t, float, int32_t, uint32_t>(), |
| 218 | + "LoadData 2dv2 only support fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | ||
| 216 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); | 219 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); |
| 217 | if ASCEND_IS_AIC { | 220 | if ASCEND_IS_AIC { |
| 218 | - if (loadDataParam.ifTranspose) { | 221 | + if constexpr (SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t>()) { |
| 219 | - load_cbuf_to_ca(dst, | 222 | + if (loadDataParam.ifTranspose) { |
| 220 | - src, | 223 | + load_cbuf_to_ca_s4(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 221 | - loadDataParam.mStartPosition, | 224 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 222 | - loadDataParam.kStartPosition, | 225 | + loadDataParam.dstStride, 1); |
| 223 | - loadDataParam.mStep, | 226 | + } else { |
| 224 | - loadDataParam.kStep, | 227 | + load_cbuf_to_ca_s4(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 225 | - loadDataParam.srcStride, | 228 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 226 | - loadDataParam.dstStride, | 229 | + loadDataParam.dstStride, 0); |
| 227 | - 1); | 230 | + } |
| 228 | } else { | 231 | } else { |
| 229 | - load_cbuf_to_ca(dst, | 232 | + if (loadDataParam.ifTranspose) { |
| 230 | - src, | 233 | + load_cbuf_to_ca(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 231 | - loadDataParam.mStartPosition, | 234 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 232 | - loadDataParam.kStartPosition, | 235 | + loadDataParam.dstStride, 1); |
| 233 | - loadDataParam.mStep, | 236 | + } else { |
| 234 | - loadDataParam.kStep, | 237 | + load_cbuf_to_ca(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 235 | - loadDataParam.srcStride, | 238 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 236 | - loadDataParam.dstStride, | 239 | + loadDataParam.dstStride, 0); |
| 237 | - 0); | 240 | + } |
| 238 | } | 241 | } |
| 239 | } | 242 | } |
| 240 | } | 243 | } |
| @@ -242,31 +245,32 @@ __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const | |||
| 242 | template <typename T> | 245 | template <typename T> |
| 243 | __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const LoadData2DParamsV2 &loadDataParam) | 246 | __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const LoadData2DParamsV2 &loadDataParam) |
| 244 | { | 247 | { |
| 245 | - static_assert(SupportType<T, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, bfloat16_t, | 248 | + static_assert( |
| 246 | - float, int32_t, uint32_t>(), | 249 | + SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, half, |
| 247 | - "LoadData 2dv2 only support uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | 250 | + bfloat16_t, float, int32_t, uint32_t>(), |
| 251 | + "LoadData 2dv2 only support fp4x2_e2m1_t, fp4x2_e1m2_t, uint8_t, int8_t, hifloat8_t, fp8_e5m2_t, fp8_e4m3fn_t, \ | ||
| 248 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); | 252 | half, bfloat16_t, float, int32_t, uint32_t on current device!"); |
| 249 | if ASCEND_IS_AIC { | 253 | if ASCEND_IS_AIC { |
| 250 | - if (loadDataParam.ifTranspose) { | 254 | + if constexpr (SupportType<T, fp4x2_e2m1_t, fp4x2_e1m2_t>()) { |
| 251 | - load_cbuf_to_cb(dst, | 255 | + if (loadDataParam.ifTranspose) { |
| 252 | - src, | 256 | + load_cbuf_to_cb_s4(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 253 | - loadDataParam.mStartPosition, | 257 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 254 | - loadDataParam.kStartPosition, | 258 | + loadDataParam.dstStride, 1); |
| 255 | - loadDataParam.mStep, | 259 | + } else { |
| 256 | - loadDataParam.kStep, | 260 | + load_cbuf_to_cb_s4(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 257 | - loadDataParam.srcStride, | 261 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 258 | - loadDataParam.dstStride, | 262 | + loadDataParam.dstStride, 0); |
| 259 | - 1); | 263 | + } |
| 260 | } else { | 264 | } else { |
| 261 | - load_cbuf_to_cb(dst, | 265 | + if (loadDataParam.ifTranspose) { |
| 262 | - src, | 266 | + load_cbuf_to_cb(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 263 | - loadDataParam.mStartPosition, | 267 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 264 | - loadDataParam.kStartPosition, | 268 | + loadDataParam.dstStride, 1); |
| 265 | - loadDataParam.mStep, | 269 | + } else { |
| 266 | - loadDataParam.kStep, | 270 | + load_cbuf_to_cb(dst, src, loadDataParam.mStartPosition, loadDataParam.kStartPosition, |
| 267 | - loadDataParam.srcStride, | 271 | + loadDataParam.mStep, loadDataParam.kStep, loadDataParam.srcStride, |
| 268 | - loadDataParam.dstStride, | 272 | + loadDataParam.dstStride, 0); |
| 269 | - 0); | 273 | + } |
| 270 | } | 274 | } |
| 271 | } | 275 | } |
| 272 | } | 276 | } |