{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":1195048,"sourceType":"datasetVersion","datasetId":680469},{"sourceId":3951115,"sourceType":"datasetVersion","datasetId":1027206}],"dockerImageVersionId":30528,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport os\nimport sys\nimport json\nimport math\nimport random\nimport cv2\ntimm_path = \"../input/timm-pytorch-image-models/pytorch-image-models-master\"\nsys.path.append(timm_path)\nimport timm\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom torch.cuda.amp import autocast\nimport seaborn as sns\nfrom sklearn import model_selection\nfrom sklearn.metrics import mean_squared_error\nfrom tqdm.notebook import tqdm\nimport random\nimport glob\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset,DataLoader\nfrom torch import optim\nfrom torchvision import transforms\nfrom transformers import  get_cosine_schedule_with_warmup\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"_uuid":"159303a1-fb2a-4ad1-a63c-971e5b737131","_cell_guid":"b063e629-f50a-4c6b-b51f-4a2fcc538140","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.272052Z","iopub.execute_input":"2024-07-03T16:03:14.272650Z","iopub.status.idle":"2024-07-03T16:03:14.284768Z","shell.execute_reply.started":"2024-07-03T16:03:14.272602Z","shell.execute_reply":"2024-07-03T16:03:14.283641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\nset_seed(42)","metadata":{"_uuid":"cf3582c0-0996-4c18-9f19-5d4b40370be7","_cell_guid":"e5878bd3-bcf7-413d-8198-7f71b4f0200e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.293511Z","iopub.execute_input":"2024-07-03T16:03:14.294205Z","iopub.status.idle":"2024-07-03T16:03:14.302276Z","shell.execute_reply.started":"2024-07-03T16:03:14.294164Z","shell.execute_reply":"2024-07-03T16:03:14.301153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/siim-isic-melanoma-classification/train.csv')","metadata":{"_uuid":"08447492-158d-407c-a54d-7aff8147c3b6","_cell_guid":"3a1c48a2-85c7-46bb-a465-7464d26c6481","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.315436Z","iopub.execute_input":"2024-07-03T16:03:14.315724Z","iopub.status.idle":"2024-07-03T16:03:14.389985Z","shell.execute_reply.started":"2024-07-03T16:03:14.315698Z","shell.execute_reply":"2024-07-03T16:03:14.388968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.isnull().sum()","metadata":{"_uuid":"3bbf8157-5b50-493b-b0e8-30b7a31307aa","_cell_guid":"b407c6dc-8637-465b-a7e6-0d976c00399b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.391684Z","iopub.execute_input":"2024-07-03T16:03:14.392001Z","iopub.status.idle":"2024-07-03T16:03:14.463788Z","shell.execute_reply.started":"2024-07-03T16:03:14.391973Z","shell.execute_reply":"2024-07-03T16:03:14.462852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['sex'] = df['sex'].map({'male': 1, 'female': 0})\ndf['sex'] = df['sex'].fillna(-1)\n\n\ndf['age_approx'] /= 90\ndf['age_approx'] = df['age_approx'].fillna(0)\n\n\ndf['n_images'] = df.patient_id.map(df.groupby(['patient_id']).image_name.count())\ndf.loc[df['patient_id'] == -1, 'n_images'] = 1","metadata":{"_uuid":"f187e19b-4173-4881-8ea9-14f75bdbb026","_cell_guid":"517a558f-08dc-4fcd-95c0-6337f4532dcc","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.465126Z","iopub.execute_input":"2024-07-03T16:03:14.465508Z","iopub.status.idle":"2024-07-03T16:03:14.499751Z","shell.execute_reply.started":"2024-07-03T16:03:14.465474Z","shell.execute_reply":"2024-07-03T16:03:14.498963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['path'] = [f\"/kaggle/input/siic-isic-224x224-images/train/{x}.png\" for x in df[\"image_name\"].values]\ndense_features = [\n    'sex', 'age_approx', 'n_images'\n]","metadata":{"_uuid":"9c366400-0b29-4f7d-a6d8-2eeb7255f6e5","_cell_guid":"fb3897d7-3c74-4193-8845-bf0179f76af4","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.502179Z","iopub.execute_input":"2024-07-03T16:03:14.502711Z","iopub.status.idle":"2024-07-03T16:03:14.517746Z","shell.execute_reply.started":"2024-07-03T16:03:14.502675Z","shell.execute_reply":"2024-07-03T16:03:14.516816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"strat_kfold = model_selection.StratifiedGroupKFold(n_splits=5, random_state=42, shuffle=True)\n\n# Create an empty 'fold' column in the DataFrame\ndf['fold'] = -1  # Initialize with a default value\n\nfor i, (_, train_index) in enumerate(strat_kfold.split(df.index, y=df['target'], groups=df['patient_id'])):\n    df.loc[df.index[train_index], 'fold'] = i\n\ndf['fold'] = df['fold'].astype('int')","metadata":{"_uuid":"45809ab1-0190-4eef-a9bd-5dd06a65617a","_cell_guid":"f10b6296-5acd-4961-83b8-5071c99fce42","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:14.518809Z","iopub.execute_input":"2024-07-03T16:03:14.519103Z","iopub.status.idle":"2024-07-03T16:03:15.576450Z","shell.execute_reply.started":"2024-07-03T16:03:14.519078Z","shell.execute_reply":"2024-07-03T16:03:15.575622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"_uuid":"5d047a4b-a94b-4c2e-b760-2e3e6973059b","_cell_guid":"c7bb0f97-c14e-4d94-987c-be3060b3be84","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.577615Z","iopub.execute_input":"2024-07-03T16:03:15.577916Z","iopub.status.idle":"2024-07-03T16:03:15.595531Z","shell.execute_reply.started":"2024-07-03T16:03:15.577890Z","shell.execute_reply":"2024-07-03T16:03:15.594582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.sex.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-07-03T16:03:15.596583Z","iopub.execute_input":"2024-07-03T16:03:15.596855Z","iopub.status.idle":"2024-07-03T16:03:15.605417Z","shell.execute_reply.started":"2024-07-03T16:03:15.596831Z","shell.execute_reply":"2024-07-03T16:03:15.604260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.n_images = df.n_images/df.n_images.max()","metadata":{"execution":{"iopub.status.busy":"2024-07-03T16:03:15.606607Z","iopub.execute_input":"2024-07-03T16:03:15.606945Z","iopub.status.idle":"2024-07-03T16:03:15.613354Z","shell.execute_reply.started":"2024-07-03T16:03:15.606918Z","shell.execute_reply":"2024-07-03T16:03:15.612137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_size = 224\ntrain_aug = A.Compose(\n    [   A.RandomResizedCrop(image_size,image_size,p= 0.8),\n        A.Resize(image_size,image_size,p=1.0),\n        A.HorizontalFlip(p=0.5),   \n        A.RandomBrightnessContrast(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=30, p=0.5),\n          \n   A.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n        ToTensorV2()\n    ]\n)\nval_aug = A.Compose(\n    [ \n     A.Resize(image_size,image_size,p=1.0),\n        A.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD),\n        ToTensorV2()\n    ]\n)","metadata":{"_uuid":"371b8c63-e79e-46d3-8350-cf10596f3c08","_cell_guid":"61fa726f-690c-40d5-8434-25d20a37ae60","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.618133Z","iopub.execute_input":"2024-07-03T16:03:15.618519Z","iopub.status.idle":"2024-07-03T16:03:15.626331Z","shell.execute_reply.started":"2024-07-03T16:03:15.618491Z","shell.execute_reply":"2024-07-03T16:03:15.625284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Op4bio_data(Dataset):\n    def __init__(self,df, augs):\n        self.df = df\n        self.augs = augs\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        \n        img_src = self.df.loc[idx,'path']\n        image = cv2.imread(img_src)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        transformed = self.augs(image=image)\n        image = transformed['image']\n        \n        meta = self.df[dense_features].iloc[idx, :].values\n        \n        label  = self.df['target'][idx]\n        \n        return image , torch.FloatTensor(meta) , label","metadata":{"_uuid":"586582cf-ad04-4a7d-8836-20284ce1de92","_cell_guid":"6a4f4db2-6cbd-4f5b-891f-2f516d60cb28","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.627777Z","iopub.execute_input":"2024-07-03T16:03:15.628187Z","iopub.status.idle":"2024-07-03T16:03:15.637879Z","shell.execute_reply.started":"2024-07-03T16:03:15.628125Z","shell.execute_reply":"2024-07-03T16:03:15.636862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t_data = Op4bio_data(df, augs = train_aug)","metadata":{"execution":{"iopub.status.busy":"2024-07-03T16:03:15.639022Z","iopub.execute_input":"2024-07-03T16:03:15.639408Z","iopub.status.idle":"2024-07-03T16:03:15.644379Z","shell.execute_reply.started":"2024-07-03T16:03:15.639374Z","shell.execute_reply":"2024-07-03T16:03:15.643388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t_data[11]","metadata":{"execution":{"iopub.status.busy":"2024-07-03T16:03:15.645456Z","iopub.execute_input":"2024-07-03T16:03:15.645766Z","iopub.status.idle":"2024-07-03T16:03:15.669342Z","shell.execute_reply.started":"2024-07-03T16:03:15.645729Z","shell.execute_reply":"2024-07-03T16:03:15.668403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 定义ViT模型\nclass ViTModel(nn.Module):\n    def __init__(self, pretrained=True):\n        super(ViTModel, self).__init__()\n        self.backbone = timm.create_model('deit_tiny_patch16_224', pretrained=pretrained, num_classes=0)\n        self.pool = nn.AdaptiveAvgPool1d(1)\n        \n        self.fc1 = torch.nn.Linear(192, 768)  # 添加一个线性层将 192 维映射到 768 维\n        self.fc2 = torch.nn.Linear(768, 64)   # 原来的 64 维映射\n        self.norm = torch.nn.BatchNorm1d(64)\n\n    def forward(self, image):\n        image = self.backbone(image)  # 提取特征\n        image = self.fc1(image)       # 映射到 768 维\n        image = self.fc2(image)       # 映射到 64 维\n        image = self.norm(image)      # 批归一化\n        return image","metadata":{"execution":{"iopub.status.busy":"2024-07-03T16:03:15.670533Z","iopub.execute_input":"2024-07-03T16:03:15.670814Z","iopub.status.idle":"2024-07-03T16:03:15.678981Z","shell.execute_reply.started":"2024-07-03T16:03:15.670789Z","shell.execute_reply":"2024-07-03T16:03:15.677935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" \nclass GTNModel(nn.Module):\n    def __init__(self, input_dim, hidden_dim, output_dim, num_heads=4, num_layers=2):\n        super(GTNModel, self).__init__()\n        # 用一个线性层将 3 维数据投影到 64 维\n        self.input_projection = nn.Linear(3, input_dim)\n        self.encoder_layer = nn.TransformerEncoderLayer(d_model=input_dim, nhead=num_heads)\n        self.encoder = nn.TransformerEncoder(self.encoder_layer, num_layers=num_layers)\n        self.fc = nn.Linear(input_dim, hidden_dim)\n        self.output_layer = nn.Linear(hidden_dim, output_dim)\n\n    def forward(self, x):\n        # 先将输入数据投影到 64 维\n        x = self.input_projection(x)\n        x = self.encoder(x)\n        x = self.fc(x)\n        x = self.output_layer(x)\n        return x\n    \n    \n\n# 定义CLIP模型\nclass CLIPModel(nn.Module):\n    def __init__(self, model1, model2):\n        super(CLIPModel, self).__init__()\n        self.model1 = model1  # ViT\n        self.model2 = model2  # GTN\n        self.t = nn.Parameter(torch.Tensor([1]))  # 可学习的参数\n\n    def forward(self, image, graph_data):\n        with autocast():  # 启用自动混合精度\n            image_features = self.model1(image)\n            graph_features = self.model2(graph_data)\n\n            image_features = F.normalize(image_features, p=2, dim=1)\n            graph_features = F.normalize(graph_features, p=2, dim=1)\n            \n            image_features = image_features.to(graph_features.dtype)\n\n            dot_product = torch.matmul(image_features, graph_features.t())\n            dot_product_weighted = torch.exp(self.t) * dot_product\n\n            return dot_product_weighted","metadata":{"_uuid":"e2a54b70-5b08-4863-838f-9638706b0bfe","_cell_guid":"b6af3f44-1baa-4e32-8b87-a2234f2c0939","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.680347Z","iopub.execute_input":"2024-07-03T16:03:15.680751Z","iopub.status.idle":"2024-07-03T16:03:15.693030Z","shell.execute_reply.started":"2024-07-03T16:03:15.680705Z","shell.execute_reply":"2024-07-03T16:03:15.692273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"_uuid":"b23c516d-0f98-481b-9055-41aea58a9031","_cell_guid":"27ff72c8-4c14-49be-9b2e-e4de36862798","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.694473Z","iopub.execute_input":"2024-07-03T16:03:15.694789Z","iopub.status.idle":"2024-07-03T16:03:15.704863Z","shell.execute_reply.started":"2024-07-03T16:03:15.694763Z","shell.execute_reply":"2024-07-03T16:03:15.703820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(train_loader,model,optimizer,criterion,e,epochs,device):\n    '''Trains the model for a single epoch and returns Loss,Accuracy, AUC for that epoch'''\n\n    losses = AverageMeter()\n    model.train()\n    global_step = 0\n    loop = tqdm(enumerate(train_loader),total = len(train_loader))\n    \n    for step,(image,tabular,l) in loop:\n        image = image.to(device)\n        tabular = tabular.to(device)\n        \n        batch_size = l.size(0)\n\n        output = model(image,tabular)\n        \n        #print(output)\n        \n        labels = torch.arange(batch_size) \n        labels = labels.to(device)\n        \n        loss1 = criterion(output, labels)\n        loss2 = criterion(output.T, labels)\n        loss = (loss1 + loss2) / 2.0\n        \n        losses.update(loss.item(), batch_size)\n\n        loss.backward()\n         # Clip gradients to avoid gradient explosion\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)  # Adjust max_norm value as needed\n        \n        \n        optimizer.step()\n        optimizer.zero_grad()\n\n        global_step += 1\n\n        loop.set_description(f\"Epoch {e+1}/{epochs}\")\n        loop.set_postfix(model_loss = loss.item() ,stage = 'train')        \n\n    return losses.avg","metadata":{"_uuid":"ceb7a311-2a5f-4cb5-b710-b6c28eaaa1cf","_cell_guid":"f6ead9da-d31a-424e-a14e-677752dfd77b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.706130Z","iopub.execute_input":"2024-07-03T16:03:15.706463Z","iopub.status.idle":"2024-07-03T16:03:15.717363Z","shell.execute_reply.started":"2024-07-03T16:03:15.706437Z","shell.execute_reply":"2024-07-03T16:03:15.716370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def val_one_epoch(loader,model,criterion,device):\n    '''Validates the model for a single epoch and returns Loss,Accuracy, AUC for that epoch'''\n    losses = AverageMeter()\n    model.eval()\n    global_step = 0\n    loop = tqdm(enumerate(loader),total = len(loader))\n    \n    for step,(image,tabular,l) in loop:\n        image = image.to(device)\n        tabular = tabular.to(device)\n        batch_size = l.size(0)\n        \n        labels = torch.arange(batch_size) \n        labels = labels.to(device)\n        \n        with torch.no_grad():\n            output = model(image,tabular)\n\n        loss1 = criterion(output, labels)\n        loss2 = criterion(output.T, labels)\n        loss = (loss1 + loss2) / 2.0        \n        losses.update(loss.item(), batch_size)\n        loop.set_postfix(model_loss = loss.item() ,stage = 'val')\n        global_step += 1\n\n    return losses.avg","metadata":{"_uuid":"350bd546-2097-46c0-bac7-4b839a96633a","_cell_guid":"6f127cfc-ecae-4de0-9f4d-0f7deb1d110b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.718658Z","iopub.execute_input":"2024-07-03T16:03:15.719014Z","iopub.status.idle":"2024-07-03T16:03:15.728129Z","shell.execute_reply.started":"2024-07-03T16:03:15.718979Z","shell.execute_reply":"2024-07-03T16:03:15.727267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit(t_loader ,v_loader, model, OUTPUT_DIR,device,optimizer):\n    \n    T_LOSS1 = []\n    V_LOSS1 = []\n    model.to(device)\n    #model.to(device)\n    \n    criterion = nn.CrossEntropyLoss() # Loss function\n    optimizer = optimizer\n    \n    epochs = 15\n    loop = range(epochs)\n    for e in loop:\n      \n        loss = train_one_epoch(t_loader,model,optimizer,criterion,e,epochs,device)\n        \n        print(f'For epoch {e+1}/{epochs}')\n        print(f'average model train_loss {loss}')\n        \n        T_LOSS1.append(loss)\n\n        val_loss = val_one_epoch(v_loader,model,criterion,device)\n        \n        print(f'average model val_loss {val_loss}')\n        \n        V_LOSS1.append(val_loss)\n\n        torch.save(model.state_dict(),OUTPUT_DIR+ f'model_val_loss {val_loss}.pth')\n\n    return T_LOSS1,V_LOSS1","metadata":{"_uuid":"e97a342b-2e7f-447a-90f8-87f23684d9e8","_cell_guid":"6f2a2a46-9dcb-4652-b559-22b1bad01f21","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.729286Z","iopub.execute_input":"2024-07-03T16:03:15.729597Z","iopub.status.idle":"2024-07-03T16:03:15.739968Z","shell.execute_reply.started":"2024-07-03T16:03:15.729571Z","shell.execute_reply":"2024-07-03T16:03:15.739192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data= df[df.fold != 0]\nval_data  = df[df.fold == 0]\n    \nt_data = Op4bio_data(train_data.reset_index(drop=True) , augs = train_aug)\nv_data = Op4bio_data(val_data.reset_index(drop=True) , augs = val_aug)\n\n\nt_loader = DataLoader(t_data, shuffle=True,\n                        num_workers=2,\n                        batch_size=16,drop_last =True)\n\nv_loader = DataLoader(v_data, shuffle=False,\n                        num_workers=2,\n                        batch_size=16,drop_last =False)","metadata":{"_uuid":"5d351733-e22e-4066-8336-90ba48fea98c","_cell_guid":"683918f5-50ee-4830-a8fb-29efcb24c08f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.741261Z","iopub.execute_input":"2024-07-03T16:03:15.742257Z","iopub.status.idle":"2024-07-03T16:03:15.759005Z","shell.execute_reply.started":"2024-07-03T16:03:15.742196Z","shell.execute_reply":"2024-07-03T16:03:15.758271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"_uuid":"e3782ec8-9e01-4ea0-8bab-bf8c6ac124ff","_cell_guid":"919fdf4b-3e76-4a0c-a743-47b75c4c2649","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.760122Z","iopub.execute_input":"2024-07-03T16:03:15.760421Z","iopub.status.idle":"2024-07-03T16:03:15.764859Z","shell.execute_reply.started":"2024-07-03T16:03:15.760396Z","shell.execute_reply":"2024-07-03T16:03:15.763872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model1 = ViTModel(pretrained=True)  # ViT模型\nmodel2 = GTNModel(input_dim=64, hidden_dim=64, output_dim=64, num_heads=4, num_layers=2)","metadata":{"_uuid":"8692c7b3-486d-4132-8130-f854454a8065","_cell_guid":"c6d42433-5f59-4a16-a8a8-3b75b44a3ba2","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:15.766035Z","iopub.execute_input":"2024-07-03T16:03:15.766591Z","iopub.status.idle":"2024-07-03T16:03:16.071927Z","shell.execute_reply.started":"2024-07-03T16:03:15.766555Z","shell.execute_reply":"2024-07-03T16:03:16.071059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"combined_model = CLIPModel(model1, model2)\noptimizer = optim.AdamW(combined_model.parameters(), lr=1e-6 , weight_decay = 1e-5 ) \nT_LOSS1, V_LOSS1= fit(t_loader ,v_loader, combined_model, OUTPUT_DIR,device,optimizer )","metadata":{"_uuid":"e7003288-813c-499f-af29-fc890505773f","_cell_guid":"7050e2ab-4583-4cf8-aac9-db22a96bff84","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:03:16.076111Z","iopub.execute_input":"2024-07-03T16:03:16.076420Z","iopub.status.idle":"2024-07-03T16:37:11.979182Z","shell.execute_reply.started":"2024-07-03T16:03:16.076394Z","shell.execute_reply":"2024-07-03T16:37:11.977846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Generate x-axis values\nepochs = len(T_LOSS1)\nx = list(range(1, epochs + 1))\n\n# Plot training loss\nplt.plot(x, T_LOSS1, label='Training Loss')\n\n# Plot validation loss\nplt.plot(x, V_LOSS1, label='Validation Loss')\n\n# Set plot labels and title\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Training and Validation Loss')\nplt.legend()\n\n# Show the plot\nplt.show()","metadata":{"_uuid":"cb74c080-fee6-4a9e-943f-afabd1d43e0c","_cell_guid":"41ebf9dd-8e4d-4bb4-a147-0852c8740488","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-07-03T16:37:11.981026Z","iopub.execute_input":"2024-07-03T16:37:11.981439Z","iopub.status.idle":"2024-07-03T16:37:12.319310Z","shell.execute_reply.started":"2024-07-03T16:37:11.981392Z","shell.execute_reply":"2024-07-03T16:37:12.318225Z"},"trusted":true},"execution_count":null,"outputs":[]}]}