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