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
3
4
5model = vggkagn(3,
6 1000,
7 groups=1,
8 degree=5,
9 dropout=0.15,
10 l1_decay=0,
11 dropout_linear=0.25,
12 width_scale=2,
13 vgg_type='VGG11v2',
14 expected_feature_shape=(1, 1),
15 affine=True
16 )
17
18model.from_pretrained('brivangl/vgg_kagn11_v2')1from torchvision.transforms import v2
2
3
4transforms_val = v2.Compose([
5 v2.ToImage(),
6 v2.Resize(256, antialias=True),
7 v2.CenterCrop(224),
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 ])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) |
|---|---|---|---|
| 59.1 | 82.29 | 99.43 | 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}