JimmyChin1998/Pytorch-Learning-File
0
1{2 "cells": [3 {4 "cell_type": "code",5 "execution_count": 11,6 "id": "42b32900-93af-4b4d-abea-eb5d7b771b9e",7 "metadata": {},8 "outputs": [9 {10 "name": "stdout",11 "output_type": "stream",12 "text": [13 "[INFO] torch/torchvision versions not as required, installing nightly versions.\n",14 "Looking in indexes: https://pypi.org/simple, https://download.pytorch.org/whl/cu113\n",15 "Requirement already satisfied: torch in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (2.5.1)\n",16 "Requirement already satisfied: torchvision in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (0.20.1)\n",17 "Requirement already satisfied: torchaudio in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (2.5.1)\n",18 "Requirement already satisfied: filelock in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (3.13.1)\n",19 "Requirement already satisfied: typing-extensions>=4.8.0 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (4.11.0)\n",20 "Requirement already satisfied: networkx in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (3.2.1)\n",21 "Requirement already satisfied: jinja2 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (3.1.4)\n",22 "Requirement already satisfied: fsspec in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (2024.10.0)\n",23 "Requirement already satisfied: setuptools in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (75.1.0)\n",24 "Requirement already satisfied: sympy==1.13.1 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torch) (1.13.1)\n",25 "Requirement already satisfied: mpmath<1.4,>=1.1.0 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from sympy==1.13.1->torch) (1.3.0)\n",26 "Requirement already satisfied: numpy in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torchvision) (1.26.4)\n",27 "Requirement already satisfied: pillow!=8.3.*,>=5.3.0 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from torchvision) (10.4.0)\n",28 "Requirement already satisfied: MarkupSafe>=2.0 in c:\\users\\user\\desktop\\pytorch\\lib\\site-packages (from jinja2->torch) (2.1.3)\n",29 "torch version: 2.5.0\n",30 "torchvision version: 0.20.0\n"31 ]32 }33 ],34 "source": [35 "# For this notebook to run with updated APIs, we need torch 1.12+ and torchvision 0.13+\n",36 "try:\n",37 " import torch\n",38 " import torchvision\n",39 " assert int(torch.__version__.split(\".\")[1]) >= 12, \"torch version should be 1.12+\"\n",40 " assert int(torchvision.__version__.split(\".\")[1]) >= 13, \"torchvision version should be 0.13+\"\n",41 " print(f\"torch version: {torch.__version__}\")\n",42 " print(f\"torchvision version: {torchvision.__version__}\")\n",43 "except:\n",44 " print(f\"[INFO] torch/torchvision versions not as required, installing nightly versions.\")\n",45 " !pip3 install -U torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cu113\n",46 " import torch\n",47 " import torchvision\n",48 " print(f\"torch version: {torch.__version__}\")\n",49 " print(f\"torchvision version: {torchvision.__version__}\")"50 ]51 },52 {53 "cell_type": "code",54 "execution_count": 12,55 "id": "78c413d7-48b6-4638-bdef-edd7077d6ef0",56 "metadata": {},57 "outputs": [58 {59 "name": "stdout",60 "output_type": "stream",61 "text": [62 "torch version: 2.5.0\n",63 "torchvision version: 0.20.0\n"64 ]65 }66 ],67 "source": [68 "import torch\n",69 "import torchvision\n",70 "print(f\"torch version: {torch.__version__}\")\n",71 "print(f\"torchvision version: {torchvision.__version__}\")"72 ]73 },74 {75 "cell_type": "code",76 "execution_count": 13,77 "id": "d9cc9181-3b9f-4391-91e7-6b8c76f870d2",78 "metadata": {},79 "outputs": [],80 "source": [81 "# Continue with regular imports\n",82 "import matplotlib.pyplot as plt\n",83 "import torch\n",84 "import torchvision\n",85 "\n",86 "from torch import nn\n",87 "from torchvision import transforms\n",88 "\n",89 "# Try to get torchinfo, install it if it doesn't work\n",90 "try:\n",91 " from torchinfo import summary\n",92 "except:\n",93 " print(\"[INFO] Couldn't find torchinfo... installing it.\")\n",94 " !pip install -q torchinfo\n",95 " from torchinfo import summary"96 ]97 },98 {99 "cell_type": "code",100 "execution_count": 14,101 "id": "07f97bae-b0c0-4e81-ac62-8cb11f5002d4",102 "metadata": {},103 "outputs": [],104 "source": [105 "import sys\n",106 "# sys.path.append(r'C:\\Users\\User\\Desktop\\Pytorch\\pytorchPractice') # 添加父路徑到系統路徑中\n",107 "sys.path.append(r'C:\\Users\\jimmychin\\Desktop\\Pytorch\\pytorchPractice') # 添加父路徑到系統路徑中\n",108 "from going_modular import data_setup, engine"109 ]110 },111 {112 "cell_type": "code",113 "execution_count": 15,114 "id": "4b5d63db-a22e-4ec4-a889-3df40bc1b7ee",115 "metadata": {},116 "outputs": [117 {118 "data": {119 "text/plain": [120 "'cuda'"121 ]122 },123 "execution_count": 15,124 "metadata": {},125 "output_type": "execute_result"126 }127 ],128 "source": [129 "# Setup device agnostic code\n",130 "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",131 "device"132 ]133 },134 {135 "cell_type": "code",136 "execution_count": 16,137 "id": "513a71fd-1479-4fae-bb5f-1c5773cd9450",138 "metadata": {},139 "outputs": [140 {141 "name": "stdout",142 "output_type": "stream",143 "text": [144 "data\\pizza_steak_sushi directory exists.\n"145 ]146 }147 ],148 "source": [149 "import os\n",150 "import zipfile\n",151 "\n",152 "from pathlib import Path\n",153 "\n",154 "import requests\n",155 "\n",156 "# Setup path to data folder\n",157 "data_path = Path(\"data/\")\n",158 "image_path = data_path / \"pizza_steak_sushi\"\n",159 "\n",160 "# If the image folder doesn't exist, download it and prepare it... \n",161 "if image_path.is_dir():\n",162 " print(f\"{image_path} directory exists.\")\n",163 "else:\n",164 " print(f\"Did not find {image_path} directory, creating one...\")\n",165 " image_path.mkdir(parents=True, exist_ok=True)\n",166 " \n",167 " # Download pizza, steak, sushi data\n",168 " with open(data_path / \"pizza_steak_sushi.zip\", \"wb\") as f:\n",169 " request = requests.get(\"https://github.com/mrdbourke/pytorch-deep-learning/raw/main/data/pizza_steak_sushi.zip\")\n",170 " print(\"Downloading pizza, steak, sushi data...\")\n",171 " f.write(request.content)\n",172 "\n",173 " # Unzip pizza, steak, sushi data\n",174 " with zipfile.ZipFile(data_path / \"pizza_steak_sushi.zip\", \"r\") as zip_ref:\n",175 " print(\"Unzipping pizza, steak, sushi data...\") \n",176 " zip_ref.extractall(image_path)\n",177 "\n",178 " # Remove .zip file\n",179 " os.remove(data_path / \"pizza_steak_sushi.zip\")"180 ]181 },182 {183 "cell_type": "code",184 "execution_count": 17,185 "id": "30882980-ba49-4c11-8262-4915daecb3f8",186 "metadata": {},187 "outputs": [],188 "source": [189 "# Setup Dirs\n",190 "train_dir = image_path / \"train\"\n",191 "test_dir = image_path / \"test\""192 ]193 },194 {195 "cell_type": "code",196 "execution_count": 18,197 "id": "b77a8391-f2ab-4afe-8bed-1c7565d86326",198 "metadata": {},199 "outputs": [],200 "source": [201 "# Create a transforms pipeline manually (required for torchvision < 0.13)\n",202 "manual_transforms = transforms.Compose([\n",203 " transforms.Resize((224, 224)), # 1. Reshape all images to 224x224 (though some models may require different sizes)\n",204 " transforms.ToTensor(), # 2. Turn image values to between 0 & 1 \n",205 " transforms.Normalize(mean=[0.485, 0.456, 0.406], # 3. A mean of [0.485, 0.456, 0.406] (across each colour channel)\n",206 " std=[0.229, 0.224, 0.225]) # 4. A standard deviation of [0.229, 0.224, 0.225] (across each colour channel),\n",207 "])"208 ]209 },210 {211 "cell_type": "code",212 "execution_count": 19,213 "id": "afa23646-69e8-4922-b29c-f80f7bc7abd7",214 "metadata": {},215 "outputs": [216 {217 "data": {218 "text/plain": [219 "(<torch.utils.data.dataloader.DataLoader at 0x27c86b37200>,\n",220 " <torch.utils.data.dataloader.DataLoader at 0x27c86b37260>,\n",221 " ['pizza', 'steak', 'sushi'])"222 ]223 },224 "execution_count": 19,225 "metadata": {},226 "output_type": "execute_result"227 }228 ],229 "source": [230 "# Create training and testing DataLoaders as well as get a list of class names\n",231 "train_dataloader, test_dataloader, class_names = data_setup.create_dataloaders(train_dir=train_dir,\n",232 " test_dir=test_dir,\n",233 " transform=manual_transforms, # resize, convert images to between 0 & 1 and normalize them\n",234 " batch_size=32) # set mini-batch size to 32\n",235 "\n",236 "train_dataloader, test_dataloader, class_names"237 ]238 },239 {240 "cell_type": "code",241 "execution_count": 20,242 "id": "720bf99f-ff8d-4ea2-b452-8aaabf80071a",243 "metadata": {},244 "outputs": [245 {246 "data": {247 "text/plain": [248 "EfficientNet_B0_Weights.IMAGENET1K_V1"249 ]250 },251 "execution_count": 20,252 "metadata": {},253 "output_type": "execute_result"254 }255 ],256 "source": [257 "# Get a set of pretrained model weights\n",258 "weights = torchvision.models.EfficientNet_B0_Weights.DEFAULT # .DEFAULT = best available weights from pretraining on ImageNet\n",259 "weights"260 ]261 },262 {263 "cell_type": "code",264 "execution_count": 21,265 "id": "b31e1eea-f43c-4afc-be64-32acad4b6e2c",266 "metadata": {},267 "outputs": [268 {269 "data": {270 "text/plain": [271 "ImageClassification(\n",272 " crop_size=[224]\n",273 " resize_size=[256]\n",274 " mean=[0.485, 0.456, 0.406]\n",275 " std=[0.229, 0.224, 0.225]\n",276 " interpolation=InterpolationMode.BICUBIC\n",277 ")"278 ]279 },280 "execution_count": 21,281 "metadata": {},282 "output_type": "execute_result"283 }284 ],285 "source": [286 "# Get the transforms used to create our pretrained weights\n",287 "auto_transforms = weights.transforms()\n",288 "auto_transforms"289 ]290 },291 {292 "cell_type": "code",293 "execution_count": 22,294 "id": "eb296a2b-d133-416c-a576-cf0bf62aaf4a",295 "metadata": {},296 "outputs": [297 {298 "data": {299 "text/plain": [300 "(<torch.utils.data.dataloader.DataLoader at 0x27c86b379b0>,\n",301 " <torch.utils.data.dataloader.DataLoader at 0x27c86b37410>,\n",302 " ['pizza', 'steak', 'sushi'])"303 ]304 },305 "execution_count": 22,306 "metadata": {},307 "output_type": "execute_result"308 }309 ],310 "source": [311 "# Create training and testing DataLoaders as well as get a list of class names\n",312 "train_dataloader, test_dataloader, class_names = data_setup.create_dataloaders(train_dir=train_dir,\n",313 " test_dir=test_dir,\n",314 " transform=auto_transforms, # perform same data transforms on our own data as the pretrained model\n",315 " batch_size=32) # set mini-batch size to 32\n",316 "\n",317 "train_dataloader, test_dataloader, class_names"318 ]319 },320 {321 "cell_type": "code",322 "execution_count": 23,323 "id": "2ae81537-7302-45db-ae48-85c8e21f22b4",324 "metadata": {},325 "outputs": [326 {327 "name": "stdout",328 "output_type": "stream",329 "text": [330 "<class 'torch.utils.data.dataloader.DataLoader'>\n",331 "<class 'list'>\n"332 ]333 }334 ],335 "source": [336 "print(type(train_dataloader))\n",337 "print(type(class_names))"338 ]339 },340 {341 "cell_type": "code",342 "execution_count": 24,343 "id": "c67ef5a6-7dbb-46b6-8687-cbec0b9289eb",344 "metadata": {},345 "outputs": [346 {347 "name": "stderr",348 "output_type": "stream",349 "text": [350 "Downloading: \"https://download.pytorch.org/models/efficientnet_b0_rwightman-7f5810bc.pth\" to C:\\Users\\User/.cache\\torch\\hub\\checkpoints\\efficientnet_b0_rwightman-7f5810bc.pth\n",351 "100%|█████████████████████████████████████████████████████████| 20.5M/20.5M [00:01<00:00, 11.4MB/s]\n"352 ]353 }354 ],355 "source": [356 "# OLD: Setup the model with pretrained weights and send it to the target device (this was prior to torchvision v0.13)\n",357 "# model = torchvision.models.efficientnet_b0(pretrained=True).to(device) # OLD method (with pretrained=True)\n",358 "\n",359 "# NEW: Setup the model with pretrained weights and send it to the target device (torchvision v0.13+)\n",360 "weights = torchvision.models.EfficientNet_B0_Weights.DEFAULT # .DEFAULT = best available weights \n",361 "model = torchvision.models.efficientnet_b0(weights=weights).to(device)\n",362 "\n",363 "#model # uncomment to output (it's very long)"364 ]365 },366 {367 "cell_type": "code",368 "execution_count": 25,369 "id": "1357403b-b623-47e3-a43a-65d55ce32c9a",370 "metadata": {},371 "outputs": [372 {373 "data": {374 "text/plain": [375 "'cuda'"376 ]377 },378 "execution_count": 25,379 "metadata": {},380 "output_type": "execute_result"381 }382 ],383 "source": [384 "device"385 ]386 },387 {388 "cell_type": "code",389 "execution_count": 26,390 "id": "b8e9c5fb-e1f4-48c2-9a79-0e6e690de4d6",391 "metadata": {},392 "outputs": [393 {394 "data": {395 "text/plain": [396 "============================================================================================================================================\n",397 "Layer (type (var_name)) Input Shape Output Shape Param # Trainable\n",398 "============================================================================================================================================\n",399 "EfficientNet (EfficientNet) [32, 3, 224, 224] [32, 1000] -- True\n",400 "├─Sequential (features) [32, 3, 224, 224] [32, 1280, 7, 7] -- True\n",401 "│ └─Conv2dNormActivation (0) [32, 3, 224, 224] [32, 32, 112, 112] -- True\n",402 "│ │ └─Conv2d (0) [32, 3, 224, 224] [32, 32, 112, 112] 864 True\n",403 "│ │ └─BatchNorm2d (1) [32, 32, 112, 112] [32, 32, 112, 112] 64 True\n",404 "│ │ └─SiLU (2) [32, 32, 112, 112] [32, 32, 112, 112] -- --\n",405 "│ └─Sequential (1) [32, 32, 112, 112] [32, 16, 112, 112] -- True\n",406 "│ │ └─MBConv (0) [32, 32, 112, 112] [32, 16, 112, 112] 1,448 True\n",407 "│ └─Sequential (2) [32, 16, 112, 112] [32, 24, 56, 56] -- True\n",408 "│ │ └─MBConv (0) [32, 16, 112, 112] [32, 24, 56, 56] 6,004 True\n",409 "│ │ └─MBConv (1) [32, 24, 56, 56] [32, 24, 56, 56] 10,710 True\n",410 "│ └─Sequential (3) [32, 24, 56, 56] [32, 40, 28, 28] -- True\n",411 "│ │ └─MBConv (0) [32, 24, 56, 56] [32, 40, 28, 28] 15,350 True\n",412 "│ │ └─MBConv (1) [32, 40, 28, 28] [32, 40, 28, 28] 31,290 True\n",413 "│ └─Sequential (4) [32, 40, 28, 28] [32, 80, 14, 14] -- True\n",414 "│ │ └─MBConv (0) [32, 40, 28, 28] [32, 80, 14, 14] 37,130 True\n",415 "│ │ └─MBConv (1) [32, 80, 14, 14] [32, 80, 14, 14] 102,900 True\n",416 "│ │ └─MBConv (2) [32, 80, 14, 14] [32, 80, 14, 14] 102,900 True\n",417 "│ └─Sequential (5) [32, 80, 14, 14] [32, 112, 14, 14] -- True\n",418 "│ │ └─MBConv (0) [32, 80, 14, 14] [32, 112, 14, 14] 126,004 True\n",419 "│ │ └─MBConv (1) [32, 112, 14, 14] [32, 112, 14, 14] 208,572 True\n",420 "│ │ └─MBConv (2) [32, 112, 14, 14] [32, 112, 14, 14] 208,572 True\n",421 "│ └─Sequential (6) [32, 112, 14, 14] [32, 192, 7, 7] -- True\n",422 "│ │ └─MBConv (0) [32, 112, 14, 14] [32, 192, 7, 7] 262,492 True\n",423 "│ │ └─MBConv (1) [32, 192, 7, 7] [32, 192, 7, 7] 587,952 True\n",424 "│ │ └─MBConv (2) [32, 192, 7, 7] [32, 192, 7, 7] 587,952 True\n",425 "│ │ └─MBConv (3) [32, 192, 7, 7] [32, 192, 7, 7] 587,952 True\n",426 "│ └─Sequential (7) [32, 192, 7, 7] [32, 320, 7, 7] -- True\n",427 "│ │ └─MBConv (0) [32, 192, 7, 7] [32, 320, 7, 7] 717,232 True\n",428 "│ └─Conv2dNormActivation (8) [32, 320, 7, 7] [32, 1280, 7, 7] -- True\n",429 "│ │ └─Conv2d (0) [32, 320, 7, 7] [32, 1280, 7, 7] 409,600 True\n",430 "│ │ └─BatchNorm2d (1) [32, 1280, 7, 7] [32, 1280, 7, 7] 2,560 True\n",431 "│ │ └─SiLU (2) [32, 1280, 7, 7] [32, 1280, 7, 7] -- --\n",432 "├─AdaptiveAvgPool2d (avgpool) [32, 1280, 7, 7] [32, 1280, 1, 1] -- --\n",433 "├─Sequential (classifier) [32, 1280] [32, 1000] -- True\n",434 "│ └─Dropout (0) [32, 1280] [32, 1280] -- --\n",435 "│ └─Linear (1) [32, 1280] [32, 1000] 1,281,000 True\n",436 "============================================================================================================================================\n",437 "Total params: 5,288,548\n",438 "Trainable params: 5,288,548\n",439 "Non-trainable params: 0\n",440 "Total mult-adds (Units.GIGABYTES): 12.35\n",441 "============================================================================================================================================\n",442 "Input size (MB): 19.27\n",443 "Forward/backward pass size (MB): 3452.35\n",444 "Params size (MB): 21.15\n",445 "Estimated Total Size (MB): 3492.77\n",446 "============================================================================================================================================"447 ]448 },449 "execution_count": 26,450 "metadata": {},451 "output_type": "execute_result"452 }453 ],454 "source": [455 "# Print a summary using torchinfo (uncomment for actual output)\n",456 "summary(model=model, \n",457 " input_size=(32, 3, 224, 224), # make sure this is \"input_size\", not \"input_shape\"\n",458 " # col_names=[\"input_size\"], # uncomment for smaller output\n",459 " col_names=[\"input_size\", \"output_size\", \"num_params\", \"trainable\"],\n",460 " col_width=20,\n",461 " row_settings=[\"var_names\"]\n",462 ") "463 ]464 },465 {466 "cell_type": "code",467 "execution_count": 27,468 "id": "182603d9-c801-4df1-84bd-6f35b123c506",469 "metadata": {},470 "outputs": [471 {472 "name": "stdout",473 "output_type": "stream",474 "text": [475 "EfficientNet(\n",476 " (features): Sequential(\n",477 " (0): Conv2dNormActivation(\n",478 " (0): Conv2d(3, 32, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n",479 " (1): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",480 " (2): SiLU(inplace=True)\n",481 " )\n",482 " (1): Sequential(\n",483 " (0): MBConv(\n",484 " (block): Sequential(\n",485 " (0): Conv2dNormActivation(\n",486 " (0): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), groups=32, bias=False)\n",487 " (1): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",488 " (2): SiLU(inplace=True)\n",489 " )\n",490 " (1): SqueezeExcitation(\n",491 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",492 " (fc1): Conv2d(32, 8, kernel_size=(1, 1), stride=(1, 1))\n",493 " (fc2): Conv2d(8, 32, kernel_size=(1, 1), stride=(1, 1))\n",494 " (activation): SiLU(inplace=True)\n",495 " (scale_activation): Sigmoid()\n",496 " )\n",497 " (2): Conv2dNormActivation(\n",498 " (0): Conv2d(32, 16, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",499 " (1): BatchNorm2d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",500 " )\n",501 " )\n",502 " (stochastic_depth): StochasticDepth(p=0.0, mode=row)\n",503 " )\n",504 " )\n",505 " (2): Sequential(\n",506 " (0): MBConv(\n",507 " (block): Sequential(\n",508 " (0): Conv2dNormActivation(\n",509 " (0): Conv2d(16, 96, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",510 " (1): BatchNorm2d(96, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",511 " (2): SiLU(inplace=True)\n",512 " )\n",513 " (1): Conv2dNormActivation(\n",514 " (0): Conv2d(96, 96, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), groups=96, bias=False)\n",515 " (1): BatchNorm2d(96, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",516 " (2): SiLU(inplace=True)\n",517 " )\n",518 " (2): SqueezeExcitation(\n",519 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",520 " (fc1): Conv2d(96, 4, kernel_size=(1, 1), stride=(1, 1))\n",521 " (fc2): Conv2d(4, 96, kernel_size=(1, 1), stride=(1, 1))\n",522 " (activation): SiLU(inplace=True)\n",523 " (scale_activation): Sigmoid()\n",524 " )\n",525 " (3): Conv2dNormActivation(\n",526 " (0): Conv2d(96, 24, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",527 " (1): BatchNorm2d(24, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",528 " )\n",529 " )\n",530 " (stochastic_depth): StochasticDepth(p=0.0125, mode=row)\n",531 " )\n",532 " (1): MBConv(\n",533 " (block): Sequential(\n",534 " (0): Conv2dNormActivation(\n",535 " (0): Conv2d(24, 144, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",536 " (1): BatchNorm2d(144, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",537 " (2): SiLU(inplace=True)\n",538 " )\n",539 " (1): Conv2dNormActivation(\n",540 " (0): Conv2d(144, 144, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), groups=144, bias=False)\n",541 " (1): BatchNorm2d(144, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",542 " (2): SiLU(inplace=True)\n",543 " )\n",544 " (2): SqueezeExcitation(\n",545 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",546 " (fc1): Conv2d(144, 6, kernel_size=(1, 1), stride=(1, 1))\n",547 " (fc2): Conv2d(6, 144, kernel_size=(1, 1), stride=(1, 1))\n",548 " (activation): SiLU(inplace=True)\n",549 " (scale_activation): Sigmoid()\n",550 " )\n",551 " (3): Conv2dNormActivation(\n",552 " (0): Conv2d(144, 24, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",553 " (1): BatchNorm2d(24, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",554 " )\n",555 " )\n",556 " (stochastic_depth): StochasticDepth(p=0.025, mode=row)\n",557 " )\n",558 " )\n",559 " (3): Sequential(\n",560 " (0): MBConv(\n",561 " (block): Sequential(\n",562 " (0): Conv2dNormActivation(\n",563 " (0): Conv2d(24, 144, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",564 " (1): BatchNorm2d(144, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",565 " (2): SiLU(inplace=True)\n",566 " )\n",567 " (1): Conv2dNormActivation(\n",568 " (0): Conv2d(144, 144, kernel_size=(5, 5), stride=(2, 2), padding=(2, 2), groups=144, bias=False)\n",569 " (1): BatchNorm2d(144, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",570 " (2): SiLU(inplace=True)\n",571 " )\n",572 " (2): SqueezeExcitation(\n",573 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",574 " (fc1): Conv2d(144, 6, kernel_size=(1, 1), stride=(1, 1))\n",575 " (fc2): Conv2d(6, 144, kernel_size=(1, 1), stride=(1, 1))\n",576 " (activation): SiLU(inplace=True)\n",577 " (scale_activation): Sigmoid()\n",578 " )\n",579 " (3): Conv2dNormActivation(\n",580 " (0): Conv2d(144, 40, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",581 " (1): BatchNorm2d(40, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",582 " )\n",583 " )\n",584 " (stochastic_depth): StochasticDepth(p=0.037500000000000006, mode=row)\n",585 " )\n",586 " (1): MBConv(\n",587 " (block): Sequential(\n",588 " (0): Conv2dNormActivation(\n",589 " (0): Conv2d(40, 240, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",590 " (1): BatchNorm2d(240, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",591 " (2): SiLU(inplace=True)\n",592 " )\n",593 " (1): Conv2dNormActivation(\n",594 " (0): Conv2d(240, 240, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=240, bias=False)\n",595 " (1): BatchNorm2d(240, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",596 " (2): SiLU(inplace=True)\n",597 " )\n",598 " (2): SqueezeExcitation(\n",599 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",600 " (fc1): Conv2d(240, 10, kernel_size=(1, 1), stride=(1, 1))\n",601 " (fc2): Conv2d(10, 240, kernel_size=(1, 1), stride=(1, 1))\n",602 " (activation): SiLU(inplace=True)\n",603 " (scale_activation): Sigmoid()\n",604 " )\n",605 " (3): Conv2dNormActivation(\n",606 " (0): Conv2d(240, 40, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",607 " (1): BatchNorm2d(40, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",608 " )\n",609 " )\n",610 " (stochastic_depth): StochasticDepth(p=0.05, mode=row)\n",611 " )\n",612 " )\n",613 " (4): Sequential(\n",614 " (0): MBConv(\n",615 " (block): Sequential(\n",616 " (0): Conv2dNormActivation(\n",617 " (0): Conv2d(40, 240, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",618 " (1): BatchNorm2d(240, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",619 " (2): SiLU(inplace=True)\n",620 " )\n",621 " (1): Conv2dNormActivation(\n",622 " (0): Conv2d(240, 240, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), groups=240, bias=False)\n",623 " (1): BatchNorm2d(240, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",624 " (2): SiLU(inplace=True)\n",625 " )\n",626 " (2): SqueezeExcitation(\n",627 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",628 " (fc1): Conv2d(240, 10, kernel_size=(1, 1), stride=(1, 1))\n",629 " (fc2): Conv2d(10, 240, kernel_size=(1, 1), stride=(1, 1))\n",630 " (activation): SiLU(inplace=True)\n",631 " (scale_activation): Sigmoid()\n",632 " )\n",633 " (3): Conv2dNormActivation(\n",634 " (0): Conv2d(240, 80, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",635 " (1): BatchNorm2d(80, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",636 " )\n",637 " )\n",638 " (stochastic_depth): StochasticDepth(p=0.0625, mode=row)\n",639 " )\n",640 " (1): MBConv(\n",641 " (block): Sequential(\n",642 " (0): Conv2dNormActivation(\n",643 " (0): Conv2d(80, 480, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",644 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",645 " (2): SiLU(inplace=True)\n",646 " )\n",647 " (1): Conv2dNormActivation(\n",648 " (0): Conv2d(480, 480, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), groups=480, bias=False)\n",649 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",650 " (2): SiLU(inplace=True)\n",651 " )\n",652 " (2): SqueezeExcitation(\n",653 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",654 " (fc1): Conv2d(480, 20, kernel_size=(1, 1), stride=(1, 1))\n",655 " (fc2): Conv2d(20, 480, kernel_size=(1, 1), stride=(1, 1))\n",656 " (activation): SiLU(inplace=True)\n",657 " (scale_activation): Sigmoid()\n",658 " )\n",659 " (3): Conv2dNormActivation(\n",660 " (0): Conv2d(480, 80, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",661 " (1): BatchNorm2d(80, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",662 " )\n",663 " )\n",664 " (stochastic_depth): StochasticDepth(p=0.07500000000000001, mode=row)\n",665 " )\n",666 " (2): MBConv(\n",667 " (block): Sequential(\n",668 " (0): Conv2dNormActivation(\n",669 " (0): Conv2d(80, 480, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",670 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",671 " (2): SiLU(inplace=True)\n",672 " )\n",673 " (1): Conv2dNormActivation(\n",674 " (0): Conv2d(480, 480, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), groups=480, bias=False)\n",675 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",676 " (2): SiLU(inplace=True)\n",677 " )\n",678 " (2): SqueezeExcitation(\n",679 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",680 " (fc1): Conv2d(480, 20, kernel_size=(1, 1), stride=(1, 1))\n",681 " (fc2): Conv2d(20, 480, kernel_size=(1, 1), stride=(1, 1))\n",682 " (activation): SiLU(inplace=True)\n",683 " (scale_activation): Sigmoid()\n",684 " )\n",685 " (3): Conv2dNormActivation(\n",686 " (0): Conv2d(480, 80, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",687 " (1): BatchNorm2d(80, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",688 " )\n",689 " )\n",690 " (stochastic_depth): StochasticDepth(p=0.08750000000000001, mode=row)\n",691 " )\n",692 " )\n",693 " (5): Sequential(\n",694 " (0): MBConv(\n",695 " (block): Sequential(\n",696 " (0): Conv2dNormActivation(\n",697 " (0): Conv2d(80, 480, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",698 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",699 " (2): SiLU(inplace=True)\n",700 " )\n",701 " (1): Conv2dNormActivation(\n",702 " (0): Conv2d(480, 480, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=480, bias=False)\n",703 " (1): BatchNorm2d(480, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",704 " (2): SiLU(inplace=True)\n",705 " )\n",706 " (2): SqueezeExcitation(\n",707 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",708 " (fc1): Conv2d(480, 20, kernel_size=(1, 1), stride=(1, 1))\n",709 " (fc2): Conv2d(20, 480, kernel_size=(1, 1), stride=(1, 1))\n",710 " (activation): SiLU(inplace=True)\n",711 " (scale_activation): Sigmoid()\n",712 " )\n",713 " (3): Conv2dNormActivation(\n",714 " (0): Conv2d(480, 112, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",715 " (1): BatchNorm2d(112, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",716 " )\n",717 " )\n",718 " (stochastic_depth): StochasticDepth(p=0.1, mode=row)\n",719 " )\n",720 " (1): MBConv(\n",721 " (block): Sequential(\n",722 " (0): Conv2dNormActivation(\n",723 " (0): Conv2d(112, 672, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",724 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",725 " (2): SiLU(inplace=True)\n",726 " )\n",727 " (1): Conv2dNormActivation(\n",728 " (0): Conv2d(672, 672, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=672, bias=False)\n",729 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",730 " (2): SiLU(inplace=True)\n",731 " )\n",732 " (2): SqueezeExcitation(\n",733 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",734 " (fc1): Conv2d(672, 28, kernel_size=(1, 1), stride=(1, 1))\n",735 " (fc2): Conv2d(28, 672, kernel_size=(1, 1), stride=(1, 1))\n",736 " (activation): SiLU(inplace=True)\n",737 " (scale_activation): Sigmoid()\n",738 " )\n",739 " (3): Conv2dNormActivation(\n",740 " (0): Conv2d(672, 112, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",741 " (1): BatchNorm2d(112, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",742 " )\n",743 " )\n",744 " (stochastic_depth): StochasticDepth(p=0.1125, mode=row)\n",745 " )\n",746 " (2): MBConv(\n",747 " (block): Sequential(\n",748 " (0): Conv2dNormActivation(\n",749 " (0): Conv2d(112, 672, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",750 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",751 " (2): SiLU(inplace=True)\n",752 " )\n",753 " (1): Conv2dNormActivation(\n",754 " (0): Conv2d(672, 672, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=672, bias=False)\n",755 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",756 " (2): SiLU(inplace=True)\n",757 " )\n",758 " (2): SqueezeExcitation(\n",759 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",760 " (fc1): Conv2d(672, 28, kernel_size=(1, 1), stride=(1, 1))\n",761 " (fc2): Conv2d(28, 672, kernel_size=(1, 1), stride=(1, 1))\n",762 " (activation): SiLU(inplace=True)\n",763 " (scale_activation): Sigmoid()\n",764 " )\n",765 " (3): Conv2dNormActivation(\n",766 " (0): Conv2d(672, 112, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",767 " (1): BatchNorm2d(112, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",768 " )\n",769 " )\n",770 " (stochastic_depth): StochasticDepth(p=0.125, mode=row)\n",771 " )\n",772 " )\n",773 " (6): Sequential(\n",774 " (0): MBConv(\n",775 " (block): Sequential(\n",776 " (0): Conv2dNormActivation(\n",777 " (0): Conv2d(112, 672, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",778 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",779 " (2): SiLU(inplace=True)\n",780 " )\n",781 " (1): Conv2dNormActivation(\n",782 " (0): Conv2d(672, 672, kernel_size=(5, 5), stride=(2, 2), padding=(2, 2), groups=672, bias=False)\n",783 " (1): BatchNorm2d(672, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",784 " (2): SiLU(inplace=True)\n",785 " )\n",786 " (2): SqueezeExcitation(\n",787 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",788 " (fc1): Conv2d(672, 28, kernel_size=(1, 1), stride=(1, 1))\n",789 " (fc2): Conv2d(28, 672, kernel_size=(1, 1), stride=(1, 1))\n",790 " (activation): SiLU(inplace=True)\n",791 " (scale_activation): Sigmoid()\n",792 " )\n",793 " (3): Conv2dNormActivation(\n",794 " (0): Conv2d(672, 192, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",795 " (1): BatchNorm2d(192, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",796 " )\n",797 " )\n",798 " (stochastic_depth): StochasticDepth(p=0.1375, mode=row)\n",799 " )\n",800 " (1): MBConv(\n",801 " (block): Sequential(\n",802 " (0): Conv2dNormActivation(\n",803 " (0): Conv2d(192, 1152, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",804 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",805 " (2): SiLU(inplace=True)\n",806 " )\n",807 " (1): Conv2dNormActivation(\n",808 " (0): Conv2d(1152, 1152, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=1152, bias=False)\n",809 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",810 " (2): SiLU(inplace=True)\n",811 " )\n",812 " (2): SqueezeExcitation(\n",813 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",814 " (fc1): Conv2d(1152, 48, kernel_size=(1, 1), stride=(1, 1))\n",815 " (fc2): Conv2d(48, 1152, kernel_size=(1, 1), stride=(1, 1))\n",816 " (activation): SiLU(inplace=True)\n",817 " (scale_activation): Sigmoid()\n",818 " )\n",819 " (3): Conv2dNormActivation(\n",820 " (0): Conv2d(1152, 192, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",821 " (1): BatchNorm2d(192, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",822 " )\n",823 " )\n",824 " (stochastic_depth): StochasticDepth(p=0.15000000000000002, mode=row)\n",825 " )\n",826 " (2): MBConv(\n",827 " (block): Sequential(\n",828 " (0): Conv2dNormActivation(\n",829 " (0): Conv2d(192, 1152, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",830 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",831 " (2): SiLU(inplace=True)\n",832 " )\n",833 " (1): Conv2dNormActivation(\n",834 " (0): Conv2d(1152, 1152, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=1152, bias=False)\n",835 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",836 " (2): SiLU(inplace=True)\n",837 " )\n",838 " (2): SqueezeExcitation(\n",839 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",840 " (fc1): Conv2d(1152, 48, kernel_size=(1, 1), stride=(1, 1))\n",841 " (fc2): Conv2d(48, 1152, kernel_size=(1, 1), stride=(1, 1))\n",842 " (activation): SiLU(inplace=True)\n",843 " (scale_activation): Sigmoid()\n",844 " )\n",845 " (3): Conv2dNormActivation(\n",846 " (0): Conv2d(1152, 192, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",847 " (1): BatchNorm2d(192, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",848 " )\n",849 " )\n",850 " (stochastic_depth): StochasticDepth(p=0.1625, mode=row)\n",851 " )\n",852 " (3): MBConv(\n",853 " (block): Sequential(\n",854 " (0): Conv2dNormActivation(\n",855 " (0): Conv2d(192, 1152, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",856 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",857 " (2): SiLU(inplace=True)\n",858 " )\n",859 " (1): Conv2dNormActivation(\n",860 " (0): Conv2d(1152, 1152, kernel_size=(5, 5), stride=(1, 1), padding=(2, 2), groups=1152, bias=False)\n",861 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",862 " (2): SiLU(inplace=True)\n",863 " )\n",864 " (2): SqueezeExcitation(\n",865 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",866 " (fc1): Conv2d(1152, 48, kernel_size=(1, 1), stride=(1, 1))\n",867 " (fc2): Conv2d(48, 1152, kernel_size=(1, 1), stride=(1, 1))\n",868 " (activation): SiLU(inplace=True)\n",869 " (scale_activation): Sigmoid()\n",870 " )\n",871 " (3): Conv2dNormActivation(\n",872 " (0): Conv2d(1152, 192, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",873 " (1): BatchNorm2d(192, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",874 " )\n",875 " )\n",876 " (stochastic_depth): StochasticDepth(p=0.17500000000000002, mode=row)\n",877 " )\n",878 " )\n",879 " (7): Sequential(\n",880 " (0): MBConv(\n",881 " (block): Sequential(\n",882 " (0): Conv2dNormActivation(\n",883 " (0): Conv2d(192, 1152, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",884 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",885 " (2): SiLU(inplace=True)\n",886 " )\n",887 " (1): Conv2dNormActivation(\n",888 " (0): Conv2d(1152, 1152, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), groups=1152, bias=False)\n",889 " (1): BatchNorm2d(1152, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",890 " (2): SiLU(inplace=True)\n",891 " )\n",892 " (2): SqueezeExcitation(\n",893 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",894 " (fc1): Conv2d(1152, 48, kernel_size=(1, 1), stride=(1, 1))\n",895 " (fc2): Conv2d(48, 1152, kernel_size=(1, 1), stride=(1, 1))\n",896 " (activation): SiLU(inplace=True)\n",897 " (scale_activation): Sigmoid()\n",898 " )\n",899 " (3): Conv2dNormActivation(\n",900 " (0): Conv2d(1152, 320, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",901 " (1): BatchNorm2d(320, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",902 " )\n",903 " )\n",904 " (stochastic_depth): StochasticDepth(p=0.1875, mode=row)\n",905 " )\n",906 " )\n",907 " (8): Conv2dNormActivation(\n",908 " (0): Conv2d(320, 1280, kernel_size=(1, 1), stride=(1, 1), bias=False)\n",909 " (1): BatchNorm2d(1280, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n",910 " (2): SiLU(inplace=True)\n",911 " )\n",912 " )\n",913 " (avgpool): AdaptiveAvgPool2d(output_size=1)\n",914 " (classifier): Sequential(\n",915 " (0): Dropout(p=0.2, inplace=True)\n",916 " (1): Linear(in_features=1280, out_features=1000, bias=True)\n",917 " )\n",918 ")\n"919 ]920 }921 ],922 "source": [923 "print(model)"924 ]925 },926 {927 "cell_type": "code",928 "execution_count": 28,929 "id": "d468c7f5-9d9e-4439-b65b-c69108169d9b",930 "metadata": {},931 "outputs": [932 {933 "name": "stdout",934 "output_type": "stream",935 "text": [936 "features.0.0.weight: requires_grad = False\n",937 "features.0.1.weight: requires_grad = False\n",938 "features.0.1.bias: requires_grad = False\n",939 "features.1.0.block.0.0.weight: requires_grad = False\n",940 "features.1.0.block.0.1.weight: requires_grad = False\n",941 "features.1.0.block.0.1.bias: requires_grad = False\n",942 "features.1.0.block.1.fc1.weight: requires_grad = False\n",943 "features.1.0.block.1.fc1.bias: requires_grad = False\n",944 "features.1.0.block.1.fc2.weight: requires_grad = False\n",945 "features.1.0.block.1.fc2.bias: requires_grad = False\n",946 "features.1.0.block.2.0.weight: requires_grad = False\n",947 "features.1.0.block.2.1.weight: requires_grad = False\n",948 "features.1.0.block.2.1.bias: requires_grad = False\n",949 "features.2.0.block.0.0.weight: requires_grad = False\n",950 "features.2.0.block.0.1.weight: requires_grad = False\n",951 "features.2.0.block.0.1.bias: requires_grad = False\n",952 "features.2.0.block.1.0.weight: requires_grad = False\n",953 "features.2.0.block.1.1.weight: requires_grad = False\n",954 "features.2.0.block.1.1.bias: requires_grad = False\n",955 "features.2.0.block.2.fc1.weight: requires_grad = False\n",956 "features.2.0.block.2.fc1.bias: requires_grad = False\n",957 "features.2.0.block.2.fc2.weight: requires_grad = False\n",958 "features.2.0.block.2.fc2.bias: requires_grad = False\n",959 "features.2.0.block.3.0.weight: requires_grad = False\n",960 "features.2.0.block.3.1.weight: requires_grad = False\n",961 "features.2.0.block.3.1.bias: requires_grad = False\n",962 "features.2.1.block.0.0.weight: requires_grad = False\n",963 "features.2.1.block.0.1.weight: requires_grad = False\n",964 "features.2.1.block.0.1.bias: requires_grad = False\n",965 "features.2.1.block.1.0.weight: requires_grad = False\n",966 "features.2.1.block.1.1.weight: requires_grad = False\n",967 "features.2.1.block.1.1.bias: requires_grad = False\n",968 "features.2.1.block.2.fc1.weight: requires_grad = False\n",969 "features.2.1.block.2.fc1.bias: requires_grad = False\n",970 "features.2.1.block.2.fc2.weight: requires_grad = False\n",971 "features.2.1.block.2.fc2.bias: requires_grad = False\n",972 "features.2.1.block.3.0.weight: requires_grad = False\n",973 "features.2.1.block.3.1.weight: requires_grad = False\n",974 "features.2.1.block.3.1.bias: requires_grad = False\n",975 "features.3.0.block.0.0.weight: requires_grad = False\n",976 "features.3.0.block.0.1.weight: requires_grad = False\n",977 "features.3.0.block.0.1.bias: requires_grad = False\n",978 "features.3.0.block.1.0.weight: requires_grad = False\n",979 "features.3.0.block.1.1.weight: requires_grad = False\n",980 "features.3.0.block.1.1.bias: requires_grad = False\n",981 "features.3.0.block.2.fc1.weight: requires_grad = False\n",982 "features.3.0.block.2.fc1.bias: requires_grad = False\n",983 "features.3.0.block.2.fc2.weight: requires_grad = False\n",984 "features.3.0.block.2.fc2.bias: requires_grad = False\n",985 "features.3.0.block.3.0.weight: requires_grad = False\n",986 "features.3.0.block.3.1.weight: requires_grad = False\n",987 "features.3.0.block.3.1.bias: requires_grad = False\n",988 "features.3.1.block.0.0.weight: requires_grad = False\n",989 "features.3.1.block.0.1.weight: requires_grad = False\n",990 "features.3.1.block.0.1.bias: requires_grad = False\n",991 "features.3.1.block.1.0.weight: requires_grad = False\n",992 "features.3.1.block.1.1.weight: requires_grad = False\n",993 "features.3.1.block.1.1.bias: requires_grad = False\n",994 "features.3.1.block.2.fc1.weight: requires_grad = False\n",995 "features.3.1.block.2.fc1.bias: requires_grad = False\n",996 "features.3.1.block.2.fc2.weight: requires_grad = False\n",997 "features.3.1.block.2.fc2.bias: requires_grad = False\n",998 "features.3.1.block.3.0.weight: requires_grad = False\n",999 "features.3.1.block.3.1.weight: requires_grad = False\n",1000 "features.3.1.block.3.1.bias: requires_grad = False\n",1001 "features.4.0.block.0.0.weight: requires_grad = False\n",1002 "features.4.0.block.0.1.weight: requires_grad = False\n",1003 "features.4.0.block.0.1.bias: requires_grad = False\n",1004 "features.4.0.block.1.0.weight: requires_grad = False\n",1005 "features.4.0.block.1.1.weight: requires_grad = False\n",1006 "features.4.0.block.1.1.bias: requires_grad = False\n",1007 "features.4.0.block.2.fc1.weight: requires_grad = False\n",1008 "features.4.0.block.2.fc1.bias: requires_grad = False\n",1009 "features.4.0.block.2.fc2.weight: requires_grad = False\n",1010 "features.4.0.block.2.fc2.bias: requires_grad = False\n",1011 "features.4.0.block.3.0.weight: requires_grad = False\n",1012 "features.4.0.block.3.1.weight: requires_grad = False\n",1013 "features.4.0.block.3.1.bias: requires_grad = False\n",1014 "features.4.1.block.0.0.weight: requires_grad = False\n",1015 "features.4.1.block.0.1.weight: requires_grad = False\n",1016 "features.4.1.block.0.1.bias: requires_grad = False\n",1017 "features.4.1.block.1.0.weight: requires_grad = False\n",1018 "features.4.1.block.1.1.weight: requires_grad = False\n",1019 "features.4.1.block.1.1.bias: requires_grad = False\n",1020 "features.4.1.block.2.fc1.weight: requires_grad = False\n",1021 "features.4.1.block.2.fc1.bias: requires_grad = False\n",1022 "features.4.1.block.2.fc2.weight: requires_grad = False\n",1023 "features.4.1.block.2.fc2.bias: requires_grad = False\n",1024 "features.4.1.block.3.0.weight: requires_grad = False\n",1025 "features.4.1.block.3.1.weight: requires_grad = False\n",1026 "features.4.1.block.3.1.bias: requires_grad = False\n",1027 "features.4.2.block.0.0.weight: requires_grad = False\n",1028 "features.4.2.block.0.1.weight: requires_grad = False\n",1029 "features.4.2.block.0.1.bias: requires_grad = False\n",1030 "features.4.2.block.1.0.weight: requires_grad = False\n",1031 "features.4.2.block.1.1.weight: requires_grad = False\n",1032 "features.4.2.block.1.1.bias: requires_grad = False\n",1033 "features.4.2.block.2.fc1.weight: requires_grad = False\n",1034 "features.4.2.block.2.fc1.bias: requires_grad = False\n",1035 "features.4.2.block.2.fc2.weight: requires_grad = False\n",1036 "features.4.2.block.2.fc2.bias: requires_grad = False\n",1037 "features.4.2.block.3.0.weight: requires_grad = False\n",1038 "features.4.2.block.3.1.weight: requires_grad = False\n",1039 "features.4.2.block.3.1.bias: requires_grad = False\n",1040 "features.5.0.block.0.0.weight: requires_grad = False\n",1041 "features.5.0.block.0.1.weight: requires_grad = False\n",1042 "features.5.0.block.0.1.bias: requires_grad = False\n",1043 "features.5.0.block.1.0.weight: requires_grad = False\n",1044 "features.5.0.block.1.1.weight: requires_grad = False\n",1045 "features.5.0.block.1.1.bias: requires_grad = False\n",1046 "features.5.0.block.2.fc1.weight: requires_grad = False\n",1047 "features.5.0.block.2.fc1.bias: requires_grad = False\n",1048 "features.5.0.block.2.fc2.weight: requires_grad = False\n",1049 "features.5.0.block.2.fc2.bias: requires_grad = False\n",1050 "features.5.0.block.3.0.weight: requires_grad = False\n",1051 "features.5.0.block.3.1.weight: requires_grad = False\n",1052 "features.5.0.block.3.1.bias: requires_grad = False\n",1053 "features.5.1.block.0.0.weight: requires_grad = False\n",1054 "features.5.1.block.0.1.weight: requires_grad = False\n",1055 "features.5.1.block.0.1.bias: requires_grad = False\n",1056 "features.5.1.block.1.0.weight: requires_grad = False\n",1057 "features.5.1.block.1.1.weight: requires_grad = False\n",1058 "features.5.1.block.1.1.bias: requires_grad = False\n",1059 "features.5.1.block.2.fc1.weight: requires_grad = False\n",1060 "features.5.1.block.2.fc1.bias: requires_grad = False\n",1061 "features.5.1.block.2.fc2.weight: requires_grad = False\n",1062 "features.5.1.block.2.fc2.bias: requires_grad = False\n",1063 "features.5.1.block.3.0.weight: requires_grad = False\n",1064 "features.5.1.block.3.1.weight: requires_grad = False\n",1065 "features.5.1.block.3.1.bias: requires_grad = False\n",1066 "features.5.2.block.0.0.weight: requires_grad = False\n",1067 "features.5.2.block.0.1.weight: requires_grad = False\n",1068 "features.5.2.block.0.1.bias: requires_grad = False\n",1069 "features.5.2.block.1.0.weight: requires_grad = False\n",1070 "features.5.2.block.1.1.weight: requires_grad = False\n",1071 "features.5.2.block.1.1.bias: requires_grad = False\n",1072 "features.5.2.block.2.fc1.weight: requires_grad = False\n",1073 "features.5.2.block.2.fc1.bias: requires_grad = False\n",1074 "features.5.2.block.2.fc2.weight: requires_grad = False\n",1075 "features.5.2.block.2.fc2.bias: requires_grad = False\n",1076 "features.5.2.block.3.0.weight: requires_grad = False\n",1077 "features.5.2.block.3.1.weight: requires_grad = False\n",1078 "features.5.2.block.3.1.bias: requires_grad = False\n",1079 "features.6.0.block.0.0.weight: requires_grad = False\n",1080 "features.6.0.block.0.1.weight: requires_grad = False\n",1081 "features.6.0.block.0.1.bias: requires_grad = False\n",1082 "features.6.0.block.1.0.weight: requires_grad = False\n",1083 "features.6.0.block.1.1.weight: requires_grad = False\n",1084 "features.6.0.block.1.1.bias: requires_grad = False\n",1085 "features.6.0.block.2.fc1.weight: requires_grad = False\n",1086 "features.6.0.block.2.fc1.bias: requires_grad = False\n",1087 "features.6.0.block.2.fc2.weight: requires_grad = False\n",1088 "features.6.0.block.2.fc2.bias: requires_grad = False\n",1089 "features.6.0.block.3.0.weight: requires_grad = False\n",1090 "features.6.0.block.3.1.weight: requires_grad = False\n",1091 "features.6.0.block.3.1.bias: requires_grad = False\n",1092 "features.6.1.block.0.0.weight: requires_grad = False\n",1093 "features.6.1.block.0.1.weight: requires_grad = False\n",1094 "features.6.1.block.0.1.bias: requires_grad = False\n",1095 "features.6.1.block.1.0.weight: requires_grad = False\n",1096 "features.6.1.block.1.1.weight: requires_grad = False\n",1097 "features.6.1.block.1.1.bias: requires_grad = False\n",1098 "features.6.1.block.2.fc1.weight: requires_grad = False\n",1099 "features.6.1.block.2.fc1.bias: requires_grad = False\n",1100 "features.6.1.block.2.fc2.weight: requires_grad = False\n",1101 "features.6.1.block.2.fc2.bias: requires_grad = False\n",1102 "features.6.1.block.3.0.weight: requires_grad = False\n",1103 "features.6.1.block.3.1.weight: requires_grad = False\n",1104 "features.6.1.block.3.1.bias: requires_grad = False\n",1105 "features.6.2.block.0.0.weight: requires_grad = False\n",1106 "features.6.2.block.0.1.weight: requires_grad = False\n",1107 "features.6.2.block.0.1.bias: requires_grad = False\n",1108 "features.6.2.block.1.0.weight: requires_grad = False\n",1109 "features.6.2.block.1.1.weight: requires_grad = False\n",1110 "features.6.2.block.1.1.bias: requires_grad = False\n",1111 "features.6.2.block.2.fc1.weight: requires_grad = False\n",1112 "features.6.2.block.2.fc1.bias: requires_grad = False\n",1113 "features.6.2.block.2.fc2.weight: requires_grad = False\n",1114 "features.6.2.block.2.fc2.bias: requires_grad = False\n",1115 "features.6.2.block.3.0.weight: requires_grad = False\n",1116 "features.6.2.block.3.1.weight: requires_grad = False\n",1117 "features.6.2.block.3.1.bias: requires_grad = False\n",1118 "features.6.3.block.0.0.weight: requires_grad = False\n",1119 "features.6.3.block.0.1.weight: requires_grad = False\n",1120 "features.6.3.block.0.1.bias: requires_grad = False\n",1121 "features.6.3.block.1.0.weight: requires_grad = False\n",1122 "features.6.3.block.1.1.weight: requires_grad = False\n",1123 "features.6.3.block.1.1.bias: requires_grad = False\n",1124 "features.6.3.block.2.fc1.weight: requires_grad = False\n",1125 "features.6.3.block.2.fc1.bias: requires_grad = False\n",1126 "features.6.3.block.2.fc2.weight: requires_grad = False\n",1127 "features.6.3.block.2.fc2.bias: requires_grad = False\n",1128 "features.6.3.block.3.0.weight: requires_grad = False\n",1129 "features.6.3.block.3.1.weight: requires_grad = False\n",1130 "features.6.3.block.3.1.bias: requires_grad = False\n",1131 "features.7.0.block.0.0.weight: requires_grad = False\n",1132 "features.7.0.block.0.1.weight: requires_grad = False\n",1133 "features.7.0.block.0.1.bias: requires_grad = False\n",1134 "features.7.0.block.1.0.weight: requires_grad = False\n",1135 "features.7.0.block.1.1.weight: requires_grad = False\n",1136 "features.7.0.block.1.1.bias: requires_grad = False\n",1137 "features.7.0.block.2.fc1.weight: requires_grad = False\n",1138 "features.7.0.block.2.fc1.bias: requires_grad = False\n",1139 "features.7.0.block.2.fc2.weight: requires_grad = False\n",1140 "features.7.0.block.2.fc2.bias: requires_grad = False\n",1141 "features.7.0.block.3.0.weight: requires_grad = False\n",1142 "features.7.0.block.3.1.weight: requires_grad = False\n",1143 "features.7.0.block.3.1.bias: requires_grad = False\n",1144 "features.8.0.weight: requires_grad = False\n",1145 "features.8.1.weight: requires_grad = False\n",1146 "features.8.1.bias: requires_grad = False\n",1147 "classifier.1.weight: requires_grad = True\n",1148 "classifier.1.bias: requires_grad = True\n"1149 ]1150 }1151 ],1152 "source": [1153 "# Freeze all base layers in the \"features\" section of the model (the feature extractor) by setting requires_grad=False\n",1154 "for param in model.features.parameters():\n",1155 " param.requires_grad = False\n",1156 "for name, param in model.named_parameters():\n",1157 " print(f\"{name}: requires_grad = {param.requires_grad}\")"1158 ]1159 },1160 {1161 "cell_type": "code",1162 "execution_count": 29,1163 "id": "6e16c5d7-5460-4ed3-967a-68b042e30dec",1164 "metadata": {},1165 "outputs": [1166 {1167 "data": {1168 "text/plain": [1169 "3"1170 ]1171 },1172 "execution_count": 29,1173 "metadata": {},1174 "output_type": "execute_result"1175 }1176 ],1177 "source": [1178 "len(class_names)"1179 ]1180 },1181 {1182 "cell_type": "code",1183 "execution_count": 30,1184 "id": "9bcb303f-d099-4ba1-b537-e5ef444fb9bf",1185 "metadata": {},1186 "outputs": [],1187 "source": [1188 "# Set the manual seeds\n",1189 "torch.manual_seed(42)\n",1190 "torch.cuda.manual_seed(42)\n",1191 "\n",1192 "# Get the length of class_names (one output unit for each class)\n",1193 "output_shape = len(class_names)\n",1194 "\n",1195 "# Recreate the classifier layer and seed it to the target device\n",1196 "model.classifier = torch.nn.Sequential(\n",1197 " torch.nn.Dropout(p=0.2, inplace=True), \n",1198 " torch.nn.Linear(in_features=1280, \n",1199 " out_features=output_shape, # same number of output units as our number of classes\n",1200 " bias=True)).to(device)"