已合并
loaddatav2 support fp4 #2018
loaddatav2 support fp4 #2018
已合并
吴洋创建于 5月8日
1 个文件变更+54-50
Mimpl/basic_api/dav_c310/kernel_operator_mm_impl.h+54-50
@@ -110,9 +110,10 @@ __aicore__ inline void LoadData2DL12L0BCal(__cb__ T* dst, __cbuf__ T* src, const
110template <typename T>110template <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#if defined(ASCENDC_CPU_DEBUG) && ASCENDC_CPU_DEBUG == 1119#if defined(ASCENDC_CPU_DEBUG) && ASCENDC_CPU_DEBUG == 1
@@ -160,9 +161,10 @@ __aicore__ inline void LoadData2DL12L0ACal(__ca__ T *dst, __cbuf__ T *src, const
160template <typename T>161template <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#if defined(ASCENDC_CPU_DEBUG) && ASCENDC_CPU_DEBUG == 1170#if defined(ASCENDC_CPU_DEBUG) && ASCENDC_CPU_DEBUG == 1
@@ -210,31 +212,32 @@ __aicore__ inline void LoadData2DL12L0BCal(__cb__ T *dst, __cbuf__ T *src, const
210template <typename T>212template <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
242template <typename T>245template <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}