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 における畳み込み型コルモゴロフ・アルノルドネットワーク
何がここにあるか
- コルモゴロフ・アルノルド表現(KAN)に基づく畳み込み層を実装する研究指向の PyTorch ライブラリ。固定された重み行列ではなく、各カーネルは学習可能な一変数関数(スプライン、多項式、ウェーブレットなど)の集合を格納する。
- 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
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト