已合并
fix(asinh): improve large-value precision #4822
m0_73979701创建于 8月18日
fix(asinh): improve large-value precision #4822
已合并
共 1 个文件变更+25-5
| @@ -49,7 +49,9 @@ constexpr float CONST_S_MAX = 3.4028235e34f; // clip 上界(REQUIREMENTS §8.1 | |||
| 49 | // 字面量截断到 FP32 实际可表示精度(mantissa 24 bit,约 7 位有效数字)。 | 49 | // 字面量截断到 FP32 实际可表示精度(mantissa 24 bit,约 7 位有效数字)。 |
| 50 | // 完整精度参考: ln(2) ≈ 0.69314718055994530941723212145818...,编译期 round 到最近 FP32。 | 50 | // 完整精度参考: ln(2) ≈ 0.69314718055994530941723212145818...,编译期 round 到最近 FP32。 |
| 51 | constexpr float CONST_LN2 = 0.6931472f; | 51 | constexpr float CONST_LN2 = 0.6931472f; |
| 52 | -constexpr float CONST_BRANCH_THRESHOLD = 0.00024414063f; // 2^-12,小参数分支阈值 | 52 | +constexpr float CONST_BRANCH_THRESHOLD = 0.00024414063f; // 2^-12,小参数分支阈值 |
| 53 | +constexpr float CONST_DIRECT_LOG_THRESHOLD = 10.0f; // 中大参数直接使用 log(u) | ||
| 54 | +constexpr float CONST_ASYMPTOTIC_THRESHOLD = 268435456.0f; // 2^28,以上使用 log(|x|) + ln(2) | ||
| 53 | 55 | ||
| 54 | // Double Buffer 固定为 2(19 步含 Log/Sqrt/Div/Compare 计算密集,双缓冲收益显著) | 56 | // Double Buffer 固定为 2(19 步含 Log/Sqrt/Div/Compare 计算密集,双缓冲收益显著) |
| 55 | static constexpr int32_t BUFFER_NUM = 2; | 57 | static constexpr int32_t BUFFER_NUM = 2; |
| @@ -313,9 +315,14 @@ __aicore__ inline void Asinh<T>::ComputeFp32Pipeline(LocalTensor<float>& xOrigFp | |||
| 313 | // Log natural 三参数版本,框架从未 InitBuffer 的剩余 UB 自动申请 tmpBuffer | 315 | // Log natural 三参数版本,框架从未 InitBuffer 的剩余 UB 自动申请 tmpBuffer |
| 314 | AscendC::Log(s, s, n); | 316 | AscendC::Log(s, s, n); |
| 315 | 317 | ||
| 316 | - // ===== Step 15: s = log(u) * r / clipped_s → 主路径 res ===== | 318 | + // ===== Step 15: 小参数保留 log1p 补偿,中大参数直接使用 log(u) ===== |
| 317 | - AscendC::Mul(s, s, r, n); // s = log(u) * r | 319 | + // 对 |x| < 10,r/(u-1) 可补偿 u=1+r 在 FP32 下的舍入;对更大输入, |
| 318 | - AscendC::Div(s, s, b, n); // s = log(u) * r / clipped_s | 320 | + // 该乘除会引入额外 1~2 ULP,因此将分子、分母同时置 1,保留 step 14 的 log(u)。 |
| 321 | + AscendC::CompareScalar(selMask, absX, CONST_DIRECT_LOG_THRESHOLD, AscendC::CMPMODE::LT, nAligned); | ||
| 322 | + AscendC::Select(r, selMask, r, CONST_ONE, AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, nAligned); | ||
| 323 | + AscendC::Select(b, selMask, b, CONST_ONE, AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, nAligned); | ||
| 324 | + AscendC::Mul(s, s, r, n); // 小参数: log(u) * r;中大参数: log(u) * 1 | ||
| 325 | + AscendC::Div(s, s, b, n); // 小参数: / (u-1);中大参数: / 1 | ||
| 319 | // bBuf 此后再次释放 | 326 | // bBuf 此后再次释放 |
| 320 | 327 | ||
| 321 | // ===== Step 16: 大参数修正 result_2 = min(res, log(|x|) + ln(2) + 1/|x|²) ===== | 328 | // ===== Step 16: 大参数修正 result_2 = min(res, log(|x|) + ln(2) + 1/|x|²) ===== |
| @@ -333,7 +340,20 @@ __aicore__ inline void Asinh<T>::ComputeFp32Pipeline(LocalTensor<float>& xOrigFp | |||
| 333 | AscendC::Div(b, b, absX, n); | 340 | AscendC::Div(b, b, absX, n); |
| 334 | AscendC::Mul(b, b, b, n); // b = 1/|x|² | 341 | AscendC::Mul(b, b, b, n); // b = 1/|x|² |
| 335 | AscendC::Add(r, r, b, n); // r = log(|x|) + ln(2) + 1/|x|² | 342 | AscendC::Add(r, r, b, n); // r = log(|x|) + ln(2) + 1/|x|² |
| 336 | - // 16d: s = min(s, r) → result_2 | 343 | + |
| 344 | + // 16d: 分段选择 correction。 | ||
| 345 | + // |x| < 10: 保持原有 min(res, correction); | ||
| 346 | + // 10 <= |x| < 2^28: 屏蔽 correction,直接保留 log(u); | ||
| 347 | + // |x| >= 2^28: 恢复渐近式,避免 u≈2|x| 溢出。 | ||
| 348 | + // bBuf 已结束 1/|x|² 的生命周期,借其保存 correction,供极大值和 NaN 路径恢复。 | ||
| 349 | + AscendC::Adds(b, r, CONST_ZERO, n); | ||
| 350 | + AscendC::CompareScalar(selMask, absX, CONST_DIRECT_LOG_THRESHOLD, AscendC::CMPMODE::LT, nAligned); | ||
| 351 | + AscendC::Select(r, selMask, r, CONST_S_MAX, AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, nAligned); | ||
| 352 | + AscendC::CompareScalar(selMask, absX, CONST_ASYMPTOTIC_THRESHOLD, AscendC::CMPMODE::GE, nAligned); | ||
| 353 | + AscendC::Select(r, selMask, b, r, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE, nAligned); | ||
| 354 | + // NaN 与自身比较为 false:恢复 NaN correction,保持原有 NaN 传播语义。 | ||
| 355 | + AscendC::Compare(selMask, absX, absX, AscendC::CMPMODE::EQ, nAligned); | ||
| 356 | + AscendC::Select(r, selMask, r, b, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE, nAligned); | ||
| 337 | AscendC::Min(s, s, r, n); | 357 | AscendC::Min(s, s, r, n); |
| 338 | 358 | ||
| 339 | // ===== Step 17: output = select(|x| < 2^-12, |x|, result_2) ===== | 359 | // ===== Step 17: output = select(|x| < 2^-12, |x|, result_2) ===== |