{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":1087.035813,"end_time":"2023-01-22T10:17:16.597092","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-01-22T09:59:09.561279","version":"2.3.4"},"colab":{"provenance":[]},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4619805,"sourceType":"datasetVersion","datasetId":2688675}],"dockerImageVersionId":30616,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n\n**This project was completed as part of our coursework at SPIT, where we focused on developing a platform for cancer detection using advanced machine learning techniques. Our goal was to leverage the power of AI to assist in early detection and diagnosis of cancer, ultimately improving patient outcomes.**\n\n","metadata":{"papermill":{"duration":0.010044,"end_time":"2023-01-22T09:59:18.486192","exception":false,"start_time":"2023-01-22T09:59:18.476148","status":"completed"},"tags":[],"id":"07810ece"}},{"cell_type":"markdown","source":"## importing and installing stuff","metadata":{"id":"Y3iW0S4HfaSb"}},{"cell_type":"code","source":"!pip install pydicom -q","metadata":{"id":"5nMAz_AQt3N0","outputId":"5a367117-0401-41e1-9838-07b1a238ccec","execution":{"iopub.status.busy":"2024-04-12T18:43:20.242683Z","iopub.execute_input":"2024-04-12T18:43:20.243431Z","iopub.status.idle":"2024-04-12T18:43:33.56614Z","shell.execute_reply.started":"2024-04-12T18:43:20.243397Z","shell.execute_reply":"2024-04-12T18:43:33.564915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install numpy==1.22.0\n","metadata":{"execution":{"iopub.status.busy":"2024-04-12T18:43:33.568527Z","iopub.execute_input":"2024-04-12T18:43:33.568864Z","iopub.status.idle":"2024-04-12T18:43:50.658876Z","shell.execute_reply.started":"2024-04-12T18:43:33.568837Z","shell.execute_reply":"2024-04-12T18:43:50.657629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pydicom\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom ignite.engine import Events, create_supervised_trainer, create_supervised_evaluator\nfrom ignite.metrics import Accuracy, Loss, RunningAverage\nfrom ignite.contrib.handlers import ProgressBar\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import models, transforms\n","metadata":{"papermill":{"duration":3.715291,"end_time":"2023-01-22T09:59:22.251424","exception":false,"start_time":"2023-01-22T09:59:18.536133","status":"completed"},"tags":[],"id":"d51b515d","execution":{"iopub.status.busy":"2024-04-12T18:43:50.660642Z","iopub.execute_input":"2024-04-12T18:43:50.661041Z","iopub.status.idle":"2024-04-12T18:43:56.617376Z","shell.execute_reply.started":"2024-04-12T18:43:50.66101Z","shell.execute_reply":"2024-04-12T18:43:56.616369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n","metadata":{"papermill":{"duration":0.005588,"end_time":"2023-01-22T09:59:18.497792","exception":false,"start_time":"2023-01-22T09:59:18.492204","status":"completed"},"tags":[],"id":"7cdc0644"}},{"cell_type":"markdown","source":"","metadata":{"papermill":{"duration":0.005403,"end_time":"2023-01-22T09:59:18.50875","exception":false,"start_time":"2023-01-22T09:59:18.503347","status":"completed"},"tags":[],"id":"0aed2f4a"}},{"cell_type":"markdown","source":"## utility functions\nWe write some utility functions beforehand. To know how to work with PyDicom funtions, read [this](https://towardsdatascience.com/introducing-pydicom-its-classes-methods-and-attributes-518c1d71162).","metadata":{"papermill":{"duration":0.005416,"end_time":"2023-01-22T09:59:18.530519","exception":false,"start_time":"2023-01-22T09:59:18.525103","status":"completed"},"tags":[],"id":"d95bbc52"}},{"cell_type":"code","source":"def read_xray(file_path, img_size=None):\n    \"\"\"\n    Read the dicom data and get the image\n    Args:\n        file_path: The path of the dicom file\n        img_size: Size of the output image\n    \"\"\"\n\n    dicom = pydicom.read_file(file_path)\n    img = dicom.pixel_array\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = np.max(img) - img\n\n    if img_size:\n        img = cv2.resize(img, img_size)\n\n    # Add channel dim at First\n    img = img[np.newaxis]\n\n    # Converting img to float32\n    img = img / np.max(img)\n    img = img.astype(\"float32\")\n\n    return img","metadata":{"papermill":{"duration":0.015247,"end_time":"2023-01-22T09:59:22.272633","exception":false,"start_time":"2023-01-22T09:59:22.257386","status":"completed"},"tags":[],"id":"221fb680","execution":{"iopub.status.busy":"2024-04-12T18:43:56.619074Z","iopub.execute_input":"2024-04-12T18:43:56.619705Z","iopub.status.idle":"2024-04-12T18:43:56.625735Z","shell.execute_reply.started":"2024-04-12T18:43:56.619676Z","shell.execute_reply":"2024-04-12T18:43:56.624742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def patchify(batch, patch_size):\n    \"\"\"\n    Patchify the batch of images\n\n    Shape:\n        batch: (b, h, w, c)\n        output: (b, nh, nw, ph, pw, c)\n    \"\"\"\n    b, c, h, w = batch.shape\n    ph, pw = patch_size\n    nh, nw = h // ph, w // pw\n\n    batch_patches = torch.reshape(batch, (b, c, nh, ph, nw, pw))\n    batch_patches = torch.permute(batch_patches, (0, 1, 2, 4, 3, 5))\n\n    return batch_patches","metadata":{"papermill":{"duration":0.015657,"end_time":"2023-01-22T09:59:22.293761","exception":false,"start_time":"2023-01-22T09:59:22.278104","status":"completed"},"tags":[],"id":"d80572d3","execution":{"iopub.status.busy":"2024-04-12T18:43:56.628265Z","iopub.execute_input":"2024-04-12T18:43:56.628561Z","iopub.status.idle":"2024-04-12T18:43:56.635502Z","shell.execute_reply.started":"2024-04-12T18:43:56.628512Z","shell.execute_reply":"2024-04-12T18:43:56.634602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We test our `patchify` function on a single image.","metadata":{"papermill":{"duration":0.00537,"end_time":"2023-01-22T09:59:22.30469","exception":false,"start_time":"2023-01-22T09:59:22.29932","status":"completed"},"tags":[],"id":"645afc03"}},{"cell_type":"code","source":"FILE_PATH = ('/kaggle/input/rsna-breast-cancer-detection/'\n             'train_images/10006/1459541791.dcm')\n\nimg = read_xray(FILE_PATH, img_size=(512, 512))\n\nbatch = torch.tensor(img[None])\npatch_size = (16, 16)\nbatch_patches = patchify(batch, patch_size)\n\npatches = batch_patches[0]\nc, nh, nw, ph, pw = patches.shape\n\nplt.figure(figsize=(5, 5))\nplt.imshow(img[0], cmap=\"gray\")\nplt.axis(\"off\")\n\nplt.figure(figsize=(5, 5))\nfor i in range(nh):\n    for j in range(nw):\n        plt.subplot(nh, nw, i * nw + j + 1)\n        plt.imshow(patches[0, i, j], cmap=\"gray\")\n        plt.axis(\"off\")","metadata":{"papermill":{"duration":48.930892,"end_time":"2023-01-22T10:00:11.241361","exception":false,"start_time":"2023-01-22T09:59:22.310469","status":"completed"},"tags":[],"id":"522d21b9","outputId":"92211244-22b1-4e50-fe3d-c0f167d522e8","execution":{"iopub.status.busy":"2024-04-12T18:43:56.636814Z","iopub.execute_input":"2024-04-12T18:43:56.637603Z","iopub.status.idle":"2024-04-12T18:44:36.736933Z","shell.execute_reply.started":"2024-04-12T18:43:56.637572Z","shell.execute_reply":"2024-04-12T18:44:36.736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mlp(in_features, hidden_units, out_features):\n    \"\"\"\n    Returns a MLP head\n    \"\"\"\n    dims = [in_features] + hidden_units + [out_features]\n    layers = []\n    for dim1, dim2 in zip(dims[:-2], dims[1:-1]):\n        layers.append(nn.Linear(dim1, dim2))\n        layers.append(nn.ReLU())\n    layers.append(nn.Linear(dims[-2], dims[-1]))\n    return nn.Sequential(*layers)","metadata":{"papermill":{"duration":0.017061,"end_time":"2023-01-22T10:00:11.2657","exception":false,"start_time":"2023-01-22T10:00:11.248639","status":"completed"},"tags":[],"id":"5b559866","execution":{"iopub.status.busy":"2024-04-12T18:44:36.738141Z","iopub.execute_input":"2024-04-12T18:44:36.738417Z","iopub.status.idle":"2024-04-12T18:44:36.744575Z","shell.execute_reply.started":"2024-04-12T18:44:36.738392Z","shell.execute_reply":"2024-04-12T18:44:36.743749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## image to sequence block\nThis Block takes a batch of image as input and returns a batch of sequences. Later on we feed this sequences into the transformer encoder.","metadata":{"papermill":{"duration":0.006314,"end_time":"2023-01-22T10:00:11.278755","exception":false,"start_time":"2023-01-22T10:00:11.272441","status":"completed"},"tags":[],"id":"6cb491e3"}},{"cell_type":"code","source":"class Img2Seq(nn.Module):\n    \"\"\"\n    This layers takes a batch of images as input and\n    returns a batch of sequences\n\n    Shape:\n        input: (b, h, w, c)\n        output: (b, s, d)\n    \"\"\"\n    def __init__(self, img_size, patch_size, n_channels, d_model):\n        super().__init__()\n        self.patch_size = patch_size\n        self.img_size = img_size\n\n        nh, nw = img_size[0] // patch_size[0], img_size[1] // patch_size[1]\n        n_tokens = nh * nw\n\n        token_dim = patch_size[0] * patch_size[1] * n_channels\n        self.linear = nn.Linear(token_dim, d_model)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, d_model))\n        self.pos_emb = nn.Parameter(torch.randn(n_tokens, d_model))\n\n    def __call__(self, batch):\n        batch = patchify(batch, self.patch_size)\n\n        b, c, nh, nw, ph, pw = batch.shape\n\n        # Flattening the patches\n        batch = torch.permute(batch, [0, 2, 3, 4, 5, 1])\n        batch = torch.reshape(batch, [b, nh * nw, ph * pw * c])\n\n        batch = self.linear(batch)\n        cls = self.cls_token.expand([b, -1, -1])\n        emb = batch + self.pos_emb\n\n        return torch.cat([cls, emb], axis=1)","metadata":{"papermill":{"duration":0.018674,"end_time":"2023-01-22T10:00:11.304195","exception":false,"start_time":"2023-01-22T10:00:11.285521","status":"completed"},"tags":[],"id":"58ae05ba","execution":{"iopub.status.busy":"2024-04-12T18:44:36.745655Z","iopub.execute_input":"2024-04-12T18:44:36.745914Z","iopub.status.idle":"2024-04-12T18:44:36.768486Z","shell.execute_reply.started":"2024-04-12T18:44:36.745892Z","shell.execute_reply":"2024-04-12T18:44:36.767449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## visual transformer module\nThis modules wraps up everything. We can divide this module into 3 parts:\n* An image to sequence encoder\n* Transformer encoder\n* Multilayer perceptron head classification\n\nWe use `torch.nn.TransformerEncoder` and `torch.nn.TransformerEncoderLayer` to implement our transformer encoder. We highly recommend to read the official documentation to [learn more about the layers](https://pytorch.org/docs/stable/generated/torch.nn.TransformerEncoder.html).\nFor the loss function, we will be using [gelu](https://arxiv.org/abs/1606.08415). You can also use this video to [learn about it](https://www.youtube.com/watch?v=FWhMkpo9yuM).","metadata":{"papermill":{"duration":0.006736,"end_time":"2023-01-22T10:00:11.317817","exception":false,"start_time":"2023-01-22T10:00:11.311081","status":"completed"},"tags":[],"id":"26d7647d"}},{"cell_type":"code","source":"class ViT(nn.Module):\n    def __init__(\n        self,\n        img_size,\n        patch_size,\n        n_channels,\n        d_model,\n        nhead,\n        dim_feedforward,\n        blocks,\n        mlp_head_units,\n        n_classes,\n    ):\n        super().__init__()\n        \"\"\"\n        Args:\n            img_size: Size of the image\n            patch_size: Size of the patch\n            n_channels: Number of image channels\n            d_model: The number of features in the transformer encoder\n            nhead: The number of heads in the multiheadattention models\n            dim_feedforward: The dimension of the feedforward network model in the encoder\n            blocks: The number of sub-encoder-layers in the encoder\n            mlp_head_units: The hidden units of mlp_head\n            n_classes: The number of output classes\n        \"\"\"\n        self.img2seq = Img2Seq(img_size, patch_size, n_channels, d_model)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model, nhead, dim_feedforward, activation=\"gelu\", batch_first=True\n        )\n        self.transformer_encoder = nn.TransformerEncoder(\n            encoder_layer, blocks\n        )\n        self.mlp = get_mlp(d_model, mlp_head_units, n_classes)\n\n        self.output = nn.Sigmoid() if n_classes == 1 else nn.Softmax()\n\n    def __call__(self, batch):\n\n        batch = self.img2seq(batch)\n        batch = self.transformer_encoder(batch)\n        batch = batch[:, 0, :]\n        batch = self.mlp(batch)\n        output = self.output(batch)\n        return output","metadata":{"papermill":{"duration":0.018775,"end_time":"2023-01-22T10:00:11.34348","exception":false,"start_time":"2023-01-22T10:00:11.324705","status":"completed"},"tags":[],"id":"00b544d2","execution":{"iopub.status.busy":"2024-04-12T18:44:36.769906Z","iopub.execute_input":"2024-04-12T18:44:36.77035Z","iopub.status.idle":"2024-04-12T18:44:36.783343Z","shell.execute_reply.started":"2024-04-12T18:44:36.770321Z","shell.execute_reply":"2024-04-12T18:44:36.782368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## training\n\nHere, we design a simple training to loop to train or `ViT` model on a subset of dataset. We use an already cropped dataset.","metadata":{"papermill":{"duration":0.006607,"end_time":"2023-01-22T10:00:11.356973","exception":false,"start_time":"2023-01-22T10:00:11.350366","status":"completed"},"tags":[],"id":"e396a0ac"}},{"cell_type":"markdown","source":"## hyperparameters setup","metadata":{"papermill":{"duration":0.007061,"end_time":"2023-01-22T10:00:11.370556","exception":false,"start_time":"2023-01-22T10:00:11.363495","status":"completed"},"tags":[],"id":"829d3a13"}},{"cell_type":"code","source":"img_size = (512, 512)\npatch_size = (16, 16)\nn_channels = 1\nd_model = 1024\nnhead = 4\ndim_feedforward = 2048\nblocks = 8\nmlp_head_units = [1024, 512]\nn_classes = 1\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"papermill":{"duration":0.137093,"end_time":"2023-01-22T10:00:11.514575","exception":false,"start_time":"2023-01-22T10:00:11.377482","status":"completed"},"tags":[],"id":"6a5ac8c3","execution":{"iopub.status.busy":"2024-04-12T18:44:36.784575Z","iopub.execute_input":"2024-04-12T18:44:36.785347Z","iopub.status.idle":"2024-04-12T18:44:36.846689Z","shell.execute_reply.started":"2024-04-12T18:44:36.785322Z","shell.execute_reply":"2024-04-12T18:44:36.845686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset loading\nWe will be using the [CLAHE](https://www.geeksforgeeks.org/clahe-histogram-eqalization-opencv/) algorithm to improve contrast between the tiled images. Apart from this, we are using the pytorch data handling modules like DataLoaded which you can read about [here](https://pytorch.org/docs/stable/data.html).","metadata":{"papermill":{"duration":0.007039,"end_time":"2023-01-22T10:00:11.528747","exception":false,"start_time":"2023-01-22T10:00:11.521708","status":"completed"},"tags":[],"id":"4587c40c"}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n\n    def __init__(self, df, img_path):\n        self.df = df\n        self.img_path = img_path\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        patient_id, image_id, cancer = self.df.iloc[idx][['patient_id', 'image_id', 'cancer']]\n        file = os.path.join(self.img_path, f'{patient_id}_{image_id}.png')\n        file = cv2.imread(file, cv2.COLOR_BGR2GRAY)\n        clahe = cv2.createCLAHE(clipLimit = 15, tileGridSize=[8, 8])\n        file = clahe.apply(file)\n        file = file / file.max()\n        X = torch.tensor(file[np.newaxis].astype('float32')).to(device)\n        y = torch.tensor([cancer]).float().to(device)\n        return X, y","metadata":{"papermill":{"duration":0.018483,"end_time":"2023-01-22T10:00:11.554179","exception":false,"start_time":"2023-01-22T10:00:11.535696","status":"completed"},"tags":[],"id":"8349a5ae","execution":{"iopub.status.busy":"2024-04-12T18:44:36.847881Z","iopub.execute_input":"2024-04-12T18:44:36.848188Z","iopub.status.idle":"2024-04-12T18:44:36.856849Z","shell.execute_reply.started":"2024-04-12T18:44:36.848163Z","shell.execute_reply":"2024-04-12T18:44:36.855946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ncounts = df['cancer'].value_counts()\ndf['weights'] = df['cancer'].apply(lambda x: 1/counts[x])\n\ntrain_df, val_df = train_test_split(df, test_size=0.3, stratify=df['cancer'])","metadata":{"papermill":{"duration":0.432908,"end_time":"2023-01-22T10:00:11.993993","exception":false,"start_time":"2023-01-22T10:00:11.561085","status":"completed"},"tags":[],"id":"c8cdd3ff","execution":{"iopub.status.busy":"2024-04-12T18:44:36.858025Z","iopub.execute_input":"2024-04-12T18:44:36.85832Z","iopub.status.idle":"2024-04-12T18:44:37.346947Z","shell.execute_reply.started":"2024-04-12T18:44:36.858296Z","shell.execute_reply":"2024-04-12T18:44:37.346055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '/kaggle/input/rsna-breast-cancer-512-pngs'\ntrain_samples = 1000\nval_samples = 500\n\ntrain_ds = RSNADataset(train_df, img_path)\nval_ds = RSNADataset(val_df, img_path)\n\ntrain_sampler = WeightedRandomSampler(train_df['weights'].values, train_samples)\ntrain_loader = DataLoader(train_ds, batch_size=8, sampler=train_sampler)\n\nval_sampler = WeightedRandomSampler(val_df['weights'].values, val_samples)\nval_loader = DataLoader(val_ds, batch_size=32, sampler=val_sampler)","metadata":{"papermill":{"duration":0.018775,"end_time":"2023-01-22T10:00:12.019912","exception":false,"start_time":"2023-01-22T10:00:12.001137","status":"completed"},"tags":[],"id":"b2c20004","execution":{"iopub.status.busy":"2024-04-12T18:44:37.348329Z","iopub.execute_input":"2024-04-12T18:44:37.34872Z","iopub.status.idle":"2024-04-12T18:44:37.355836Z","shell.execute_reply.started":"2024-04-12T18:44:37.348683Z","shell.execute_reply":"2024-04-12T18:44:37.354729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model creation and training start","metadata":{"papermill":{"duration":0.007287,"end_time":"2023-01-22T10:00:12.034465","exception":false,"start_time":"2023-01-22T10:00:12.027178","status":"completed"},"tags":[],"id":"87621c6e"}},{"cell_type":"code","source":"model = ViT(\n    img_size = (512, 512),\n    patch_size = (16, 16),\n    n_channels = 1,\n    d_model = 1024,\n    nhead = 4,\n    dim_feedforward = 1024,\n    blocks = 8,\n    mlp_head_units = [512, 512],\n    n_classes = 1,\n).to(device)\n\noptimizer = Adam(model.parameters())\ncriterion = nn.BCELoss()\n\ntrainer = create_supervised_trainer(model, optimizer, criterion, device=device)\nval_metrics = {\n    \"bce\": Loss(criterion)\n}\nevaluator = create_supervised_evaluator(model, metrics=val_metrics, device=device)","metadata":{"papermill":{"duration":3.667453,"end_time":"2023-01-22T10:00:15.709219","exception":false,"start_time":"2023-01-22T10:00:12.041766","status":"completed"},"tags":[],"id":"44c7f2df","execution":{"iopub.status.busy":"2024-04-12T18:44:37.360557Z","iopub.execute_input":"2024-04-12T18:44:37.360844Z","iopub.status.idle":"2024-04-12T18:44:37.767348Z","shell.execute_reply.started":"2024-04-12T18:44:37.36082Z","shell.execute_reply":"2024-04-12T18:44:37.766577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"log_interval = 10\nmax_epochs = 5\nbest_loss = float('inf')\n\nRunningAverage(output_transform=lambda x: x).attach(trainer, 'loss')\n\npbar = ProgressBar()\npbar.attach(trainer, ['loss'])\n\n@trainer.on(Events.EPOCH_COMPLETED)\ndef log_validation_results(trainer):\n    global best_loss\n    evaluator.run(val_loader)\n    loss = evaluator.state.metrics['bce']\n    if loss < best_loss:\n        best_loss = loss\n        torch.save(model.state_dict(), 'best_model_vit.pt')\n    print(f\"Validation Results - Epoch: {trainer.state.epoch} Avg loss: {loss:.2f}\")\n\noutput_state = trainer.run(train_loader, max_epochs=max_epochs)","metadata":{"papermill":{"duration":848.375644,"end_time":"2023-01-22T10:14:24.091936","exception":false,"start_time":"2023-01-22T10:00:15.716292","status":"completed"},"tags":[],"id":"081e32ba","execution":{"iopub.status.busy":"2024-04-12T18:48:55.126094Z","iopub.execute_input":"2024-04-12T18:48:55.126483Z","iopub.status.idle":"2024-04-12T18:59:25.313204Z","shell.execute_reply.started":"2024-04-12T18:48:55.126456Z","shell.execute_reply":"2024-04-12T18:59:25.312046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModelForImageClassification\n\nmodel = AutoModelForImageClassification.from_pretrained('google/vit-base-patch16-224')\nmodel.save_pretrained(\"my_model\")","metadata":{"execution":{"iopub.status.busy":"2024-04-12T19:22:42.568339Z","iopub.execute_input":"2024-04-12T19:22:42.568754Z","iopub.status.idle":"2024-04-12T19:22:43.952976Z","shell.execute_reply.started":"2024-04-12T19:22:42.568722Z","shell.execute_reply":"2024-04-12T19:22:43.952027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import TFAutoModelForImageClassification\n\ntf_model = TFAutoModelForImageClassification.from_pretrained(\"my_model\", from_pt=True)\ntf_model.save_pretrained(\"my_model_tf\")","metadata":{"execution":{"iopub.status.busy":"2024-04-12T19:22:49.667301Z","iopub.execute_input":"2024-04-12T19:22:49.66782Z","iopub.status.idle":"2024-04-12T19:22:51.673051Z","shell.execute_reply.started":"2024-04-12T19:22:49.667787Z","shell.execute_reply":"2024-04-12T19:22:51.672145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n# List files in the directory\nfiles = os.listdir(\"my_model_tf\")\nprint(files)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-12T19:23:14.668384Z","iopub.execute_input":"2024-04-12T19:23:14.669222Z","iopub.status.idle":"2024-04-12T19:23:14.674617Z","shell.execute_reply.started":"2024-04-12T19:23:14.669189Z","shell.execute_reply":"2024-04-12T19:23:14.67362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\n\n# Define the path to the directory containing the .h5 file\ndirectory_path = \"my_model_tf\"\n\n# Define the name of the .h5 file\nfile_name = \"tf_model.h5\"\n\n# Define the full path to the .h5 file\nfile_path = os.path.join(directory_path, file_name)\n\n# Create a download link for the file\nFileLink(file_path)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-12T19:08:58.338565Z","iopub.execute_input":"2024-04-12T19:08:58.339281Z","iopub.status.idle":"2024-04-12T19:08:58.346324Z","shell.execute_reply.started":"2024-04-12T19:08:58.339247Z","shell.execute_reply":"2024-04-12T19:08:58.345382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport tensorflow as tf\nfrom transformers import ViTForImageClassification, ViTConfig\n\n# Load the ViTForImageClassification model\nvit_model = ViTForImageClassification.from_pretrained('google/vit-base-patch16-224')\nvit_config = ViTConfig.from_pretrained('google/vit-base-patch16-224')\n\n# Load your TensorFlow model\nyour_tf_model = tf.keras.models.load_model(\"my_saved_model_tf\")\n\n# Get configurations\nvit_config_dict = vit_config.to_dict()\nyour_tf_config_dict = your_tf_model.get_config()\n\n# Compare configurations\nif vit_config_dict == your_tf_config_dict:\n    print(\"Your TensorFlow model is compatible with ViTForImageClassification.\")\nelse:\n    print(\"Your TensorFlow model is not fully compatible with ViTForImageClassification.\")\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-12T19:33:12.755676Z","iopub.execute_input":"2024-04-12T19:33:12.756108Z","iopub.status.idle":"2024-04-12T19:33:22.57558Z","shell.execute_reply.started":"2024-04-12T19:33:12.756076Z","shell.execute_reply":"2024-04-12T19:33:22.574493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Team Members:-\n\nSPIT Minor Project\n\n- **Nikhil Chaudhari**: \n\n- **Atharva Patil**: \n\n- **Madhura Kanfade**: \n\n\n","metadata":{}}]}