Views
No views yet

git clone https://github.com/IvanDrokin/torch-conv-kan.git
cd torch-conv-kan
pip install -r requirements.txt1import torch
2from models import vggkagn
3model = vggkagn(3,
4 1000,
5 groups=1,
6 degree=5,
7 dropout=0.15,
8 l1_decay=0,
9 dropout_linear=0.25,
10 width_scale=2,
11 vgg_type='VGG11v4',
12 expected_feature_shape=(1, 1),
13 affine=True
14 )
15model.from_pretrained('brivangl/vgg_kagn11_v4')1from torchvision.transforms import v2
2transforms_val = v2.Compose([
3 v2.ToImage(),
4 v2.Resize(256, antialias=True),
5 v2.CenterCrop(224),
6 v2.ToDtype(torch.float32, scale=True),
7 v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
8 ])1{'learning_rate': 0.0009, 'adam_beta1': 0.9, 'adam_beta2': 0.999, 'adam_weight_decay': 5e-06,
2'adam_epsilon': 1e-08, 'lr_warmup_steps': 7500, 'lr_power': 0.3, 'lr_end': 1e-07, 'set_grads_to_none': False}1transforms_train = v2.Compose([
2 v2.ToImage(),
3 v2.RandomHorizontalFlip(p=0.5),
4 v2.RandomResizedCrop(224, antialias=True),
5 v2.RandomChoice([v2.AutoAugment(AutoAugmentPolicy.CIFAR10),
6 v2.AutoAugment(AutoAugmentPolicy.IMAGENET)
7 ]),
8 v2.ToDtype(torch.float32, scale=True),
9 v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
10])| Accuracy, top1 | Accuracy, top5 | AUC (ovo) | AUC (ovr) |
|---|---|---|---|
| 61.17 | 83.26 | 99.42 | 99.43 |
1@misc{torch-conv-kan,
2 author = {Ivan Drokin},
3 title = {Torch Conv KAN},
4 year = {2024},
5 publisher = {GitHub},
6 journal = {GitHub repository},
7 howpublished = {\url{https://github.com/IvanDrokin/torch-conv-kan}}
8}