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