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
  • 專案
  • 專案
  • 專案
  • 專案