17 KiB
17 KiB
In [1]:
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import matplotlib.pyplot as plt
from pathlib import Path
import warnings
warnings.filterwarnings('ignore')
from htb_ai_library.core import set_reproducibility
from htb_ai_library.data import get_mnist_loaders
from htb_ai_library.models import SimpleLeNet
from htb_ai_library.training import train_model
from htb_ai_library.utils import save_model, load_model
from htb_ai_library.visualization import use_htb_style
use_htb_style()
set_reproducibility(1337)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {device}")
output_dir = Path('output')
output_dir.mkdir(exist_ok=True)[31m---------------------------------------------------------------------------[39m [31mRuntimeError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[1][39m[32m, line 10[39m [32m 7[39m [38;5;28;01mimport[39;00m[38;5;250m [39m[34;01mwarnings[39;00m [32m 8[39m warnings.filterwarnings([33m'[39m[33mignore[39m[33m'[39m) [32m---> [39m[32m10[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mhtb_ai_library[39;00m[34;01m.[39;00m[34;01mcore[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m set_reproducibility [32m 11[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mhtb_ai_library[39;00m[34;01m.[39;00m[34;01mdata[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m get_mnist_loaders [32m 12[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mhtb_ai_library[39;00m[34;01m.[39;00m[34;01mmodels[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m SimpleLeNet [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/htb_ai_library/__init__.py:16[39m [32m 6[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mcore[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m ( [32m 7[39m set_reproducibility, [32m 8[39m save_model, [32m (...)[39m[32m 12[39m HTB_PALETTE, get_color, get_color_palette [32m 13[39m ) [32m 15[39m [38;5;66;03m# data subpackage[39;00m [32m---> [39m[32m16[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mdata[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m ( [32m 17[39m get_mnist_loaders, [32m 18[39m download_sms_spam_dataset, [32m 19[39m mnist_denormalize, [32m 20[39m cifar_normalize, [32m 21[39m load_adult_census, [32m 22[39m get_cifar10_loaders, [32m 23[39m get_cifar10_transform, [32m 24[39m ) [32m 26[39m [38;5;66;03m# models subpackage[39;00m [32m 27[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mmodels[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m ( [32m 28[39m SimpleLeNet, [32m 29[39m SimpleCNN, [32m (...)[39m[32m 34[39m AttackModel, [32m 35[39m ) [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/htb_ai_library/data/__init__.py:5[39m [32m 1[39m [33;03m"""[39;00m [32m 2[39m [33;03mData subpackage providing loaders and dataset utilities.[39;00m [32m 3[39m [33;03m"""[39;00m [32m----> [39m[32m5[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mmnist[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m get_mnist_loaders, mnist_denormalize [32m 6[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01msms[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m download_sms_spam_dataset [32m 7[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mtransforms[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m cifar_normalize [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/htb_ai_library/data/mnist.py:11[39m [32m 9[39m [38;5;28;01mimport[39;00m[38;5;250m [39m[34;01mtorch[39;00m [32m 10[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mtorch[39;00m[34;01m.[39;00m[34;01mutils[39;00m[34;01m.[39;00m[34;01mdata[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m DataLoader [32m---> [39m[32m11[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mtorchvision[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m datasets, transforms [32m 14[39m [38;5;28;01mdef[39;00m[38;5;250m [39m[34mget_mnist_loaders[39m( [32m 15[39m batch_size: [38;5;28mint[39m = [32m128[39m, [32m 16[39m data_dir: [38;5;28mstr[39m = [33m"[39m[33m./data[39m[33m"[39m, [32m (...)[39m[32m 19[39m seed: Optional[[38;5;28mint[39m] = [32m1337[39m, [32m 20[39m ) -> Tuple[DataLoader, DataLoader]: [32m 21[39m [38;5;250m [39m[33;03m"""[39;00m [32m 22[39m [33;03m Create MNIST data loaders with optional normalization and deterministic shuffling.[39;00m [32m 23[39m [32m (...)[39m[32m 39[39m [33;03m Training and test data loaders.[39;00m [32m 40[39m [33;03m """[39;00m [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/torchvision/__init__.py:10[39m [32m 7[39m [38;5;66;03m# Don't re-order these, we need to load the _C extension (done when importing[39;00m [32m 8[39m [38;5;66;03m# .extensions) before entering _meta_registrations.[39;00m [32m 9[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mextension[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m _HAS_OPS [38;5;66;03m# usort:skip[39;00m [32m---> [39m[32m10[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01mtorchvision[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m _meta_registrations, datasets, io, models, ops, transforms, utils [38;5;66;03m# usort:skip[39;00m [32m 12[39m [38;5;28;01mtry[39;00m: [32m 13[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01m.[39;00m[34;01mversion[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m __version__ [38;5;66;03m# noqa: F401[39;00m [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/torchvision/_meta_registrations.py:163[39m [32m 153[39m torch._check( [32m 154[39m grad.dtype == rois.dtype, [32m 155[39m [38;5;28;01mlambda[39;00m: ( [32m (...)[39m[32m 158[39m ), [32m 159[39m ) [32m 160[39m [38;5;28;01mreturn[39;00m grad.new_empty((batch_size, channels, height, width)) [32m--> [39m[32m163[39m [38;5;129;43m@torch[39;49m[43m.[49m[43mlibrary[49m[43m.[49m[43mregister_fake[49m[43m([49m[33;43m"[39;49m[33;43mtorchvision::nms[39;49m[33;43m"[39;49m[43m)[49m [32m 164[39m [38;5;28;43;01mdef[39;49;00m[38;5;250;43m [39;49m[34;43mmeta_nms[39;49m[43m([49m[43mdets[49m[43m,[49m[43m [49m[43mscores[49m[43m,[49m[43m [49m[43miou_threshold[49m[43m)[49m[43m:[49m [32m 165[39m [43m [49m[43mtorch[49m[43m.[49m[43m_check[49m[43m([49m[43mdets[49m[43m.[49m[43mdim[49m[43m([49m[43m)[49m[43m [49m[43m==[49m[43m [49m[32;43m2[39;49m[43m,[49m[43m [49m[38;5;28;43;01mlambda[39;49;00m[43m:[49m[43m [49m[33;43mf[39;49m[33;43m"[39;49m[33;43mboxes should be a 2d tensor, got [39;49m[38;5;132;43;01m{[39;49;00m[43mdets[49m[43m.[49m[43mdim[49m[43m([49m[43m)[49m[38;5;132;43;01m}[39;49;00m[33;43mD[39;49m[33;43m"[39;49m[43m)[49m [32m 166[39m [43m [49m[43mtorch[49m[43m.[49m[43m_check[49m[43m([49m[43mdets[49m[43m.[49m[43msize[49m[43m([49m[32;43m1[39;49m[43m)[49m[43m [49m[43m==[49m[43m [49m[32;43m4[39;49m[43m,[49m[43m [49m[38;5;28;43;01mlambda[39;49;00m[43m:[49m[43m [49m[33;43mf[39;49m[33;43m"[39;49m[33;43mboxes should have 4 elements in dimension 1, got [39;49m[38;5;132;43;01m{[39;49;00m[43mdets[49m[43m.[49m[43msize[49m[43m([49m[32;43m1[39;49m[43m)[49m[38;5;132;43;01m}[39;49;00m[33;43m"[39;49m[43m)[49m [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/torch/library.py:1073[39m, in [36mregister_fake.<locals>.register[39m[34m(func)[39m [32m 1071[39m [38;5;28;01melse[39;00m: [32m 1072[39m use_lib = lib [32m-> [39m[32m1073[39m [43muse_lib[49m[43m.[49m[43m_register_fake[49m[43m([49m [32m 1074[39m [43m [49m[43mop_name[49m[43m,[49m[43m [49m[43mfunc[49m[43m,[49m[43m [49m[43m_stacklevel[49m[43m=[49m[43mstacklevel[49m[43m [49m[43m+[49m[43m [49m[32;43m1[39;49m[43m,[49m[43m [49m[43mallow_override[49m[43m=[49m[43mallow_override[49m [32m 1075[39m [43m[49m[43m)[49m [32m 1076[39m [38;5;28;01mreturn[39;00m func [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/torch/library.py:203[39m, in [36mLibrary._register_fake[39m[34m(self, op_name, fn, _stacklevel, allow_override)[39m [32m 200[39m [38;5;28;01melse[39;00m: [32m 201[39m func_to_register = fn [32m--> [39m[32m203[39m handle = [43mentry[49m[43m.[49m[43mfake_impl[49m[43m.[49m[43mregister[49m[43m([49m [32m 204[39m [43m [49m[43mfunc_to_register[49m[43m,[49m[43m [49m[43msource[49m[43m,[49m[43m [49m[43mlib[49m[43m=[49m[38;5;28;43mself[39;49m[43m,[49m[43m [49m[43mallow_override[49m[43m=[49m[43mallow_override[49m [32m 205[39m [43m[49m[43m)[49m [32m 206[39m [38;5;28mself[39m._registration_handles.append(handle) [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/torch/_library/fake_impl.py:50[39m, in [36mFakeImplHolder.register[39m[34m(self, func, source, lib, allow_override)[39m [32m 44[39m [38;5;28;01mif[39;00m [38;5;28mself[39m.kernel [38;5;129;01mis[39;00m [38;5;129;01mnot[39;00m [38;5;28;01mNone[39;00m: [32m 45[39m [38;5;28;01mraise[39;00m [38;5;167;01mRuntimeError[39;00m( [32m 46[39m [33mf[39m[33m"[39m[33mregister_fake(...): the operator [39m[38;5;132;01m{[39;00m[38;5;28mself[39m.qualname[38;5;132;01m}[39;00m[33m [39m[33m"[39m [32m 47[39m [33mf[39m[33m"[39m[33malready has an fake impl registered at [39m[33m"[39m [32m 48[39m [33mf[39m[33m"[39m[38;5;132;01m{[39;00m[38;5;28mself[39m.kernel.source[38;5;132;01m}[39;00m[33m.[39m[33m"[39m [32m 49[39m ) [32m---> [39m[32m50[39m [38;5;28;01mif[39;00m [43mtorch[49m[43m.[49m[43m_C[49m[43m.[49m[43m_dispatch_has_kernel_for_dispatch_key[49m[43m([49m[38;5;28;43mself[39;49m[43m.[49m[43mqualname[49m[43m,[49m[43m [49m[33;43m"[39;49m[33;43mMeta[39;49m[33;43m"[39;49m[43m)[49m: [32m 51[39m [38;5;28;01mraise[39;00m [38;5;167;01mRuntimeError[39;00m( [32m 52[39m [33mf[39m[33m"[39m[33mregister_fake(...): the operator [39m[38;5;132;01m{[39;00m[38;5;28mself[39m.qualname[38;5;132;01m}[39;00m[33m [39m[33m"[39m [32m 53[39m [33mf[39m[33m"[39m[33malready has an DispatchKey::Meta implementation via a [39m[33m"[39m [32m (...)[39m[32m 56[39m [33mf[39m[33m"[39m[33mregister_fake.[39m[33m"[39m [32m 57[39m ) [32m 59[39m [38;5;28;01mif[39;00m torch._C._dispatch_has_kernel_for_dispatch_key( [32m 60[39m [38;5;28mself[39m.qualname, [33m"[39m[33mCompositeImplicitAutograd[39m[33m"[39m [32m 61[39m ): [31mRuntimeError[39m: operator torchvision::nms does not exist
In [ ]: