:原理、代码对比与源码级实现)
深入解析 pykan 的稀疏初始化sparse_init原理、代码对比与源码级实现【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan本篇文章聚焦 pykan 中KAN模型的稀疏初始化机制即构造参数sparse_init。文章以官方文档 docs/Interp/Interp_11_sparse_init.rst 为核心脉络从一张 5×5×5×1 网络的初始化对比实验出发逐步剖析sparse_initTrue/False对网络连接结构的影响并下沉到 kan/KANLayer.py、kan/MultKAN.py 与 kan/utils.py 的源码说明稀疏掩码mask是如何生成、如何在初始化阶段施加到激活函数上、以及它在整个训练生命周期中保持固定的实现细节。读完本文你将能够独立复现该对比实验并理解sparse_init这一开关在代码层面的真实作用边界。实验入口5×5×5×1 网络的两种初始化对比在 docs/Interp/Interp_11_sparse_init.ipynb 中官方通过两段几乎完全一致的代码展示了sparse_init的唯一差异同一组模型超参数下一个以默认方式密集初始化创建另一个以稀疏方式创建。首先看密集初始化版本from kan import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) model KAN([5,5,5,1], sparse_initFalse, devicedevice) x torch.rand(100,5).to(device) model.get_act(x) model.plot()再来看稀疏初始化版本除sparse_initTrue外其他代码完全一致from kan import * device torch.device(cuda if torch.cuda.is_available() else cpu) print(device) model KAN([5,5,5,1], sparse_initTrue, devicedevice) x torch.rand(100,5).to(device) model.get_act(x) model.plot()两段代码的运行流程完全相同先用torch.device(cuda if torch.cuda.is_available() else cpu)自动选择计算设备在官方文档的运行输出中为cuda然后以宽度列表[5,5,5,1]构造一个 3 层4 个宽度层的 KAN 模型随后用torch.rand(100,5)生成 100 个 5 维随机样本作为输入调用model.get_act(x)驱动一次前向传播同时收集中间激活数据最后调用model.plot()将网络结构与激活函数可视化。构造模型时控制台会输出checkpoint directory created: ./model与saving model version 0.0这是 pykan 默认的自动保存auto_saveTrue行为模型初始状态会被保存为state_id0的检查点。对比如下两张图可以直观看到两种初始化在连接模式上的差异sparse_initFalse的网络中层与层之间保持稠密的连接而sparse_initTrue的网络中大量连接在初始化阶段即被“剪断”只有少数关键路径保留下来整体呈现明显的稀疏带状结构。sparse_init 参数在构造链中的传递路径sparse_init并不是KAN类自身新发明的开关而是从模型构造函数一路透传到最底层的单个 KAN 层。这条参数传递链可以从源码中完整还原模型层kan/MultKAN.py 中MultKAN.__init__的签名包含sparse_initFalse默认关闭其 docstring 明确说明sparse initialization (True) or normal dense initialization. Default: False。逐层传递在构造每一层激活函数时kan/MultKAN.py#L214 将sparse_init原样传入KANLayersp_batch KANLayer(in_dimwidth_in[l], out_dimwidth_out[l1], numgrid_l, kk_l, noise_scalenoise_scale, scale_base_muscale_base_mu, scale_base_sigmascale_base_sigma, scale_sp1., base_funbase_fun, grid_epsgrid_eps, grid_rangegrid_range, sp_trainablesp_trainable, sb_trainablesb_trainable, sparse_initsparse_init)层级实现kan/KANLayer.py#L105-L108 根据该开关决定掩码的来源if sparse_init: self.mask torch.nn.Parameter(sparse_mask(in_dim, out_dim)).requires_grad_(False) else: self.mask torch.nn.Parameter(torch.ones(in_dim, out_dim)).requires_grad_(False)从源码看两种模式下mask都被注册为torch.nn.Parameter且设置了requires_grad_(False)——即掩码本身永远不可训练。区别仅仅在于密集模式下掩码是全 1 矩阵所有连接等效保留稀疏模式下掩码由sparse_mask函数按规则生成只保留部分连接。稀疏掩码的生成算法最近邻连接规则稀疏掩码的真正实现位于 kan/utils.py#L268-L284 的sparse_mask函数def sparse_mask(in_dim, out_dim): get sparse mask in_coord torch.arange(in_dim) * 1/in_dim 1/(2*in_dim) out_coord torch.arange(out_dim) * 1/out_dim 1/(2*out_dim) dist_mat torch.abs(out_coord[:,None] - in_coord[None,:]) in_nearest torch.argmin(dist_mat, dim0) in_connection torch.stack([torch.arange(in_dim), in_nearest]).permute(1,0) out_nearest torch.argmin(dist_mat, dim1) out_connection torch.stack([out_nearest, torch.arange(out_dim)]).permute(1,0) all_connection torch.cat([in_connection, out_connection], dim0) mask torch.zeros(in_dim, out_dim) mask[all_connection[:,0], all_connection[:,1]] 1. return mask从算法结构看它实现的是一个基于坐标最近邻的带状banded稀疏连接模式输入神经元与输出神经元分别被投影到[0,1]区间内的归一化坐标每个神经元取该区间等分点的中点即in_coord与out_coord构造输入-输出之间的绝对距离矩阵dist_mat每个输入神经元找到距离自己最近的输出神经元in_nearest按行取argmin建立一条连接每个输出神经元同样找到距离自己最近的输入神经元out_nearest按列取argmin建立一条连接两者合并去重后将对应位置的掩码置为 1。可以推断该算法保证稀疏掩码中每个输入至少连接到一个输出、每个输出至少连接到一个输入从而避免出现“孤立神经元”同时把连接数从稠密的in_dim × out_dim压缩到in_dim out_dim量级。由于输入/输出按顺序编号且坐标单调排列最终保留下来的连接在拓扑上近似于一条“对角线”带状结构——这正是两张对比图中稀疏网络呈现对角线式连接的原因。当层宽不相等如 5→1 的最后一层时多个输入会共享同一个最近邻输出连接进一步汇聚。掩码如何作用于前向计算mask 的施加位置仅生成掩码还不够它必须真正参与前向计算才能改变网络行为。在 kan/KANLayer.py#L156-L166 的forward中可以看到掩码的作用位置base self.base_fun(x) # (batch, in_dim) y coef2curve(x_evalx, gridself.grid, coefself.coef, kself.k) postspline y.clone().permute(0,2,1) y self.scale_base[None,:,:] * base[:,:,None] self.scale_sp[None,:,:] * y y self.mask[None,:,:] * y postacts y.clone().permute(0,2,1) y torch.sum(y, dim1)前向过程可拆解为四步残差基函数项self.scale_base * base(x)与 B 样条项self.scale_sp * y由coef2curve依据系数coef与网格grid计算相加构成完整的激活输出掩码逐元素相乘y self.mask[None,:,:] * y被掩码置 0 的位置上对应输入-输出对的激活贡献被强制归零记录施加掩码后的postacts用于可视化与激活提取沿输入维度求和得到该层的输出y。此外掩码的影响在初始化阶段就已经渗透到可训练参数中kan/KANLayer.py#L112 在初始化scale_sp样条幅度时直接乘以self.maskself.scale_sp torch.nn.Parameter(torch.ones(in_dim, out_dim) * scale_sp * 1 / np.sqrt(in_dim) * self.mask).requires_grad_(sp_trainable)这意味着稀疏化后被掩码的连接不仅在前向时被屏蔽其样条幅度参数也被初始化为 0配合mask的持续屏蔽这些连接在整个训练过程中都不会产生梯度意义上的有效贡献——除非后续通过其他机制如剪枝/生长的重连流程显式修改掩码。需要说明的是mask本身requires_grad_(False)因此常规fit/训练循环不会改变它sparse_init描述的是“初始化阶段”的结构选择而非训练过程中的动态稀疏化。使用方式与适用边界结合官方文档与源码可以给出以下实操结论如何开启在构造KAN([...], sparse_initTrue, ...)时传入即可无需其他配套代码。由于KAN继承自MultKAN见 kan/MultKAN.py该参数同时适用于KAN与MultKAN两种入口默认行为sparse_init默认False即默认采用全连接式的密集初始化与大多数神经网络库的直觉一致与剪枝的区别sparse_init是在初始化时刻通过固定掩码直接限定初始拓扑而 pykan 文档中另一主题pruning是在训练过程中依据激活重要性动态移除连接两者在时间点、机制上均不同与网格/宽度参数的关系sparse_init独立于grid、k样条阶数、noise_scale等参数它只决定掩码的形态不改变样条本身的分辨率或阶数可视化验证构造模型后先get_act(x)再plot()即可直接观察掩码效果本文开头的两段代码即为最小可复现样例无需额外依赖。从源码结构可以进一步推断稀疏初始化保留了层间“每个神经元至少一个连接”的连通性保证这在多任务、多输出场景如宽度列表尾部出现较窄层下可避免神经元失联。至于稀疏初始化是否带来精度或训练效率的收益官方文档未给出性能数据本文也不作推测读者可在自己的数据集上通过sparse_initTrue/False的对照实验自行验证。小结sparse_init是 pykan 提供的一个轻量却影响深远的初始化开关它通过 kan/utils.py 中的sparse_mask最近邻算法生成带状稀疏掩码在 kan/KANLayer.py 的初始化阶段施加于scale_sp并固定为不可训练参数随后在每次前向传播中持续屏蔽非保留连接。理解这条从KAN([5,5,5,1], sparse_initTrue)到self.mask[None,:,:] * y的完整链路有助于你在实际建模中正确选择初始化策略也能为后续阅读 pykan 的剪枝、符号化等高级可解释性主题打下基础。【免费下载链接】pykanKolmogorov Arnold Networks项目地址: https://gitcode.com/GitHub_Trending/pyk/pykan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考