IvanDrokin/torch-conv-kan
This project is dedicated to the implementation and research of Kolmogorov-Arnold convolutional networks. The repository includes implementations of 1D, 2D, and 3D convolutions with different kernels, ResNet-like and DenseNet-like models, training code based on accelerate/PyTorch, as well as scripts for experiments with CIFAR-10 and Tiny ImageNet.
TorchConv‑KAN – PyTorch 中的卷积型柯尔莫哥洛夫-阿诺德网络
这是什么
- 一个研究导向的 PyTorch 库,实现了基于柯尔莫哥洛夫-阿诺德表示(KAN)的卷积层。与固定权重矩阵不同,每个卷积核存储一组可学习的一元函数(样条、多项式、小波等)。
- 提供不断增长的变体集合——KAN、Fast‑KAN、KALN、KAGN、ChebyKAN、WavKAN、JacobiKAN、Bernstein‑KAN、ReLU‑KAN,以及瓶颈版本——适用于 1‑D、2‑D 和 3‑D 卷积。
- 将这些层打包成熟悉的 CNN 主干网络(ResNet‑like、DenseNet‑like、VGG‑like、U‑Net/U2‑Net),可轻松将 Conv‑KAN 块插入现有架构中。
主要特性
- 层动物园 – 即用型
KANConv*、FastKANConv*、KALNConv*、KAGNConv*、WavKANConv*、JacobiKANConv*、BernsteinKANConv*、ReLUKANConv*及其瓶颈等价物。 - 模型动物园 – ResKANet、DenseKANet、VGG‑KAN、UKANet/U2KANet,以及基于这些层的新型 ConvNeXt‑style 块。
- 预训练检查点 – 多个 VGG‑KAN 变体(如 VGG‑KAGN‑11‑BN 仅用 7.25 M 参数即达到 68.5 % top‑1)和近期 ConvNeXt‑KAGN 模型的 ImageNet‑1k 权重。
- 训练工具 – 支持 MNIST、CIFAR‑10/100、Tiny‑Imagenet 和 ImageNet‑1k 的训练脚本;与 🤗 Accelerate、Hydra 配置和 Weights & Biases 日志集成。
- 量化与 PEFT – 支持 KAN 基模型的后训练量化和参数高效微调(PEFT)。
- 超参数搜索 – 提供 Ray‑Tune 包装器用于自动调参;可选 LBFGS 优化器。
典型用例
import torch, torch.nn as nn
from kan_convs import KANConv2DLayer
class SimpleConvKAN(nn.Module):
def __init__(self, layer_sizes, num_classes=10, in_ch=1, spline_order=3):
super().__init__()
self.features = nn.Sequential(
KANConv2DLayer(in_ch, layer_sizes[0], spline_order, kernel_size=3, padding=1),
KANConv2DLayer(layer_sizes[0], layer_sizes[1], spline_order, kernel_size=3, stride=2, padding=1),
KANConv2DLayer(layer_sizes[1], layer_sizes[2], spline_order, kernel_size=3, stride=2, padding=1),
KANConv2DLayer(layer_sizes[2], layer_sizes[3], spline_order, kernel_size=3, padding=1),
nn.AdaptiveAvgPool2d(1),
)
self.head = nn.Linear(layer_sizes[3], num_classes)
self.drop = nn.Dropout(0.25)
def forward(self, x):
x = self.features(x)
x = torch.flatten(x, 1)
x = self.drop(x)
return self.head(x)
运行提供的训练脚本(如 python mnist_conv.py)可在 MNIST、CIFAR‑10/100 上训练/评估,或切换至基于 Accelerate 的脚本以处理更大数据集。
安装与快速入门
git clone https://github.com/IvanDrokin/torch-conv-kan.git
cd torch-conv-kan
pip install -r requirements.txt # PyTorch + CUDA, accelerate, hydra, wandb, ray[tune]
# 可选:wandb login # 用于实验追踪
accelerate launch cifar.py # 使用默认配置在 CIFAR‑10 上训练 ResKANet
项目状态
- 持续开发中(2024年5月至7月每日更新,另含2026年7月更新)。
- 核心研究论文已发布:Kolmogorov‑Arnold Convolutions: Design Principles and Empirical Studies (arXiv 2407.01092)。
- 基准测试仍处于初步阶段;作者指出,CIFAR‑10/100 上的性能落后于传统 CNN,且部分变体(如 ChebyKAN)存在稳定性问题。
谁可能用得上
- 探索卷积核内可学习激活函数的研究人员。
- 希望在视觉模型中实验 KAN 风格层但不想从头构建的实践者。
- 对非标准卷积架构的量化或 PEFT 感兴趣的人。
引用 若使用代码或基线结果,请引用随附的 arXiv 论文(BibTeX 已在仓库中提供)。
以上所有信息均直接取自仓库的 README;未添加任何外部假设。
相关
- Dispatch
- 项目
- 项目
- 项目
- 项目