已合并
dynamic_mx_quant zero block with maxLowBound fix #7894
季骏创建于 7月24日
dynamic_mx_quant zero block with maxLowBound fix #7894
已合并
共 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; // 尾数与指数加1 | 366 | 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/scale | 371 | 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 | ||