ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

wincnn 源码逐行解读:如何用 SymPy 精确算术自动推导 Winograd 卷积变换矩阵

wincnn 源码逐行解读:如何用 SymPy 精确算术自动推导 Winograd 卷积变换矩阵 wincnn 源码逐行解读如何用 SymPy 精确算术自动推导 Winograd 卷积变换矩阵【免费下载链接】wincnnWinograd minimal convolution algorithm generator for convolutional neural networks.项目地址: https://gitcode.com/gh_mirrors/wi/wincnnwincnn是一个专为卷积神经网络CNN生成最小 Winograd 卷积算法的 Python 模块。它完全构建在 SymPy 符号计算引擎之上利用 SymPy 的精确算术有理数而非浮点数自动推导出 F(m, r) 卷积所需的 AT、G、BT 三张变换矩阵并通过符号化验证保证结果零误差。如果你不想手推拉格朗日插值系数这篇文章就是完整的阅读指南。一、先懂原理Winograd 最小卷积在做什么直接卷积 F(2,3)输出 2 个点、核长 3需要6 次乘法而 Winograd 变换把卷积变成逐点相乘Y AT · ((G·g) ⊙ (BT·d))只需4 次乘法——这是数学上可证明的最少次数。其中AT / A把输出信号采样到插值点Vandermonde 矩阵G滤波器的变换元素是有理数如 1/2、1/90由拉格朗日插值系数决定BT把输入数据变换到同一插值点空间。 原理细节与 F(4x4, 3x3) 的运算量对比见 FAQ.md 的 How do Winograds fast convolution algorithms work? 一节。二、为什么必须用精确算术浮点数会毁掉整个推导推导 G、BT 矩阵需要反复做符号矩阵求逆源码中f ** (-1)。插值系数天然是分数算法矩阵中出现的系数F(2,3)1/2F(4,3)1/4、-1/6、1/24F(6,3)1/90、16/45、-21/4如果这些系数用 float 参与矩阵运算舍入误差会沿着符号化简一路污染最终结果——你得到的将是一张近似正确的矩阵部署后输出会逐像素漂移且无法定位。wincnn 的解法很简单全程使用 SymPy 的符号对象。整数保持整数分数用sympy.Rational表示矩阵元素永远是精确的代数式。这就是 README 中 F(6,3) 示例特意使用Rational(1,2)作为插值点的原因from sympy import Rational wincnn.showCookToomFilter((0, 1, -1, 2, -2, Rational(1,2), -Rational(1,2)), 6, 3)推导出的G矩阵里出现1/90、16/45这样的精确分数——这是浮点运算不可能碰巧产生的干净结果。三、快速上手安装与第一个变换安装只需一行要求 Python ≥ 3.8、SymPy ≥ 1.9pip install wincnn或从源码安装git clone https://gitcode.com/gh_mirrors/wi/wincnn cd wincnn pip install .最小示例——用插值点 (0, 1, -1) 推导 F(2,3) 变换import wincnn wincnn.showCookToomFilter((0, 1, -1), 2, 3)输出节选自 README.mdAT ⎡1 1 1 0⎤ ⎣0 1 -1 1⎦ G ⎡ 1 0 0 ⎤ ⎢1/2 1/2 1/2⎥ ⎢1/2 -1/2 1/2⎥ ⎣ 0 0 1 ⎦ BT ⎡1 0 -1 0⎤ ⎢0 1 1 0⎥ ⎢0 -1 1 0⎥ ⎣0 -1 0 1⎦ AT*((G*g)(BT*d)) ⎡d[0]⋅g[0] d[1]⋅g[1] d[2]⋅g[2]⎤ ⎣d[1]⋅g[0] d[2]⋅g[1] d[3]⋅g[2]⎦最后一行是自动符号验证把三张矩阵代入AT·((G·g)⊙(BT·d))并化简得到的恰好就是 F(2,3) 卷积的定义式。由于全程符号运算这个验证是严格成立的而不是误差 1e-6。四、源码逐段解读整个项目只有 wincnn.py 一个核心文件约 252 行所有逻辑都能一览无余。下面按执行顺序逐段拆解。4.1 符号工具箱SymPy 导入第 1–13 行from sympy import IndexedBase, Matrix, Poly, simplify, symbols, zeros, pprint每个导入都有明确分工Matrix做符号矩阵运算Poly/symbols构造多项式IndexedBase生成d[i]、g[i]这类带下标的符号simplify负责化简验证式pprint输出漂亮的矩阵排版。4.2 求值矩阵 At 与 A第 16–23 行def At(a, m, n): return Matrix(m, n, lambda i, j: a[i] ** j)At 就是 Vandermonde 矩阵第 i 行是第 i 个插值点a[i]的 0~n-1 次幂。它回答的问题是把多项式在各个插值点上求值。def A(a, m, n): return At(a, m - 1, n).row_insert( m - 1, Matrix(1, n, lambda i, j: 1 if j n - 1 else 0) )A 在 At 的底部插入一行哨兵[0 … 0 1]。这个技巧把多项式最高次系数直接抄送出来保证变换矩阵可逆——这正是 F(2,3) 的 AT 里出现末尾 0/1 列的原因。4.3 伴随矩阵 T第 26–29 行def T(a, n): return Matrix( Matrix.eye(n).col_insert(n, Matrix(n, 1, lambda i, j: -(a[i] ** n))) )T 是多项式递推的伴随矩阵给定前 n 次幂在插值点的值它能算出第 n 次幂的值最后一列-(a[i]**n)由插值点满足的范德蒙关系推出。它在推导 BT 时起作用BT 的每一行本质上是某插值点处的多项式系数。4.4 拉格朗日基Lx、F、L第 32–76 行def Lx(a, n): x symbols(x) return Matrix( n, 1, lambda i, j: Poly( reduce(operator.mul, ((x - a[k] if k ! i else 1) for k in range(0, n)), 1 ).expand(basicTrue), x ).as_expr(), )Lx的第 i 个元素就是第 i 个拉格朗日基多项式∏(x − a[k]) / ∏(a[i] − a[k])的分子用reduce(operator.mul, ...)连乘构造再用Poly(...).expand(basicTrue)以精确系数展开——注意这里没有任何浮点参与。def L(a, n): x symbols(x) lx Lx(a, n) f F(a, n) return Matrix(n, n, lambda i, j: lx[i, 0].coeff(x, j) / f[i]).TL把每个基多项式除以分母f[i]并按幂次拆成系数行.coeff(x, j)转置后得到点值 → 系数的插值矩阵。这里的除法lx / f[i]作用于 SymPy 有理数所以1/2、-1/6这类分数永远精确。4.5 数据变换 Bt / B第 79–86 行def Bt(a, n): return L(a, n) * T(a, n)Bt L·T先用伴随矩阵 T 补齐高次项再用插值矩阵 L 还原系数——这就是 BT 矩阵的来源纯符号乘法无求逆。4.6 分数放哪fractionsIn 的四种选择第 89–92 行FractionsInG 0 FractionsInA 1 FractionsInB 2 FractionsInF 3这是 wincnn 很实用的一招**分数系数可以任意搬运**到四个矩阵中的某一个。工程上通常希望某张矩阵只有整数或 0/1方便手写优化代码把分数集中到另一张矩阵即可。4.7 核心函数 cookToomFilter第 95–141 行alpha n r - 1 f FdiagPlus1(a, alpha) if f[0, 0] 0: f[0, :] * -1alpha n r - 1即输出长度 核长 − 1也就是所需插值点数量FdiagPlus1构造对角分数矩阵对角线是各插值点的拉格朗日分母符号规整若首元素为负整行取负让输出更整洁。if fractionsIn FractionsInG: AT A(a, alpha, n).T G (A(a, alpha, r).T * f ** (-1)).T BT f * B(a, alpha).T默认模式FractionsInGAT直接取求值矩阵转置——纯整数G(Aᵀ · f⁻¹)ᵀ——唯一一次符号矩阵求逆分数全部落在这里所以 G 里出现 1/2、1/90 这类有理数BTf · Bᵀ——用分数对角阵吸收插值分母后得到以整数为主的数据变换。其余三个elif分支只是把f ** (-1)挪到 AT、BT 或独立输出 f数学上完全等价供使用者按部署需求挑选。4.8 符号验证filterVerify第 144–162 行di IndexedBase(d); gi IndexedBase(g) d Matrix(alpha, 1, lambda i, j: di[i]) g Matrix(r, 1, lambda i, j: gi[i]) V BT * d U G * g M U.multiply_elementwise(V) Y simplify(AT * M)六行代码复刻了完整的卷积流程BT·d数据变换 →G·g滤波变换 →逐元素相乘multiply_elementwise→AT逆变换 →simplify化简。只要输出的每一项都形如d[k]⋅g[j]且与卷积定义逐项吻合变换矩阵就是正确的——这正是 README 中AT*((G*g)(BT*d)) …那一段的生成逻辑。4.9 打印函数与对偶原理第 185–252 行showCookToomFilter依次pprint打印三张矩阵并调用filterVerify验证而showCookToomConvolution第 218 行起只做一件事B BT.transpose() A AT.transpose()这就是 README 所说的Transformation Principle把 FIR 滤波形式的数据/逆变换矩阵交换并转置立刻得到线性卷积含边缘的变换无需重新推导。五、测试如何守护推导正确性测试文件 tests/test_wincnn.py 的策略非常硬核将 F(2,3)、F(4,3)、F(6,3) 三组 AT / G / BT 矩阵硬编码为符号期望值全部是Rational精确分数逐一断言相等test_filter_verify与test_convolution_verify则断言simplify(Y − 期望卷积式) 零矩阵——用符号减法代替数值容差。这意味着任何人修改推导逻辑导致任何一个分数分子分母变化测试都会立刻失败不存在近似通过的灰色地带。六、进阶阅读与选型建议想解决的问题去哪里看Winograd 为何比 FFT 卷积省乘法1 次/点 vs 约 1.5 次/点FAQ.md 第 1 问变换开销大总运算量真的减少吗FAQ.md 第 2 问F(4x4,3x3) 运算量核算支持 strided / dilated 卷积吗FAQ.md 后两问抽取-求和分解法基于中国剩余定理的更一般 Winograd 算法论文补充材料 2464-supp.pdf最后给三条实用建议插值点从 (0, 1, −1, 2, −2, 1/2, −1/2) 里选足够且数值行为温和变换矩阵随规模增大会变病态系数如 −21/4小尺寸 3×3 卷积是 Cook-Toom 算法的最佳用武之地部署前务必用fractionsIn挑一种分数分布让目标硬件上乘法次数最少、加法树最浅。wincnn 的全部智慧浓缩在一句话里把数值计算交给 SymPy 符号引擎用精确算术让变换矩阵的推导从手算核对变成一键生成 自动验证——这也是它不到 260 行代码就能覆盖整个推导流程的原因。【免费下载链接】wincnnWinograd minimal convolution algorithm generator for convolutional neural networks.项目地址: https://gitcode.com/gh_mirrors/wi/wincnn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表