近似双曲正切函数

Hacker News Top 论文

摘要

本文梳理了多种快速双曲正切近似方法——泰勒展开、Padé 逼近、样条曲线及位级技巧,面向神经网络与实时音频场景。

暂无内容
查看原文
查看缓存全文

缓存时间: 2026/04/23 00:25

# 近似双曲正切 来源:https://jtomschroeder.com/blog/approximating-tanh/ 快速 tanh 近似方法综述:Taylor 级数、Padé 逼近、样条,以及 K-TanH、Schraudolph 等位运算技巧 双曲正切函数 \\( tanh \\) 把任意实数光滑地映射到区间 \\(-1,1\\),呈 S 形曲线。这一特性使其在神经网络里充当激活函数时既能引入非线性又能把输出限幅,在音频处理里则能提供自然的软削波饱和效果。无论哪种场景,速度都至关重要:神经网络的一次前向传播就可能调用数百万次 \\( tanh \\),音频实时处理则要求 44.1 kHz 甚至更高的采样率。标准库实现的高精度往往意味着更多计算,而精心定制的近似算法可以显著提速。本文梳理几种常见思路:传统多项式法(Taylor、Padé、样条),以及利用 IEEE-754 浮点格式的奇技淫巧,后者几乎不费劲就能跑得非常快。 ## 近似方法一览 先走马观花看一遍对 \\( tanh \\) 做快速近似的常用套路。 ### Taylor 级数 还记得微积分课吗?Taylor 级数用逐阶导数把函数展成无穷多项式,截断前几项即可得到轻量级近似。 ```rust pub fn tanhf(x: f32) -> f32 { // 当 |x| 过大时多项式开始发散,直接截断到 ±1 if x.abs() > 1.365 { return 1f32.copysign(x); } let t1 = x; let t2 = x.powi(3) * (1. / 3.); let t3 = x.powi(5) * (2. / 15.); let t4 = x.powi(7) * (17. / 315.); let t5 = x.powi(9) * (62. / 2835.); let t6 = x.powi(11) * (1382. / 155925.); t1 - t2 + t3 - t4 + t5 - t6 } ``` ### Padé 逼近 与 Taylor 类似,Padé 用两个多项式相除,精度更高但要多做一次除法。下面改编自 JUCE 的 FastMathApproximations,数学上称为 "[7/6] Padé",分子 7 次、分母 6 次。 ```rust pub fn tanhf(x: f32) -> f32 { // 该逼近仅在 [-5,5] 区间保证误差可控 if x.abs() > 5. { return 1f32.copysign(x); } let x2 = x * x; let numerator = x * (135135. + x2 * (17325. + x2 * (378. + x2))); let denominator = 135135. + x2 * (62370. + x2 * (3150. + 28. * x2)); numerator / denominator } ``` ### 样条 把函数切成多段,每段配一个三次多项式。下面例子来自论文《Efficiently inaccurate approximation of hyperbolic tangent used as transfer function in artificial neural networks》(Simos, Tsitouras)。作者把 [0,18] 切成 3 段,系数用 MATLAB 事先拟合好,速度优先、精度其次。 ```rust pub fn tanhf3(xin: f32) -> f32 { const N1: f32 = 0.371025186672900; const N2: f32 = 2.572153900248530; const N3: f32 = 18.; match xin.abs() { x if x <= N1 => { -3.695076086125492e-1 * x.powi(3) + 1.987219343897867e-2 * x.powi(2) + x } x if x <= N2 => { let n = x - N1; 5.928356367224758e-2 * n.powi(3) - 3.914176949486042e-1 * n.powi(2) + 8.621472609449146e-1 * n + 3.548881072496229e-1 } x if x <= N3 => { let n = x - N2; -3.347599023061577e-6 * n.powi(3) + 5.456777761558641e-5 * n.powi(2) + 7.066442941005233e-4 * n + 9.884026213740197e-1 } _ => 1., }.copysign(xin) } ``` ## 借格式作弊 看完数学派,再来见识"格式流":直接拿 IEEE-754 浮点的位模式开刀。 32 位浮点的二进制布局:符号 1 位、指数 8 位、尾数 23 位。 名义值公式: $$ (-1)^s · 2^{E} · (1 + M/2^{p}) $$ ### K-TanH 论文《K-TanH: Efficient TanH For Deep Learning》提出纯整数运算 + 512 bit 查找表的硬件友好算法。核心思想:把指数低 2 位与尾数高 3 位拼成 5 bit 索引,查表得到新的指数、右移位数和偏置,然后现场拼装结果。 ```rust pub fn tanhf(x: f32) -> f32 { const T1: f32 = 0.25; const T2: f32 = 3.75; let xa = x.abs(); if xa < T1 { x } else if xa > T2 { 1f32.copysign(x) } else { let xb = x.to_bits(); let mi = (xb >> 16) & 0b0111_1111; let so = xb & 0x8000_0000; let t = (xb >> 20) & 0b11_111; let (et, rt, bt) = unpack(LUT[t as usize]); let eo = (et as u32) << 23; let mo = (((mi >> rt) as i32 + bt as i32) as u32) << 16; f32::from_bits(so | eo | mo) } } ``` 查找表仅 32 项,可塞进 AVX512 的 512 bit 寄存器,一次并行查多条。论文还提到对 bfloat16 更友好,尾数短、表更小,深度学习精度足够。 ### Schraudolph 指数逼近 1999 年,Nicol Schraudolph 在《A Fast, Compact Approximation of the Exponential Function》里把浮点位模式当整数玩,几行代码近似 \\( e^x \\)。思路与 Quake 的 fast inv-sqrt 异曲同工:利用指数域本身就像对数这一事实。 原始 C++ 版(double)核心: ```cpp i = (int)(EXP_A * y) + (1072693248 - EXP_C); ``` Rust 单精度移植: ```rust pub fn expf(y: f32) -> f32 { use core::f32::consts::LN_2; const BIAS: i16 = f32::MAX_EXP as i16 - 1; const MANTISSA_BITS: i16 = f32::MANTISSA_DIGITS as i16 - 1; const OFFSET_BITS: i16 = i16::BITS as i16; const X: i16 = 1 << (MANTISSA_BITS - OFFSET_BITS); const A: f32 = X as f32 / LN_2; const B: i16 = X * BIAS; const C: i16 = 8; // 调参 const D: i16 = B - C; unsafe { let y = (A * y).to_int_unchecked::<i16>() + D; core::mem::transmute( #[cfg(target_endian = "little")] [0, y], #[cfg(target_endian = "big")] [y, 0], ) } } ``` 有了超快 `expf`,就能按定义算 tanh: ```rust pub fn tanhf(x: f32) -> f32 { let y = expf(2. * x); (y - 1.) / (y + 1.) } ``` #### Schraudolph-NG:误差对消版 2018 年作者本人在 Stack Overflow 留言:把 `exp(x/2) / exp(-x/2)` 拼回去,两项的分段线性误差高度相关,一除就能抵消大半,白捡精度。 ```rust pub fn expf(x: f32) -> f32 { /* … 常量同上 … */ let num = f32::from_bits(((A/2.) * x + B) as u32); let den = f32::from_bits(((-A/2.) * x + B) as u32); num / den } ``` ARM NEON 版可一次算俩指数再相除,SIMD 友好。 #### 延伸阅读 - Martin Leitner-Ankerl: Optimized Exponential Functions for Java - Martin Leitner-Ankerl: Optimized pow() approximation for Java, C / C++, and C# - ekmett/approximate/cbits/fast.c - gingerBill/gb/gb_math.h

相似文章

关于神经网络的显式超表达逼近

arXiv cs.LG

本文研究了固定架构神经网络的显式参数-误差权衡逼近,利用中国剩余定理作为构造性编码机制,并获得了Lipschitz和Hölder光滑函数的显式界。

分叉附近的状态空间NTK坍缩

arXiv cs.LG

本文发展了动力模型分叉附近梯度下降的局部理论,表明状态空间神经正切核坍缩为秩一算子,主导学习动力学,使优化有效低维且可从规范形式预测。

通过负微分电阻网络桥接函数逼近与器件物理

Reddit r/singularity

本文介绍了KANalogue,一种完全模拟的Kolmogorov-Arnold网络实现,它利用负微分电阻器件直接在硬件中执行可学习的非线性函数。在MNIST、FashionMNIST和CIFAR-10数据集上,该方案以少于模拟MLP的参数实现了具有竞争力的精度。