{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":7944294,"datasetId":4670999,"databundleVersionId":8053257},{"sourceType":"datasetVersion","sourceId":8155987,"datasetId":4824526,"databundleVersionId":8277486}],"dockerImageVersionId":30683,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"papermill":{"default_parameters":{},"duration":17324.565701,"end_time":"2024-03-26T23:26:42.769088","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-03-26T18:37:58.203387","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install monai thop pydicom\n# !pip install gdcm pylibjpeg pylibjpeg-libjpeg","metadata":{"papermill":{"duration":14.580562,"end_time":"2024-03-26T18:38:15.480097","exception":false,"start_time":"2024-03-26T18:38:00.899535","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:42:01.303491Z","iopub.execute_input":"2024-04-18T12:42:01.304283Z","iopub.status.idle":"2024-04-18T12:42:16.70535Z","shell.execute_reply.started":"2024-04-18T12:42:01.304252Z","shell.execute_reply":"2024-04-18T12:42:16.704316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install thop","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:42:16.707697Z","iopub.execute_input":"2024-04-18T12:42:16.708139Z","iopub.status.idle":"2024-04-18T12:42:29.604349Z","shell.execute_reply.started":"2024-04-18T12:42:16.708101Z","shell.execute_reply":"2024-04-18T12:42:29.603159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python -c \"import monai\" || pip install -q \"monai-weekly[pillow, tqdm]\"\n!python -c \"import matplotlib\" || pip install -q matplotlib\n%matplotlib inline","metadata":{"papermill":{"duration":45.198596,"end_time":"2024-03-26T18:39:00.69051","exception":false,"start_time":"2024-03-26T18:38:15.491914","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:42:29.606054Z","iopub.execute_input":"2024-04-18T12:42:29.60681Z","iopub.status.idle":"2024-04-18T12:43:18.451028Z","shell.execute_reply.started":"2024-04-18T12:42:29.60677Z","shell.execute_reply":"2024-04-18T12:43:18.449746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\nimport tempfile\nimport matplotlib.pyplot as plt\nimport PIL\nimport torch\nimport numpy as np\nfrom sklearn.metrics import classification_report\n\nfrom monai.apps import download_and_extract\nfrom monai.config import print_config\nfrom monai.data import decollate_batch, DataLoader\nfrom monai.metrics import ROCAUCMetric\nfrom monai.networks.nets import DenseNet121\nfrom monai.transforms import (\n    Activations,\n    EnsureChannelFirst,\n    AsDiscrete,\n    Compose,\n    LoadImage,\n    RandFlip,\n    RandRotate,\n    RandZoom,\n    ScaleIntensity,\n)\nfrom monai.utils import set_determinism\n\nprint_config()","metadata":{"papermill":{"duration":29.259606,"end_time":"2024-03-26T18:39:29.962134","exception":false,"start_time":"2024-03-26T18:39:00.702528","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:43:59.464052Z","iopub.execute_input":"2024-04-18T12:43:59.464405Z","iopub.status.idle":"2024-04-18T12:44:25.468478Z","shell.execute_reply.started":"2024-04-18T12:43:59.464378Z","shell.execute_reply":"2024-04-18T12:44:25.467438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport pydicom\nfrom torchvision import transforms\nimport pickle\nimport cv2\n\nfrom sklearn.model_selection import train_test_split\n\n# Timing utility\nfrom timeit import default_timer as timer\nfrom tqdm import tqdm\n\nimport torch\nfrom sklearn.metrics import roc_auc_score, accuracy_score, precision_score, recall_score, f1_score, roc_curve, auc\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.metrics import multilabel_confusion_matrix\nimport seaborn as sn\n\nimport warnings\nwarnings.filterwarnings('ignore', category=FutureWarning)\nwarnings.filterwarnings('ignore', category=UserWarning)\n\n\nprint(\"all imported\")\n\nset_determinism(seed=0)","metadata":{"papermill":{"duration":0.217627,"end_time":"2024-03-26T18:39:30.195448","exception":false,"start_time":"2024-03-26T18:39:29.977821","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:25.470311Z","iopub.execute_input":"2024-04-18T12:44:25.470973Z","iopub.status.idle":"2024-04-18T12:44:25.484572Z","shell.execute_reply.started":"2024-04-18T12:44:25.470944Z","shell.execute_reply":"2024-04-18T12:44:25.483551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nnum_classes = 15\nIMAGE_SIZE = (224, 224)\nin_channels = 1\n# IMAGE_SIZE[0], IMAGE_SIZE[1]\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:44:25.485918Z","iopub.execute_input":"2024-04-18T12:44:25.486235Z","iopub.status.idle":"2024-04-18T12:44:25.553994Z","shell.execute_reply.started":"2024-04-18T12:44:25.486212Z","shell.execute_reply":"2024-04-18T12:44:25.552974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from thop import profile\nfrom thop import clever_format\nimport torch\n\ndef display_params_flops(model):\n    #params\n    num_params = sum(p.numel() for p in model.parameters())\n    num_params_millions = num_params / 1e6\n    print(f\"Number of parameters in millions: {num_params_millions:.2f} M\")\n\n    num_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    num_params_millions = num_params / 1e6\n    print(f\"Number of trainable parameters in millions: {num_params_millions:.2f} M\")\n\n\n    #FLOPS\n    input_size = (1, in_channels, IMAGE_SIZE[0], IMAGE_SIZE[1])  \n\n\n    # Move the model to GPU if available\n    if torch.cuda.is_available():\n        model = model.cuda()\n\n    # Use thop.profile to count FLOPs\n    input_tensor = torch.randn(*input_size)\n    if torch.cuda.is_available():\n        input_tensor = input_tensor.cuda()\n    flops, params = profile(model, inputs=(input_tensor,))\n\n    # Convert FLOPs to gigaFLOPs and format the results\n    flops, params = clever_format([flops, params], \"%.2f\")\n    print(f\"FLOPs: {flops}, Params: {params}\")","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:44:25.556641Z","iopub.execute_input":"2024-04-18T12:44:25.557368Z","iopub.status.idle":"2024-04-18T12:44:25.566821Z","shell.execute_reply.started":"2024-04-18T12:44:25.557316Z","shell.execute_reply":"2024-04-18T12:44:25.565983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nfrom torch import nn\nfrom timm import create_model\n\n\nclass CoatNetModel(nn.Module):\n    def __init__(self, num_classes, fine_tune=False):\n        super(CoatNetModel, self).__init__()\n        self.coatnet = create_model(\n            'timm/coatnet_3_rw_224.sw_in12k', # accepts only 224x224 images\n            pretrained=True,\n            num_classes=num_classes,\n            in_chans=in_channels\n        )\n        \n        if not fine_tune:\n            for param in self.coatnet.parameters():\n                param.requires_grad = False\n            \n            for param in self.coatnet.head.parameters():\n                param.requires_grad = True\n        \n\n    def forward(self, x):\n        x = self.coatnet(x)\n        return x\n    \n    \n    \ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = CoatNetModel(num_classes, fine_tune=True)\nmodel.to(device)\n\n# print(model)\nprint()\n\nx = torch.randn(1, in_channels, IMAGE_SIZE[0], IMAGE_SIZE[1]).to(device)\n\n\noutput = model(x)\nprint(\"Model output's shape:\", output.shape)\n# print(output) # logits\ndisplay_params_flops(model)","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:44:25.567923Z","iopub.execute_input":"2024-04-18T12:44:25.568177Z","iopub.status.idle":"2024-04-18T12:44:34.992024Z","shell.execute_reply.started":"2024-04-18T12:44:25.568155Z","shell.execute_reply":"2024-04-18T12:44:34.991033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"diseases = ['Aortic enlargement', 'Atelectasis', 'Calcification', 'Cardiomegaly',\n 'Consolidation', 'ILD', 'Infiltration', 'Lung Opacity', 'Nodule/Mass',\n 'Other lesion', 'Pleural effusion', 'Pleural thickening', 'Pneumothorax',\n 'Pulmonary fibrosis']\n\n# decided on the basis of frequency of occurence of individual diseases in images.\n\n# Drop columns not in the list\ncolumns_to_keep = diseases.copy()\ncolumns_to_keep.append('image_id')\n\nprint(diseases)\nprint(columns_to_keep)","metadata":{"papermill":{"duration":0.02171,"end_time":"2024-03-26T18:39:30.230676","exception":false,"start_time":"2024-03-26T18:39:30.208966","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:34.993411Z","iopub.execute_input":"2024-04-18T12:44:34.993695Z","iopub.status.idle":"2024-04-18T12:44:34.999881Z","shell.execute_reply.started":"2024-04-18T12:44:34.993671Z","shell.execute_reply":"2024-04-18T12:44:34.998826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# some helper functions:-\n\ndef delete_columns(df, columns_to_keep=columns_to_keep, root_folder='/train'):\n    columns_to_drop = [col for col in df.columns if col not in columns_to_keep]\n    print(\"COLUMNS to drop:\", columns_to_drop)\n    df.drop(columns=columns_to_drop, inplace=True)\n    print(f\"Deleted columns from {root_folder.split('/')[-1]} folder\")\n    return df\n\ndef remove_all_zeros(df, diseases=diseases, root_folder='/train'):\n    # Remove rows where all values are 0 in the disease labels\n    df = df[(df[diseases] != 0).any(axis=1)]\n    df.reset_index(drop=True, inplace=True)\n    return df\n\ndef add_file_path_column(df, root_folder='/train'):\n    df['file_path'] = df['image_id'].apply(lambda x: os.path.join(root_folder, f\"{x}.npy\")) # adjust dicom or png\n    print(f\"Added file_path column to {root_folder.split('/')[-1]} folder\")\n    return df\n\n","metadata":{"papermill":{"duration":0.023135,"end_time":"2024-03-26T18:39:30.267193","exception":false,"start_time":"2024-03-26T18:39:30.244058","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:52.695041Z","iopub.execute_input":"2024-04-18T12:44:52.695802Z","iopub.status.idle":"2024-04-18T12:44:52.7035Z","shell.execute_reply.started":"2024-04-18T12:44:52.69577Z","shell.execute_reply":"2024-04-18T12:44:52.702559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\n# Load the original CSV file\ntrain_data = pd.read_csv(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv\")\n\n\n# Extract unique pairs of class_id and class_name\nclass_mapping = train_data[['class_id', 'class_name']].drop_duplicates()\n\n# Sort the mapping by class_id\nclass_mapping = class_mapping.sort_values(by='class_id')\n\n# Display the sorted mapping without the index\nprint(class_mapping.to_string(index=False))\n\ndisease_labels = class_mapping['class_name'].values\nprint(\"Disease LABELS :\", disease_labels)\nprint('-'*100)\n\nprint(\"Now converting into one hot vector format\")\n### Convert the csv into image_id -> 0,0,0, ..... 1., 0, 1 (15 labels format)\n\n# Convert class_id to string to enable one-hot encoding\ntrain_data['class_id'] = train_data['class_id'].astype(str)\n\n# Perform one-hot encoding to convert class IDs into binary columns\none_hot_encoded = pd.get_dummies(train_data['class_id'])\n\n# Convert boolean values to integers (0 for False, 1 for True)\none_hot_encoded = one_hot_encoded.astype(int)\n\n# Concatenate one-hot encoded columns with original DataFrame\ntrain_data = pd.concat([train_data, one_hot_encoded], axis=1)\n# print(train_data.columns)\n\n# Group by image ID and aggregate the one-hot encoded class columns for each class separately\ntrain_data = train_data.groupby('image_id').agg({\n#     'width': 'first',\n#     'height': 'first',\n    '0': 'max', '1': 'max', '2': 'max', '3': 'max', '4': 'max',\n    '5': 'max', '6': 'max', '7': 'max', '8': 'max', '9': 'max',\n    '10': 'max', '11': 'max', '12': 'max', '13': 'max', '14': 'max'\n}).reset_index()\n\n\n\nprint(\"UNIQUE images IN THE DATASET:\", len(train_data.image_id.unique()))\nprint(\"num rows in dataset  :\", len(train_data))\n# train_data.head()\n\n# Count occurrences of each class\nclass_counts = train_data.iloc[:, 1:].sum()\nprint(\"TOTAL Individual class counts:-\\n\", class_counts)\nprint('-'*100)\n\n\n# the 14 th class is no finding...\n\n\n# Save the new DataFrame to a new CSV file\n# train_data.to_csv(\"one_hot_vector_data.csv\", index=False)\n\n\n# Check if 'no-finding' class is marked while any other abnormality class is also marked\n# for index, row in train_data.iterrows():\n#     if row['14'] == 1 and row.iloc[1:].sum() > 1:\n#         print(\"Issue found in row:\", index)\n#         print(\"Other abnormality classes also marked in this row.\")\n\n\n\n# Rename the one-hot encoded columns\ntrain_data = train_data.rename(columns=dict(zip(train_data.columns[1:], disease_labels)))\nprint(\"Renamed the columns from numbers to disease_label names\")\n\n\n### Removing the 'no finding label' ###\ntrain_data = delete_columns(train_data, columns_to_keep=columns_to_keep, root_folder='train')\ntrain_data = remove_all_zeros(train_data, diseases=diseases, root_folder='train')\n\n\n\n## Adding file_path column\ntrain_data = add_file_path_column(train_data, root_folder='/kaggle/input/vinbigdata-512-voi-clahe/vinbigdata-512-resized-clahe2-8,8_train')\n\n# Split 10% of the training data as validation data, with random_state=42, for reproducibility\ntrain_data, val_data = train_test_split(train_data, test_size=0.2, random_state=42)\nval_data, test_data = train_test_split(val_data, test_size=0.5,random_state=42)\n\n\n# reducing the size (comment this)\n# train_data = train_data.head(500)\n# val_data = val_data.head(500)\n# test_data = test_data.head(500)\n\n# Resetting the index, very important\ntrain_data.reset_index(drop=True, inplace=True)\nval_data.reset_index(drop=True, inplace=True)\ntest_data.reset_index(drop=True, inplace=True)\n\n### ADJUST diseases & disease_labels ###\n# collecting & storing the labels separately\n\n#adjust diseases or disease_labels\ntrain_paths = train_data['file_path'].values\ntrain_labels = train_data[diseases].values\n\nval_paths = val_data['file_path'].values\nval_labels = val_data[diseases].values\n\ntest_paths = test_data['file_path'].values\ntest_labels = test_data[diseases].values\n\n\n\nprint(\"length of train:\", len(train_data), len(train_paths), len(train_labels))\nprint(\"length of val:\", len(val_data), len(val_data), len(val_data))\nprint(\"length of test:\", len(test_data), len(test_data), len(test_data))\n\n\nprint(\"Unique image_ids in train:\", len(train_data['image_id'].unique()))\nprint(\"Unique image_ids in val:\", len(val_data['image_id'].unique()))\nprint(\"Unique image_ids in test:\", len(test_data['image_id'].unique()))\n\n# train_data.info() except image_id all the other columns are of int datatype\n\n","metadata":{"papermill":{"duration":0.392919,"end_time":"2024-03-26T18:39:30.673449","exception":false,"start_time":"2024-03-26T18:39:30.28053","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:54.197227Z","iopub.execute_input":"2024-04-18T12:44:54.197572Z","iopub.status.idle":"2024-04-18T12:44:54.54654Z","shell.execute_reply.started":"2024-04-18T12:44:54.197547Z","shell.execute_reply":"2024-04-18T12:44:54.545507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Count occurrences of each class\nclass_counts = train_data.iloc[:, 1:-1].sum()\nprint(\"TRAIN Individual class counts:-\\n\", class_counts)\n\n# Count occurrences of each class\nclass_counts = val_data.iloc[:, 1:-1].sum()\nprint(\"VAL Individual class counts:-\\n\", class_counts)\n\n# Count occurrences of each class\nclass_counts = test_data.iloc[:, 1:-1].sum()\nprint(\"TEST Individual class counts:-\\n\", class_counts)","metadata":{"papermill":{"duration":0.02778,"end_time":"2024-03-26T18:39:30.715166","exception":false,"start_time":"2024-03-26T18:39:30.687386","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:55.38552Z","iopub.execute_input":"2024-04-18T12:44:55.385906Z","iopub.status.idle":"2024-04-18T12:44:55.397555Z","shell.execute_reply.started":"2024-04-18T12:44:55.385876Z","shell.execute_reply":"2024-04-18T12:44:55.396479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"papermill":{"duration":0.034975,"end_time":"2024-03-26T18:39:30.763917","exception":false,"start_time":"2024-03-26T18:39:30.728942","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:56.402036Z","iopub.execute_input":"2024-04-18T12:44:56.402406Z","iopub.status.idle":"2024-04-18T12:44:56.421939Z","shell.execute_reply.started":"2024-04-18T12:44:56.402375Z","shell.execute_reply":"2024-04-18T12:44:56.420774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_data.head()","metadata":{"papermill":{"duration":0.030583,"end_time":"2024-03-26T18:39:30.808584","exception":false,"start_time":"2024-03-26T18:39:30.778001","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:57.437501Z","iopub.execute_input":"2024-04-18T12:44:57.43825Z","iopub.status.idle":"2024-04-18T12:44:57.453837Z","shell.execute_reply.started":"2024-04-18T12:44:57.438219Z","shell.execute_reply":"2024-04-18T12:44:57.452759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data.head()","metadata":{"papermill":{"duration":0.031117,"end_time":"2024-03-26T18:39:30.854238","exception":false,"start_time":"2024-03-26T18:39:30.823121","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:44:58.374225Z","iopub.execute_input":"2024-04-18T12:44:58.374985Z","iopub.status.idle":"2024-04-18T12:44:58.389797Z","shell.execute_reply.started":"2024-04-18T12:44:58.374953Z","shell.execute_reply":"2024-04-18T12:44:58.388897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualising the images","metadata":{"papermill":{"duration":0.014869,"end_time":"2024-03-26T18:39:30.884307","exception":false,"start_time":"2024-03-26T18:39:30.869438","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_image_files = [\n    os.path.join('/', train_paths[i]) for i in range(len(train_paths))\n]\nprint(len(train_image_files))\n\nval_image_files = [\n    os.path.join('/', val_paths[i]) for i in range(len(val_paths))\n]\nprint(len(val_image_files))\n\ntest_image_files = [\n    os.path.join('/', test_paths[i]) for i in range(len(test_paths))\n]\nprint(len(test_image_files))","metadata":{"papermill":{"duration":0.049028,"end_time":"2024-03-26T18:39:30.948027","exception":false,"start_time":"2024-03-26T18:39:30.898999","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:01.968774Z","iopub.execute_input":"2024-04-18T12:45:01.96915Z","iopub.status.idle":"2024-04-18T12:45:01.985349Z","shell.execute_reply.started":"2024-04-18T12:45:01.969121Z","shell.execute_reply":"2024-04-18T12:45:01.984388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\n\ndef read_xray(path, voi_lut = True, fix_monochrome = True, apply_clahe=True, clipLimit=2.0, tileGridSize=(8,8)):\n    dicom = pydicom.read_file(path)\n    \n    # VOI LUT (if available by DICOM device) is used to transform raw DICOM data to \"human-friendly\" view\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n    \n    # depending on this value, X-ray may look inverted - fix that:\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n        \n    data = data - np.min(data)\n    data = data / np.max(data)\n#     data = data.astype(np.float32)\n#     data = (data * 255.0).astype(np.float32) # no need for this I think \n\n    if apply_clahe:\n        data = apply_clahe_to_image(data, clipLimit=clipLimit, tileGridSize=tileGridSize)\n        \n    return data\n\n\n\ndef apply_clahe_to_image(image, clipLimit=2.0, tileGridSize=(8,8)):\n    # Convert image to uint16\n    image = (image * 65535).astype(np.uint16)\n    \n    # Create a CLAHE object\n    clahe = cv2.createCLAHE(clipLimit=clipLimit, tileGridSize=tileGridSize)\n    \n    # Apply CLAHE\n    clahe_image = clahe.apply(image)\n    \n    # Convert image back to float32 in range [0, 1]\n    clahe_image = clahe_image.astype(np.float32) / 65535.0\n    \n    return clahe_image\n\n","metadata":{"papermill":{"duration":0.030672,"end_time":"2024-03-26T18:39:30.993912","exception":false,"start_time":"2024-03-26T18:39:30.96324","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:02.062826Z","iopub.execute_input":"2024-04-18T12:45:02.063093Z","iopub.status.idle":"2024-04-18T12:45:02.077216Z","shell.execute_reply.started":"2024-04-18T12:45:02.06307Z","shell.execute_reply":"2024-04-18T12:45:02.07642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_images(images, titles, rows, cols):\n    fig, axes = plt.subplots(rows, cols, figsize=(9, 9))\n    for i, (image, title) in enumerate(zip(images, titles)):\n        ax = axes[i // cols, i % cols]\n        ax.imshow(image, cmap='gray')\n        ax.set_title(title)\n        ax.axis('off')\n    plt.show()\n\n\npath = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/0005e8e3701dfb1dd93d53e2ff537b6e.dicom'\n# path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/0007d316f756b3fa0baea2ff514ce945.dicom'\n# path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/003cfe5ce5c0ec5163138eb3b740e328.dicom'\n# path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/0076d6a1e3139927fd62459c54276c3c.dicom'\n# path = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/0101ad90f31ddb8fb24e9935a3dac9db.dicom'\n\n\nclipLimits = [2.0, 3.0]\ntileGridSizes = [(6,6),(8,8)]\nimages = []\ntitles = []\n\n\nfor clipLimit in clipLimits:\n    for tileSize in tileGridSizes:\n        images.append(read_xray(path, clipLimit=clipLimit, tileGridSize=tileSize))\n        titles.append(f\"clipLimit={clipLimit}, tileGridSize={tileSize}\")\n    \ntitles = [f\"CL={clipLimit}, GS={tileSize}\" for clipLimit in clipLimits for tileSize in tileGridSizes]\nplot_images(images, titles, len(clipLimits), len(tileGridSizes))","metadata":{"papermill":{"duration":4.125995,"end_time":"2024-03-26T18:39:35.134647","exception":false,"start_time":"2024-03-26T18:39:31.008652","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:07.132908Z","iopub.execute_input":"2024-04-18T12:45:07.13328Z","iopub.status.idle":"2024-04-18T12:45:11.289761Z","shell.execute_reply.started":"2024-04-18T12:45:07.133252Z","shell.execute_reply":"2024-04-18T12:45:11.28876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = read_xray(path, voi_lut=False, apply_clahe=False)\nplt.figure(figsize = (8,8))\n# plt.imshow(img, cmap='bone')\nplt.imshow(img, cmap='gray')","metadata":{"papermill":{"duration":1.362619,"end_time":"2024-03-26T18:39:36.51676","exception":false,"start_time":"2024-03-26T18:39:35.154141","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:11.291298Z","iopub.execute_input":"2024-04-18T12:45:11.291993Z","iopub.status.idle":"2024-04-18T12:45:12.673634Z","shell.execute_reply.started":"2024-04-18T12:45:11.291963Z","shell.execute_reply":"2024-04-18T12:45:12.672757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = read_xray(path, voi_lut=True, apply_clahe=False)\nplt.figure(figsize = (8,8))\n# plt.imshow(img, cmap='bone')\nplt.imshow(img, cmap='gray')","metadata":{"papermill":{"duration":1.486622,"end_time":"2024-03-26T18:39:38.025894","exception":false,"start_time":"2024-03-26T18:39:36.539272","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:12.906999Z","iopub.execute_input":"2024-04-18T12:45:12.907801Z","iopub.status.idle":"2024-04-18T12:45:14.450676Z","shell.execute_reply.started":"2024-04-18T12:45:12.907697Z","shell.execute_reply":"2024-04-18T12:45:14.449748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# the best configuration\nimg = read_xray(path, apply_clahe=True, clipLimit=2.0, tileGridSize=(8,8))\nplt.figure(figsize = (8,8))\n# plt.imshow(img, cmap='bone')\nplt.imshow(img, cmap='gray')","metadata":{"papermill":{"duration":1.495401,"end_time":"2024-03-26T18:39:39.547959","exception":false,"start_time":"2024-03-26T18:39:38.052558","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:14.452328Z","iopub.execute_input":"2024-04-18T12:45:14.45262Z","iopub.status.idle":"2024-04-18T12:45:16.019536Z","shell.execute_reply.started":"2024-04-18T12:45:14.452594Z","shell.execute_reply":"2024-04-18T12:45:16.01867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### The preprocessed saved images:-","metadata":{"papermill":{"duration":0.030494,"end_time":"2024-03-26T18:39:39.610165","exception":false,"start_time":"2024-03-26T18:39:39.579671","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import pydicom\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\n\n# Assuming train_image_files is a list of DICOM file paths\n# Example:\n# train_image_files = [\"path/to/dicom/file1.dcm\", \"path/to/dicom/file2.dcm\", ...]\n\nfor i, k in enumerate(np.random.randint(len(train_image_files), size=9)):\n\n\n#     print(\"Image path\", train_image_files[k])\n    img_mat = np.load(train_image_files[k])\n\n#     img_mat = cv2.imread(train_image_files[k])\n# #     img_mat = cv2.resize(img_mat, (224, 224))\n#     img_black = cv2.cvtColor(img_mat, cv2.COLOR_BGR2GRAY)\n    float64_image_array = img_mat.astype(np.float64)\n\n    \n    print(\"image_shape:\", img_mat.shape)\n    \n    # Visualization using Matplotlib\n    plt.subplot(3, 3, i + 1)\n    plt.imshow(float64_image_array, cmap='gray')  # Assuming grayscale DICOM images\n    plt.title(f\"Image {i + 1}\")\n    plt.axis('off')\n\nplt.show()\n","metadata":{"papermill":{"duration":0.877317,"end_time":"2024-03-26T18:39:40.516851","exception":false,"start_time":"2024-03-26T18:39:39.639534","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:17.269444Z","iopub.execute_input":"2024-04-18T12:45:17.270283Z","iopub.status.idle":"2024-04-18T12:45:18.234938Z","shell.execute_reply.started":"2024-04-18T12:45:17.270244Z","shell.execute_reply":"2024-04-18T12:45:18.233921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 4\nIMAGE_SIZE = (224,224) # the size to which the image would be resized before passing to the model\n# num_classes = len(diseases)\nnum_classes = len(train_data.columns) - 2 \n## ADJUST num_classes accordingly \n\nprint(\"BATCH_SIZE: \", BATCH_SIZE) \nprint(\"IMAGE_SIZE: \", IMAGE_SIZE) \nprint(\"num_classes: \", num_classes) ","metadata":{"papermill":{"duration":0.041129,"end_time":"2024-03-26T18:39:40.592071","exception":false,"start_time":"2024-03-26T18:39:40.550942","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:18.750177Z","iopub.execute_input":"2024-04-18T12:45:18.75054Z","iopub.status.idle":"2024-04-18T12:45:18.757178Z","shell.execute_reply.started":"2024-04-18T12:45:18.750509Z","shell.execute_reply":"2024-04-18T12:45:18.756216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.transforms import Compose, Lambda, EnsureChannelFirst, ScaleIntensity,\\\n RandRotate, RandFlip, RandZoom, RandSpatialCrop, RandRotate90, ResizeWithPadOrCrop, Resize\n\ndef load_tensor(data):\n    # Convert the data to a supported data type (e.g., float32)\n    return data.astype(np.float32)\n\n\ntrain_transforms = Compose(\n    [\n#         LoadImage(image_only=True), \n        Lambda(load_tensor),\n        EnsureChannelFirst(channel_dim=0),\n#         RandSpatialCrop(IMAGE_SIZE[0],IMAGE_SIZE[1]), random_size=False), # crops randomly ; loss of information\n        Resize(spatial_size=(IMAGE_SIZE[0],IMAGE_SIZE[1])),\n#         ResizeWithPadOrCrop(spatial_size=IMAGE_SIZE), # this will centrally crop : loss of information\n#         RandRotate90(prob=0.5, spatial_axes=(0, 1)),\n        ScaleIntensity(),\n#         RandRotate(range_x=np.pi / 12, prob=0.5, keep_size=True),\n#         RandFlip(spatial_axis=0, prob=0.5),\n        RandZoom(min_zoom=0.9, max_zoom=1.1, prob=0.5),\n    ]\n)\n\nval_transforms = Compose(\n    [\n#         LoadImage(image_only=True),\n        Lambda(load_tensor),\n        EnsureChannelFirst(channel_dim=0),\n#         ResizeWithPadOrCrop(spatial_size=IMAGE_SIZE), # this will centrally crop : loss of information\n        Resize(spatial_size=(IMAGE_SIZE[0],IMAGE_SIZE[1])),\n\n        ScaleIntensity(),\n\n    ]\n)\n\n# y_pred_trans = Compose([Activations(softmax=True)])\n# y_trans = Compose([AsDiscrete(to_onehot=num_class)])\n\ny_pred_trans = Compose([Activations(sigmoid=True)]) # for multi-label classfication\ny_trans = Compose([AsDiscrete(threshold_values=True)])\n","metadata":{"papermill":{"duration":0.047102,"end_time":"2024-03-26T18:39:40.670556","exception":false,"start_time":"2024-03-26T18:39:40.623454","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:20.537579Z","iopub.execute_input":"2024-04-18T12:45:20.538492Z","iopub.status.idle":"2024-04-18T12:45:20.551403Z","shell.execute_reply.started":"2024-04-18T12:45:20.538457Z","shell.execute_reply":"2024-04-18T12:45:20.55034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VINDR_BigData_Dataset(torch.utils.data.Dataset):\n    def __init__(self, paths, labels, in_chans=1, transforms=None):\n        self.paths = paths\n        self.labels = labels\n        self.transforms = transforms\n        self.target_size = IMAGE_SIZE\n        self.in_chans = in_chans\n        \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, index):\n#         print(index)\n        img_path = self.paths[index]\n#         print(img_path, index)\n        labels = self.labels[index]\n#         image = read_xray(img_path, voi_lut=False, fix_monochrome=False, normalize=False)\n        \n#         img_mat = cv2.imread(img_path)\n#         img_black = cv2.cvtColor(img_mat, cv2.COLOR_BGR2GRAY)\n        \n        image = np.load(img_path)\n        \n        # Expand the dimensions to add a channel dimension\n        image = np.expand_dims(image, axis=0)\n        \n        if self.transforms:\n            image = self.transforms(image)\n        \n        return image, labels\n\n\ntrain_ds = VINDR_BigData_Dataset(train_paths, train_labels, in_chans=1, transforms=train_transforms)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=4)\n\nval_ds = VINDR_BigData_Dataset(val_paths, val_labels, in_chans=1, transforms=val_transforms)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, num_workers=4)\n\ntest_ds = VINDR_BigData_Dataset(test_paths, test_labels, in_chans=1, transforms=val_transforms)\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, num_workers=4)\n\nprint(\"No. of TRAIN batches:\", len(train_loader))\nprint(\"No. of VAL batches:\", len(val_loader))\nprint(\"No. of TEST batches:\", len(test_loader))\n","metadata":{"papermill":{"duration":0.04539,"end_time":"2024-03-26T18:39:40.747279","exception":false,"start_time":"2024-03-26T18:39:40.701889","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:21.90494Z","iopub.execute_input":"2024-04-18T12:45:21.905754Z","iopub.status.idle":"2024-04-18T12:45:21.917745Z","shell.execute_reply.started":"2024-04-18T12:45:21.905688Z","shell.execute_reply":"2024-04-18T12:45:21.916735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from thop import profile\nfrom thop import clever_format\nimport torch\n\ndef display_params_flops(model):\n    #params\n    num_params = sum(p.numel() for p in model.parameters())\n    num_params_millions = num_params / 1e6\n    print(f\"Number of parameters in millions: {num_params_millions:.2f} M\")\n\n    num_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    num_params_millions = num_params / 1e6\n    print(f\"Number of trainable parameters in millions: {num_params_millions:.2f} M\")\n\n\n    #FLOPS\n    input_size = (1, 1, IMAGE_SIZE[0], IMAGE_SIZE[1])  \n\n\n    # Move the model to GPU if available\n    if torch.cuda.is_available():\n        model = model.cuda()\n\n    # Use thop.profile to count FLOPs\n    input_tensor = torch.randn(*input_size)\n    if torch.cuda.is_available():\n        input_tensor = input_tensor.cuda()\n    flops, params = profile(model, inputs=(input_tensor,))\n\n    # Convert FLOPs to gigaFLOPs and format the results\n    flops, params = clever_format([flops, params], \"%.2f\")\n    print(f\"FLOPs: {flops}, Params: {params}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:23.02418Z","iopub.execute_input":"2024-04-18T12:45:23.025127Z","iopub.status.idle":"2024-04-18T12:45:23.033926Z","shell.execute_reply.started":"2024-04-18T12:45:23.025072Z","shell.execute_reply":"2024-04-18T12:45:23.032965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Definition","metadata":{"papermill":{"duration":0.030913,"end_time":"2024-03-26T18:39:40.809626","exception":false,"start_time":"2024-03-26T18:39:40.778713","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# _features.py\n\n\"\"\" PyTorch Feature Extraction Helpers\n\nA collection of classes, functions, modules to help extract features from models\nand provide a common interface for describing them.\n\nThe return_layers, module re-writing idea inspired by torchvision IntermediateLayerGetter\nhttps://github.com/pytorch/vision/blob/d88d8961ae51507d0cb680329d985b1488b1b76b/torchvision/models/_utils.py\n\nHacked together by / Copyright 2020 Ross Wightman\n\"\"\"\nfrom collections import OrderedDict, defaultdict\nfrom copy import deepcopy\nfrom functools import partial\nfrom typing import Dict, List, Optional, Sequence, Set, Tuple, Union\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.checkpoint import checkpoint\n\nfrom timm.layers import Format\n\n\n__all__ = [\n    'FeatureInfo', 'FeatureHooks', 'FeatureDictNet', 'FeatureListNet', 'FeatureHookNet', 'FeatureGetterNet',\n    'feature_take_indices'\n]\n\n\ndef _take_indices(\n        num_blocks: int,\n        n: Optional[Union[int, List[int], Tuple[int]]],\n) -> Tuple[Set[int], int]:\n    if isinstance(n, int):\n        assert n >= 0\n        take_indices = {x for x in range(num_blocks - n, num_blocks)}\n    else:\n        take_indices = {num_blocks + idx if idx < 0 else idx for idx in n}\n    return take_indices, max(take_indices)\n\n\ndef _take_indices_jit(\n        num_blocks: int,\n        n: Union[int, List[int], Tuple[int]],\n) -> Tuple[List[int], int]:\n    if isinstance(n, int):\n        assert n >= 0\n        take_indices = [num_blocks - n + i for i in range(n)]\n    elif isinstance(n, tuple):\n        # splitting this up is silly, but needed for torchscript type resolution of n\n        take_indices = [num_blocks + idx if idx < 0 else idx for idx in n]\n    else:\n        take_indices = [num_blocks + idx if idx < 0 else idx for idx in n]\n    return take_indices, max(take_indices)\n\n\ndef feature_take_indices(\n        num_blocks: int,\n        indices: Optional[Union[int, List[int], Tuple[int]]] = None,\n) -> Tuple[List[int], int]:\n    if indices is None:\n        indices = num_blocks  # all blocks if None\n    if torch.jit.is_scripting():\n        return _take_indices_jit(num_blocks, indices)\n    else:\n        # NOTE non-jit returns Set[int] instead of List[int] but torchscript can't handle that anno\n        return _take_indices(num_blocks, indices)\n\n\ndef _out_indices_as_tuple(x: Union[int, Tuple[int, ...]]) -> Tuple[int, ...]:\n    if isinstance(x, int):\n        # if indices is an int, take last N features\n        return tuple(range(-x, 0))\n    return tuple(x)\n\n\nOutIndicesT = Union[int, Tuple[int, ...]]\n\n\nclass FeatureInfo:\n\n    def __init__(\n            self,\n            feature_info: List[Dict],\n            out_indices: OutIndicesT,\n    ):\n        out_indices = _out_indices_as_tuple(out_indices)\n        prev_reduction = 1\n        for i, fi in enumerate(feature_info):\n            # sanity check the mandatory fields, there may be additional fields depending on the model\n            assert 'num_chs' in fi and fi['num_chs'] > 0\n            assert 'reduction' in fi and fi['reduction'] >= prev_reduction\n            prev_reduction = fi['reduction']\n            assert 'module' in fi\n            fi.setdefault('index', i)\n        self.out_indices = out_indices\n        self.info = feature_info\n\n    def from_other(self, out_indices: OutIndicesT):\n        out_indices = _out_indices_as_tuple(out_indices)\n        return FeatureInfo(deepcopy(self.info), out_indices)\n\n    def get(self, key: str, idx: Optional[Union[int, List[int]]] = None):\n        \"\"\" Get value by key at specified index (indices)\n        if idx == None, returns value for key at each output index\n        if idx is an integer, return value for that feature module index (ignoring output indices)\n        if idx is a list/tuple, return value for each module index (ignoring output indices)\n        \"\"\"\n        if idx is None:\n            return [self.info[i][key] for i in self.out_indices]\n        if isinstance(idx, (tuple, list)):\n            return [self.info[i][key] for i in idx]\n        else:\n            return self.info[idx][key]\n\n    def get_dicts(self, keys: Optional[List[str]] = None, idx: Optional[Union[int, List[int]]] = None):\n        \"\"\" return info dicts for specified keys (or all if None) at specified indices (or out_indices if None)\n        \"\"\"\n        if idx is None:\n            if keys is None:\n                return [self.info[i] for i in self.out_indices]\n            else:\n                return [{k: self.info[i][k] for k in keys} for i in self.out_indices]\n        if isinstance(idx, (tuple, list)):\n            return [self.info[i] if keys is None else {k: self.info[i][k] for k in keys} for i in idx]\n        else:\n            return self.info[idx] if keys is None else {k: self.info[idx][k] for k in keys}\n\n    def channels(self, idx: Optional[Union[int, List[int]]] = None):\n        \"\"\" feature channels accessor\n        \"\"\"\n        return self.get('num_chs', idx)\n\n    def reduction(self, idx: Optional[Union[int, List[int]]] = None):\n        \"\"\" feature reduction (output stride) accessor\n        \"\"\"\n        return self.get('reduction', idx)\n\n    def module_name(self, idx: Optional[Union[int, List[int]]] = None):\n        \"\"\" feature module name accessor\n        \"\"\"\n        return self.get('module', idx)\n\n    def __getitem__(self, item):\n        return self.info[item]\n\n    def __len__(self):\n        return len(self.info)\n\n\nclass FeatureHooks:\n    \"\"\" Feature Hook Helper\n\n    This module helps with the setup and extraction of hooks for extracting features from\n    internal nodes in a model by node name.\n\n    FIXME This works well in eager Python but needs redesign for torchscript.\n    \"\"\"\n\n    def __init__(\n            self,\n            hooks: Sequence[str],\n            named_modules: dict,\n            out_map: Sequence[Union[int, str]] = None,\n            default_hook_type: str = 'forward',\n    ):\n        # setup feature hooks\n        self._feature_outputs = defaultdict(OrderedDict)\n        modules = {k: v for k, v in named_modules}\n        for i, h in enumerate(hooks):\n            hook_name = h['module']\n            m = modules[hook_name]\n            hook_id = out_map[i] if out_map else hook_name\n            hook_fn = partial(self._collect_output_hook, hook_id)\n            hook_type = h.get('hook_type', default_hook_type)\n            if hook_type == 'forward_pre':\n                m.register_forward_pre_hook(hook_fn)\n            elif hook_type == 'forward':\n                m.register_forward_hook(hook_fn)\n            else:\n                assert False, \"Unsupported hook type\"\n\n    def _collect_output_hook(self, hook_id, *args):\n        x = args[-1]  # tensor we want is last argument, output for fwd, input for fwd_pre\n        if isinstance(x, tuple):\n            x = x[0]  # unwrap input tuple\n        self._feature_outputs[x.device][hook_id] = x\n\n    def get_output(self, device) -> Dict[str, torch.tensor]:\n        output = self._feature_outputs[device]\n        self._feature_outputs[device] = OrderedDict()  # clear after reading\n        return output\n\n\ndef _module_list(module, flatten_sequential=False):\n    # a yield/iter would be better for this but wouldn't be compatible with torchscript\n    ml = []\n    for name, module in module.named_children():\n        if flatten_sequential and isinstance(module, nn.Sequential):\n            # first level of Sequential containers is flattened into containing model\n            for child_name, child_module in module.named_children():\n                combined = [name, child_name]\n                ml.append(('_'.join(combined), '.'.join(combined), child_module))\n        else:\n            ml.append((name, name, module))\n    return ml\n\n\ndef _get_feature_info(net, out_indices: OutIndicesT):\n    feature_info = getattr(net, 'feature_info')\n    if isinstance(feature_info, FeatureInfo):\n        return feature_info.from_other(out_indices)\n    elif isinstance(feature_info, (list, tuple)):\n        return FeatureInfo(net.feature_info, out_indices)\n    else:\n        assert False, \"Provided feature_info is not valid\"\n\n\ndef _get_return_layers(feature_info, out_map):\n    module_names = feature_info.module_name()\n    return_layers = {}\n    for i, name in enumerate(module_names):\n        return_layers[name] = out_map[i] if out_map is not None else feature_info.out_indices[i]\n    return return_layers\n\n\nclass FeatureDictNet(nn.ModuleDict):\n    \"\"\" Feature extractor with OrderedDict return\n\n    Wrap a model and extract features as specified by the out indices, the network is\n    partially re-built from contained modules.\n\n    There is a strong assumption that the modules have been registered into the model in the same\n    order as they are used. There should be no reuse of the same nn.Module more than once, including\n    trivial modules like `self.relu = nn.ReLU`.\n\n    Only submodules that are directly assigned to the model class (`model.feature1`) or at most\n    one Sequential container deep (`model.features.1`, with flatten_sequent=True) can be captured.\n    All Sequential containers that are directly assigned to the original model will have their\n    modules assigned to this module with the name `model.features.1` being changed to `model.features_1`\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            out_indices: OutIndicesT = (0, 1, 2, 3, 4),\n            out_map: Sequence[Union[int, str]] = None,\n            output_fmt: str = 'NCHW',\n            feature_concat: bool = False,\n            flatten_sequential: bool = False,\n    ):\n        \"\"\"\n        Args:\n            model: Model from which to extract features.\n            out_indices: Output indices of the model features to extract.\n            out_map: Return id mapping for each output index, otherwise str(index) is used.\n            feature_concat: Concatenate intermediate features that are lists or tuples instead of selecting\n                first element e.g. `x[0]`\n            flatten_sequential: Flatten first two-levels of sequential modules in model (re-writes model modules)\n        \"\"\"\n        super(FeatureDictNet, self).__init__()\n        self.feature_info = _get_feature_info(model, out_indices)\n        self.output_fmt = Format(output_fmt)\n        self.concat = feature_concat\n        self.grad_checkpointing = False\n        self.return_layers = {}\n\n        return_layers = _get_return_layers(self.feature_info, out_map)\n        modules = _module_list(model, flatten_sequential=flatten_sequential)\n        remaining = set(return_layers.keys())\n        layers = OrderedDict()\n        for new_name, old_name, module in modules:\n            layers[new_name] = module\n            if old_name in remaining:\n                # return id has to be consistently str type for torchscript\n                self.return_layers[new_name] = str(return_layers[old_name])\n                remaining.remove(old_name)\n            if not remaining:\n                break\n        assert not remaining and len(self.return_layers) == len(return_layers), \\\n            f'Return layers ({remaining}) are not present in model'\n        self.update(layers)\n\n    def set_grad_checkpointing(self, enable: bool = True):\n        self.grad_checkpointing = enable\n\n    def _collect(self, x) -> (Dict[str, torch.Tensor]):\n        out = OrderedDict()\n        for i, (name, module) in enumerate(self.items()):\n            if self.grad_checkpointing and not torch.jit.is_scripting():\n                # Skipping checkpoint of first module because need a gradient at input\n                # Skipping last because networks with in-place ops might fail w/ checkpointing enabled\n                # NOTE: first_or_last module could be static, but recalc in is_scripting guard to avoid jit issues\n                first_or_last_module = i == 0 or i == max(len(self) - 1, 0)\n                x = module(x) if first_or_last_module else checkpoint(module, x)\n            else:\n                x = module(x)\n\n            if name in self.return_layers:\n                out_id = self.return_layers[name]\n                if isinstance(x, (tuple, list)):\n                    # If model tap is a tuple or list, concat or select first element\n                    # FIXME this may need to be more generic / flexible for some nets\n                    out[out_id] = torch.cat(x, 1) if self.concat else x[0]\n                else:\n                    out[out_id] = x\n        return out\n\n    def forward(self, x) -> Dict[str, torch.Tensor]:\n        return self._collect(x)\n\n\nclass FeatureListNet(FeatureDictNet):\n    \"\"\" Feature extractor with list return\n\n    A specialization of FeatureDictNet that always returns features as a list (values() of dict).\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            out_indices: OutIndicesT = (0, 1, 2, 3, 4),\n            output_fmt: str = 'NCHW',\n            feature_concat: bool = False,\n            flatten_sequential: bool = False,\n    ):\n        \"\"\"\n        Args:\n            model: Model from which to extract features.\n            out_indices: Output indices of the model features to extract.\n            feature_concat: Concatenate intermediate features that are lists or tuples instead of selecting\n                first element e.g. `x[0]`\n            flatten_sequential: Flatten first two-levels of sequential modules in model (re-writes model modules)\n        \"\"\"\n        super().__init__(\n            model,\n            out_indices=out_indices,\n            output_fmt=output_fmt,\n            feature_concat=feature_concat,\n            flatten_sequential=flatten_sequential,\n        )\n\n    def forward(self, x) -> (List[torch.Tensor]):\n        return list(self._collect(x).values())\n\n\nclass FeatureHookNet(nn.ModuleDict):\n    \"\"\" FeatureHookNet\n\n    Wrap a model and extract features specified by the out indices using forward/forward-pre hooks.\n\n    If `no_rewrite` is True, features are extracted via hooks without modifying the underlying\n    network in any way.\n\n    If `no_rewrite` is False, the model will be re-written as in the\n    FeatureList/FeatureDict case by folding first to second (Sequential only) level modules into this one.\n\n    FIXME this does not currently work with Torchscript, see FeatureHooks class\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            out_indices: OutIndicesT = (0, 1, 2, 3, 4),\n            out_map: Optional[Sequence[Union[int, str]]] = None,\n            return_dict: bool = False,\n            output_fmt: str = 'NCHW',\n            no_rewrite: bool = False,\n            flatten_sequential: bool = False,\n            default_hook_type: str = 'forward',\n    ):\n        \"\"\"\n\n        Args:\n            model: Model from which to extract features.\n            out_indices: Output indices of the model features to extract.\n            out_map: Return id mapping for each output index, otherwise str(index) is used.\n            return_dict: Output features as a dict.\n            no_rewrite: Enforce that model is not re-written if True, ie no modules are removed / changed.\n                flatten_sequential arg must also be False if this is set True.\n            flatten_sequential: Re-write modules by flattening first two levels of nn.Sequential containers.\n            default_hook_type: The default hook type to use if not specified in model.feature_info.\n        \"\"\"\n        super().__init__()\n        assert not torch.jit.is_scripting()\n        self.feature_info = _get_feature_info(model, out_indices)\n        self.return_dict = return_dict\n        self.output_fmt = Format(output_fmt)\n        self.grad_checkpointing = False\n\n        layers = OrderedDict()\n        hooks = []\n        if no_rewrite:\n            assert not flatten_sequential\n            if hasattr(model, 'reset_classifier'):  # make sure classifier is removed?\n                model.reset_classifier(0)\n            layers['body'] = model\n            hooks.extend(self.feature_info.get_dicts())\n        else:\n            modules = _module_list(model, flatten_sequential=flatten_sequential)\n            remaining = {\n                f['module']: f['hook_type'] if 'hook_type' in f else default_hook_type\n                for f in self.feature_info.get_dicts()\n            }\n            for new_name, old_name, module in modules:\n                layers[new_name] = module\n                for fn, fm in module.named_modules(prefix=old_name):\n                    if fn in remaining:\n                        hooks.append(dict(module=fn, hook_type=remaining[fn]))\n                        del remaining[fn]\n                if not remaining:\n                    break\n            assert not remaining, f'Return layers ({remaining}) are not present in model'\n        self.update(layers)\n        self.hooks = FeatureHooks(hooks, model.named_modules(), out_map=out_map)\n\n    def set_grad_checkpointing(self, enable: bool = True):\n        self.grad_checkpointing = enable\n\n    def forward(self, x):\n        for i, (name, module) in enumerate(self.items()):\n            if self.grad_checkpointing and not torch.jit.is_scripting():\n                # Skipping checkpoint of first module because need a gradient at input\n                # Skipping last because networks with in-place ops might fail w/ checkpointing enabled\n                # NOTE: first_or_last module could be static, but recalc in is_scripting guard to avoid jit issues\n                first_or_last_module = i == 0 or i == max(len(self) - 1, 0)\n                x = module(x) if first_or_last_module else checkpoint(module, x)\n            else:\n                x = module(x)\n        out = self.hooks.get_output(x.device)\n        return out if self.return_dict else list(out.values())\n\n\nclass FeatureGetterNet(nn.ModuleDict):\n    \"\"\" FeatureGetterNet\n\n    Wrap models with a feature getter method, like 'get_intermediate_layers'\n\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            out_indices: OutIndicesT = 4,\n            out_map: Optional[Sequence[Union[int, str]]] = None,\n            return_dict: bool = False,\n            output_fmt: str = 'NCHW',\n            norm: bool = False,\n            prune: bool = True,\n    ):\n        \"\"\"\n\n        Args:\n            model: Model to wrap.\n            out_indices: Indices of features to extract.\n            out_map: Remap feature names for dict output (WIP, not supported).\n            return_dict: Return features as dictionary instead of list (WIP, not supported).\n            norm: Apply final model norm to all output features (if possible).\n        \"\"\"\n        super().__init__()\n        if prune and hasattr(model, 'prune_intermediate_layers'):\n            # replace out_indices after they've been normalized, -ve indices will be invalid after prune\n            out_indices = model.prune_intermediate_layers(\n                out_indices,\n                prune_norm=not norm,\n            )\n            out_indices = list(out_indices)\n        self.feature_info = _get_feature_info(model, out_indices)\n        self.model = model\n        self.out_indices = out_indices\n        self.out_map = out_map\n        self.return_dict = return_dict\n        self.output_fmt = output_fmt\n        self.norm = norm\n\n    def forward(self, x):\n        features = self.model.forward_intermediates(\n            x,\n            indices=self.out_indices,\n            norm=self.norm,\n            output_fmt=self.output_fmt,\n            intermediates_only=True,\n        )\n        return features","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:25.434084Z","iopub.execute_input":"2024-04-18T12:45:25.434416Z","iopub.status.idle":"2024-04-18T12:45:25.511058Z","shell.execute_reply.started":"2024-04-18T12:45:25.434392Z","shell.execute_reply":"2024-04-18T12:45:25.510194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# _features_fx.py\n\n\"\"\" PyTorch FX Based Feature Extraction Helpers\nUsing https://pytorch.org/vision/stable/feature_extraction.html\n\"\"\"\nfrom typing import Callable, Dict, List, Optional, Union, Tuple, Type\n\nimport torch\nfrom torch import nn\n\n# from ._features import _get_feature_info, _get_return_layers\n\ntry:\n    from torchvision.models.feature_extraction import create_feature_extractor as _create_feature_extractor\n    has_fx_feature_extraction = True\nexcept ImportError:\n    has_fx_feature_extraction = False\n\n# Layers we went to treat as leaf modules\nfrom timm.layers import Conv2dSame, ScaledStdConv2dSame, CondConv2d, StdConv2dSame\nfrom timm.layers.non_local_attn import BilinearAttnTransform\nfrom timm.layers.pool2d_same import MaxPool2dSame, AvgPool2dSame\nfrom timm.layers.norm_act import (\n    BatchNormAct2d,\n    SyncBatchNormAct,\n    FrozenBatchNormAct2d,\n    GroupNormAct,\n    GroupNorm1Act,\n    LayerNormAct,\n    LayerNormAct2d\n)\n\n__all__ = ['register_notrace_module', 'is_notrace_module', 'get_notrace_modules',\n           'register_notrace_function', 'is_notrace_function', 'get_notrace_functions',\n           'create_feature_extractor', 'FeatureGraphNet', 'GraphExtractNet']\n\n\n# NOTE: By default, any modules from timm.models.layers that we want to treat as leaf modules go here\n# BUT modules from timm.models should use the registration mechanism below\n_leaf_modules = {\n    BilinearAttnTransform,  # reason: flow control t <= 1\n    # Reason: get_same_padding has a max which raises a control flow error\n    Conv2dSame, MaxPool2dSame, ScaledStdConv2dSame, StdConv2dSame, AvgPool2dSame,\n    CondConv2d,  # reason: TypeError: F.conv2d received Proxy in groups=self.groups * B (because B = x.shape[0]),\n    BatchNormAct2d,\n    SyncBatchNormAct,\n    FrozenBatchNormAct2d,\n    GroupNormAct,\n    GroupNorm1Act,\n    LayerNormAct,\n    LayerNormAct2d,\n}\n\ntry:\n    from timm.layers import InplaceAbn\n    _leaf_modules.add(InplaceAbn)\nexcept ImportError:\n    pass\n\n\ndef register_notrace_module(module: Type[nn.Module]):\n    \"\"\"\n    Any module not under timm.models.layers should get this decorator if we don't want to trace through it.\n    \"\"\"\n    _leaf_modules.add(module)\n    return module\n\n\ndef is_notrace_module(module: Type[nn.Module]):\n    return module in _leaf_modules\n\n\ndef get_notrace_modules():\n    return list(_leaf_modules)\n\n\n# Functions we want to autowrap (treat them as leaves)\n_autowrap_functions = set()\n\n\ndef register_notrace_function(func: Callable):\n    \"\"\"\n    Decorator for functions which ought not to be traced through\n    \"\"\"\n    _autowrap_functions.add(func)\n    return func\n\n\ndef is_notrace_function(func: Callable):\n    return func in _autowrap_functions\n\n\ndef get_notrace_functions():\n    return list(_autowrap_functions)\n\n\ndef create_feature_extractor(model: nn.Module, return_nodes: Union[Dict[str, str], List[str]]):\n    assert has_fx_feature_extraction, 'Please update to PyTorch 1.10+, torchvision 0.11+ for FX feature extraction'\n    return _create_feature_extractor(\n        model, return_nodes,\n        tracer_kwargs={'leaf_modules': list(_leaf_modules), 'autowrap_functions': list(_autowrap_functions)}\n    )\n\n\nclass FeatureGraphNet(nn.Module):\n    \"\"\" A FX Graph based feature extractor that works with the model feature_info metadata\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            out_indices: Tuple[int, ...],\n            out_map: Optional[Dict] = None,\n    ):\n        super().__init__()\n        assert has_fx_feature_extraction, 'Please update to PyTorch 1.10+, torchvision 0.11+ for FX feature extraction'\n        self.feature_info = _get_feature_info(model, out_indices)\n        if out_map is not None:\n            assert len(out_map) == len(out_indices)\n        return_nodes = _get_return_layers(self.feature_info, out_map)\n        self.graph_module = create_feature_extractor(model, return_nodes)\n\n    def forward(self, x):\n        return list(self.graph_module(x).values())\n\n\nclass GraphExtractNet(nn.Module):\n    \"\"\" A standalone feature extraction wrapper that maps dict -> list or single tensor\n    NOTE:\n      * one can use feature_extractor directly if dictionary output is desired\n      * unlike FeatureGraphNet, this is intended to be used standalone and not with model feature_info\n      metadata for builtin feature extraction mode\n      * create_feature_extractor can be used directly if dictionary output is acceptable\n\n    Args:\n        model: model to extract features from\n        return_nodes: node names to return features from (dict or list)\n        squeeze_out: if only one output, and output in list format, flatten to single tensor\n    \"\"\"\n    def __init__(\n            self,\n            model: nn.Module,\n            return_nodes: Union[Dict[str, str], List[str]],\n            squeeze_out: bool = True,\n    ):\n        super().__init__()\n        self.squeeze_out = squeeze_out\n        self.graph_module = create_feature_extractor(model, return_nodes)\n\n    def forward(self, x) -> Union[List[torch.Tensor], torch.Tensor]:\n        out = list(self.graph_module(x).values())\n        if self.squeeze_out and len(out) == 1:\n            return out[0]\n        return out","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:26.135532Z","iopub.execute_input":"2024-04-18T12:45:26.135997Z","iopub.status.idle":"2024-04-18T12:45:26.159544Z","shell.execute_reply.started":"2024-04-18T12:45:26.135965Z","shell.execute_reply":"2024-04-18T12:45:26.15849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ._builder.py\n\n\nimport dataclasses\nimport logging\nimport os\nfrom copy import deepcopy\nfrom typing import Optional, Dict, Callable, Any, Tuple\n\nfrom torch import nn as nn\nfrom torch.hub import load_state_dict_from_url\n\n# from timm.models._features import FeatureListNet, FeatureDictNet, FeatureHookNet, FeatureGetterNet\nfrom timm.models._features_fx import FeatureGraphNet\nfrom timm.models._helpers import load_state_dict\nfrom timm.models._hub import has_hf_hub, download_cached_file, check_cached_file, load_state_dict_from_hf\nfrom timm.models._manipulate import adapt_input_conv\nfrom timm.models._pretrained import PretrainedCfg\nfrom timm.models._prune import adapt_model_from_file\nfrom timm.models._registry import get_pretrained_cfg\n\n_logger = logging.getLogger(__name__)\n\n# Global variables for rarely used pretrained checkpoint download progress and hash check.\n# Use set_pretrained_download_progress / set_pretrained_check_hash functions to toggle.\n_DOWNLOAD_PROGRESS = False\n_CHECK_HASH = False\n_USE_OLD_CACHE = int(os.environ.get('TIMM_USE_OLD_CACHE', 0)) > 0\n\n__all__ = ['set_pretrained_download_progress', 'set_pretrained_check_hash', 'load_custom_pretrained', 'load_pretrained',\n           'pretrained_cfg_for_features', 'resolve_pretrained_cfg', 'build_model_with_cfg']\n\n\ndef _resolve_pretrained_source(pretrained_cfg):\n    cfg_source = pretrained_cfg.get('source', '')\n    pretrained_url = pretrained_cfg.get('url', None)\n    pretrained_file = pretrained_cfg.get('file', None)\n    pretrained_sd = pretrained_cfg.get('state_dict', None)\n    hf_hub_id = pretrained_cfg.get('hf_hub_id', None)\n\n    # resolve where to load pretrained weights from\n    load_from = ''\n    pretrained_loc = ''\n    if cfg_source == 'hf-hub' and has_hf_hub(necessary=True):\n        # hf-hub specified as source via model identifier\n        load_from = 'hf-hub'\n        assert hf_hub_id\n        pretrained_loc = hf_hub_id\n    else:\n        # default source == timm or unspecified\n        if pretrained_sd:\n            # direct state_dict pass through is the highest priority\n            load_from = 'state_dict'\n            pretrained_loc = pretrained_sd\n            assert isinstance(pretrained_loc, dict)\n        elif pretrained_file:\n            # file load override is the second-highest priority if set\n            load_from = 'file'\n            pretrained_loc = pretrained_file\n        else:\n            old_cache_valid = False\n            if _USE_OLD_CACHE:\n                # prioritized old cached weights if exists and env var enabled\n                old_cache_valid = check_cached_file(pretrained_url) if pretrained_url else False\n            if not old_cache_valid and hf_hub_id and has_hf_hub(necessary=True):\n                # hf-hub available as alternate weight source in default_cfg\n                load_from = 'hf-hub'\n                pretrained_loc = hf_hub_id\n            elif pretrained_url:\n                load_from = 'url'\n                pretrained_loc = pretrained_url\n\n    if load_from == 'hf-hub' and pretrained_cfg.get('hf_hub_filename', None):\n        # if a filename override is set, return tuple for location w/ (hub_id, filename)\n        pretrained_loc = pretrained_loc, pretrained_cfg['hf_hub_filename']\n    return load_from, pretrained_loc\n\n\ndef set_pretrained_download_progress(enable=True):\n    \"\"\" Set download progress for pretrained weights on/off (globally). \"\"\"\n    global _DOWNLOAD_PROGRESS\n    _DOWNLOAD_PROGRESS = enable\n\n\ndef set_pretrained_check_hash(enable=True):\n    \"\"\" Set hash checking for pretrained weights on/off (globally). \"\"\"\n    global _CHECK_HASH\n    _CHECK_HASH = enable\n\n\ndef load_custom_pretrained(\n        model: nn.Module,\n        pretrained_cfg: Optional[Dict] = None,\n        load_fn: Optional[Callable] = None,\n):\n    r\"\"\"Loads a custom (read non .pth) weight file\n\n    Downloads checkpoint file into cache-dir like torch.hub based loaders, but calls\n    a passed in custom load fun, or the `load_pretrained` model member fn.\n\n    If the object is already present in `model_dir`, it's deserialized and returned.\n    The default value of `model_dir` is ``<hub_dir>/checkpoints`` where\n    `hub_dir` is the directory returned by :func:`~torch.hub.get_dir`.\n\n    Args:\n        model: The instantiated model to load weights into\n        pretrained_cfg (dict): Default pretrained model cfg\n        load_fn: An external standalone fn that loads weights into provided model, otherwise a fn named\n            'laod_pretrained' on the model will be called if it exists\n    \"\"\"\n    pretrained_cfg = pretrained_cfg or getattr(model, 'pretrained_cfg', None)\n    if not pretrained_cfg:\n        _logger.warning(\"Invalid pretrained config, cannot load weights.\")\n        return\n\n    load_from, pretrained_loc = _resolve_pretrained_source(pretrained_cfg)\n    if not load_from:\n        _logger.warning(\"No pretrained weights exist for this model. Using random initialization.\")\n        return\n    if load_from == 'hf-hub':\n        _logger.warning(\"Hugging Face hub not currently supported for custom load pretrained models.\")\n    elif load_from == 'url':\n        pretrained_loc = download_cached_file(\n            pretrained_loc,\n            check_hash=_CHECK_HASH,\n            progress=_DOWNLOAD_PROGRESS,\n        )\n\n    if load_fn is not None:\n        load_fn(model, pretrained_loc)\n    elif hasattr(model, 'load_pretrained'):\n        model.load_pretrained(pretrained_loc)\n    else:\n        _logger.warning(\"Valid function to load pretrained weights is not available, using random initialization.\")\n\n\ndef load_pretrained(\n        model: nn.Module,\n        pretrained_cfg: Optional[Dict] = None,\n        num_classes: int = 1000,\n        in_chans: int = 3,\n        filter_fn: Optional[Callable] = None,\n        strict: bool = True,\n):\n    \"\"\" Load pretrained checkpoint\n\n    Args:\n        model (nn.Module) : PyTorch model module\n        pretrained_cfg (Optional[Dict]): configuration for pretrained weights / target dataset\n        num_classes (int): num_classes for target model\n        in_chans (int): in_chans for target model\n        filter_fn (Optional[Callable]): state_dict filter fn for load (takes state_dict, model as args)\n        strict (bool): strict load of checkpoint\n\n    \"\"\"\n    pretrained_cfg = pretrained_cfg or getattr(model, 'pretrained_cfg', None)\n    if not pretrained_cfg:\n        raise RuntimeError(\"Invalid pretrained config, cannot load weights. Use `pretrained=False` for random init.\")\n\n    load_from, pretrained_loc = _resolve_pretrained_source(pretrained_cfg)\n    if load_from == 'state_dict':\n        _logger.info(f'Loading pretrained weights from state dict')\n        state_dict = pretrained_loc  # pretrained_loc is the actual state dict for this override\n    elif load_from == 'file':\n        _logger.info(f'Loading pretrained weights from file ({pretrained_loc})')\n        if pretrained_cfg.get('custom_load', False):\n            model.load_pretrained(pretrained_loc)\n            return\n        else:\n            state_dict = load_state_dict(pretrained_loc)\n    elif load_from == 'url':\n        _logger.info(f'Loading pretrained weights from url ({pretrained_loc})')\n        if pretrained_cfg.get('custom_load', False):\n            pretrained_loc = download_cached_file(\n                pretrained_loc,\n                progress=_DOWNLOAD_PROGRESS,\n                check_hash=_CHECK_HASH,\n            )\n            model.load_pretrained(pretrained_loc)\n            return\n        else:\n            state_dict = load_state_dict_from_url(\n                pretrained_loc,\n                map_location='cpu',\n                progress=_DOWNLOAD_PROGRESS,\n                check_hash=_CHECK_HASH,\n            )\n    elif load_from == 'hf-hub':\n        _logger.info(f'Loading pretrained weights from Hugging Face hub ({pretrained_loc})')\n        if isinstance(pretrained_loc, (list, tuple)):\n            state_dict = load_state_dict_from_hf(*pretrained_loc)\n        else:\n            state_dict = load_state_dict_from_hf(pretrained_loc)\n    else:\n        model_name = pretrained_cfg.get('architecture', 'this model')\n        raise RuntimeError(f\"No pretrained weights exist for {model_name}. Use `pretrained=False` for random init.\")\n\n    if filter_fn is not None:\n        try:\n            state_dict = filter_fn(state_dict, model)\n        except TypeError as e:\n            # for backwards compat with filter fn that take one arg\n            state_dict = filter_fn(state_dict)\n\n    input_convs = pretrained_cfg.get('first_conv', None)\n    if input_convs is not None and in_chans != 3:\n        if isinstance(input_convs, str):\n            input_convs = (input_convs,)\n        for input_conv_name in input_convs:\n            weight_name = input_conv_name + '.weight'\n            try:\n                state_dict[weight_name] = adapt_input_conv(in_chans, state_dict[weight_name])\n                _logger.info(\n                    f'Converted input conv {input_conv_name} pretrained weights from 3 to {in_chans} channel(s)')\n            except NotImplementedError as e:\n                del state_dict[weight_name]\n                strict = False\n                _logger.warning(\n                    f'Unable to convert pretrained {input_conv_name} weights, using random init for this layer.')\n\n    classifiers = pretrained_cfg.get('classifier', None)\n    label_offset = pretrained_cfg.get('label_offset', 0)\n    if classifiers is not None:\n        if isinstance(classifiers, str):\n            classifiers = (classifiers,)\n        if num_classes != pretrained_cfg['num_classes']:\n            for classifier_name in classifiers:\n                # completely discard fully connected if model num_classes doesn't match pretrained weights\n                state_dict.pop(classifier_name + '.weight', None)\n                state_dict.pop(classifier_name + '.bias', None)\n            strict = False\n        elif label_offset > 0:\n            for classifier_name in classifiers:\n                # special case for pretrained weights with an extra background class in pretrained weights\n                classifier_weight = state_dict[classifier_name + '.weight']\n                state_dict[classifier_name + '.weight'] = classifier_weight[label_offset:]\n                classifier_bias = state_dict[classifier_name + '.bias']\n                state_dict[classifier_name + '.bias'] = classifier_bias[label_offset:]\n\n    load_result = model.load_state_dict(state_dict, strict=strict)\n    if load_result.missing_keys:\n        _logger.info(\n            f'Missing keys ({\", \".join(load_result.missing_keys)}) discovered while loading pretrained weights.'\n            f' This is expected if model is being adapted.')\n    if load_result.unexpected_keys:\n        _logger.warning(\n            f'Unexpected keys ({\", \".join(load_result.unexpected_keys)}) found while loading pretrained weights.'\n            f' This may be expected if model is being adapted.')\n\n\ndef pretrained_cfg_for_features(pretrained_cfg):\n    pretrained_cfg = deepcopy(pretrained_cfg)\n    # remove default pretrained cfg fields that don't have much relevance for feature backbone\n    to_remove = ('num_classes', 'classifier', 'global_pool')  # add default final pool size?\n    for tr in to_remove:\n        pretrained_cfg.pop(tr, None)\n    return pretrained_cfg\n\n\ndef _filter_kwargs(kwargs, names):\n    if not kwargs or not names:\n        return\n    for n in names:\n        kwargs.pop(n, None)\n\n\ndef _update_default_model_kwargs(pretrained_cfg, kwargs, kwargs_filter):\n    \"\"\" Update the default_cfg and kwargs before passing to model\n\n    Args:\n        pretrained_cfg: input pretrained cfg (updated in-place)\n        kwargs: keyword args passed to model build fn (updated in-place)\n        kwargs_filter: keyword arg keys that must be removed before model __init__\n    \"\"\"\n    # Set model __init__ args that can be determined by default_cfg (if not already passed as kwargs)\n    default_kwarg_names = ('num_classes', 'global_pool', 'in_chans')\n    if pretrained_cfg.get('fixed_input_size', False):\n        # if fixed_input_size exists and is True, model takes an img_size arg that fixes its input size\n        default_kwarg_names += ('img_size',)\n\n    for n in default_kwarg_names:\n        # for legacy reasons, model __init__args uses img_size + in_chans as separate args while\n        # pretrained_cfg has one input_size=(C, H ,W) entry\n        if n == 'img_size':\n            input_size = pretrained_cfg.get('input_size', None)\n            if input_size is not None:\n                assert len(input_size) == 3\n                kwargs.setdefault(n, input_size[-2:])\n        elif n == 'in_chans':\n            input_size = pretrained_cfg.get('input_size', None)\n            if input_size is not None:\n                assert len(input_size) == 3\n                kwargs.setdefault(n, input_size[0])\n        elif n == 'num_classes':\n            default_val = pretrained_cfg.get(n, None)\n            # if default is < 0, don't pass through to model\n            if default_val is not None and default_val >= 0:\n                kwargs.setdefault(n, pretrained_cfg[n])\n        else:\n            default_val = pretrained_cfg.get(n, None)\n            if default_val is not None:\n                kwargs.setdefault(n, pretrained_cfg[n])\n\n    # Filter keyword args for task specific model variants (some 'features only' models, etc.)\n    _filter_kwargs(kwargs, names=kwargs_filter)\n\n\ndef resolve_pretrained_cfg(\n        variant: str,\n        pretrained_cfg=None,\n        pretrained_cfg_overlay=None,\n) -> PretrainedCfg:\n    model_with_tag = variant\n    pretrained_tag = None\n    if pretrained_cfg:\n        if isinstance(pretrained_cfg, dict):\n            # pretrained_cfg dict passed as arg, validate by converting to PretrainedCfg\n            pretrained_cfg = PretrainedCfg(**pretrained_cfg)\n        elif isinstance(pretrained_cfg, str):\n            pretrained_tag = pretrained_cfg\n            pretrained_cfg = None\n\n    # fallback to looking up pretrained cfg in model registry by variant identifier\n    if not pretrained_cfg:\n        if pretrained_tag:\n            model_with_tag = '.'.join([variant, pretrained_tag])\n        pretrained_cfg = get_pretrained_cfg(model_with_tag)\n\n    if not pretrained_cfg:\n        _logger.warning(\n            f\"No pretrained configuration specified for {model_with_tag} model. Using a default.\"\n            f\" Please add a config to the model pretrained_cfg registry or pass explicitly.\")\n        pretrained_cfg = PretrainedCfg()  # instance with defaults\n\n    pretrained_cfg_overlay = pretrained_cfg_overlay or {}\n    if not pretrained_cfg.architecture:\n        pretrained_cfg_overlay.setdefault('architecture', variant)\n    pretrained_cfg = dataclasses.replace(pretrained_cfg, **pretrained_cfg_overlay)\n\n    return pretrained_cfg\n\n\ndef build_model_with_cfg(\n        model_cls: Callable,\n        variant: str,\n        pretrained: bool,\n        pretrained_cfg: Optional[Dict] = None,\n        pretrained_cfg_overlay: Optional[Dict] = None,\n        model_cfg: Optional[Any] = None,\n        feature_cfg: Optional[Dict] = None,\n        pretrained_strict: bool = True,\n        pretrained_filter_fn: Optional[Callable] = None,\n        kwargs_filter: Optional[Tuple[str]] = None,\n        **kwargs,\n):\n    \"\"\" Build model with specified default_cfg and optional model_cfg\n\n    This helper fn aids in the construction of a model including:\n      * handling default_cfg and associated pretrained weight loading\n      * passing through optional model_cfg for models with config based arch spec\n      * features_only model adaptation\n      * pruning config / model adaptation\n\n    Args:\n        model_cls (nn.Module): model class\n        variant (str): model variant name\n        pretrained (bool): load pretrained weights\n        pretrained_cfg (dict): model's pretrained weight/task config\n        model_cfg (Optional[Dict]): model's architecture config\n        feature_cfg (Optional[Dict]: feature extraction adapter config\n        pretrained_strict (bool): load pretrained weights strictly\n        pretrained_filter_fn (Optional[Callable]): filter callable for pretrained weights\n        kwargs_filter (Optional[Tuple]): kwargs to filter before passing to model\n        **kwargs: model args passed through to model __init__\n    \"\"\"\n    pruned = kwargs.pop('pruned', False)\n    features = False\n    feature_cfg = feature_cfg or {}\n\n    # resolve and update model pretrained config and model kwargs\n    pretrained_cfg = resolve_pretrained_cfg(\n        variant,\n        pretrained_cfg=pretrained_cfg,\n        pretrained_cfg_overlay=pretrained_cfg_overlay\n    )\n\n    # FIXME converting back to dict, PretrainedCfg use should be propagated further, but not into model\n    pretrained_cfg = pretrained_cfg.to_dict()\n\n    _update_default_model_kwargs(pretrained_cfg, kwargs, kwargs_filter)\n\n    # Setup for feature extraction wrapper done at end of this fn\n    if kwargs.pop('features_only', False):\n        features = True\n        feature_cfg.setdefault('out_indices', (0, 1, 2, 3, 4))\n        if 'out_indices' in kwargs:\n            feature_cfg['out_indices'] = kwargs.pop('out_indices')\n\n    # Instantiate the model\n    if model_cfg is None:\n        model = model_cls(**kwargs)\n    else:\n        model = model_cls(cfg=model_cfg, **kwargs)\n    model.pretrained_cfg = pretrained_cfg\n    model.default_cfg = model.pretrained_cfg  # alias for backwards compat\n\n    if pruned:\n        model = adapt_model_from_file(model, variant)\n\n    # For classification models, check class attr, then kwargs, then default to 1k, otherwise 0 for feats\n    num_classes_pretrained = 0 if features else getattr(model, 'num_classes', kwargs.get('num_classes', 1000))\n    if pretrained:\n        load_pretrained(\n            model,\n            pretrained_cfg=pretrained_cfg,\n            num_classes=num_classes_pretrained,\n            in_chans=kwargs.get('in_chans', 3),\n            filter_fn=pretrained_filter_fn,\n            strict=pretrained_strict,\n        )\n\n    # Wrap the model in a feature extraction module if enabled\n    if features:\n        feature_cls = FeatureListNet\n        output_fmt = getattr(model, 'output_fmt', None)\n        if output_fmt is not None:\n            feature_cfg.setdefault('output_fmt', output_fmt)\n        if 'feature_cls' in feature_cfg:\n            feature_cls = feature_cfg.pop('feature_cls')\n            if isinstance(feature_cls, str):\n                feature_cls = feature_cls.lower()\n                if 'hook' in feature_cls:\n                    feature_cls = FeatureHookNet\n                elif feature_cls == 'dict':\n                    feature_cls = FeatureDictNet\n                elif feature_cls == 'fx':\n                    feature_cls = FeatureGraphNet\n                elif feature_cls == 'getter':\n                    feature_cls = FeatureGetterNet\n                else:\n                    assert False, f'Unknown feature class {feature_cls}'\n        model = feature_cls(model, **feature_cfg)\n        model.pretrained_cfg = pretrained_cfg_for_features(pretrained_cfg)  # add back pretrained cfg\n        model.default_cfg = model.pretrained_cfg  # alias for rename backwards compat (default_cfg -> pretrained_cfg)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:27.157831Z","iopub.execute_input":"2024-04-18T12:45:27.158154Z","iopub.status.idle":"2024-04-18T12:45:27.214568Z","shell.execute_reply.started":"2024-04-18T12:45:27.158129Z","shell.execute_reply":"2024-04-18T12:45:27.213493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# _manipulate.py\n\nimport collections.abc\nimport math\nimport re\nfrom collections import defaultdict\nfrom itertools import chain\nfrom typing import Any, Callable, Dict, Iterator, Tuple, Type, Union\n\nimport torch\nfrom torch import nn as nn\nfrom torch.utils.checkpoint import checkpoint\n\n__all__ = ['model_parameters', 'named_apply', 'named_modules', 'named_modules_with_params', 'adapt_input_conv',\n           'group_with_matcher', 'group_modules', 'group_parameters', 'flatten_modules', 'checkpoint_seq']\n\n\ndef model_parameters(model: nn.Module, exclude_head: bool = False):\n    if exclude_head:\n        # FIXME this a bit of a quick and dirty hack to skip classifier head params based on ordering\n        return [p for p in model.parameters()][:-2]\n    else:\n        return model.parameters()\n\n\ndef named_apply(\n        fn: Callable,\n        module: nn.Module, name='',\n        depth_first: bool = True,\n        include_root: bool = False,\n) -> nn.Module:\n    if not depth_first and include_root:\n        fn(module=module, name=name)\n    for child_name, child_module in module.named_children():\n        child_name = '.'.join((name, child_name)) if name else child_name\n        named_apply(fn=fn, module=child_module, name=child_name, depth_first=depth_first, include_root=True)\n    if depth_first and include_root:\n        fn(module=module, name=name)\n    return module\n\n\ndef named_modules(\n        module: nn.Module,\n        name: str = '',\n        depth_first: bool = True,\n        include_root: bool = False,\n):\n    if not depth_first and include_root:\n        yield name, module\n    for child_name, child_module in module.named_children():\n        child_name = '.'.join((name, child_name)) if name else child_name\n        yield from named_modules(\n            module=child_module, name=child_name, depth_first=depth_first, include_root=True)\n    if depth_first and include_root:\n        yield name, module\n\n\ndef named_modules_with_params(\n        module: nn.Module,\n        name: str = '',\n        depth_first: bool = True,\n        include_root: bool = False,\n):\n    if module._parameters and not depth_first and include_root:\n        yield name, module\n    for child_name, child_module in module.named_children():\n        child_name = '.'.join((name, child_name)) if name else child_name\n        yield from named_modules_with_params(\n            module=child_module, name=child_name, depth_first=depth_first, include_root=True)\n    if module._parameters and depth_first and include_root:\n        yield name, module\n\n\nMATCH_PREV_GROUP = (99999,)\n\n\ndef group_with_matcher(\n        named_objects: Iterator[Tuple[str, Any]],\n        group_matcher: Union[Dict, Callable],\n        return_values: bool = False,\n        reverse: bool = False\n):\n    if isinstance(group_matcher, dict):\n        # dictionary matcher contains a dict of raw-string regex expr that must be compiled\n        compiled = []\n        for group_ordinal, (group_name, mspec) in enumerate(group_matcher.items()):\n            if mspec is None:\n                continue\n            # map all matching specifications into 3-tuple (compiled re, prefix, suffix)\n            if isinstance(mspec, (tuple, list)):\n                # multi-entry match specifications require each sub-spec to be a 2-tuple (re, suffix)\n                for sspec in mspec:\n                    compiled += [(re.compile(sspec[0]), (group_ordinal,), sspec[1])]\n            else:\n                compiled += [(re.compile(mspec), (group_ordinal,), None)]\n        group_matcher = compiled\n\n    def _get_grouping(name):\n        if isinstance(group_matcher, (list, tuple)):\n            for match_fn, prefix, suffix in group_matcher:\n                r = match_fn.match(name)\n                if r:\n                    parts = (prefix, r.groups(), suffix)\n                    # map all tuple elem to int for numeric sort, filter out None entries\n                    return tuple(map(float, chain.from_iterable(filter(None, parts))))\n            return float('inf'),  # un-matched layers (neck, head) mapped to largest ordinal\n        else:\n            ord = group_matcher(name)\n            if not isinstance(ord, collections.abc.Iterable):\n                return ord,\n            return tuple(ord)\n\n    # map layers into groups via ordinals (ints or tuples of ints) from matcher\n    grouping = defaultdict(list)\n    for k, v in named_objects:\n        grouping[_get_grouping(k)].append(v if return_values else k)\n\n    # remap to integers\n    layer_id_to_param = defaultdict(list)\n    lid = -1\n    for k in sorted(filter(lambda x: x is not None, grouping.keys())):\n        if lid < 0 or k[-1] != MATCH_PREV_GROUP[0]:\n            lid += 1\n        layer_id_to_param[lid].extend(grouping[k])\n\n    if reverse:\n        assert not return_values, \"reverse mapping only sensible for name output\"\n        # output reverse mapping\n        param_to_layer_id = {}\n        for lid, lm in layer_id_to_param.items():\n            for n in lm:\n                param_to_layer_id[n] = lid\n        return param_to_layer_id\n\n    return layer_id_to_param\n\n\ndef group_parameters(\n        module: nn.Module,\n        group_matcher,\n        return_values: bool = False,\n        reverse: bool = False,\n):\n    return group_with_matcher(\n        module.named_parameters(), group_matcher, return_values=return_values, reverse=reverse)\n\n\ndef group_modules(\n        module: nn.Module,\n        group_matcher,\n        return_values: bool = False,\n        reverse: bool = False,\n):\n    return group_with_matcher(\n        named_modules_with_params(module), group_matcher, return_values=return_values, reverse=reverse)\n\n\ndef flatten_modules(\n        named_modules: Iterator[Tuple[str, nn.Module]],\n        depth: int = 1,\n        prefix: Union[str, Tuple[str, ...]] = '',\n        module_types: Union[str, Tuple[Type[nn.Module]]] = 'sequential',\n):\n    prefix_is_tuple = isinstance(prefix, tuple)\n    if isinstance(module_types, str):\n        if module_types == 'container':\n            module_types = (nn.Sequential, nn.ModuleList, nn.ModuleDict)\n        else:\n            module_types = (nn.Sequential,)\n    for name, module in named_modules:\n        if depth and isinstance(module, module_types):\n            yield from flatten_modules(\n                module.named_children(),\n                depth - 1,\n                prefix=(name,) if prefix_is_tuple else name,\n                module_types=module_types,\n            )\n        else:\n            if prefix_is_tuple:\n                name = prefix + (name,)\n                yield name, module\n            else:\n                if prefix:\n                    name = '.'.join([prefix, name])\n                yield name, module\n\n\ndef checkpoint_seq(\n        functions,\n        x,\n        every=1,\n        flatten=False,\n        skip_last=False,\n        preserve_rng_state=True\n):\n    r\"\"\"A helper function for checkpointing sequential models.\n\n    Sequential models execute a list of modules/functions in order\n    (sequentially). Therefore, we can divide such a sequence into segments\n    and checkpoint each segment. All segments except run in :func:`torch.no_grad`\n    manner, i.e., not storing the intermediate activations. The inputs of each\n    checkpointed segment will be saved for re-running the segment in the backward pass.\n\n    See :func:`~torch.utils.checkpoint.checkpoint` on how checkpointing works.\n\n    .. warning::\n        Checkpointing currently only supports :func:`torch.autograd.backward`\n        and only if its `inputs` argument is not passed. :func:`torch.autograd.grad`\n        is not supported.\n\n    .. warning:\n        At least one of the inputs needs to have :code:`requires_grad=True` if\n        grads are needed for model inputs, otherwise the checkpointed part of the\n        model won't have gradients.\n\n    Args:\n        functions: A :class:`torch.nn.Sequential` or the list of modules or functions to run sequentially.\n        x: A Tensor that is input to :attr:`functions`\n        every: checkpoint every-n functions (default: 1)\n        flatten (bool): flatten nn.Sequential of nn.Sequentials\n        skip_last (bool): skip checkpointing the last function in the sequence if True\n        preserve_rng_state (bool, optional, default=True):  Omit stashing and restoring\n            the RNG state during each checkpoint.\n\n    Returns:\n        Output of running :attr:`functions` sequentially on :attr:`*inputs`\n\n    Example:\n        >>> model = nn.Sequential(...)\n        >>> input_var = checkpoint_seq(model, input_var, every=2)\n    \"\"\"\n    def run_function(start, end, functions):\n        def forward(_x):\n            for j in range(start, end + 1):\n                _x = functions[j](_x)\n            return _x\n        return forward\n\n    if isinstance(functions, torch.nn.Sequential):\n        functions = functions.children()\n    if flatten:\n        functions = chain.from_iterable(functions)\n    if not isinstance(functions, (tuple, list)):\n        functions = tuple(functions)\n\n    num_checkpointed = len(functions)\n    if skip_last:\n        num_checkpointed -= 1\n    end = -1\n    for start in range(0, num_checkpointed, every):\n        end = min(start + every - 1, num_checkpointed - 1)\n        x = checkpoint(run_function(start, end, functions), x, preserve_rng_state=preserve_rng_state)\n    if skip_last:\n        return run_function(end + 1, len(functions) - 1, functions)(x)\n    return x\n\n\ndef adapt_input_conv(in_chans, conv_weight):\n    conv_type = conv_weight.dtype\n    conv_weight = conv_weight.float()  # Some weights are in torch.half, ensure it's float for sum on CPU\n    O, I, J, K = conv_weight.shape\n    if in_chans == 1:\n        if I > 3:\n            assert conv_weight.shape[1] % 3 == 0\n            # For models with space2depth stems\n            conv_weight = conv_weight.reshape(O, I // 3, 3, J, K)\n            conv_weight = conv_weight.sum(dim=2, keepdim=False)\n        else:\n            conv_weight = conv_weight.sum(dim=1, keepdim=True)\n    elif in_chans != 3:\n        if I != 3:\n            raise NotImplementedError('Weight format not supported by conversion.')\n        else:\n            # NOTE this strategy should be better than random init, but there could be other combinations of\n            # the original RGB input layer weights that'd work better for specific cases.\n            repeat = int(math.ceil(in_chans / 3))\n            conv_weight = conv_weight.repeat(1, repeat, 1, 1)[:, :in_chans, :, :]\n            conv_weight *= (3 / float(in_chans))\n    conv_weight = conv_weight.to(conv_type)\n    return conv_weight","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:27.899816Z","iopub.execute_input":"2024-04-18T12:45:27.900157Z","iopub.status.idle":"2024-04-18T12:45:27.944069Z","shell.execute_reply.started":"2024-04-18T12:45:27.900131Z","shell.execute_reply":"2024-04-18T12:45:27.943166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nfrom collections import deque, defaultdict\nfrom dataclasses import dataclass, field, replace, asdict\nfrom typing import Any, Deque, Dict, Tuple, Optional, Union\n\n\n__all__ = ['PretrainedCfg', 'filter_pretrained_cfg', 'DefaultCfg']\n\n\n@dataclass\nclass PretrainedCfg:\n    \"\"\"\n    \"\"\"\n    # weight source locations\n    url: Optional[Union[str, Tuple[str, str]]] = None  # remote URL\n    file: Optional[str] = None  # local / shared filesystem path\n    state_dict: Optional[Dict[str, Any]] = None  # in-memory state dict\n    hf_hub_id: Optional[str] = None  # Hugging Face Hub model id ('organization/model')\n    hf_hub_filename: Optional[str] = None  # Hugging Face Hub filename (overrides default)\n\n    source: Optional[str] = None  # source of cfg / weight location used (url, file, hf-hub)\n    architecture: Optional[str] = None  # architecture variant can be set when not implicit\n    tag: Optional[str] = None  # pretrained tag of source\n    custom_load: bool = False  # use custom model specific model.load_pretrained() (ie for npz files)\n\n    # input / data config\n    input_size: Tuple[int, int, int] = (3, 224, 224)\n    test_input_size: Optional[Tuple[int, int, int]] = None\n    min_input_size: Optional[Tuple[int, int, int]] = None\n    fixed_input_size: bool = False\n    interpolation: str = 'bicubic'\n    crop_pct: float = 0.875\n    test_crop_pct: Optional[float] = None\n    crop_mode: str = 'center'\n    mean: Tuple[float, ...] = (0.485, 0.456, 0.406)\n    std: Tuple[float, ...] = (0.229, 0.224, 0.225)\n\n    # head / classifier config and meta-data\n    num_classes: int = 1000\n    label_offset: Optional[int] = None\n    label_names: Optional[Tuple[str]] = None\n    label_descriptions: Optional[Dict[str, str]] = None\n\n    # model attributes that vary with above or required for pretrained adaptation\n    pool_size: Optional[Tuple[int, ...]] = None\n    test_pool_size: Optional[Tuple[int, ...]] = None\n    first_conv: Optional[str] = None\n    classifier: Optional[str] = None\n\n    license: Optional[str] = None\n    description: Optional[str] = None\n    origin_url: Optional[str] = None\n    paper_name: Optional[str] = None\n    paper_ids: Optional[Union[str, Tuple[str]]] = None\n    notes: Optional[Tuple[str]] = None\n\n    @property\n    def has_weights(self):\n        return self.url or self.file or self.hf_hub_id\n\n    def to_dict(self, remove_source=False, remove_null=True):\n        return filter_pretrained_cfg(\n            asdict(self),\n            remove_source=remove_source,\n            remove_null=remove_null\n        )\n\n\ndef filter_pretrained_cfg(cfg, remove_source=False, remove_null=True):\n    filtered_cfg = {}\n    keep_null = {'pool_size', 'first_conv', 'classifier'}  # always keep these keys, even if none\n    for k, v in cfg.items():\n        if remove_source and k in {'url', 'file', 'hf_hub_id', 'hf_hub_id', 'hf_hub_filename', 'source'}:\n            continue\n        if remove_null and v is None and k not in keep_null:\n            continue\n        filtered_cfg[k] = v\n    return filtered_cfg\n\n\n@dataclass\nclass DefaultCfg:\n    tags: Deque[str] = field(default_factory=deque)  # priority queue of tags (first is default)\n    cfgs: Dict[str, PretrainedCfg] = field(default_factory=dict)  # pretrained cfgs by tag\n    is_pretrained: bool = False  # at least one of the configs has a pretrained source set\n\n    @property\n    def default(self):\n        return self.cfgs[self.tags[0]]\n\n    @property\n    def default_with_tag(self):\n        tag = self.tags[0]\n        return tag, self.cfgs[tag]","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:28.757927Z","iopub.execute_input":"2024-04-18T12:45:28.758298Z","iopub.status.idle":"2024-04-18T12:45:28.779783Z","shell.execute_reply.started":"2024-04-18T12:45:28.758268Z","shell.execute_reply":"2024-04-18T12:45:28.778754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# _registry.py\n\n\n\"\"\" Model Registry\nHacked together by / Copyright 2020 Ross Wightman\n\"\"\"\n\nimport fnmatch\nimport re\nimport sys\nimport warnings\nfrom collections import defaultdict, deque\nfrom copy import deepcopy\nfrom dataclasses import replace\nfrom typing import Any, Callable, Dict, Iterable, List, Optional, Set, Sequence, Union, Tuple\n\n# from ._pretrained import PretrainedCfg, DefaultCfg\n\n__all__ = [\n    'split_model_name_tag', 'get_arch_name', 'register_model', 'generate_default_cfgs',\n    'list_models', 'list_pretrained', 'is_model', 'model_entrypoint', 'list_modules', 'is_model_in_modules',\n    'get_pretrained_cfg_value', 'is_model_pretrained'\n]\n\n_module_to_models: Dict[str, Set[str]] = defaultdict(set)  # dict of sets to check membership of model in module\n_model_to_module: Dict[str, str] = {}  # mapping of model names to module names\n_model_entrypoints: Dict[str, Callable[..., Any]] = {}  # mapping of model names to architecture entrypoint fns\n_model_has_pretrained: Set[str] = set()  # set of model names that have pretrained weight url present\n_model_default_cfgs: Dict[str, PretrainedCfg] = {}  # central repo for model arch -> default cfg objects\n_model_pretrained_cfgs: Dict[str, PretrainedCfg] = {}  # central repo for model arch.tag -> pretrained cfgs\n_model_with_tags: Dict[str, List[str]] = defaultdict(list)  # shortcut to map each model arch to all model + tag names\n_module_to_deprecated_models: Dict[str, Dict[str, Optional[str]]] = defaultdict(dict)\n_deprecated_models: Dict[str, Optional[str]] = {}\n\n\ndef split_model_name_tag(model_name: str, no_tag: str = '') -> Tuple[str, str]:\n    model_name, *tag_list = model_name.split('.', 1)\n    tag = tag_list[0] if tag_list else no_tag\n    return model_name, tag\n\n\ndef get_arch_name(model_name: str) -> str:\n    return split_model_name_tag(model_name)[0]\n\n\ndef generate_default_cfgs(cfgs: Dict[str, Union[Dict[str, Any], PretrainedCfg]]):\n    out = defaultdict(DefaultCfg)\n    default_set = set()  # no tag and tags ending with * are prioritized as default\n\n    for k, v in cfgs.items():\n        if isinstance(v, dict):\n            v = PretrainedCfg(**v)\n        has_weights = v.has_weights\n\n        model, tag = split_model_name_tag(k)\n        is_default_set = model in default_set\n        priority = (has_weights and not tag) or (tag.endswith('*') and not is_default_set)\n        tag = tag.strip('*')\n\n        default_cfg = out[model]\n\n        if priority:\n            default_cfg.tags.appendleft(tag)\n            default_set.add(model)\n        elif has_weights and not default_cfg.is_pretrained:\n            default_cfg.tags.appendleft(tag)\n        else:\n            default_cfg.tags.append(tag)\n\n        if has_weights:\n            default_cfg.is_pretrained = True\n\n        default_cfg.cfgs[tag] = v\n\n    return out\n\n\ndef register_model(fn: Callable[..., Any]) -> Callable[..., Any]:\n    # lookup containing module\n    mod = sys.modules[fn.__module__]\n    module_name_split = fn.__module__.split('.')\n    module_name = module_name_split[-1] if len(module_name_split) else ''\n\n    # add model to __all__ in module\n    model_name = fn.__name__\n    if hasattr(mod, '__all__'):\n        mod.__all__.append(model_name)\n    else:\n        mod.__all__ = [model_name]  # type: ignore\n\n    # add entries to registry dict/sets\n    if model_name in _model_entrypoints:\n        warnings.warn(\n            f'Overwriting {model_name} in registry with {fn.__module__}.{model_name}. This is because the name being '\n            'registered conflicts with an existing name. Please check if this is not expected.',\n            stacklevel=2,\n        )\n    _model_entrypoints[model_name] = fn\n    _model_to_module[model_name] = module_name\n    _module_to_models[module_name].add(model_name)\n    if hasattr(mod, 'default_cfgs') and model_name in mod.default_cfgs:\n        # this will catch all models that have entrypoint matching cfg key, but miss any aliasing\n        # entrypoints or non-matching combos\n        default_cfg = mod.default_cfgs[model_name]\n        if not isinstance(default_cfg, DefaultCfg):\n            # new style default cfg dataclass w/ multiple entries per model-arch\n            assert isinstance(default_cfg, dict)\n            # old style cfg dict per model-arch\n            pretrained_cfg = PretrainedCfg(**default_cfg)\n            default_cfg = DefaultCfg(tags=deque(['']), cfgs={'': pretrained_cfg})\n\n        for tag_idx, tag in enumerate(default_cfg.tags):\n            is_default = tag_idx == 0\n            pretrained_cfg = default_cfg.cfgs[tag]\n            model_name_tag = '.'.join([model_name, tag]) if tag else model_name\n            replace_items = dict(architecture=model_name, tag=tag if tag else None)\n            if pretrained_cfg.hf_hub_id and pretrained_cfg.hf_hub_id == 'timm/':\n                # auto-complete hub name w/ architecture.tag\n                replace_items['hf_hub_id'] = pretrained_cfg.hf_hub_id + model_name_tag\n            pretrained_cfg = replace(pretrained_cfg, **replace_items)\n\n            if is_default:\n                _model_pretrained_cfgs[model_name] = pretrained_cfg\n                if pretrained_cfg.has_weights:\n                    # add tagless entry if it's default and has weights\n                    _model_has_pretrained.add(model_name)\n\n            if tag:\n                _model_pretrained_cfgs[model_name_tag] = pretrained_cfg\n                if pretrained_cfg.has_weights:\n                    # add model w/ tag if tag is valid\n                    _model_has_pretrained.add(model_name_tag)\n                _model_with_tags[model_name].append(model_name_tag)\n            else:\n                _model_with_tags[model_name].append(model_name)  # has empty tag (to slowly remove these instances)\n\n        _model_default_cfgs[model_name] = default_cfg\n\n    return fn\n\n\ndef _deprecated_model_shim(deprecated_name: str, current_fn: Callable = None, current_tag: str = ''):\n    def _fn(pretrained=False, **kwargs):\n        assert current_fn is not None,  f'Model {deprecated_name} has been removed with no replacement.'\n        current_name = '.'.join([current_fn.__name__, current_tag]) if current_tag else current_fn.__name__\n        warnings.warn(f'Mapping deprecated model name {deprecated_name} to current {current_name}.', stacklevel=2)\n        pretrained_cfg = kwargs.pop('pretrained_cfg', None)\n        return current_fn(pretrained=pretrained, pretrained_cfg=pretrained_cfg or current_tag, **kwargs)\n    return _fn\n\n\ndef register_model_deprecations(module_name: str, deprecation_map: Dict[str, Optional[str]]):\n    mod = sys.modules[module_name]\n    module_name_split = module_name.split('.')\n    module_name = module_name_split[-1] if len(module_name_split) else ''\n\n    for deprecated, current in deprecation_map.items():\n        if hasattr(mod, '__all__'):\n            mod.__all__.append(deprecated)\n        current_fn = None\n        current_tag = ''\n        if current:\n            current_name, current_tag = split_model_name_tag(current)\n            current_fn = getattr(mod, current_name)\n        deprecated_entrypoint_fn = _deprecated_model_shim(deprecated, current_fn, current_tag)\n        setattr(mod, deprecated, deprecated_entrypoint_fn)\n        _model_entrypoints[deprecated] = deprecated_entrypoint_fn\n        _model_to_module[deprecated] = module_name\n        _module_to_models[module_name].add(deprecated)\n        _deprecated_models[deprecated] = current\n        _module_to_deprecated_models[module_name][deprecated] = current\n\n\ndef _natural_key(string_: str) -> List[Union[int, str]]:\n    \"\"\"See https://blog.codinghorror.com/sorting-for-humans-natural-sort-order/\"\"\"\n    return [int(s) if s.isdigit() else s for s in re.split(r'(\\d+)', string_.lower())]\n\n\ndef _expand_filter(filter: str):\n    \"\"\" expand a 'base_filter' to 'base_filter.*' if no tag portion\"\"\"\n    filter_base, filter_tag = split_model_name_tag(filter)\n    if not filter_tag:\n        return ['.'.join([filter_base, '*']), filter]\n    else:\n        return [filter]\n\n\ndef list_models(\n        filter: Union[str, List[str]] = '',\n        module: str = '',\n        pretrained: bool = False,\n        exclude_filters: Union[str, List[str]] = '',\n        name_matches_cfg: bool = False,\n        include_tags: Optional[bool] = None,\n) -> List[str]:\n    \"\"\" Return list of available model names, sorted alphabetically\n\n    Args:\n        filter - Wildcard filter string that works with fnmatch\n        module - Limit model selection to a specific submodule (ie 'vision_transformer')\n        pretrained - Include only models with valid pretrained weights if True\n        exclude_filters - Wildcard filters to exclude models after including them with filter\n        name_matches_cfg - Include only models w/ model_name matching default_cfg name (excludes some aliases)\n        include_tags - Include pretrained tags in model names (model.tag). If None, defaults\n            set to True when pretrained=True else False (default: None)\n\n    Returns:\n        models - The sorted list of models\n\n    Example:\n        model_list('gluon_resnet*') -- returns all models starting with 'gluon_resnet'\n        model_list('*resnext*, 'resnet') -- returns all models with 'resnext' in 'resnet' module\n    \"\"\"\n    if filter:\n        include_filters = filter if isinstance(filter, (tuple, list)) else [filter]\n    else:\n        include_filters = []\n\n    if include_tags is None:\n        # FIXME should this be default behaviour? or default to include_tags=True?\n        include_tags = pretrained\n\n    all_models: Set[str] = _module_to_models[module] if module else set(_model_entrypoints.keys())\n    all_models = all_models - _deprecated_models.keys()  # remove deprecated models from listings\n\n    if include_tags:\n        # expand model names to include names w/ pretrained tags\n        models_with_tags: Set[str] = set()\n        for m in all_models:\n            models_with_tags.update(_model_with_tags[m])\n        all_models = models_with_tags\n        # expand include and exclude filters to include a '.*' for proper match if no tags in filter\n        include_filters = [ef for f in include_filters for ef in _expand_filter(f)]\n        exclude_filters = [ef for f in exclude_filters for ef in _expand_filter(f)]\n\n    if include_filters:\n        models: Set[str] = set()\n        for f in include_filters:\n            include_models = fnmatch.filter(all_models, f)  # include these models\n            if len(include_models):\n                models = models.union(include_models)\n    else:\n        models = all_models\n\n    if exclude_filters:\n        if not isinstance(exclude_filters, (tuple, list)):\n            exclude_filters = [exclude_filters]\n        for xf in exclude_filters:\n            exclude_models = fnmatch.filter(models, xf)  # exclude these models\n            if len(exclude_models):\n                models = models.difference(exclude_models)\n\n    if pretrained:\n        models = _model_has_pretrained.intersection(models)\n\n    if name_matches_cfg:\n        models = set(_model_pretrained_cfgs).intersection(models)\n\n    return sorted(models, key=_natural_key)\n\n\ndef list_pretrained(\n        filter: Union[str, List[str]] = '',\n        exclude_filters: str = '',\n) -> List[str]:\n    return list_models(\n        filter=filter,\n        pretrained=True,\n        exclude_filters=exclude_filters,\n        include_tags=True,\n    )\n\n\ndef get_deprecated_models(module: str = '') -> Dict[str, str]:\n    all_deprecated = _module_to_deprecated_models[module] if module else _deprecated_models\n    return deepcopy(all_deprecated)\n\n\ndef is_model(model_name: str) -> bool:\n    \"\"\" Check if a model name exists\n    \"\"\"\n    arch_name = get_arch_name(model_name)\n    return arch_name in _model_entrypoints\n\n\ndef model_entrypoint(model_name: str, module_filter: Optional[str] = None) -> Callable[..., Any]:\n    \"\"\"Fetch a model entrypoint for specified model name\n    \"\"\"\n    arch_name = get_arch_name(model_name)\n    if module_filter and arch_name not in _module_to_models.get(module_filter, {}):\n        raise RuntimeError(f'Model ({model_name} not found in module {module_filter}.')\n    return _model_entrypoints[arch_name]\n\n\ndef list_modules() -> List[str]:\n    \"\"\" Return list of module names that contain models / model entrypoints\n    \"\"\"\n    modules = _module_to_models.keys()\n    return sorted(modules)\n\n\ndef is_model_in_modules(\n        model_name: str, module_names: Union[Tuple[str, ...], List[str], Set[str]]\n) -> bool:\n    \"\"\"Check if a model exists within a subset of modules\n\n    Args:\n        model_name - name of model to check\n        module_names - names of modules to search in\n    \"\"\"\n    arch_name = get_arch_name(model_name)\n    assert isinstance(module_names, (tuple, list, set))\n    return any(arch_name in _module_to_models[n] for n in module_names)\n\n\ndef is_model_pretrained(model_name: str) -> bool:\n    return model_name in _model_has_pretrained\n\n\ndef get_pretrained_cfg(model_name: str, allow_unregistered: bool = True) -> Optional[PretrainedCfg]:\n    if model_name in _model_pretrained_cfgs:\n        return deepcopy(_model_pretrained_cfgs[model_name])\n    arch_name, tag = split_model_name_tag(model_name)\n    if arch_name in _model_default_cfgs:\n        # if model arch exists, but the tag is wrong, error out\n        raise RuntimeError(f'Invalid pretrained tag ({tag}) for {arch_name}.')\n    if allow_unregistered:\n        # if model arch doesn't exist, it has no pretrained_cfg registered, allow a default to be created\n        return None\n    raise RuntimeError(f'Model architecture ({arch_name}) has no pretrained cfg registered.')\n\n\ndef get_pretrained_cfg_value(model_name: str, cfg_key: str) -> Optional[Any]:\n    \"\"\" Get a specific model default_cfg value by key. None if key doesn't exist.\n    \"\"\"\n    cfg = get_pretrained_cfg(model_name, allow_unregistered=False)\n    return getattr(cfg, cfg_key, None)","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:29.906682Z","iopub.execute_input":"2024-04-18T12:45:29.907018Z","iopub.status.idle":"2024-04-18T12:45:29.960523Z","shell.execute_reply.started":"2024-04-18T12:45:29.906993Z","shell.execute_reply":"2024-04-18T12:45:29.959474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\" MaxVit and CoAtNet Vision Transformer - CNN Hybrids in PyTorch\n\nThis is a from-scratch implementation of both CoAtNet and MaxVit in PyTorch.\n\n99% of the implementation was done from papers, however last minute some adjustments were made\nbased on the (as yet unfinished?) public code release https://github.com/google-research/maxvit\n\nThere are multiple sets of models defined for both architectures. Typically, names with a\n `_rw` suffix are my own original configs prior to referencing https://github.com/google-research/maxvit.\nThese configs work well and appear to be a bit faster / lower resource than the paper.\n\nThe models without extra prefix / suffix' (coatnet_0_224, maxvit_tiny_224, etc), are intended to\nmatch paper, BUT, without any official pretrained weights it's difficult to confirm a 100% match.\n\nPapers:\n\nMaxViT: Multi-Axis Vision Transformer - https://arxiv.org/abs/2204.01697\n@article{tu2022maxvit,\n  title={MaxViT: Multi-Axis Vision Transformer},\n  author={Tu, Zhengzhong and Talebi, Hossein and Zhang, Han and Yang, Feng and Milanfar, Peyman and Bovik, Alan and Li, Yinxiao},\n  journal={ECCV},\n  year={2022},\n}\n\nCoAtNet: Marrying Convolution and Attention for All Data Sizes - https://arxiv.org/abs/2106.04803\n@article{DBLP:journals/corr/abs-2106-04803,\n  author    = {Zihang Dai and Hanxiao Liu and Quoc V. Le and Mingxing Tan},\n  title     = {CoAtNet: Marrying Convolution and Attention for All Data Sizes},\n  journal   = {CoRR},\n  volume    = {abs/2106.04803},\n  year      = {2021}\n}\n\nHacked together by / Copyright 2022, Ross Wightman\n\"\"\"\n\nimport math\nfrom collections import OrderedDict\nfrom dataclasses import dataclass, replace, field\nfrom functools import partial\nfrom typing import Callable, Optional, Union, Tuple, List\n\nimport torch\nfrom torch import nn\nfrom torch.jit import Final\n\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nfrom timm.layers import Mlp, ConvMlp, DropPath, LayerNorm, ClassifierHead, NormMlpClassifierHead\nfrom timm.layers import create_attn, get_act_layer, get_norm_layer, get_norm_act_layer, create_conv2d, create_pool2d\nfrom timm.layers import trunc_normal_tf_, to_2tuple, extend_tuple, make_divisible, _assert\nfrom timm.layers import RelPosMlp, RelPosBias, RelPosBiasTf, use_fused_attn, resize_rel_pos_bias_table\n# from ._builder import build_model_with_cfg\n# from ._features_fx import register_notrace_function\n# from ._manipulate import named_apply, checkpoint_seq\n# from ._registry import generate_default_cfgs, register_model\n\n__all__ = ['MaxxVitCfg', 'MaxxVitConvCfg', 'MaxxVitTransformerCfg', 'MaxxVit']\n\n\n@dataclass\nclass MaxxVitTransformerCfg:\n    dim_head: int = 32\n    head_first: bool = True  # head ordering in qkv channel dim\n    expand_ratio: float = 4.0\n    expand_first: bool = True\n    shortcut_bias: bool = True\n    attn_bias: bool = True\n    attn_drop: float = 0.\n    proj_drop: float = 0.\n    pool_type: str = 'avg2'\n    rel_pos_type: str = 'bias'\n    rel_pos_dim: int = 512  # for relative position types w/ MLP\n    partition_ratio: int = 32\n    window_size: Optional[Tuple[int, int]] = None\n    grid_size: Optional[Tuple[int, int]] = None\n    no_block_attn: bool = False  # disable window block attention for maxvit (ie only grid)\n    use_nchw_attn: bool = False  # for MaxViT variants (not used for CoAt), keep tensors in NCHW order\n    init_values: Optional[float] = None\n    act_layer: str = 'gelu'\n    norm_layer: str = 'layernorm2d'\n    norm_layer_cl: str = 'layernorm'\n    norm_eps: float = 1e-6\n\n    def __post_init__(self):\n        if self.grid_size is not None:\n            self.grid_size = to_2tuple(self.grid_size)\n        if self.window_size is not None:\n            self.window_size = to_2tuple(self.window_size)\n            if self.grid_size is None:\n                self.grid_size = self.window_size\n\n\n@dataclass\nclass MaxxVitConvCfg:\n    block_type: str = 'mbconv'\n    expand_ratio: float = 4.0\n    expand_output: bool = True  # calculate expansion channels from output (vs input chs)\n    kernel_size: int = 3\n    group_size: int = 1  # 1 == depthwise\n    pre_norm_act: bool = False  # activation after pre-norm\n    output_bias: bool = True  # bias for shortcut + final 1x1 projection conv\n    stride_mode: str = 'dw'  # stride done via one of 'pool', '1x1', 'dw'\n    pool_type: str = 'avg2'\n    downsample_pool_type: str = 'avg2'\n    padding: str = ''\n    attn_early: bool = False  # apply attn between conv2 and norm2, instead of after norm2\n    attn_layer: str = 'se'\n    attn_act_layer: str = 'silu'\n    attn_ratio: float = 0.25\n    init_values: Optional[float] = 1e-6  # for ConvNeXt block, ignored by MBConv\n    act_layer: str = 'gelu'\n    norm_layer: str = ''\n    norm_layer_cl: str = ''\n    norm_eps: Optional[float] = None\n\n    def __post_init__(self):\n        # mbconv vs convnext blocks have different defaults, set in post_init to avoid explicit config args\n        assert self.block_type in ('mbconv', 'convnext')\n        use_mbconv = self.block_type == 'mbconv'\n        if not self.norm_layer:\n            self.norm_layer = 'batchnorm2d' if use_mbconv else 'layernorm2d'\n        if not self.norm_layer_cl and not use_mbconv:\n            self.norm_layer_cl = 'layernorm'\n        if self.norm_eps is None:\n            self.norm_eps = 1e-5 if use_mbconv else 1e-6\n        self.downsample_pool_type = self.downsample_pool_type or self.pool_type\n\n\n@dataclass\nclass MaxxVitCfg:\n    embed_dim: Tuple[int, ...] = (96, 192, 384, 768)\n    depths: Tuple[int, ...] = (2, 3, 5, 2)\n    block_type: Tuple[Union[str, Tuple[str, ...]], ...] = ('C', 'C', 'T', 'T')\n    stem_width: Union[int, Tuple[int, int]] = 64\n    stem_bias: bool = False\n    conv_cfg: MaxxVitConvCfg = field(default_factory=MaxxVitConvCfg)\n    transformer_cfg: MaxxVitTransformerCfg = field(default_factory=MaxxVitTransformerCfg)\n    head_hidden_size: int = None\n    weight_init: str = 'vit_eff'\n\n\nclass Attention2d(nn.Module):\n    fused_attn: Final[bool]\n\n    \"\"\" multi-head attention for 2D NCHW tensors\"\"\"\n    def __init__(\n            self,\n            dim: int,\n            dim_out: Optional[int] = None,\n            dim_head: int = 32,\n            bias: bool = True,\n            expand_first: bool = True,\n            head_first: bool = True,\n            rel_pos_cls: Callable = None,\n            attn_drop: float = 0.,\n            proj_drop: float = 0.\n    ):\n        super().__init__()\n        dim_out = dim_out or dim\n        dim_attn = dim_out if expand_first else dim\n        self.num_heads = dim_attn // dim_head\n        self.dim_head = dim_head\n        self.head_first = head_first\n        self.scale = dim_head ** -0.5\n        self.fused_attn = use_fused_attn()\n\n        self.qkv = nn.Conv2d(dim, dim_attn * 3, 1, bias=bias)\n        self.rel_pos = rel_pos_cls(num_heads=self.num_heads) if rel_pos_cls else None\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Conv2d(dim_attn, dim_out, 1, bias=bias)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x, shared_rel_pos: Optional[torch.Tensor] = None):\n#         print(f\"Inside Attention2d forward: {x.shape}\")\n        B, C, H, W = x.shape\n\n#         y = x.clone()\n#         qkv_temp = self.qkv(y)\n\n#         print(f\"After self.qkv {qkv_temp.shape}\")\n        \n        if self.head_first:\n            q, k, v = self.qkv(x).view(B, self.num_heads, self.dim_head * 3, -1).chunk(3, dim=2)\n#             print(\"self.head_first=True\")\n            \n        else:\n            q, k, v = self.qkv(x).reshape(B, 3, self.num_heads, self.dim_head, -1).unbind(1)\n#             print(\"self.head_first=False\")\n            \n        \n#         print(f\"q, k, v: {q.shape}, {k.shape}, {v.shape}\")\n\n\n        if self.fused_attn:\n            # here\n#             print(\"self.fused_attn = True\")\n            attn_bias = None\n            if self.rel_pos is not None:\n                #here\n                attn_bias = self.rel_pos.get_bias()\n#                 print(\"self.rel_pos is not None\")\n                \n            elif shared_rel_pos is not None:\n                attn_bias = shared_rel_pos\n#                 print(\"shared_rel_pos is not None\")\n            \n            \n#             if attn_bias is not None:    \n#                 print(f\"attn_bias: {attn_bias.shape}\")\n#             else:\n#                 print(\"attn_bias is none\")\n            \n#             print(f\"Now applying x = torch.nn.functional.scaled_dot_product_attention(\\\n#                 q.transpose(-1, -2).contiguous(),\\\n#                 k.transpose(-1, -2).contiguous(),\\\n#                 v.transpose(-1, -2).contiguous(),\\\n#                 attn_mask=attn_bias,\\\n#                 dropout_p=self.attn_drop.p if self.training else 0.,\\\n#             ).transpose(-1, -2).reshape(B, -1, H, W)\")\n\n            x = torch.nn.functional.scaled_dot_product_attention(\n                q.transpose(-1, -2).contiguous(),\n                k.transpose(-1, -2).contiguous(),\n                v.transpose(-1, -2).contiguous(),\n                attn_mask=attn_bias,\n                dropout_p=self.attn_drop.p if self.training else 0.,\n            ).transpose(-1, -2).reshape(B, -1, H, W)\n        else:\n#             print(\"self.fused_attn = False\")\n            q = q * self.scale\n            attn = q.transpose(-2, -1) @ k\n            if self.rel_pos is not None:\n                attn = self.rel_pos(attn)\n            elif shared_rel_pos is not None:\n                attn = attn + shared_rel_pos\n            attn = attn.softmax(dim=-1)\n            attn = self.attn_drop(attn)\n            x = (v @ attn.transpose(-2, -1)).view(B, -1, H, W)\n        \n        \n#         print(f\"After attention mechanism: {x.shape}\")\n        x = self.proj(x)\n#         print(f\"After self.proj: {x.shape}\")\n        x = self.proj_drop(x)\n#         print(f\"After self.proj_drop: {x.shape}\")\n#         print('*'*100)\n        return x\n\n\n### PROPOSED ATTENTION MECHANISM ###\nimport torch.nn.functional as F\n\nclass Attention2dPyramidal(nn.Module):\n    fused_attn: Final[bool]\n        \n    def __init__(self, num_levels: int, \n            dim: int,\n            dim_out: Optional[int] = None,\n            dim_head: int = 32,\n            bias: bool = True,\n            expand_first: bool = True,\n            head_first: bool = True,\n            rel_pos_cls: Callable = None,\n            attn_drop: float = 0.,\n            proj_drop: float = 0.):\n        super().__init__()\n        self.num_levels = num_levels\n        self.attention_modules = nn.ModuleList([\n            Attention2d(dim, dim_out, dim_head, expand_first, bias, rel_pos_cls, attn_drop, proj_drop) \n            for _ in range(num_levels)\n        ])\n        self.attn_bias = bias\n        self.rel_pos_cls = rel_pos_cls\n        \n\n    def forward(self, x, shared_rel_pos: Optional[torch.Tensor] = None):\n#         print(f\"In Attention2dPyramidal forward num_levels={self.num_levels} and x={x.shape}:-\")\n    \n        # Divide input into pyramidal levels\n        pyramidal_levels = [x]\n        for i in range(1, self.num_levels):\n            downsample_factor = 2 * i\n            level_input = F.avg_pool2d(x, kernel_size=downsample_factor, stride=downsample_factor)\n#             print(f\"Shape after downsampling(avgpool2D) {i}th level={level_input.shape}\")\n            pyramidal_levels.append(level_input)\n\n        # Compute attention at each level\n        attention_outputs = []\n        for level, level_input in enumerate(pyramidal_levels):\n            attention_output = self.attention_modules[level](level_input)\n#             print(f\"Shape after applying attention at {level}th level={attention_output.shape}\")\n            attention_outputs.append(attention_output)\n\n        \n        # Upsample and combine attention outputs starting from the last level\n        combined_attention_output = attention_outputs[-1]\n        for i in range(self.num_levels - 2, -1, -1): \n            # Interpolate the attention output of the current level to match the spatial dimensions of the level above it\n            combined_attention_output = F.interpolate(combined_attention_output, size=attention_outputs[i].shape[2:], mode='bilinear', align_corners=False)\n#             print(f\"Shape after applying interpolation at {i+1}th level={combined_attention_output.shape}\")\n\n            # Combine the attention output of the current level with the level above it\n            combined_attention_output += attention_outputs[i]\n#             print(f\"Shape after combining {i+1} and {i}th level={combined_attention_output.shape}\")\n\n\n        # Resize the final combined attention output to match the spatial dimensions of the input at the first level\n        combined_attention_output = F.interpolate(combined_attention_output, size=pyramidal_levels[0].shape[2:], mode='bilinear', align_corners=False)\n#         print(f\"Final shape after (interpolation_last) Attention2dPyramidal : {combined_attention_output.shape}\")\n#         print(\"$\"*100)\n        return combined_attention_output\n\n    def match_spatial_dimensions(self, tensor1, tensor2):\n        \"\"\"\n        Pad or crop tensor1 to match the spatial dimensions of tensor2.\n        \"\"\"\n        if tensor1.shape[2:] != tensor2.shape[2:]:\n            diff_h = tensor2.shape[2] - tensor1.shape[2]\n            diff_w = tensor2.shape[3] - tensor1.shape[3]\n            pad_left = diff_w // 2\n            pad_right = diff_w - pad_left\n            pad_top = diff_h // 2\n            pad_bottom = diff_h - pad_top\n            tensor1 = F.pad(tensor1, (pad_left, pad_right, pad_top, pad_bottom))\n        return tensor1\n    \n\nclass AttentionCl(nn.Module):\n    \"\"\" Channels-last multi-head attention (B, ..., C) \"\"\"\n    fused_attn: Final[bool]\n\n    def __init__(\n            self,\n            dim: int,\n            dim_out: Optional[int] = None,\n            dim_head: int = 32,\n            bias: bool = True,\n            expand_first: bool = True,\n            head_first: bool = True,\n            rel_pos_cls: Callable = None,\n            attn_drop: float = 0.,\n            proj_drop: float = 0.\n    ):\n        super().__init__()\n        dim_out = dim_out or dim\n        dim_attn = dim_out if expand_first and dim_out > dim else dim\n        assert dim_attn % dim_head == 0, 'attn dim should be divisible by head_dim'\n        self.num_heads = dim_attn // dim_head\n        self.dim_head = dim_head\n        self.head_first = head_first\n        self.scale = dim_head ** -0.5\n        self.fused_attn = use_fused_attn()\n\n        self.qkv = nn.Linear(dim, dim_attn * 3, bias=bias)\n        self.rel_pos = rel_pos_cls(num_heads=self.num_heads) if rel_pos_cls else None\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim_attn, dim_out, bias=bias)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x, shared_rel_pos: Optional[torch.Tensor] = None):\n        B = x.shape[0]\n        restore_shape = x.shape[:-1]\n\n        if self.head_first:\n            q, k, v = self.qkv(x).view(B, -1, self.num_heads, self.dim_head * 3).transpose(1, 2).chunk(3, dim=3)\n        else:\n            q, k, v = self.qkv(x).reshape(B, -1, 3, self.num_heads, self.dim_head).transpose(1, 3).unbind(2)\n\n        if self.fused_attn:\n            attn_bias = None\n            if self.rel_pos is not None:\n                attn_bias = self.rel_pos.get_bias()\n            elif shared_rel_pos is not None:\n                attn_bias = shared_rel_pos\n\n            x = torch.nn.functional.scaled_dot_product_attention(\n                q, k, v,\n                attn_mask=attn_bias,\n                dropout_p=self.attn_drop.p if self.training else 0.,\n            )\n        else:\n            q = q * self.scale\n            attn = q @ k.transpose(-2, -1)\n            if self.rel_pos is not None:\n                attn = self.rel_pos(attn, shared_rel_pos=shared_rel_pos)\n            elif shared_rel_pos is not None:\n                attn = attn + shared_rel_pos\n            attn = attn.softmax(dim=-1)\n            attn = self.attn_drop(attn)\n            x = attn @ v\n\n        x = x.transpose(1, 2).reshape(restore_shape + (-1,))\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n\nclass LayerScale(nn.Module):\n    def __init__(self, dim, init_values=1e-5, inplace=False):\n        super().__init__()\n        self.inplace = inplace\n        self.gamma = nn.Parameter(init_values * torch.ones(dim))\n\n    def forward(self, x):\n        gamma = self.gamma\n        return x.mul_(gamma) if self.inplace else x * gamma\n\n\nclass LayerScale2d(nn.Module):\n    def __init__(self, dim, init_values=1e-5, inplace=False):\n        super().__init__()\n        self.inplace = inplace\n        self.gamma = nn.Parameter(init_values * torch.ones(dim))\n\n    def forward(self, x):\n        gamma = self.gamma.view(1, -1, 1, 1)\n        return x.mul_(gamma) if self.inplace else x * gamma\n\n\nclass Downsample2d(nn.Module):\n    \"\"\" A downsample pooling module supporting several maxpool and avgpool modes\n    * 'max' - MaxPool2d w/ kernel_size 3, stride 2, padding 1\n    * 'max2' - MaxPool2d w/ kernel_size = stride = 2\n    * 'avg' - AvgPool2d w/ kernel_size 3, stride 2, padding 1\n    * 'avg2' - AvgPool2d w/ kernel_size = stride = 2\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            dim_out: int,\n            pool_type: str = 'avg2',\n            padding: str = '',\n            bias: bool = True,\n    ):\n        super().__init__()\n        assert pool_type in ('max', 'max2', 'avg', 'avg2')\n        if pool_type == 'max':\n            self.pool = create_pool2d('max', kernel_size=3, stride=2, padding=padding or 1)\n        elif pool_type == 'max2':\n            self.pool = create_pool2d('max', 2, padding=padding or 0)  # kernel_size == stride == 2\n        elif pool_type == 'avg':\n            self.pool = create_pool2d(\n                'avg', kernel_size=3, stride=2, count_include_pad=False, padding=padding or 1)\n        else:\n            self.pool = create_pool2d('avg', 2, padding=padding or 0)\n\n        if dim != dim_out:\n            self.expand = nn.Conv2d(dim, dim_out, 1, bias=bias)\n        else:\n            self.expand = nn.Identity()\n\n    def forward(self, x):\n#         print(f\"In Downsample2D forward:- {x.shape}\")\n        x = self.pool(x)  # spatial downsample\n#         print(f\"After self.pool: {x.shape}\")\n        x = self.expand(x)  # expand chs\n#         print(f\"After self.expand: {x.shape}\")\n        return x\n\n\ndef _init_transformer(module, name, scheme=''):\n    if isinstance(module, (nn.Conv2d, nn.Linear)):\n        if scheme == 'normal':\n            nn.init.normal_(module.weight, std=.02)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        elif scheme == 'trunc_normal':\n            trunc_normal_tf_(module.weight, std=.02)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        elif scheme == 'xavier_normal':\n            nn.init.xavier_normal_(module.weight)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        else:\n            # vit like\n            nn.init.xavier_uniform_(module.weight)\n            if module.bias is not None:\n                if 'mlp' in name:\n                    nn.init.normal_(module.bias, std=1e-6)\n                else:\n                    nn.init.zeros_(module.bias)\n\n\nclass TransformerBlock2d(nn.Module):\n    \"\"\" Transformer block with 2D downsampling\n    '2D' NCHW tensor layout\n\n    Some gains can be seen on GPU using a 1D / CL block, BUT w/ the need to switch back/forth to NCHW\n    for spatial pooling, the benefit is minimal so ended up using just this variant for CoAt configs.\n\n    This impl was faster on TPU w/ PT XLA than the 1D experiment.\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            dim_out: int,\n            stride: int = 1,\n            rel_pos_cls: Callable = None,\n            cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path: float = 0.,\n    ):\n        super().__init__()\n        norm_layer = partial(get_norm_layer(cfg.norm_layer), eps=cfg.norm_eps)\n        act_layer = get_act_layer(cfg.act_layer)\n\n        if stride == 2:\n            self.shortcut = Downsample2d(dim, dim_out, pool_type=cfg.pool_type, bias=cfg.shortcut_bias)\n            self.norm1 = nn.Sequential(OrderedDict([\n                ('norm', norm_layer(dim)),\n                ('down', Downsample2d(dim, dim, pool_type=cfg.pool_type)),\n            ]))\n        else:\n            assert dim == dim_out\n            self.shortcut = nn.Identity()\n            self.norm1 = norm_layer(dim)\n\n        ### MODIFIED HERE    \n        self.attn = Attention2dPyramidal(\n            4, # num_levels\n            dim,\n            dim_out,\n            dim_head=cfg.dim_head,\n            expand_first=cfg.expand_first,\n            bias=cfg.attn_bias,\n            rel_pos_cls=rel_pos_cls,\n            attn_drop=cfg.attn_drop,\n            proj_drop=cfg.proj_drop\n        )\n        self.ls1 = LayerScale2d(dim_out, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        self.norm2 = norm_layer(dim_out)\n        self.mlp = ConvMlp(\n            in_features=dim_out,\n            hidden_features=int(dim_out * cfg.expand_ratio),\n            act_layer=act_layer,\n            drop=cfg.proj_drop)\n        self.ls2 = LayerScale2d(dim_out, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def init_weights(self, scheme=''):\n        named_apply(partial(_init_transformer, scheme=scheme), self)\n\n    def forward(self, x, shared_rel_pos: Optional[torch.Tensor] = None):\n#         print(f\"Inside TransformerBlock2d forward: {x.shape}\")\n        \n#         y = x.clone()\n        \n#         y_shortcut = y.clone()\n        \n        \n#         y = self.norm1(y)\n#         print(f\"After self.norm1 {y.shape}\")\n        \n#         y = self.attn(y, shared_rel_pos=shared_rel_pos)\n#         print(f\"After self.attn(y, shared_rel_pos=shared_rel_pos): {y.shape} \")\n        \n#         y = self.ls1(y)\n#         print(f\"After self.ls1(y): {y.shape}\")\n        \n#         y = self.drop_path1(y)\n#         print(f\"After self.drop_path1(y): {y.shape}\")\n        \n#         y_shortcut = self.shortcut(y_shortcut)\n#         print(f\"After self.shortcut(y_shortcut): {y_shortcut.shape}\")\n        \n#         y = y_shortcut + y\n#         print(f\"After y_shortcut + y : {y.shape}\")\n        \n        \n        x = self.shortcut(x) + self.drop_path1(self.ls1(self.attn(self.norm1(x), shared_rel_pos=shared_rel_pos)))\n#         print(f\"after self.shortcut(x) + self.drop_path1(self.ls1(self.attn(self.norm1(x), shared_rel_pos=shared_rel_pos))) : {x.shape}\")\n        \n        \n        \n#         y = x.clone()\n        \n#         y_shortcut = x.clone()\n        \n#         y = self.norm2(y)\n#         print(f\"After self.norm2(y): {y.shape}\")\n        \n#         y = self.mlp(y)\n#         print(f\"After self.mlp: {y.shape}\")\n        \n#         y = self.ls2(y)\n#         print(f\"After self.ls2: {y.shape}\")\n\n#         y = self.drop_path2(y)\n#         print(f\"After self.drop_path2: {y.shape}\")\n\n        \n#         y = y_shortcut + y\n#         print(f\"After y_shortcut + y: {y.shape}\")\n        \n        x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n#         print(f\"After x + self.drop_path2(self.ls2(self.mlp(self.norm2(x)))): {x.shape}\")\n        \n        return x\n\n\ndef _init_conv(module, name, scheme=''):\n    if isinstance(module, nn.Conv2d):\n        if scheme == 'normal':\n            nn.init.normal_(module.weight, std=.02)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        elif scheme == 'trunc_normal':\n            trunc_normal_tf_(module.weight, std=.02)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        elif scheme == 'xavier_normal':\n            nn.init.xavier_normal_(module.weight)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        else:\n            # efficientnet like\n            fan_out = module.kernel_size[0] * module.kernel_size[1] * module.out_channels\n            fan_out //= module.groups\n            nn.init.normal_(module.weight, 0, math.sqrt(2.0 / fan_out))\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n\n\ndef num_groups(group_size, channels):\n    if not group_size:  # 0 or None\n        return 1  # normal conv with 1 group\n    else:\n        # NOTE group_size == 1 -> depthwise conv\n        assert channels % group_size == 0\n        return channels // group_size\n\n\nclass MbConvBlock(nn.Module):\n    \"\"\" Pre-Norm Conv Block - 1x1 - kxk - 1x1, w/ inverted bottleneck (expand)\n    \"\"\"\n    def __init__(\n            self,\n            in_chs: int,\n            out_chs: int,\n            stride: int = 1,\n            dilation: Tuple[int, int] = (1, 1),\n            cfg: MaxxVitConvCfg = MaxxVitConvCfg(),\n            drop_path: float = 0.\n    ):\n        super(MbConvBlock, self).__init__()\n        norm_act_layer = partial(get_norm_act_layer(cfg.norm_layer, cfg.act_layer), eps=cfg.norm_eps)\n        mid_chs = make_divisible((out_chs if cfg.expand_output else in_chs) * cfg.expand_ratio)\n        groups = num_groups(cfg.group_size, mid_chs)\n\n        if stride == 2:\n            self.shortcut = Downsample2d(\n                in_chs, out_chs, pool_type=cfg.pool_type, bias=cfg.output_bias, padding=cfg.padding)\n        else:\n            self.shortcut = nn.Identity()\n\n        assert cfg.stride_mode in ('pool', '1x1', 'dw')\n        stride_pool, stride_1, stride_2 = 1, 1, 1\n        if cfg.stride_mode == 'pool':\n            # NOTE this is not described in paper, experiment to find faster option that doesn't stride in 1x1\n            stride_pool, dilation_2 = stride, dilation[1]\n            # FIXME handle dilation of avg pool\n        elif cfg.stride_mode == '1x1':\n            # NOTE I don't like this option described in paper, 1x1 w/ stride throws info away\n            stride_1, dilation_2 = stride, dilation[1]\n        else:\n            stride_2, dilation_2 = stride, dilation[0]\n\n        self.pre_norm = norm_act_layer(in_chs, apply_act=cfg.pre_norm_act)\n        if stride_pool > 1:\n            self.down = Downsample2d(in_chs, in_chs, pool_type=cfg.downsample_pool_type, padding=cfg.padding)\n        else:\n            self.down = nn.Identity()\n        self.conv1_1x1 = create_conv2d(in_chs, mid_chs, 1, stride=stride_1)\n        self.norm1 = norm_act_layer(mid_chs)\n\n        self.conv2_kxk = create_conv2d(\n            mid_chs, mid_chs, cfg.kernel_size,\n            stride=stride_2, dilation=dilation_2, groups=groups, padding=cfg.padding)\n\n        attn_kwargs = {}\n        if isinstance(cfg.attn_layer, str):\n            if cfg.attn_layer == 'se' or cfg.attn_layer == 'eca':\n                attn_kwargs['act_layer'] = cfg.attn_act_layer\n                attn_kwargs['rd_channels'] = int(cfg.attn_ratio * (out_chs if cfg.expand_output else mid_chs))\n\n        # two different orderings for SE and norm2 (due to some weights and trials using SE before norm2)\n        if cfg.attn_early:\n            self.se_early = create_attn(cfg.attn_layer, mid_chs, **attn_kwargs)\n            self.norm2 = norm_act_layer(mid_chs)\n            self.se = None\n        else:\n            self.se_early = None\n            self.norm2 = norm_act_layer(mid_chs)\n            self.se = create_attn(cfg.attn_layer, mid_chs, **attn_kwargs)\n\n        self.conv3_1x1 = create_conv2d(mid_chs, out_chs, 1, bias=cfg.output_bias)\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def init_weights(self, scheme=''):\n        named_apply(partial(_init_conv, scheme=scheme), self)\n\n    def forward(self, x):\n#         print('-'*100)\n#         print(f\"Inside MbConvBlock forward:{x.shape}\")\n        shortcut = self.shortcut(x)\n#         print(f\"After self.shortcut:{x.shape}\")\n\n        x = self.pre_norm(x)\n#         print(f\"After self.pre_norm:{x.shape}\")\n        x = self.down(x)\n#         print(f\"After self.down:{x.shape}\")\n\n\n        # 1x1 expansion conv & norm-act\n        x = self.conv1_1x1(x)\n#         print(f\"After self.conv1_1x1 {x.shape}\")\n        x = self.norm1(x)\n#         print(f\"After self.norm1 {x.shape}\")\n\n\n        # depthwise / grouped 3x3 conv w/ SE (or other) channel attention & norm-act\n        x = self.conv2_kxk(x)\n#         print(f\"After self.conv2_kxk {x.shape}\")\n        if self.se_early is not None:\n            x = self.se_early(x)\n#             print(f\"self.se_early is not NONE. After self.se_early {x.shape}\")\n\n        x = self.norm2(x)\n#         print(f\"After self.norm2 {x.shape}\")\n\n        if self.se is not None:\n            # here\n            x = self.se(x)\n#             print(f\"self.se is not NONE. After self.se: {x.shape}\")\n\n        # 1x1 linear projection to output width\n        x = self.conv3_1x1(x)\n#         print(f\"After self.conv3_1x1: {x.shape}\")\n\n        x = self.drop_path(x) + shortcut\n#         print(f\"After self.drop_path(x) + shortcut: {x.shape}\")\n#         print('-'*100)\n        return x\n\n\nclass ConvNeXtBlock(nn.Module):\n    \"\"\" ConvNeXt Block\n    \"\"\"\n\n    def __init__(\n            self,\n            in_chs: int,\n            out_chs: Optional[int] = None,\n            kernel_size: int = 7,\n            stride: int = 1,\n            dilation: Tuple[int, int] = (1, 1),\n            cfg: MaxxVitConvCfg = MaxxVitConvCfg(),\n            conv_mlp: bool = True,\n            drop_path: float = 0.\n    ):\n        super().__init__()\n        out_chs = out_chs or in_chs\n        act_layer = get_act_layer(cfg.act_layer)\n        if conv_mlp:\n            norm_layer = partial(get_norm_layer(cfg.norm_layer), eps=cfg.norm_eps)\n            mlp_layer = ConvMlp\n        else:\n            assert 'layernorm' in cfg.norm_layer\n            norm_layer = LayerNorm\n            mlp_layer = Mlp\n        self.use_conv_mlp = conv_mlp\n\n        if stride == 2:\n            self.shortcut = Downsample2d(in_chs, out_chs)\n        elif in_chs != out_chs:\n            self.shortcut = nn.Conv2d(in_chs, out_chs, kernel_size=1, bias=cfg.output_bias)\n        else:\n            self.shortcut = nn.Identity()\n\n        assert cfg.stride_mode in ('pool', 'dw')\n        stride_pool, stride_dw = 1, 1\n        # FIXME handle dilation?\n        if cfg.stride_mode == 'pool':\n            stride_pool = stride\n        else:\n            stride_dw = stride\n\n        if stride_pool == 2:\n            self.down = Downsample2d(in_chs, in_chs, pool_type=cfg.downsample_pool_type)\n        else:\n            self.down = nn.Identity()\n\n        self.conv_dw = create_conv2d(\n            in_chs, out_chs, kernel_size=kernel_size, stride=stride_dw, dilation=dilation[1],\n            depthwise=True, bias=cfg.output_bias)\n        self.norm = norm_layer(out_chs)\n        self.mlp = mlp_layer(out_chs, int(cfg.expand_ratio * out_chs), bias=cfg.output_bias, act_layer=act_layer)\n        if conv_mlp:\n            self.ls = LayerScale2d(out_chs, cfg.init_values) if cfg.init_values else nn.Identity()\n        else:\n            self.ls = LayerScale(out_chs, cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def forward(self, x):\n        shortcut = self.shortcut(x)\n        x = self.down(x)\n        x = self.conv_dw(x)\n        if self.use_conv_mlp:\n            x = self.norm(x)\n            x = self.mlp(x)\n            x = self.ls(x)\n        else:\n            x = x.permute(0, 2, 3, 1)\n            x = self.norm(x)\n            x = self.mlp(x)\n            x = self.ls(x)\n            x = x.permute(0, 3, 1, 2)\n\n        x = self.drop_path(x) + shortcut\n        return x\n\n\ndef window_partition(x, window_size: List[int]):\n    B, H, W, C = x.shape\n    _assert(H % window_size[0] == 0, f'height ({H}) must be divisible by window ({window_size[0]})')\n    _assert(W % window_size[1] == 0, '')\n    x = x.view(B, H // window_size[0], window_size[0], W // window_size[1], window_size[1], C)\n    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size[0], window_size[1], C)\n    return windows\n\n\n@register_notrace_function  # reason: int argument is a Proxy\ndef window_reverse(windows, window_size: List[int], img_size: List[int]):\n    H, W = img_size\n    C = windows.shape[-1]\n    x = windows.view(-1, H // window_size[0], W // window_size[1], window_size[0], window_size[1], C)\n    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, H, W, C)\n    return x\n\n\ndef grid_partition(x, grid_size: List[int]):\n    B, H, W, C = x.shape\n    _assert(H % grid_size[0] == 0, f'height {H} must be divisible by grid {grid_size[0]}')\n    _assert(W % grid_size[1] == 0, '')\n    x = x.view(B, grid_size[0], H // grid_size[0], grid_size[1], W // grid_size[1], C)\n    windows = x.permute(0, 2, 4, 1, 3, 5).contiguous().view(-1, grid_size[0], grid_size[1], C)\n    return windows\n\n\n@register_notrace_function  # reason: int argument is a Proxy\ndef grid_reverse(windows, grid_size: List[int], img_size: List[int]):\n    H, W = img_size\n    C = windows.shape[-1]\n    x = windows.view(-1, H // grid_size[0], W // grid_size[1], grid_size[0], grid_size[1], C)\n    x = x.permute(0, 3, 1, 4, 2, 5).contiguous().view(-1, H, W, C)\n    return x\n\n\ndef get_rel_pos_cls(cfg: MaxxVitTransformerCfg, window_size):\n    rel_pos_cls = None\n    if cfg.rel_pos_type == 'mlp':\n        rel_pos_cls = partial(RelPosMlp, window_size=window_size, hidden_dim=cfg.rel_pos_dim)\n    elif cfg.rel_pos_type == 'bias':\n        rel_pos_cls = partial(RelPosBias, window_size=window_size)\n    elif cfg.rel_pos_type == 'bias_tf':\n        rel_pos_cls = partial(RelPosBiasTf, window_size=window_size)\n    return rel_pos_cls\n\n\nclass PartitionAttentionCl(nn.Module):\n    \"\"\" Grid or Block partition + Attn + FFN.\n    NxC 'channels last' tensor layout.\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            partition_type: str = 'block',\n            cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path: float = 0.,\n    ):\n        super().__init__()\n        norm_layer = partial(get_norm_layer(cfg.norm_layer_cl), eps=cfg.norm_eps)  # NOTE this block is channels-last\n        act_layer = get_act_layer(cfg.act_layer)\n\n        self.partition_block = partition_type == 'block'\n        self.partition_size = to_2tuple(cfg.window_size if self.partition_block else cfg.grid_size)\n        rel_pos_cls = get_rel_pos_cls(cfg, self.partition_size)\n\n        self.norm1 = norm_layer(dim)\n        self.attn = AttentionCl(\n            dim,\n            dim,\n            dim_head=cfg.dim_head,\n            bias=cfg.attn_bias,\n            head_first=cfg.head_first,\n            rel_pos_cls=rel_pos_cls,\n            attn_drop=cfg.attn_drop,\n            proj_drop=cfg.proj_drop,\n        )\n        self.ls1 = LayerScale(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        self.norm2 = norm_layer(dim)\n        self.mlp = Mlp(\n            in_features=dim,\n            hidden_features=int(dim * cfg.expand_ratio),\n            act_layer=act_layer,\n            drop=cfg.proj_drop)\n        self.ls2 = LayerScale(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def _partition_attn(self, x):\n        img_size = x.shape[1:3]\n        if self.partition_block:\n            partitioned = window_partition(x, self.partition_size)\n        else:\n            partitioned = grid_partition(x, self.partition_size)\n\n        partitioned = self.attn(partitioned)\n\n        if self.partition_block:\n            x = window_reverse(partitioned, self.partition_size, img_size)\n        else:\n            x = grid_reverse(partitioned, self.partition_size, img_size)\n        return x\n\n    def forward(self, x):\n        x = x + self.drop_path1(self.ls1(self._partition_attn(self.norm1(x))))\n        x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n        return x\n\n\nclass ParallelPartitionAttention(nn.Module):\n    \"\"\" Experimental. Grid and Block partition + single FFN\n    NxC tensor layout.\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path: float = 0.,\n    ):\n        super().__init__()\n        assert dim % 2 == 0\n        norm_layer = partial(get_norm_layer(cfg.norm_layer_cl), eps=cfg.norm_eps)  # NOTE this block is channels-last\n        act_layer = get_act_layer(cfg.act_layer)\n\n        assert cfg.window_size == cfg.grid_size\n        self.partition_size = to_2tuple(cfg.window_size)\n        rel_pos_cls = get_rel_pos_cls(cfg, self.partition_size)\n\n        self.norm1 = norm_layer(dim)\n        self.attn_block = AttentionCl(\n            dim,\n            dim // 2,\n            dim_head=cfg.dim_head,\n            bias=cfg.attn_bias,\n            head_first=cfg.head_first,\n            rel_pos_cls=rel_pos_cls,\n            attn_drop=cfg.attn_drop,\n            proj_drop=cfg.proj_drop,\n        )\n        self.attn_grid = AttentionCl(\n            dim,\n            dim // 2,\n            dim_head=cfg.dim_head,\n            bias=cfg.attn_bias,\n            head_first=cfg.head_first,\n            rel_pos_cls=rel_pos_cls,\n            attn_drop=cfg.attn_drop,\n            proj_drop=cfg.proj_drop,\n        )\n        self.ls1 = LayerScale(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        self.norm2 = norm_layer(dim)\n        self.mlp = Mlp(\n            in_features=dim,\n            hidden_features=int(dim * cfg.expand_ratio),\n            out_features=dim,\n            act_layer=act_layer,\n            drop=cfg.proj_drop)\n        self.ls2 = LayerScale(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def _partition_attn(self, x):\n        img_size = x.shape[1:3]\n\n        partitioned_block = window_partition(x, self.partition_size)\n        partitioned_block = self.attn_block(partitioned_block)\n        x_window = window_reverse(partitioned_block, self.partition_size, img_size)\n\n        partitioned_grid = grid_partition(x, self.partition_size)\n        partitioned_grid = self.attn_grid(partitioned_grid)\n        x_grid = grid_reverse(partitioned_grid, self.partition_size, img_size)\n\n        return torch.cat([x_window, x_grid], dim=-1)\n\n    def forward(self, x):\n        x = x + self.drop_path1(self.ls1(self._partition_attn(self.norm1(x))))\n        x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n        return x\n\n\ndef window_partition_nchw(x, window_size: List[int]):\n    B, C, H, W = x.shape\n    _assert(H % window_size[0] == 0, f'height ({H}) must be divisible by window ({window_size[0]})')\n    _assert(W % window_size[1] == 0, '')\n    x = x.view(B, C, H // window_size[0], window_size[0], W // window_size[1], window_size[1])\n    windows = x.permute(0, 2, 4, 1, 3, 5).contiguous().view(-1, C, window_size[0], window_size[1])\n    return windows\n\n\n@register_notrace_function  # reason: int argument is a Proxy\ndef window_reverse_nchw(windows, window_size: List[int], img_size: List[int]):\n    H, W = img_size\n    C = windows.shape[1]\n    x = windows.view(-1, H // window_size[0], W // window_size[1], C, window_size[0], window_size[1])\n    x = x.permute(0, 3, 1, 4, 2, 5).contiguous().view(-1, C, H, W)\n    return x\n\n\ndef grid_partition_nchw(x, grid_size: List[int]):\n    B, C, H, W = x.shape\n    _assert(H % grid_size[0] == 0, f'height {H} must be divisible by grid {grid_size[0]}')\n    _assert(W % grid_size[1] == 0, '')\n    x = x.view(B, C, grid_size[0], H // grid_size[0], grid_size[1], W // grid_size[1])\n    windows = x.permute(0, 3, 5, 1, 2, 4).contiguous().view(-1, C, grid_size[0], grid_size[1])\n    return windows\n\n\n@register_notrace_function  # reason: int argument is a Proxy\ndef grid_reverse_nchw(windows, grid_size: List[int], img_size: List[int]):\n    H, W = img_size\n    C = windows.shape[1]\n    x = windows.view(-1, H // grid_size[0], W // grid_size[1], C, grid_size[0], grid_size[1])\n    x = x.permute(0, 3, 4, 1, 5, 2).contiguous().view(-1, C, H, W)\n    return x\n\n\nclass PartitionAttention2d(nn.Module):\n    \"\"\" Grid or Block partition + Attn + FFN\n\n    '2D' NCHW tensor layout.\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            partition_type: str = 'block',\n            cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path: float = 0.,\n    ):\n        super().__init__()\n        norm_layer = partial(get_norm_layer(cfg.norm_layer), eps=cfg.norm_eps)  # NOTE this block is channels-last\n        act_layer = get_act_layer(cfg.act_layer)\n\n        self.partition_block = partition_type == 'block'\n        self.partition_size = to_2tuple(cfg.window_size if self.partition_block else cfg.grid_size)\n        rel_pos_cls = get_rel_pos_cls(cfg, self.partition_size)\n\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention2d(\n            dim,\n            dim,\n            dim_head=cfg.dim_head,\n            bias=cfg.attn_bias,\n            head_first=cfg.head_first,\n            rel_pos_cls=rel_pos_cls,\n            attn_drop=cfg.attn_drop,\n            proj_drop=cfg.proj_drop,\n        )\n        self.ls1 = LayerScale2d(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        self.norm2 = norm_layer(dim)\n        self.mlp = ConvMlp(\n            in_features=dim,\n            hidden_features=int(dim * cfg.expand_ratio),\n            act_layer=act_layer,\n            drop=cfg.proj_drop)\n        self.ls2 = LayerScale2d(dim, init_values=cfg.init_values) if cfg.init_values else nn.Identity()\n        self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def _partition_attn(self, x):\n        img_size = x.shape[-2:]\n        if self.partition_block:\n            partitioned = window_partition_nchw(x, self.partition_size)\n        else:\n            partitioned = grid_partition_nchw(x, self.partition_size)\n\n        partitioned = self.attn(partitioned)\n\n        if self.partition_block:\n            x = window_reverse_nchw(partitioned, self.partition_size, img_size)\n        else:\n            x = grid_reverse_nchw(partitioned, self.partition_size, img_size)\n        return x\n\n    def forward(self, x):\n        x = x + self.drop_path1(self.ls1(self._partition_attn(self.norm1(x))))\n        x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n        return x\n\n\nclass MaxxVitBlock(nn.Module):\n    \"\"\" MaxVit conv, window partition + FFN , grid partition + FFN\n    \"\"\"\n\n    def __init__(\n            self,\n            dim: int,\n            dim_out: int,\n            stride: int = 1,\n            conv_cfg: MaxxVitConvCfg = MaxxVitConvCfg(),\n            transformer_cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path: float = 0.,\n    ):\n        super().__init__()\n        self.nchw_attn = transformer_cfg.use_nchw_attn\n\n        conv_cls = ConvNeXtBlock if conv_cfg.block_type == 'convnext' else MbConvBlock\n        self.conv = conv_cls(dim, dim_out, stride=stride, cfg=conv_cfg, drop_path=drop_path)\n\n        attn_kwargs = dict(dim=dim_out, cfg=transformer_cfg, drop_path=drop_path)\n        partition_layer = PartitionAttention2d if self.nchw_attn else PartitionAttentionCl\n        self.attn_block = None if transformer_cfg.no_block_attn else partition_layer(**attn_kwargs)\n        self.attn_grid = partition_layer(partition_type='grid', **attn_kwargs)\n\n    def init_weights(self, scheme=''):\n        if self.attn_block is not None:\n            named_apply(partial(_init_transformer, scheme=scheme), self.attn_block)\n        named_apply(partial(_init_transformer, scheme=scheme), self.attn_grid)\n        named_apply(partial(_init_conv, scheme=scheme), self.conv)\n\n    def forward(self, x):\n        # NCHW format\n        x = self.conv(x)\n\n        if not self.nchw_attn:\n            x = x.permute(0, 2, 3, 1)  # to NHWC (channels-last)\n        if self.attn_block is not None:\n            x = self.attn_block(x)\n        x = self.attn_grid(x)\n        if not self.nchw_attn:\n            x = x.permute(0, 3, 1, 2)  # back to NCHW\n        return x\n\n\nclass ParallelMaxxVitBlock(nn.Module):\n    \"\"\" MaxVit block with parallel cat(window + grid), one FF\n    Experimental timm block.\n    \"\"\"\n\n    def __init__(\n            self,\n            dim,\n            dim_out,\n            stride=1,\n            num_conv=2,\n            conv_cfg: MaxxVitConvCfg = MaxxVitConvCfg(),\n            transformer_cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            drop_path=0.,\n    ):\n        super().__init__()\n\n        conv_cls = ConvNeXtBlock if conv_cfg.block_type == 'convnext' else MbConvBlock\n        if num_conv > 1:\n            convs = [conv_cls(dim, dim_out, stride=stride, cfg=conv_cfg, drop_path=drop_path)]\n            convs += [conv_cls(dim_out, dim_out, cfg=conv_cfg, drop_path=drop_path)] * (num_conv - 1)\n            self.conv = nn.Sequential(*convs)\n        else:\n            self.conv = conv_cls(dim, dim_out, stride=stride, cfg=conv_cfg, drop_path=drop_path)\n        self.attn = ParallelPartitionAttention(dim=dim_out, cfg=transformer_cfg, drop_path=drop_path)\n\n    def init_weights(self, scheme=''):\n        named_apply(partial(_init_transformer, scheme=scheme), self.attn)\n        named_apply(partial(_init_conv, scheme=scheme), self.conv)\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = x.permute(0, 2, 3, 1)\n        x = self.attn(x)\n        x = x.permute(0, 3, 1, 2)\n        return x\n\n\nclass MaxxVitStage(nn.Module):\n    def __init__(\n            self,\n            in_chs: int,\n            out_chs: int,\n            stride: int = 2,\n            depth: int = 4,\n            feat_size: Tuple[int, int] = (14, 14),\n            block_types: Union[str, Tuple[str]] = 'C',\n            transformer_cfg: MaxxVitTransformerCfg = MaxxVitTransformerCfg(),\n            conv_cfg: MaxxVitConvCfg = MaxxVitConvCfg(),\n            drop_path: Union[float, List[float]] = 0.,\n    ):\n        super().__init__()\n        self.grad_checkpointing = False\n\n        block_types = extend_tuple(block_types, depth)\n        blocks = []\n        for i, t in enumerate(block_types):\n            block_stride = stride if i == 0 else 1\n            assert t in ('C', 'T', 'M', 'PM')\n            if t == 'C':\n                conv_cls = ConvNeXtBlock if conv_cfg.block_type == 'convnext' else MbConvBlock\n                blocks += [conv_cls(\n                    in_chs,\n                    out_chs,\n                    stride=block_stride,\n                    cfg=conv_cfg,\n                    drop_path=drop_path[i],\n                )]\n            elif t == 'T':\n                rel_pos_cls = get_rel_pos_cls(transformer_cfg, feat_size)\n                blocks += [TransformerBlock2d(\n                    in_chs,\n                    out_chs,\n                    stride=block_stride,\n                    rel_pos_cls=rel_pos_cls,\n                    cfg=transformer_cfg,\n                    drop_path=drop_path[i],\n                )]\n            elif t == 'M':\n                blocks += [MaxxVitBlock(\n                    in_chs,\n                    out_chs,\n                    stride=block_stride,\n                    conv_cfg=conv_cfg,\n                    transformer_cfg=transformer_cfg,\n                    drop_path=drop_path[i],\n                )]\n            elif t == 'PM':\n                blocks += [ParallelMaxxVitBlock(\n                    in_chs,\n                    out_chs,\n                    stride=block_stride,\n                    conv_cfg=conv_cfg,\n                    transformer_cfg=transformer_cfg,\n                    drop_path=drop_path[i],\n                )]\n            in_chs = out_chs\n        self.blocks = nn.Sequential(*blocks)\n\n    def forward(self, x):\n#         print(f\"Inside MaxxVitStage x.shape: {x.shape}\")\n        if self.grad_checkpointing and not torch.jit.is_scripting():\n            x = checkpoint_seq(self.blocks, x)\n        else:\n            x = self.blocks(x)\n#         print(\"After self.blocks\", x.shape)\n        return x\n\n\nclass Stem(nn.Module):\n\n    def __init__(\n            self,\n            in_chs: int,\n            out_chs: int,\n            kernel_size: int = 3,\n            padding: str = '',\n            bias: bool = False,\n            act_layer: str = 'gelu',\n            norm_layer: str = 'batchnorm2d',\n            norm_eps: float = 1e-5,\n    ):\n        super().__init__()\n        if not isinstance(out_chs, (list, tuple)):\n            out_chs = to_2tuple(out_chs)\n\n        norm_act_layer = partial(get_norm_act_layer(norm_layer, act_layer), eps=norm_eps)\n        self.in_chs = in_chs # added this line\n        self.out_chs = out_chs[-1]\n        self.stride = 2\n\n        self.conv1 = create_conv2d(in_chs, out_chs[0], kernel_size, stride=2, padding=padding, bias=bias)\n        self.norm1 = norm_act_layer(out_chs[0])\n        self.conv2 = create_conv2d(out_chs[0], out_chs[1], kernel_size, stride=1, padding=padding, bias=bias)\n\n    def init_weights(self, scheme=''):\n        named_apply(partial(_init_conv, scheme=scheme), self)\n\n    def forward(self, x):\n#         print(\"Inside Stem class\", x.shape)\n#         print(f\"in_chs:{self.in_chs} and out_chs:{self.out_chs}\")\n        x = self.conv1(x)\n#         print(\"after self.conv1\", x.shape)\n        x = self.norm1(x)\n#         print(\"after self.norm1\", x.shape)\n        x = self.conv2(x)\n#         print(\"after self.conv2\", x.shape)\n        \n        return x\n\n\ndef cfg_window_size(cfg: MaxxVitTransformerCfg, img_size: Tuple[int, int]):\n    if cfg.window_size is not None:\n        assert cfg.grid_size\n        return cfg\n    partition_size = img_size[0] // cfg.partition_ratio, img_size[1] // cfg.partition_ratio\n    cfg = replace(cfg, window_size=partition_size, grid_size=partition_size)\n    return cfg\n\n\ndef _overlay_kwargs(cfg: MaxxVitCfg, **kwargs):\n    transformer_kwargs = {}\n    conv_kwargs = {}\n    base_kwargs = {}\n    for k, v in kwargs.items():\n        if k.startswith('transformer_'):\n            transformer_kwargs[k.replace('transformer_', '')] = v\n        elif k.startswith('conv_'):\n            conv_kwargs[k.replace('conv_', '')] = v\n        else:\n            base_kwargs[k] = v\n    cfg = replace(\n        cfg,\n        transformer_cfg=replace(cfg.transformer_cfg, **transformer_kwargs),\n        conv_cfg=replace(cfg.conv_cfg, **conv_kwargs),\n        **base_kwargs\n    )\n    return cfg\n\n\nclass MaxxVit(nn.Module):\n    \"\"\" CoaTNet + MaxVit base model.\n\n    Highly configurable for different block compositions, tensor layouts, pooling types.\n    \"\"\"\n\n    def __init__(\n            self,\n            cfg: MaxxVitCfg,\n            img_size: Union[int, Tuple[int, int]] = 224,\n            in_chans: int = 3,\n            num_classes: int = 1000,\n            global_pool: str = 'avg',\n            drop_rate: float = 0.,\n            drop_path_rate: float = 0.,\n            **kwargs,\n    ):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        if kwargs:\n            cfg = _overlay_kwargs(cfg, **kwargs)\n        transformer_cfg = cfg_window_size(cfg.transformer_cfg, img_size)\n        self.num_classes = num_classes\n        self.global_pool = global_pool\n        self.num_features = self.embed_dim = cfg.embed_dim[-1]\n        self.drop_rate = drop_rate\n        self.grad_checkpointing = False\n        self.feature_info = []\n\n        self.stem = Stem(\n            in_chs=in_chans,\n            out_chs=cfg.stem_width,\n            padding=cfg.conv_cfg.padding,\n            bias=cfg.stem_bias,\n            act_layer=cfg.conv_cfg.act_layer,\n            norm_layer=cfg.conv_cfg.norm_layer,\n            norm_eps=cfg.conv_cfg.norm_eps,\n        )\n        stride = self.stem.stride\n        self.feature_info += [dict(num_chs=self.stem.out_chs, reduction=2, module='stem')]\n        feat_size = tuple([i // s for i, s in zip(img_size, to_2tuple(stride))])\n\n        num_stages = len(cfg.embed_dim)\n        assert len(cfg.depths) == num_stages\n        dpr = [x.tolist() for x in torch.linspace(0, drop_path_rate, sum(cfg.depths)).split(cfg.depths)]\n        in_chs = self.stem.out_chs\n        stages = []\n        for i in range(num_stages):\n            stage_stride = 2\n            out_chs = cfg.embed_dim[i]\n            feat_size = tuple([(r - 1) // stage_stride + 1 for r in feat_size])\n            stages += [MaxxVitStage(\n                in_chs,\n                out_chs,\n                depth=cfg.depths[i],\n                block_types=cfg.block_type[i],\n                conv_cfg=cfg.conv_cfg,\n                transformer_cfg=transformer_cfg,\n                feat_size=feat_size,\n                drop_path=dpr[i],\n            )]\n            stride *= stage_stride\n            in_chs = out_chs\n            self.feature_info += [dict(num_chs=out_chs, reduction=stride, module=f'stages.{i}')]\n        self.stages = nn.Sequential(*stages)\n\n        final_norm_layer = partial(get_norm_layer(cfg.transformer_cfg.norm_layer), eps=cfg.transformer_cfg.norm_eps)\n        self.head_hidden_size = cfg.head_hidden_size\n        if self.head_hidden_size:\n            self.norm = nn.Identity()\n            self.head = NormMlpClassifierHead(\n                self.num_features,\n                num_classes,\n                hidden_size=self.head_hidden_size,\n                pool_type=global_pool,\n                drop_rate=drop_rate,\n                norm_layer=final_norm_layer,\n            )\n        else:\n            # standard classifier head w/ norm, pooling, fc classifier\n            self.norm = final_norm_layer(self.num_features)\n            self.head = ClassifierHead(self.num_features, num_classes, pool_type=global_pool, drop_rate=drop_rate)\n\n        # Weight init (default PyTorch init works well for AdamW if scheme not set)\n        assert cfg.weight_init in ('', 'normal', 'trunc_normal', 'xavier_normal', 'vit_eff')\n        if cfg.weight_init:\n            named_apply(partial(self._init_weights, scheme=cfg.weight_init), self)\n\n    def _init_weights(self, module, name, scheme=''):\n        if hasattr(module, 'init_weights'):\n            try:\n                module.init_weights(scheme=scheme)\n            except TypeError:\n                module.init_weights()\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {\n            k for k, _ in self.named_parameters()\n            if any(n in k for n in [\"relative_position_bias_table\", \"rel_pos.mlp\"])}\n\n    @torch.jit.ignore\n    def group_matcher(self, coarse=False):\n        matcher = dict(\n            stem=r'^stem',  # stem and embed\n            blocks=[(r'^stages\\.(\\d+)', None), (r'^norm', (99999,))]\n        )\n        return matcher\n\n    @torch.jit.ignore\n    def set_grad_checkpointing(self, enable=True):\n        for s in self.stages:\n            s.grad_checkpointing = enable\n\n    @torch.jit.ignore\n    def get_classifier(self):\n        return self.head.fc\n\n    def reset_classifier(self, num_classes, global_pool=None):\n        self.num_classes = num_classes\n        self.head.reset(num_classes, global_pool)\n\n    def forward_features(self, x):\n#         print(f\"In forward_features: {x.shape}\")\n        x = self.stem(x)\n#         print(f\"After self.stem: {x.shape}\")\n        x = self.stages(x)\n#         print(f\"After self.stages: {x.shape}\")\n        x = self.norm(x)\n#         print(f\"After self.norm: {x.shape}\")\n        return x\n\n    def forward_head(self, x, pre_logits: bool = False):\n#         print(f\"In forward_head function: {x.shape}\")\n#         print(f\"pre_logits is {pre_logits}\")\n        return self.head(x, pre_logits=pre_logits) if pre_logits else self.head(x)\n\n    def forward(self, x):\n#         print(f\"In forward (main): {x.shape}\")\n        x = self.forward_features(x)\n#         print(f\"After forward_features (main): {x.shape}\")\n        x = self.forward_head(x)\n#         print(f\"After forward_head (main): {x.shape}\")\n        return x\n\n\ndef _rw_coat_cfg(\n        stride_mode='pool',\n        pool_type='avg2',\n        conv_output_bias=False,\n        conv_attn_early=False,\n        conv_attn_act_layer='relu',\n        conv_norm_layer='',\n        transformer_shortcut_bias=True,\n        transformer_norm_layer='layernorm2d',\n        transformer_norm_layer_cl='layernorm',\n        init_values=None,\n        rel_pos_type='bias',\n        rel_pos_dim=512,\n):\n    # 'RW' timm variant models were created and trained before seeing https://github.com/google-research/maxvit\n    # Common differences for initial timm models:\n    # - pre-norm layer in MZBConv included an activation after norm\n    # - mbconv expansion calculated from input instead of output chs\n    # - mbconv shortcut and final 1x1 conv did not have a bias\n    # - SE act layer was relu, not silu\n    # - mbconv uses silu in timm, not gelu\n    # - expansion in attention block done via output proj, not input proj\n    # Variable differences (evolved over training initial models):\n    # - avg pool with kernel_size=2 favoured downsampling (instead of maxpool for coat)\n    # - SE attention was between conv2 and norm/act\n    # - default to avg pool for mbconv downsample instead of 1x1 or dw conv\n    # - transformer block shortcut has no bias\n    return dict(\n        conv_cfg=MaxxVitConvCfg(\n            stride_mode=stride_mode,\n            pool_type=pool_type,\n            pre_norm_act=True,\n            expand_output=False,\n            output_bias=conv_output_bias,\n            attn_early=conv_attn_early,\n            attn_act_layer=conv_attn_act_layer,\n            act_layer='silu',\n            norm_layer=conv_norm_layer,\n        ),\n        transformer_cfg=MaxxVitTransformerCfg(\n            expand_first=False,\n            shortcut_bias=transformer_shortcut_bias,\n            pool_type=pool_type,\n            init_values=init_values,\n            norm_layer=transformer_norm_layer,\n            norm_layer_cl=transformer_norm_layer_cl,\n            rel_pos_type=rel_pos_type,\n            rel_pos_dim=rel_pos_dim,\n        ),\n    )\n\n\ndef _rw_max_cfg(\n        stride_mode='dw',\n        pool_type='avg2',\n        conv_output_bias=False,\n        conv_attn_ratio=1 / 16,\n        conv_norm_layer='',\n        transformer_norm_layer='layernorm2d',\n        transformer_norm_layer_cl='layernorm',\n        window_size=None,\n        dim_head=32,\n        init_values=None,\n        rel_pos_type='bias',\n        rel_pos_dim=512,\n):\n    # 'RW' timm variant models were created and trained before seeing https://github.com/google-research/maxvit\n    # Differences of initial timm models:\n    # - mbconv expansion calculated from input instead of output chs\n    # - mbconv shortcut and final 1x1 conv did not have a bias\n    # - mbconv uses silu in timm, not gelu\n    # - expansion in attention block done via output proj, not input proj\n    return dict(\n        conv_cfg=MaxxVitConvCfg(\n            stride_mode=stride_mode,\n            pool_type=pool_type,\n            expand_output=False,\n            output_bias=conv_output_bias,\n            attn_ratio=conv_attn_ratio,\n            act_layer='silu',\n            norm_layer=conv_norm_layer,\n        ),\n        transformer_cfg=MaxxVitTransformerCfg(\n            expand_first=False,\n            pool_type=pool_type,\n            dim_head=dim_head,\n            window_size=window_size,\n            init_values=init_values,\n            norm_layer=transformer_norm_layer,\n            norm_layer_cl=transformer_norm_layer_cl,\n            rel_pos_type=rel_pos_type,\n            rel_pos_dim=rel_pos_dim,\n        ),\n    )\n\n\ndef _next_cfg(\n        stride_mode='dw',\n        pool_type='avg2',\n        conv_norm_layer='layernorm2d',\n        conv_norm_layer_cl='layernorm',\n        transformer_norm_layer='layernorm2d',\n        transformer_norm_layer_cl='layernorm',\n        window_size=None,\n        no_block_attn=False,\n        init_values=1e-6,\n        rel_pos_type='mlp',  # MLP by default for maxxvit\n        rel_pos_dim=512,\n):\n    # For experimental models with convnext instead of mbconv\n    init_values = to_2tuple(init_values)\n    return dict(\n        conv_cfg=MaxxVitConvCfg(\n            block_type='convnext',\n            stride_mode=stride_mode,\n            pool_type=pool_type,\n            expand_output=False,\n            init_values=init_values[0],\n            norm_layer=conv_norm_layer,\n            norm_layer_cl=conv_norm_layer_cl,\n        ),\n        transformer_cfg=MaxxVitTransformerCfg(\n            expand_first=False,\n            pool_type=pool_type,\n            window_size=window_size,\n            no_block_attn=no_block_attn,  # enabled for MaxxViT-V2\n            init_values=init_values[1],\n            norm_layer=transformer_norm_layer,\n            norm_layer_cl=transformer_norm_layer_cl,\n            rel_pos_type=rel_pos_type,\n            rel_pos_dim=rel_pos_dim,\n        ),\n    )\n\n\ndef _tf_cfg():\n    return dict(\n        conv_cfg=MaxxVitConvCfg(\n            norm_eps=1e-3,\n            act_layer='gelu_tanh',\n            padding='same',\n        ),\n        transformer_cfg=MaxxVitTransformerCfg(\n            norm_eps=1e-5,\n            act_layer='gelu_tanh',\n            head_first=False,  # heads are interleaved (q_nh, q_hdim, k_nh, q_hdim, ....)\n            rel_pos_type='bias_tf',\n        ),\n    )\n\n\nmodel_cfgs = dict(\n    # timm specific CoAtNet configs\n    coatnet_pico_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 3, 5, 2),\n        stem_width=(32, 64),\n        **_rw_max_cfg(  # using newer max defaults here\n            conv_output_bias=True,\n            conv_attn_ratio=0.25,\n        ),\n    ),\n    coatnet_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(3, 4, 6, 3),\n        stem_width=(32, 64),\n        **_rw_max_cfg(  # using newer max defaults here\n            stride_mode='pool',\n            conv_output_bias=True,\n            conv_attn_ratio=0.25,\n        ),\n    ),\n    coatnet_0_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 3, 7, 2),  # deeper than paper '0' model\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            conv_attn_early=True,\n            transformer_shortcut_bias=False,\n        ),\n    ),\n    coatnet_1_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_early=True,\n            transformer_shortcut_bias=False,\n        )\n    ),\n    coatnet_2_rw=MaxxVitCfg(\n        embed_dim=(128, 256, 512, 1024),\n        depths=(2, 6, 14, 2),\n        stem_width=(64, 128),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_act_layer='silu',\n            #init_values=1e-6,\n        ),\n    ),\n    coatnet_3_rw=MaxxVitCfg(\n        embed_dim=(192, 384, 768, 1536),\n        depths=(2, 6, 14, 2),\n        stem_width=(96, 192),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_act_layer='silu',\n            init_values=1e-6,\n        ),\n    ),\n\n    # Experimental CoAtNet configs w/ ImageNet-1k train (different norm layers, MLP rel-pos)\n    coatnet_bn_0_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 3, 7, 2),  # deeper than paper '0' model\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_early=True,\n            transformer_shortcut_bias=False,\n            transformer_norm_layer='batchnorm2d',\n        )\n    ),\n    coatnet_rmlp_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(3, 4, 6, 3),\n        stem_width=(32, 64),\n        **_rw_max_cfg(\n            conv_output_bias=True,\n            conv_attn_ratio=0.25,\n            rel_pos_type='mlp',\n            rel_pos_dim=384,\n        ),\n    ),\n    coatnet_rmlp_0_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 3, 7, 2),  # deeper than paper '0' model\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            rel_pos_type='mlp',\n        ),\n    ),\n    coatnet_rmlp_1_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            pool_type='max',\n            conv_attn_early=True,\n            transformer_shortcut_bias=False,\n            rel_pos_type='mlp',\n            rel_pos_dim=384,  # was supposed to be 512, woops\n        ),\n    ),\n    coatnet_rmlp_1_rw2=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        stem_width=(32, 64),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            rel_pos_type='mlp',\n            rel_pos_dim=512,  # was supposed to be 512, woops\n        ),\n    ),\n    coatnet_rmlp_2_rw=MaxxVitCfg(\n        embed_dim=(128, 256, 512, 1024),\n        depths=(2, 6, 14, 2),\n        stem_width=(64, 128),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_act_layer='silu',\n            init_values=1e-6,\n            rel_pos_type='mlp'\n        ),\n    ),\n    coatnet_rmlp_3_rw=MaxxVitCfg(\n        embed_dim=(192, 384, 768, 1536),\n        depths=(2, 6, 14, 2),\n        stem_width=(96, 192),\n        **_rw_coat_cfg(\n            stride_mode='dw',\n            conv_attn_act_layer='silu',\n            init_values=1e-6,\n            rel_pos_type='mlp'\n        ),\n    ),\n\n    coatnet_nano_cc=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(3, 4, 6, 3),\n        stem_width=(32, 64),\n        block_type=('C', 'C', ('C', 'T'), ('C', 'T')),\n        **_rw_coat_cfg(),\n    ),\n    coatnext_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(3, 4, 6, 3),\n        stem_width=(32, 64),\n        weight_init='normal',\n        **_next_cfg(\n            rel_pos_type='bias',\n            init_values=(1e-5, None)\n        ),\n    ),\n\n    # Trying to be like the CoAtNet paper configs\n    coatnet_0=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 3, 5, 2),\n        stem_width=64,\n        head_hidden_size=768,\n    ),\n    coatnet_1=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        stem_width=64,\n        head_hidden_size=768,\n    ),\n    coatnet_2=MaxxVitCfg(\n        embed_dim=(128, 256, 512, 1024),\n        depths=(2, 6, 14, 2),\n        stem_width=128,\n        head_hidden_size=1024,\n    ),\n    coatnet_3=MaxxVitCfg(\n        embed_dim=(192, 384, 768, 1536),\n        depths=(2, 6, 14, 2),\n        stem_width=192,\n        head_hidden_size=1536,\n    ),\n    coatnet_4=MaxxVitCfg(\n        embed_dim=(192, 384, 768, 1536),\n        depths=(2, 12, 28, 2),\n        stem_width=192,\n        head_hidden_size=1536,\n    ),\n    coatnet_5=MaxxVitCfg(\n        embed_dim=(256, 512, 1280, 2048),\n        depths=(2, 12, 28, 2),\n        stem_width=192,\n        head_hidden_size=2048,\n    ),\n\n    # Experimental MaxVit configs\n    maxvit_pico_rw=MaxxVitCfg(\n        embed_dim=(32, 64, 128, 256),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(24, 32),\n        **_rw_max_cfg(),\n    ),\n    maxvit_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(1, 2, 3, 1),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(),\n    ),\n    maxvit_tiny_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(),\n    ),\n    maxvit_tiny_pm=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 2, 5, 2),\n        block_type=('PM',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(),\n    ),\n\n    maxvit_rmlp_pico_rw=MaxxVitCfg(\n        embed_dim=(32, 64, 128, 256),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(24, 32),\n        **_rw_max_cfg(rel_pos_type='mlp'),\n    ),\n    maxvit_rmlp_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(1, 2, 3, 1),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(rel_pos_type='mlp'),\n    ),\n    maxvit_rmlp_tiny_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(rel_pos_type='mlp'),\n    ),\n    maxvit_rmlp_small_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_rw_max_cfg(\n            rel_pos_type='mlp',\n            init_values=1e-6,\n        ),\n    ),\n    maxvit_rmlp_base_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        head_hidden_size=768,\n        **_rw_max_cfg(\n            rel_pos_type='mlp',\n        ),\n    ),\n\n    maxxvit_rmlp_nano_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(1, 2, 3, 1),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        weight_init='normal',\n        **_next_cfg(),\n    ),\n    maxxvit_rmlp_tiny_rw=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(32, 64),\n        **_next_cfg(),\n    ),\n    maxxvit_rmlp_small_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=(48, 96),\n        **_next_cfg(),\n    ),\n\n    maxxvitv2_nano_rw=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(1, 2, 3, 1),\n        block_type=('M',) * 4,\n        stem_width=(48, 96),\n        weight_init='normal',\n        **_next_cfg(\n            no_block_attn=True,\n            rel_pos_type='bias',\n        ),\n    ),\n    maxxvitv2_rmlp_base_rw=MaxxVitCfg(\n        embed_dim=(128, 256, 512, 1024),\n        depths=(2, 6, 12, 2),\n        block_type=('M',) * 4,\n        stem_width=(64, 128),\n        **_next_cfg(\n            no_block_attn=True,\n        ),\n    ),\n    maxxvitv2_rmlp_large_rw=MaxxVitCfg(\n        embed_dim=(160, 320, 640, 1280),\n        depths=(2, 6, 16, 2),\n        block_type=('M',) * 4,\n        stem_width=(80, 160),\n        head_hidden_size=1280,\n        **_next_cfg(\n            no_block_attn=True,\n        ),\n    ),\n\n    # Trying to be like the MaxViT paper configs\n    maxvit_tiny_tf=MaxxVitCfg(\n        embed_dim=(64, 128, 256, 512),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=64,\n        stem_bias=True,\n        head_hidden_size=512,\n        **_tf_cfg(),\n    ),\n    maxvit_small_tf=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 2, 5, 2),\n        block_type=('M',) * 4,\n        stem_width=64,\n        stem_bias=True,\n        head_hidden_size=768,\n        **_tf_cfg(),\n    ),\n    maxvit_base_tf=MaxxVitCfg(\n        embed_dim=(96, 192, 384, 768),\n        depths=(2, 6, 14, 2),\n        block_type=('M',) * 4,\n        stem_width=64,\n        stem_bias=True,\n        head_hidden_size=768,\n        **_tf_cfg(),\n    ),\n    maxvit_large_tf=MaxxVitCfg(\n        embed_dim=(128, 256, 512, 1024),\n        depths=(2, 6, 14, 2),\n        block_type=('M',) * 4,\n        stem_width=128,\n        stem_bias=True,\n        head_hidden_size=1024,\n        **_tf_cfg(),\n    ),\n    maxvit_xlarge_tf=MaxxVitCfg(\n        embed_dim=(192, 384, 768, 1536),\n        depths=(2, 6, 14, 2),\n        block_type=('M',) * 4,\n        stem_width=192,\n        stem_bias=True,\n        head_hidden_size=1536,\n        **_tf_cfg(),\n    ),\n)\n\n\ndef checkpoint_filter_fn(state_dict, model: nn.Module):\n    model_state_dict = model.state_dict()\n    out_dict = {}\n    for k, v in state_dict.items():\n        if k.endswith('relative_position_bias_table'):\n            m = model.get_submodule(k[:-29])\n            if v.shape != m.relative_position_bias_table.shape or m.window_size[0] != m.window_size[1]:\n                v = resize_rel_pos_bias_table(\n                    v,\n                    new_window_size=m.window_size,\n                    new_bias_shape=m.relative_position_bias_table.shape,\n                )\n\n        if k in model_state_dict and v.ndim != model_state_dict[k].ndim and v.numel() == model_state_dict[k].numel():\n            # adapt between conv2d / linear layers\n            assert v.ndim in (2, 4)\n            v = v.reshape(model_state_dict[k].shape)\n        out_dict[k] = v\n    return out_dict\n\n\ndef _create_maxxvit(variant, cfg_variant=None, pretrained=False, **kwargs):\n    if cfg_variant is None:\n        if variant in model_cfgs:\n            cfg_variant = variant\n        else:\n            cfg_variant = '_'.join(variant.split('_')[:-1])\n    return build_model_with_cfg(\n        MaxxVit, variant, pretrained,\n        model_cfg=model_cfgs[cfg_variant],\n        feature_cfg=dict(flatten_sequential=True),\n        pretrained_filter_fn=checkpoint_filter_fn,\n        **kwargs)\n\n\ndef _cfg(url='', **kwargs):\n    return {\n        'url': url, 'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': (7, 7),\n        'crop_pct': 0.95, 'interpolation': 'bicubic',\n        'mean': (0.5, 0.5, 0.5), 'std': (0.5, 0.5, 0.5),\n        'first_conv': 'stem.conv1', 'classifier': 'head.fc',\n        'fixed_input_size': True,\n        **kwargs\n    }\n\n\ndefault_cfgs = generate_default_cfgs({\n    # timm specific CoAtNet configs, ImageNet-1k pretrain, fixed rel-pos\n    'coatnet_pico_rw_224.untrained': _cfg(url=''),\n    'coatnet_nano_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_nano_rw_224_sw-f53093b4.pth',\n        crop_pct=0.9),\n    'coatnet_0_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_0_rw_224_sw-a6439706.pth'),\n    'coatnet_1_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_1_rw_224_sw-5cae1ea8.pth'\n    ),\n\n    # timm specific CoAtNet configs, ImageNet-12k pretrain w/ 1k fine-tune, fixed rel-pos\n    'coatnet_2_rw_224.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    #'coatnet_3_rw_224.untrained': _cfg(url=''),\n\n    # Experimental CoAtNet configs w/ ImageNet-12k pretrain -> 1k fine-tune (different norm layers, MLP rel-pos)\n    'coatnet_rmlp_1_rw2_224.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'coatnet_rmlp_2_rw_224.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n\n    # Experimental CoAtNet configs w/ ImageNet-1k train (different norm layers, MLP rel-pos)\n    'coatnet_bn_0_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_bn_0_rw_224_sw-c228e218.pth',\n        mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD,\n        crop_pct=0.95),\n    'coatnet_rmlp_nano_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_nano_rw_224_sw-bd1d51b3.pth',\n        crop_pct=0.9),\n    'coatnet_rmlp_0_rw_224.untrained': _cfg(url=''),\n    'coatnet_rmlp_1_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_1_rw_224_sw-9051e6c3.pth'),\n    'coatnet_rmlp_2_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnet_rmlp_2_rw_224_sw-5ccfac55.pth'),\n    'coatnet_rmlp_3_rw_224.untrained': _cfg(url=''),\n    'coatnet_nano_cc_224.untrained': _cfg(url=''),\n    'coatnext_nano_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/coatnext_nano_rw_224_ad-22cb71c2.pth',\n        crop_pct=0.9),\n\n    # ImagenNet-12k pretrain CoAtNet\n    'coatnet_2_rw_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821),\n    'coatnet_3_rw_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821),\n    'coatnet_rmlp_1_rw2_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821),\n    'coatnet_rmlp_2_rw_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821),\n\n    # Trying to be like the CoAtNet paper configs (will adapt if 'tf' weights are ever released)\n    'coatnet_0_224.untrained': _cfg(url=''),\n    'coatnet_1_224.untrained': _cfg(url=''),\n    'coatnet_2_224.untrained': _cfg(url=''),\n    'coatnet_3_224.untrained': _cfg(url=''),\n    'coatnet_4_224.untrained': _cfg(url=''),\n    'coatnet_5_224.untrained': _cfg(url=''),\n\n    # timm specific MaxVit configs, ImageNet-1k pretrain or untrained\n    'maxvit_pico_rw_256.untrained': _cfg(url='', input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_nano_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_nano_rw_256_sw-fb127241.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_tiny_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_tiny_rw_224_sw-7d0dffeb.pth'),\n    'maxvit_tiny_rw_256.untrained': _cfg(\n        url='',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_tiny_pm_256.untrained': _cfg(url='', input_size=(3, 256, 256), pool_size=(8, 8)),\n\n    # timm specific MaxVit w/ MLP rel-pos, ImageNet-1k pretrain\n    'maxvit_rmlp_pico_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_pico_rw_256_sw-8d82f2c6.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_rmlp_nano_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_nano_rw_256_sw-c17bb0d6.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_rmlp_tiny_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_tiny_rw_256_sw-bbef0ff5.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxvit_rmlp_small_rw_224.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxvit_rmlp_small_rw_224_sw-6ef0ae4f.pth',\n        crop_pct=0.9,\n    ),\n    'maxvit_rmlp_small_rw_256.untrained': _cfg(\n        url='',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n\n    # timm specific MaxVit w/ ImageNet-12k pretrain and 1k fine-tune\n    'maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n    ),\n    'maxvit_rmlp_base_rw_384.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n\n    # timm specific MaxVit w/ ImageNet-12k pretrain\n    'maxvit_rmlp_base_rw_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821,\n    ),\n\n    # timm MaxxViT configs (ConvNeXt conv blocks mixed with MaxVit transformer blocks)\n    'maxxvit_rmlp_nano_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxxvit_rmlp_nano_rw_256_sw-0325d459.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxxvit_rmlp_tiny_rw_256.untrained': _cfg(url='', input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxxvit_rmlp_small_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights-maxx/maxxvit_rmlp_small_rw_256_sw-37e217ff.pth',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n\n    # timm MaxxViT-V2 configs (ConvNeXt conv blocks mixed with MaxVit transformer blocks, more width, no block attn)\n    'maxxvitv2_nano_rw_256.sw_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 256, 256), pool_size=(8, 8)),\n    'maxxvitv2_rmlp_base_rw_224.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'maxxvitv2_rmlp_base_rw_384.sw_in12k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxxvitv2_rmlp_large_rw_224.untrained': _cfg(url=''),\n\n    'maxxvitv2_rmlp_base_rw_224.sw_in12k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=11821),\n\n    # MaxViT models ported from official Tensorflow impl\n    'maxvit_tiny_tf_224.in1k': _cfg(\n        hf_hub_id='timm/',\n        mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n    'maxvit_tiny_tf_384.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_tiny_tf_512.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_small_tf_224.in1k': _cfg(\n        hf_hub_id='timm/',\n        mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n    'maxvit_small_tf_384.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_small_tf_512.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_base_tf_224.in1k': _cfg(\n        hf_hub_id='timm/',\n        mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n    'maxvit_base_tf_384.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_base_tf_512.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_large_tf_224.in1k': _cfg(\n        hf_hub_id='timm/',\n        mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n    'maxvit_large_tf_384.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_large_tf_512.in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n\n    'maxvit_base_tf_224.in21k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=21843),\n    'maxvit_base_tf_384.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_base_tf_512.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_large_tf_224.in21k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=21843),\n    'maxvit_large_tf_384.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_large_tf_512.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_xlarge_tf_224.in21k': _cfg(\n        hf_hub_id='timm/',\n        num_classes=21843),\n    'maxvit_xlarge_tf_384.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, crop_mode='squash'),\n    'maxvit_xlarge_tf_512.in21k_ft_in1k': _cfg(\n        hf_hub_id='timm/',\n        input_size=(3, 512, 512), pool_size=(16, 16), crop_pct=1.0, crop_mode='squash'),\n})\n\n\n@register_model\ndef coatnet_pico_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_pico_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_nano_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_nano_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_0_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_0_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_1_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_1_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_2_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_2_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_3_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_3_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_bn_0_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_bn_0_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_nano_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_nano_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_0_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_0_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_1_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_1_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_1_rw2_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_1_rw2_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_2_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_2_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_2_rw_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_2_rw_384', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_rmlp_3_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_rmlp_3_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_nano_cc_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_nano_cc_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnext_nano_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnext_nano_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_0_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_0_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_1_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_1_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_2_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_2_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_3_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_3_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_4_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_4_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef coatnet_5_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('coatnet_5_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_pico_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_pico_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_nano_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_nano_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_pico_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_pico_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_nano_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_nano_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_tiny_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_tiny_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_small_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_small_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_small_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_small_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_base_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_base_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_rmlp_base_rw_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_rmlp_base_rw_384', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_pm_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_pm_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvit_rmlp_nano_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvit_rmlp_nano_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvit_rmlp_tiny_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvit_rmlp_tiny_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvit_rmlp_small_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvit_rmlp_small_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvitv2_nano_rw_256(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvitv2_nano_rw_256', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvitv2_rmlp_base_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvitv2_rmlp_base_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvitv2_rmlp_base_rw_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvitv2_rmlp_base_rw_384', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxxvitv2_rmlp_large_rw_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxxvitv2_rmlp_large_rw_224', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_tf_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_tf_224', 'maxvit_tiny_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_tf_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_tf_384', 'maxvit_tiny_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_tiny_tf_512(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_tiny_tf_512', 'maxvit_tiny_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_small_tf_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_small_tf_224', 'maxvit_small_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_small_tf_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_small_tf_384', 'maxvit_small_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_small_tf_512(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_small_tf_512', 'maxvit_small_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_base_tf_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_base_tf_224', 'maxvit_base_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_base_tf_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_base_tf_384', 'maxvit_base_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_base_tf_512(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_base_tf_512', 'maxvit_base_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_large_tf_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_large_tf_224', 'maxvit_large_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_large_tf_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_large_tf_384', 'maxvit_large_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_large_tf_512(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_large_tf_512', 'maxvit_large_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_xlarge_tf_224(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_xlarge_tf_224', 'maxvit_xlarge_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_xlarge_tf_384(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_xlarge_tf_384', 'maxvit_xlarge_tf', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef maxvit_xlarge_tf_512(pretrained=False, **kwargs) -> MaxxVit:\n    return _create_maxxvit('maxvit_xlarge_tf_512', 'maxvit_xlarge_tf', pretrained=pretrained, **kwargs)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = coatnet_3_rw_224(pretrained=False, in_chans=1, num_classes=num_classes)\n\nmodel.to(device)\n# print(model)\nprint()\n\nx = torch.randn(1, 1, IMAGE_SIZE[0], IMAGE_SIZE[1]).to(device)\n\n\noutput = model(x)\nprint(\"Model output's shape:\", output.shape)\n# print(output) # logits\n# display_params_flops(model)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:31.087489Z","iopub.execute_input":"2024-04-18T12:45:31.087875Z","iopub.status.idle":"2024-04-18T12:45:37.637705Z","shell.execute_reply.started":"2024-04-18T12:45:31.087848Z","shell.execute_reply":"2024-04-18T12:45:37.63676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = coatnet_3_rw_224(pretrained=False, in_chans=1, num_classes=num_classes) # change fine_tune as required\nmodel.to(device)\n\n\nloss_function = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), 3e-5) # adjust learning rate \n# reduce lr during fine tuning","metadata":{"papermill":{"duration":1.973611,"end_time":"2024-03-26T18:39:49.325262","exception":false,"start_time":"2024-03-26T18:39:47.351651","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:37.639581Z","iopub.execute_input":"2024-04-18T12:45:37.639974Z","iopub.status.idle":"2024-04-18T12:45:43.348611Z","shell.execute_reply.started":"2024-04-18T12:45:37.63993Z","shell.execute_reply":"2024-04-18T12:45:43.347553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\n# Assuming your model outputs logits\nlogits = torch.randn(3, num_classes)  # Example logits\ntargets = torch.randint(0, 2, (3, num_classes))  # Example targets (binary, one-hot encoded)\nprint(logits)\nprint(targets)\n\n# Define the loss function\ncriterion = nn.BCEWithLogitsLoss()\n\n# Calculate the loss\nloss = criterion(logits, targets.float())\nprint(\"loss:\", loss)","metadata":{"papermill":{"duration":0.068876,"end_time":"2024-03-26T18:39:49.430899","exception":false,"start_time":"2024-03-26T18:39:49.362023","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:43.349801Z","iopub.execute_input":"2024-04-18T12:45:43.350105Z","iopub.status.idle":"2024-04-18T12:45:43.37642Z","shell.execute_reply.started":"2024-04-18T12:45:43.350081Z","shell.execute_reply":"2024-04-18T12:45:43.375515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, loss_function=loss_function, optimizer=optimizer, num_epochs=20,\n          ranOnce=False, epochs_ran=0, model_path='model.pth', history_path='history.csv',\n         save_interval=1):\n    '''\n    returns the 'history' dataframe.\n    \n    This function trains the model for a fixed number of epochs,\n    saving the checkpoints regularly.\n    It monitors time taken per epoch and total time elapsed.\n    It tracks loss, roc_auc, f1, accuracy, precision, recall for both train & validation data.\n    It also saves the best model on loss, roc_auc, f1 and accuracy.\n    '''\n    \n    \n    initial_epoch = 0\n    \n    valid_loss_min = np.Inf\n    valid_max_accuracy = 0\n    valid_max_auc = 0\n    valid_max_precision = 0\n    valid_max_recall = 0\n    valid_max_f1 = 0\n    \n    \n    \n    if ranOnce:\n        if epochs_ran <= 0:\n            print(\"Mention the no. of epochs run by the model already for which you have the weights\")\n            return \n\n        history = pd.read_csv(history_path)\n        history=history.head(epochs_ran)\n        model.load_state_dict(torch.load(model_path))\n        '''\n        It has the following columns:-\n        epoch_number, train_loss, val_loss, train_auc, val_auc, train_accuracy, val_accuracy,\n        train_f1, val_f1, train_precision, val_precision, train_recall, val_recall,\n        time_current_epoch, total_time_elapsed\n        '''\n        initial_epoch = len(history)\n        \n        valid_loss_min = history['val_loss'].min()\n        valid_max_accuracy = history['val_accuracy'].max()\n        valid_max_auc = history['val_auc'].max()\n        valid_max_precision = history['val_precision'].max()\n        valid_max_recall = history['val_recall'].max()\n        valid_max_f1 = history['val_f1'].max()\n        \n        \n        print(f\"Model was already trained fo {initial_epoch} epochs,\\\n    with minimum loss: {valid_loss_min}, max accuracy: {valid_max_accuracy},\\\n    max auc: {valid_max_auc}, max precision: {valid_max_precision}, \\\n    max recall: {valid_max_recall}, max f1: {valid_max_f1}\")\n        \n    else:\n        print(\"Starting afresh!\")\n        history = pd.DataFrame()\n#         history = pd.DataFrame(columns=['epoch_number', 'train_loss', 'train_accuracy', 'train_f1',\\\n#             'train_precision', 'train_recall', 'train_auc', 'train_auc_scores',\\\n#             'val_loss', 'val_accuracy', 'val_f1', 'val_precision', 'val_recall', 'val_auc','val_auc_scores',\\\n#             'time_current_epoch'])\n        \n    # Main loop\n    for epoch in range(initial_epoch+1, initial_epoch + num_epochs+1):\n        \n        history_list = [] # stores the history for the current epoch in a list \n        \n        train_labels_all = []\n        train_predictions_all = []\n        train_scores_all = []\n        \n        train_loss = 0.0\n        correct_train = 0\n        total_train = 0\n        \n         # Set to training\n        model.train()\n        start = timer()\n        \n        # Training loop\n        for batch_data in tqdm(train_loader):\n            inputs, labels = batch_data[0].to(device), batch_data[1].float().to(device)  \n            \n            # Clear gradients\n            optimizer.zero_grad()\n\n            outputs = model(inputs)\n#             print(\"train_batch_outputs:\", outputs)\n#             print(\"train_batch_labels:\", labels)\n            \n            outputs = outputs.float()\n            labels = labels.float()\n\n            # Loss and backpropagation of gradients\n            loss = loss_function(outputs, labels)\n\n            loss.backward()\n            optimizer.step()  # Update the parameters\n            \n            train_loss += loss.item()\n            \n            # Calculate training accuracy\n            scores = torch.sigmoid(outputs)\n            predictions = torch.sigmoid(outputs) > 0.5\n            total_train += labels.size(0) * labels.size(1)\n            \n            correct_train += (predictions == labels).sum().item()\n            \n#             print(\"train_batch_predictions:\", predictions)\n#             print(\"train_batch_scores:\", scores)\n            \n            train_scores_all.extend(scores.detach().cpu().numpy())\n            train_labels_all.extend(labels.cpu().numpy())\n            train_predictions_all.extend(predictions.cpu().numpy())\n            \n#             train_labels_all.extend(labels)\n#             train_predictions_all.extend(predictions)\n            \n    \n        train_loss /= len(train_loader)\n        accuracy_train = correct_train / total_train\n        \n        print(f\"Current epoch {epoch}/{initial_epoch + num_epochs}\")\n        print(f\"Training Loss: {train_loss:.4f}, Accuracy: {accuracy_train:.4f}\")\n        print(\"correct:\", correct_train, \" out of \", total_train)\n\n        train_predictions_all = np.array(train_predictions_all).astype(float)\n        train_labels_all = np.array(train_labels_all).astype(float)\n        \n#         print(classification_report(train_labels_all, train_predictions_all, target_names=diseases))\n        \n        cm = multilabel_confusion_matrix(train_predictions_all, train_labels_all)\n        # we will use macro-averaging strategy.\n        accuracy_arr = []\n        precision_arr = []\n        recall_arr = []\n        f1_arr = []\n        \n        for i in range(num_classes):\n#             print(cm[i])\n#             print(cm[i].sum())\n            \n#             cfm_plot = sn.heatmap(cm[i], annot=False)\n             # TP + TN / TP + TN  + FP + FN\n            accuracy = (cm[i][0][0]+ cm[i][1][1])/cm[i].sum()\n\n            # TP / TP + FP\n            precision = cm[i][1][1]/(cm[i][1][0]+cm[i][1][1])\n\n            # TP / TP + FN\n            recall = cm[i][1][1]/(cm[i][0][1]+cm[i][1][1]) # sensitivity\n            f1 = (2*precision*recall)/(precision+recall)\n#             print(diseases[i],\": \",round(accuracy*100,2),\"%\")\n#             print(\"Precision: \",round(precision,2))\n#             print(\"Recall:\", round(recall,2))\n#             print(\"F1-Score:\", round(f1,2))\n#             print('==========================================================')\n\n\n            accuracy_arr.append(accuracy)\n            precision_arr.append(precision)\n            recall_arr.append(recall)\n            f1_arr.append(f1)\n        \n        \n        accuracy_arr = np.nan_to_num(accuracy_arr, nan=0)\n        precision_arr = np.nan_to_num(precision_arr, nan=0)\n        recall_arr = np.nan_to_num(recall_arr, nan=0)\n        f1_arr = np.nan_to_num(f1_arr, nan=0)\n        \n        accuracy_macro_train = round(sum(accuracy_arr) / len(accuracy_arr), 4)\n        precision_macro_train = round(sum(precision_arr) / len(precision_arr), 4)\n        recall_macro_train = round(sum(recall_arr) / len(recall_arr), 4)\n        f1_macro_train = round(sum(f1_arr) / len(f1_arr), 4)\n        roc_auc_macro_train = round(roc_auc_score(train_labels_all, train_scores_all, average='macro'), 4)\n\n        print(\"MACRO-averged metrics\", end=':- ')\n        print(f\"accuracy: {accuracy_macro_train}, precision: {precision_macro_train}\", end=', ')\n        print(f\"recall: {recall_macro_train}, f1: {f1_macro_train}, ROC_AUC: {roc_auc_macro_train}\")\n\n        \n        #VALIDATION START:->\n        val_labels_all = []\n        val_predictions_all = []\n        val_scores_all = []\n        \n        val_loss = 0.0\n        correct_val = 0\n        total_val = 0\n        \n        # Don't need to keep track of gradients\n        with torch.no_grad():\n            # Set to evaluation mode\n            model.eval()\n        \n            #Validation loop\n            for batch_data in tqdm(val_loader):\n                inputs, labels = batch_data[0].to(device), batch_data[1].float().to(device)  \n\n                outputs = model(inputs)\n                outputs = outputs.float()\n                labels = labels.float()\n                \n                loss = loss_function(outputs, labels)                \n                val_loss += loss.item()\n\n                scores = torch.sigmoid(outputs)\n                predictions = torch.sigmoid(outputs) > 0.5\n                total_val += labels.size(0) * labels.size(1)\n                \n                correct_val += (predictions == labels).sum().item()\n                \n                val_scores_all.extend(scores.detach().cpu().numpy())\n                val_labels_all.extend(labels.cpu().numpy())\n                val_predictions_all.extend(predictions.cpu().numpy())\n                \n\n            val_loss /= len(val_loader)\n            accuracy_val = correct_val / total_val\n\n            print(f\"Validation Loss: {val_loss:.4f}, Accuracy: {accuracy_val:.4f}\")\n            print(\"correct:\", correct_val, \" out of \", total_val)\n\n            val_predictions_all = np.array(val_predictions_all).astype(float)\n            val_labels_all = np.array(val_labels_all).astype(float)\n\n    #         print(classification_report(val_labels_all, val_predictions_all, target_names=diseases))\n\n            cm = multilabel_confusion_matrix(val_predictions_all, val_labels_all)\n            # we will use macro-averaging strategy.\n            accuracy_arr = []\n            precision_arr = []\n            recall_arr = []\n            f1_arr = []\n\n            for i in range(num_classes):\n    #             print(cm[i])\n    #             print(cm[i].sum())\n\n    #             cfm_plot = sn.heatmap(cm[i], annot=False)\n#                 TP + TN / TP + TN  + FP + FN\n                accuracy = (cm[i][0][0]+ cm[i][1][1])/cm[i].sum()\n\n                # TP / TP + FP\n                precision = cm[i][1][1]/(cm[i][1][0]+cm[i][1][1])\n\n                # TP / TP + FN\n                recall = cm[i][1][1]/(cm[i][0][1]+cm[i][1][1]) # sensitivity\n                f1 = (2*precision*recall)/(precision+recall)\n    #             print(diseases[i],\": \",round(accuracy*100,2),\"%\")\n    #             print(\"Precision: \",round(precision,2))\n    #             print(\"Recall:\", round(recall,2))\n    #             print(\"F1-Score:\", round(f1,2))\n    #             print('==========================================================')\n\n        \n                accuracy_arr.append(accuracy)\n                precision_arr.append(precision)\n                recall_arr.append(recall)\n                f1_arr.append(f1)\n            \n            accuracy_arr = np.nan_to_num(accuracy_arr, nan=0)\n            precision_arr = np.nan_to_num(precision_arr, nan=0)\n            recall_arr = np.nan_to_num(recall_arr, nan=0)\n            f1_arr = np.nan_to_num(f1_arr, nan=0)\n            \n            accuracy_macro_val = round(sum(accuracy_arr) / len(accuracy_arr), 4)\n            precision_macro_val = round(sum(precision_arr) / len(precision_arr), 4)\n            recall_macro_val = round(sum(recall_arr) / len(recall_arr), 4)\n            f1_macro_val = round(sum(f1_arr) / len(f1_arr), 4)\n            roc_auc_macro_val = round(roc_auc_score(val_labels_all, val_scores_all, average='macro'), 4)\n            \n            print(\"MACRO-averged metrics\", end=':- ')\n            print(f\"accuracy: {accuracy_macro_val}, precision: {precision_macro_val}\", end=', ')\n            print(f\"recall: {recall_macro_val}, f1: {f1_macro_val}, ROC_AUC: {roc_auc_macro_val}\")\n\n            time_this_epoch = timer()-start\n            print(f\"Time_for_this_epoch: {(time_this_epoch):.4f} seconds\")\n            print(\"-\"*120)\n\n        \n#         Add values to the history DataFrame\n        history_list.append({\n            'epoch_number': epoch,\n            'train_loss': train_loss,\n            'train_accuracy': accuracy_macro_train,\n            'train_f1': f1_macro_train,\n            'train_precision': precision_macro_train,\n            'train_recall': recall_macro_train,\n            'train_auc': roc_auc_macro_train,\n            \n            'val_loss': val_loss,\n            'val_accuracy': accuracy_macro_val,\n            'val_f1': f1_macro_val,\n            'val_precision': precision_macro_val,\n            'val_recall': recall_macro_val,\n            'val_auc': roc_auc_macro_val,\n            \n            # Add other metrics as needed\n            'time_current_epoch': time_this_epoch\n        })\n        \n        # Convert the list of dictionaries to a DataFrame\n        epoch_history = pd.DataFrame(history_list)\n\n        # Concatenate the new DataFrame with the existing history DataFrame\n        history = pd.concat([history, epoch_history], ignore_index=True)\n        \n        history.to_csv('history.csv', index=False)\n#         print(\"history_list:\", history_list)\n        print('history should be saved')\n        \n        \n#         ### HANDLING THE MODEL SAVING MECHANISM\n# #       # Save the model with the best accuracy\n        if accuracy_macro_val > valid_max_accuracy:\n            valid_max_accuracy = accuracy_macro_val\n            best_model_accuracy = model.state_dict()\n            torch.save(best_model_accuracy, 'best_model_accuracy.pth')\n\n\n#         # Save the model with the best AUC\n        if roc_auc_macro_val > valid_max_auc:\n            valid_max_auc = roc_auc_macro_val\n            best_model_auc = model.state_dict()\n            torch.save(best_model_auc, 'best_model_auc.pth')\n\n\n#         # Save the model with the best validation loss\n        if val_loss < valid_loss_min:\n            valid_loss_min = val_loss\n            best_model_val_loss = model.state_dict()\n            torch.save(best_model_val_loss, 'best_model_val_loss.pth')\n            \n        # best precision\n        if precision_macro_val > valid_max_precision:\n            valid_max_precision = precision_macro_val\n            best_model_precision = model.state_dict()\n            torch.save(best_model_precision, 'best_model_precision.pth')\n        \n        # best recall \n        if recall_macro_val > valid_max_recall :\n            valid_max_recall  = recall_macro_val\n            best_model_recall = model.state_dict()\n            torch.save(best_model_recall, 'best_model_recall.pth')\n        \n        # best f1\n        if f1_macro_val > valid_max_f1 :\n            valid_max_f1  = f1_macro_val\n            best_model_f1 = model.state_dict()\n            torch.save(best_model_f1, 'best_model_f1.pth')\n            \n\n#         # Saving model every 'save_interval' number of epochs\n        if epoch % save_interval == 0:\n            print(f\"Saving model at epoch number: {epoch}\")\n            torch.save(model.state_dict(), f\"model_{epoch}.pth\")\n        \n\n        \n    return history","metadata":{"papermill":{"duration":0.083086,"end_time":"2024-03-26T18:39:49.549144","exception":false,"start_time":"2024-03-26T18:39:49.466058","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:43.378671Z","iopub.execute_input":"2024-04-18T12:45:43.379052Z","iopub.status.idle":"2024-04-18T12:45:43.43007Z","shell.execute_reply.started":"2024-04-18T12:45:43.379022Z","shell.execute_reply":"2024-04-18T12:45:43.429094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" history = train(model, loss_function=loss_function, optimizer=optimizer, num_epochs=10,\n          ranOnce=False,\n         save_interval=10)","metadata":{"papermill":{"duration":17107.713737,"end_time":"2024-03-26T23:24:57.298152","exception":false,"start_time":"2024-03-26T18:39:49.584415","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-30T18:02:54.019398Z","iopub.execute_input":"2024-03-30T18:02:54.020126Z","iopub.status.idle":"2024-03-30T18:11:29.479364Z","shell.execute_reply.started":"2024-03-30T18:02:54.020093Z","shell.execute_reply":"2024-03-30T18:11:29.478381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Fine Tuning ( replace the model name )","metadata":{}},{"cell_type":"code","source":"model = coatnet_3_rw_224(pretrained=False, in_chans=1, num_classes=num_classes) # change fine_tune as required\nmodel.to(device)\n\n\nloss_function = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), 7e-6) # adjust learning rate \n# reduce lr during fine tuning","metadata":{"execution":{"iopub.status.busy":"2024-04-18T12:45:43.431185Z","iopub.execute_input":"2024-04-18T12:45:43.431567Z","iopub.status.idle":"2024-04-18T12:45:49.144209Z","shell.execute_reply.started":"2024-04-18T12:45:43.431535Z","shell.execute_reply":"2024-04-18T12:45:49.143407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# give the correct model & history paths & \"epochs_ran\" must be matching with model_{no.}\n\nhistory = train(model, loss_function=loss_function, optimizer=optimizer, num_epochs=90,\n          ranOnce=True, model_path='/kaggle/working/model_10.pth', history_path='/kaggle/working/history.csv',\n         epochs_ran=10, save_interval=10)","metadata":{"execution":{"iopub.status.busy":"2024-03-30T18:11:47.619291Z","iopub.execute_input":"2024-03-30T18:11:47.620184Z","iopub.status.idle":"2024-03-30T18:30:51.008322Z","shell.execute_reply.started":"2024-03-30T18:11:47.620143Z","shell.execute_reply":"2024-03-30T18:30:51.007331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing Phase","metadata":{"papermill":{"duration":4.183496,"end_time":"2024-03-26T23:25:05.680441","exception":false,"start_time":"2024-03-26T23:25:01.496945","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# set it to '/kaggle/working/history.csv' while running by SAVE and COMMIT\nhistory_path = '/kaggle/input/weights-coatnet-vinbigdata-ml-balanced/history.csv'\n# history_path = '/kaggle/input/20-03-2024-vinbigdata-multilabel-weights/history.csv'\nhistory = pd.read_csv(history_path)","metadata":{"papermill":{"duration":4.336153,"end_time":"2024-03-26T23:25:14.282245","exception":false,"start_time":"2024-03-26T23:25:09.946092","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:50.181311Z","iopub.execute_input":"2024-04-18T12:45:50.181896Z","iopub.status.idle":"2024-04-18T12:45:50.19251Z","shell.execute_reply.started":"2024-04-18T12:45:50.181866Z","shell.execute_reply":"2024-04-18T12:45:50.191535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n\n# Create subplots\nfig, axes = plt.subplots(3, 2, figsize=(15, 15))\n\n# Plot accuracy\naxes[0, 0].plot(history['epoch_number'], history['val_accuracy'], label='Validation Accuracy', color='blue')\naxes[0, 0].plot(history['epoch_number'], history['train_accuracy'], label='Training Accuracy', color='orange')\naxes[0, 0].set_xlabel('Epoch')\naxes[0, 0].set_ylabel('Accuracy')\naxes[0, 0].set_title('Validation Accuracy vs Training Accuracy')\naxes[0, 0].legend()\n\n# Plot AUC\naxes[0, 1].plot(history['epoch_number'], history['val_auc'], label='Validation AUC', color='green')\naxes[0, 1].plot(history['epoch_number'], history['train_auc'], label='Training AUC', color='red')\naxes[0, 1].set_xlabel('Epoch')\naxes[0, 1].set_ylabel('AUC')\naxes[0, 1].set_title('Validation AUC vs Training AUC')\naxes[0, 1].legend()\n\n# Plot loss\naxes[1, 0].plot(history['epoch_number'], history['val_loss'], label='Validation Loss', color='purple')\naxes[1, 0].plot(history['epoch_number'], history['train_loss'], label='Training Loss', color='brown')\naxes[1, 0].set_xlabel('Epoch')\naxes[1, 0].set_ylabel('Loss')\naxes[1, 0].set_title('Validation Loss vs Training Loss')\naxes[1, 0].legend()\n\n# Plot precision\naxes[1, 1].plot(history['epoch_number'], history['val_precision'], label='Validation Precision', color='cyan')\naxes[1, 1].plot(history['epoch_number'], history['train_precision'], label='Training Precision', color='magenta')\naxes[1, 1].set_xlabel('Epoch')\naxes[1, 1].set_ylabel('Precision')\naxes[1, 1].set_title('Validation Precision vs Training Precision')\naxes[1, 1].legend()\n\n# Plot recall\naxes[2, 0].plot(history['epoch_number'], history['val_recall'], label='Validation Recall', color='yellow')\naxes[2, 0].plot(history['epoch_number'], history['train_recall'], label='Training Recall', color='green')\naxes[2, 0].set_xlabel('Epoch')\naxes[2, 0].set_ylabel('Recall')\naxes[2, 0].set_title('Validation Recall vs Training Recall')\naxes[2, 0].legend()\n\n# Plot F1 score\naxes[2, 1].plot(history['epoch_number'], history['val_f1'], label='Validation F1 Score', color='orange')\naxes[2, 1].plot(history['epoch_number'], history['train_f1'], label='Training F1 Score', color='blue')\naxes[2, 1].set_xlabel('Epoch')\naxes[2, 1].set_ylabel('F1 Score')\naxes[2, 1].set_title('Validation F1 Score vs Training F1 Score')\naxes[2, 1].legend()\n\n# Adjust layout\nplt.tight_layout()\n\n# Show the plot\nplt.show()","metadata":{"papermill":{"duration":5.948189,"end_time":"2024-03-26T23:25:24.548768","exception":false,"start_time":"2024-03-26T23:25:18.600579","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:50.826231Z","iopub.execute_input":"2024-04-18T12:45:50.826608Z","iopub.status.idle":"2024-04-18T12:45:52.44608Z","shell.execute_reply.started":"2024-04-18T12:45:50.826577Z","shell.execute_reply":"2024-04-18T12:45:52.444845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = '/kaggle/input/weights-coatnet-vinbigdata-ml-balanced/best_model_precision.pth'\n# model_path = '/kaggle/input/20-03-2024-vinbigdata-multilabel-weights/best_model_auc.pth'\nmodel.load_state_dict(torch.load(model_path))","metadata":{"papermill":{"duration":4.665444,"end_time":"2024-03-26T23:25:33.446574","exception":false,"start_time":"2024-03-26T23:25:28.78113","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:45:52.447993Z","iopub.execute_input":"2024-04-18T12:45:52.448626Z","iopub.status.idle":"2024-04-18T12:46:03.581147Z","shell.execute_reply.started":"2024-04-18T12:45:52.448592Z","shell.execute_reply":"2024-04-18T12:46:03.580153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_scores_all = []\ntest_labels_all = []\ntest_predictions_all = []\n\ncorrect_test = 0\ntotal_test = 0\n\n\n# Don't need to keep track of gradients\nwith torch.no_grad():\n    # Set to evaluation mode\n    model.eval()\n\n    start = timer()\n    #Test loop\n    for batch_data in tqdm(test_loader):\n        inputs, labels = batch_data[0].to(device), batch_data[1].float().to(device)  \n\n        outputs = model(inputs)\n        outputs = outputs.float()\n        labels = labels.float()\n\n\n        scores = torch.sigmoid(outputs)\n        predictions = torch.sigmoid(outputs) > 0.5\n        total_test += labels.size(0) * labels.size(1)\n\n        correct_test += (predictions == labels).sum().item()\n\n        test_scores_all.extend(scores.detach().cpu().numpy())\n        test_labels_all.extend(labels.cpu().numpy())\n        test_predictions_all.extend(predictions.cpu().numpy())\n\n    accuracy_test = correct_test / total_test\n\n    print(f\"Accuracy: {accuracy_test:.4f}\")\n    print(\"correct:\", correct_test, \" out of \", total_test)\n\n    test_predictions_all = np.array(test_predictions_all).astype(float)\n    test_labels_all = np.array(test_labels_all).astype(float)\n\n    print(classification_report(test_labels_all, test_predictions_all, target_names=diseases)) # adjust the target_names\n\n    cm = multilabel_confusion_matrix(test_predictions_all, test_labels_all)\n    # we will use macro-averaging strategy.\n    accuracy_arr = []\n    precision_arr = []\n    recall_arr = []\n    f1_arr = []\n\n    for i in range(num_classes):\n        print(cm[i])\n        print(cm[i].sum())\n\n#         cfm_plot = sn.heatmap(cm[i], annot=False)\n         # TP + TN / TP + TN  + FP + FN\n        accuracy = (cm[i][0][0]+ cm[i][1][1])/cm[i].sum()\n\n        # TP / TP + FP\n        precision = cm[i][1][1]/(cm[i][1][0]+cm[i][1][1])\n\n        # TP / TP + FN\n        recall = cm[i][1][1]/(cm[i][0][1]+cm[i][1][1]) # sensitivity\n        f1 = (2*precision*recall)/(precision+recall)\n        print(disease_labels[i],\": \",round(accuracy*100,2),\"%\") ### disease or disease_labels\n        print(\"Precision: \",round(precision,2))\n        print(\"Recall:\", round(recall,2))\n        print(\"F1-Score:\", round(f1,2))\n        print('==========================================================')\n\n\n        accuracy_arr.append(accuracy)\n        precision_arr.append(precision)\n        recall_arr.append(recall)\n        f1_arr.append(f1)\n    \n    accuracy_arr = np.nan_to_num(accuracy_arr, nan=0)\n    precision_arr = np.nan_to_num(precision_arr, nan=0)\n    recall_arr = np.nan_to_num(recall_arr, nan=0)\n    f1_arr = np.nan_to_num(f1_arr, nan=0)\n            \n    accuracy_arr = np.nan_to_num(accuracy_arr, nan=0)\n    precision_arr = np.nan_to_num(precision_arr, nan=0)\n    recall_arr = np.nan_to_num(recall_arr, nan=0)\n    f1_arr = np.nan_to_num(f1_arr, nan=0)\n    \n    accuracy_macro_test = round(sum(accuracy_arr) / len(accuracy_arr), 4)\n    precision_macro_test = round(sum(precision_arr) / len(precision_arr), 4)\n    recall_macro_test = round(sum(recall_arr) / len(recall_arr), 4)\n    f1_macro_test = round(sum(f1_arr) / len(f1_arr), 4)\n    roc_auc_macro_test = round(roc_auc_score(test_labels_all, test_scores_all, average='macro'), 4)\n\n    print(\"MACRO-averged metrics\", end=':- ')\n    print(f\"accuracy: {accuracy_macro_test}, precision: {precision_macro_test}\", end=', ')\n    print(f\"recall: {recall_macro_test}, f1: {f1_macro_test}, ROC_AUC: {roc_auc_macro_test}\")\n\n    infer_time = timer()-start\n    print(f\"INFERENCE TIME: {(infer_time):.4f} seconds\")\n    print(\"-\"*120)","metadata":{"papermill":{"duration":32.202575,"end_time":"2024-03-26T23:26:09.900952","exception":false,"start_time":"2024-03-26T23:25:37.698377","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-18T12:46:03.582826Z","iopub.execute_input":"2024-04-18T12:46:03.583142Z","iopub.status.idle":"2024-04-18T12:46:16.819516Z","shell.execute_reply.started":"2024-04-18T12:46:03.583116Z","shell.execute_reply":"2024-04-18T12:46:16.818227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":4.224946,"end_time":"2024-03-26T23:26:18.511279","exception":false,"start_time":"2024-03-26T23:26:14.286333","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":4.29589,"end_time":"2024-03-26T23:26:26.991716","exception":false,"start_time":"2024-03-26T23:26:22.695826","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":4.188715,"end_time":"2024-03-26T23:26:35.381248","exception":false,"start_time":"2024-03-26T23:26:31.192533","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}