已合并
dynamic_mx_quant zero block with maxLowBound fix #7894
dynamic_mx_quant zero block with maxLowBound fix #7894
已合并
季骏创建于 7月24日
共 2 个文件变更+21-0
@@ -365,6 +365,7 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
365 Reg::RegTensor<uint16_t> maxEleFP16;365 Reg::RegTensor<uint16_t> maxEleFP16;
366 Reg::RegTensor<uint32_t> manAbs0FP32; // 尾数与指数加1366 Reg::RegTensor<uint32_t> manAbs0FP32; // 尾数与指数加1
367 Reg::RegTensor<uint32_t> manAbs1FP32;367 Reg::RegTensor<uint32_t> manAbs1FP32;
368+ Reg::RegTensor<uint32_t> zeroU32;
368 Reg::RegTensor<uint32_t> mxScale0FP32; // 指数369 Reg::RegTensor<uint32_t> mxScale0FP32; // 指数
369 Reg::RegTensor<uint32_t> mxScale1FP32;370 Reg::RegTensor<uint32_t> mxScale1FP32;
370 Reg::RegTensor<uint16_t> reversedShareExpBF16; // 1/scale371 Reg::RegTensor<uint16_t> reversedShareExpBF16; // 1/scale
@@ -385,6 +386,7 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
385 Reg::MaskReg p1;386 Reg::MaskReg p1;
386 Reg::MaskReg p2;387 Reg::MaskReg p2;
387 Reg::MaskReg p3;388 Reg::MaskReg p3;
389+ Reg::MaskReg pZeroBlock;
388 Reg::MaskReg pregAll8 = Reg::CreateMask<uint8_t, Reg::MaskPattern::ALL>();390 Reg::MaskReg pregAll8 = Reg::CreateMask<uint8_t, Reg::MaskPattern::ALL>();
389 Reg::MaskReg pregAll16 = Reg::CreateMask<uint16_t, Reg::MaskPattern::ALL>();391 Reg::MaskReg pregAll16 = Reg::CreateMask<uint16_t, Reg::MaskPattern::ALL>();
390 Reg::MaskReg pregAll32 = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>();392 Reg::MaskReg pregAll32 = Reg::CreateMask<uint32_t, Reg::MaskPattern::ALL>();
@@ -408,6 +410,7 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
408 Reg::Duplicate(maxEleU16, BF16_MAX_EXP);410 Reg::Duplicate(maxEleU16, BF16_MAX_EXP);
409 Reg::Duplicate(biasU16, BF16_EXP_BIAS);411 Reg::Duplicate(biasU16, BF16_EXP_BIAS);
410 Reg::Duplicate(zeroU16, 0);412 Reg::Duplicate(zeroU16, 0);
413+ Reg::Duplicate(zeroU32, FP32_NUMBER_ZERO);
411 Reg::Duplicate(nanU16, BF16_NAN_CUSTOM);414 Reg::Duplicate(nanU16, BF16_NAN_CUSTOM);
412 Reg::Duplicate(specialExpU16, BF16_SPECIAL_EXP_THRESHOLD);415 Reg::Duplicate(specialExpU16, BF16_SPECIAL_EXP_THRESHOLD);
413 if constexpr (IsSame<xDtype, float>::value) {416 if constexpr (IsSame<xDtype, float>::value) {
@@ -437,6 +440,8 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
437 }440 }
438 441 
439 if constexpr (calcMode == MODE_ONE) {442 if constexpr (calcMode == MODE_ONE) {
443+ Reg::CompareScalar<uint32_t, CMPMODE::NE>(pZeroBlock, (Reg::RegTensor<uint32_t>&)manAbs0FP32,
444+ FP32_NUMBER_ZERO, pregAll32);
440 Reg::Maxs((Reg::RegTensor<float>&)manAbs0FP32, (Reg::RegTensor<float>&)manAbs0FP32, maxLowBound_,445 Reg::Maxs((Reg::RegTensor<float>&)manAbs0FP32, (Reg::RegTensor<float>&)manAbs0FP32, maxLowBound_,
441 pregAll32);446 pregAll32);
442 Reg::Mul((Reg::RegTensor<float>&)manAbs0FP32, (Reg::RegTensor<float>&)manAbs0FP32,447 Reg::Mul((Reg::RegTensor<float>&)manAbs0FP32, (Reg::RegTensor<float>&)manAbs0FP32,
@@ -462,10 +467,15 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
462 Reg::Adds(manAbs0FP32, mxScale0FP32, 1, p0);467 Reg::Adds(manAbs0FP32, mxScale0FP32, 1, p0);
463 Reg::Select(mxScale0FP32, manAbs0FP32, mxScale0FP32, p0);468 Reg::Select(mxScale0FP32, manAbs0FP32, mxScale0FP32, p0);
464 469 
470+ if constexpr (calcMode == MODE_ONE) {
471+ Reg::Select<uint32_t>(mxScale0FP32, mxScale0FP32, zeroU32, pZeroBlock);
472+ }
473+ 
465 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>((Reg::RegTensor<uint16_t>&)mxScale0BF16, mxScale0FP32);474 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>((Reg::RegTensor<uint16_t>&)mxScale0BF16, mxScale0FP32);
466 475 
467 if constexpr (!IsSame<xDtype, float>::value) {476 if constexpr (!IsSame<xDtype, float>::value) {
468 if constexpr (calcMode == MODE_ONE) {477 if constexpr (calcMode == MODE_ONE) {
478+ Reg::CompareScalar<uint32_t, CMPMODE::NE>(pZeroBlock, manAbs1FP32, FP32_NUMBER_ZERO, pregAll32);
469 Reg::Maxs((Reg::RegTensor<float>&)manAbs1FP32, (Reg::RegTensor<float>&)manAbs1FP32, maxLowBound_,479 Reg::Maxs((Reg::RegTensor<float>&)manAbs1FP32, (Reg::RegTensor<float>&)manAbs1FP32, maxLowBound_,
470 pregAll32);480 pregAll32);
471 Reg::Mul((Reg::RegTensor<float>&)manAbs1FP32, (Reg::RegTensor<float>&)manAbs1FP32,481 Reg::Mul((Reg::RegTensor<float>&)manAbs1FP32, (Reg::RegTensor<float>&)manAbs1FP32,
@@ -491,6 +501,10 @@ DynamicMxQuantNotTailAxisOptimizeLargeTail<xDtype, yDtype, roundMode, calcMode>:
491 Reg::Adds(manAbs1FP32, mxScale1FP32, 1, p2);501 Reg::Adds(manAbs1FP32, mxScale1FP32, 1, p2);
492 Reg::Select(mxScale1FP32, manAbs1FP32, mxScale1FP32, p2);502 Reg::Select(mxScale1FP32, manAbs1FP32, mxScale1FP32, p2);
493 503 
504+ if constexpr (calcMode == MODE_ONE) {
505+ Reg::Select<uint32_t>(mxScale1FP32, mxScale1FP32, zeroU32, pZeroBlock);
506+ }
507+ 
494 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>((Reg::RegTensor<uint16_t>&)mxScale1BF16,508 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>((Reg::RegTensor<uint16_t>&)mxScale1BF16,
495 mxScale1FP32);509 mxScale1FP32);
496 510 
@@ -83,6 +83,7 @@ private:
83 Reg::RegTensor<uint32_t> absForXFP32;83 Reg::RegTensor<uint32_t> absForXFP32;
84 Reg::RegTensor<uint32_t> manForFP32;84 Reg::RegTensor<uint32_t> manForFP32;
85 Reg::RegTensor<uint32_t> oneU32;85 Reg::RegTensor<uint32_t> oneU32;
86+ Reg::RegTensor<uint32_t> zeroU32;
86 Reg::RegTensor<int32_t> negZeroI32;87 Reg::RegTensor<int32_t> negZeroI32;
87 Reg::RegTensor<uint16_t> bf16NegInfU16;88 Reg::RegTensor<uint16_t> bf16NegInfU16;
88 Reg::RegTensor<uint16_t> tgtMaxExpU16;89 Reg::RegTensor<uint16_t> tgtMaxExpU16;
@@ -546,6 +547,7 @@ DynamicMxQuantNotTailAxisOptimizeSmallTail<xDtype, yDtype, roundMode, calcMode>:
546 Reg::Duplicate(regs.manForFP32, FP32_MX_MAN_MASK);547 Reg::Duplicate(regs.manForFP32, FP32_MX_MAN_MASK);
547 Reg::Duplicate(regs.dstTypeMaxReg, invDstTypeMax_);548 Reg::Duplicate(regs.dstTypeMaxReg, invDstTypeMax_);
548 Reg::Duplicate(regs.oneU32, 1);549 Reg::Duplicate(regs.oneU32, 1);
550+ Reg::Duplicate(regs.zeroU32, FP32_NUMBER_ZERO);
549 551 
550 if constexpr (IsSame<DTYPE_X, half>::value) {552 if constexpr (IsSame<DTYPE_X, half>::value) {
551 if constexpr (IsSame<DTYPE_Y, fp4x2_e2m1_t>::value) {553 if constexpr (IsSame<DTYPE_Y, fp4x2_e2m1_t>::value) {
@@ -803,6 +805,7 @@ DynamicMxQuantNotTailAxisOptimizeSmallTail<xDtype, yDtype, roundMode, calcMode>:
803 Reg::RegTensor<uint32_t>& mxScaleAdd1U32, Reg::RegTensor<uint16_t>& mxScale)805 Reg::RegTensor<uint32_t>& mxScaleAdd1U32, Reg::RegTensor<uint16_t>& mxScale)
804{806{
805 if constexpr (calcMode == MODE_ONE) {807 if constexpr (calcMode == MODE_ONE) {
808+ Reg::CompareScalar<uint32_t, CMPMODE::NE>(auxRegs.maxLowBoundMask, absMaxU32, FP32_NUMBER_ZERO, auxRegs.p1);
806 if constexpr (IsSame<xDtype, float>::value) {809 if constexpr (IsSame<xDtype, float>::value) {
807 if constexpr (canMaxLowBound) {810 if constexpr (canMaxLowBound) {
808 Reg::MaskReg maskAll = Reg::CreateMask<float, Reg::MaskPattern::ALL>();811 Reg::MaskReg maskAll = Reg::CreateMask<float, Reg::MaskPattern::ALL>();
@@ -846,6 +849,10 @@ DynamicMxQuantNotTailAxisOptimizeSmallTail<xDtype, yDtype, roundMode, calcMode>:
846 Reg::Adds(mxScaleAdd1U32, mxScaleU32, 1, auxRegs.p4);849 Reg::Adds(mxScaleAdd1U32, mxScaleU32, 1, auxRegs.p4);
847 Reg::Select(mxScaleU32, mxScaleAdd1U32, mxScaleU32, auxRegs.p4);850 Reg::Select(mxScaleU32, mxScaleAdd1U32, mxScaleU32, auxRegs.p4);
848 851 
852+ if constexpr (calcMode == MODE_ONE) {
853+ Reg::Select<uint32_t>(mxScaleU32, mxScaleU32, auxRegs.zeroU32, auxRegs.maxLowBoundMask);
854+ }
855+ 
849 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>(mxScale, mxScaleU32);856 Reg::Pack<uint16_t, uint32_t, Reg::HighLowPart::LOWEST>(mxScale, mxScaleU32);
850}857}
851 858