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
  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트