16 KiB
16 KiB
In [1]:
# no torch no example womp womp
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from pathlib import Path
import warnings
warnings.filterwarnings('ignore')
# Import utilities from HTB Evasion Library
from htb_ai_library.utils import (
set_reproducibility,
save_model,
load_model,
HTB_GREEN,
NODE_BLACK,
HACKER_GREY,
WHITE,
AZURE,
NUGGET_YELLOW,
MALWARE_RED,
VIVID_PURPLE,
AQUAMARINE,
)
from htb_ai_library.data import get_mnist_loaders
from htb_ai_library.models import MNISTClassifierWithDropout
from htb_ai_library.training import train_model, evaluate_accuracy
from htb_ai_library.visualization import use_htb_style
# Apply HTB theme globally to all plots
use_htb_style()
# Set reproducibility
set_reproducibility(1337)
# Configure device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
if device.type == "cuda":
print(f"GPU: {torch.cuda.get_device_name(0)}")[31m---------------------------------------------------------------------------[39m [31mRuntimeError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[1][39m[32m, line 5[39m [32m 3[39m [38;5;28;01mimport[39;00m[38;5;250m [39m[34;01mtorch[39;00m[34;01m.[39;00m[34;01mnn[39;00m[34;01m.[39;00m[34;01mfunctional[39;00m[38;5;250m [39m[38;5;28;01mas[39;00m[38;5;250m [39m[34;01mF[39;00m [32m 4[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[32m5[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 6[39m [38;5;28;01mimport[39;00m[38;5;250m [39m[34;01mnumpy[39;00m[38;5;250m [39m[38;5;28;01mas[39;00m[38;5;250m [39m[34;01mnp[39;00m [32m 7[39m [38;5;28;01mimport[39;00m[38;5;250m [39m[34;01mmatplotlib[39;00m[34;01m.[39;00m[34;01mpyplot[39;00m[38;5;250m [39m[38;5;28;01mas[39;00m[38;5;250m [39m[34;01mplt[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 [2]:
# Get data loaders using library function
train_loader, test_loader = get_mnist_loaders(batch_size=128)
print(f"Training samples: {len(train_loader.dataset)}")
print(f"Test samples: {len(test_loader.dataset)}")
# Create output directory for saving models and results
output_dir = Path("output")
output_dir.mkdir(exist_ok=True)
# Define model checkpoint path in output directory
model_path = output_dir / "mnist_target.pth"
# Initialize model using MNISTClassifierWithDropout from library
model = MNISTClassifierWithDropout(num_classes=10).to(device)
# Check if trained model exists, otherwise train from scratch
if model_path.exists():
print(f"\nLoading existing model from {model_path}")
model = load_model(model, model_path, device)
else:
print(f"\nNo existing model found. Training new model...")
model = train_model(model, train_loader, test_loader, epochs=5, device=device)
print(f"Saving trained model to {model_path}")
save_model(model, model_path)
# Evaluate the trained model
accuracy = evaluate_accuracy(model, test_loader, device)
print(f"\nTest accuracy: {accuracy:.2f}%")
[31m---------------------------------------------------------------------------[39m [31mNameError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[2][39m[32m, line 2[39m [32m 1[39m [38;5;66;03m# Get data loaders using library function[39;00m [32m----> [39m[32m2[39m train_loader, test_loader = [43mget_mnist_loaders[49m(batch_size=[32m128[39m) [32m 3[39m [38;5;28mprint[39m([33mf[39m[33m"[39m[33mTraining samples: [39m[38;5;132;01m{[39;00m[38;5;28mlen[39m(train_loader.dataset)[38;5;132;01m}[39;00m[33m"[39m) [32m 4[39m [38;5;28mprint[39m([33mf[39m[33m"[39m[33mTest samples: [39m[38;5;132;01m{[39;00m[38;5;28mlen[39m(test_loader.dataset)[38;5;132;01m}[39;00m[33m"[39m) [31mNameError[39m: name 'get_mnist_loaders' is not defined
In [ ]: