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_bn(
6 3,
7 1000,
8 groups=1,
9 degree=5,
10 dropout= 0.05,
11 l1_decay=0,
12 width_scale=2,
13 affine=True,
14 norm_layer=nn.BatchNorm2d,
15 expected_feature_shape=(1, 1),
16 vgg_type='VGG11v4',
17 last_attention=True,
18 sa_inner_projection=None
19)
20
21model.from_pretrained('brivangl/vgg_kagn_bn11sa_v4')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) |
|---|---|---|---|
| 70.684 | 89.462 | 99.624 | 99.624 |
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}