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
  • 项目
  • 项目
  • 项目
  • 项目