{"metadata":{"kernelspec":{"display_name":"d2l","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.9.18"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA Lumbar Spine Challenge\nAuthors: Golowper","metadata":{}},{"cell_type":"markdown","source":"### Overview\n\n三类疾病\n\n1. Foraminal narrowing (on either the left or right foramen at a specified level).\n2. Subarticular stenosis (on either the left or right side at a specified level).\n3. Canal stenosis (only at a specified level).\n\n我将其命名为如下三种，变量命名规则与此相同：\n\n1. canal : 腰椎管狭窄 \n    - eg: canal_df\n2. foraminal : 神经孔狭窄 \n    - eg: foraminal_df\n3. subarticular : 关节突间狭窄 \n    - eg: subarticular_df\n","metadata":{}},{"cell_type":"markdown","source":"### Aim\n\nFor each of the conditions, you'll need to predict whether the degree of compression is normal/mild, moderate, or severe. \n\nYou can refer to the example test submission `sample_submission.csv` to get a better idea for what we're looking for in terms of output. \n\n 特定等级（正常/轻度，中度，重度） 3种\n脊柱的特定水平 (l1_l2，l2_l3，l3_l4，l4_l5，l5_s1) 5种\n病症（脊柱管狭窄，左侧神经根管狭窄，右侧神经根管狭窄，左侧关节下狭窄，右侧关节下狭窄）5种\n\n#### 对 `row_id = study_id + 病症 + 脊柱特定水平` 的 `特定等级` 给出 0-1 的打分\n","metadata":{}},{"cell_type":"markdown","source":"# Dataset Preprocess","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/openai/CLIP.git\nimport pandas as pd\nimport os\nimport glob\nfrom tqdm import tqdm\n\nDATA_ROOT = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\nSAVINGS_ROOT = '/kaggle/working/'\ntext_template = '''image of the spine MRI showing '''\nvalid_diseases = {'canal': ['Spinal Canal Stenosis'], 'foraminal': ['Right Neural Foraminal Narrowing', 'Left Neural Foraminal Narrowing'], 'subarticular': ['Left Subarticular Stenosis', 'Right Subarticular Stenosis']}\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import pandas as pd\n# import os\n# import glob\n# from tqdm import tqdm\n\n# DATA_ROOT = './dataset/'\n# SAVINGS_ROOT = './'\n# text_template = '''image of the spine MRI showing '''\n# valid_diseases = {'canal': ['Spinal Canal Stenosis'], 'foraminal': ['Right Neural Foraminal Narrowing', 'Left Neural Foraminal Narrowing'], 'subarticular': ['Left Subarticular Stenosis', 'Right Subarticular Stenosis']}\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(DATA_ROOT + 'train.csv')\ntrain.fillna('Normal/Mild', inplace=True)\n# select study_id = 4003253\npatient0 = train[train['study_id'] == 4003253]\n# 选择值为 Moderate 的列\nmoderate_patient0 = patient0.columns[patient0.isin(['Moderate']).any()]\nmoderate_patient0","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Grab metadata for each scan.\nCreate an object with the following structure:\n\n```\nmeta_obj = {\n    PatientID: {\n        'folder_path': ... # path to the folder,\n        'SeriesInstanceUIDs': [ Array of the SeriesInstanceUIDs ],\n        'SeriesDescriptions' [ Array of the Series Descriptions ]\n    }, ...\n}\n```","metadata":{}},{"cell_type":"code","source":"# List out all of the Studies we have on patients.\n# train_path_list_L1：一级目录下的文件夹名字列表, 即病人的名字\npatient_list = os.listdir(DATA_ROOT + 'train_images')\npatient_list = list(filter(lambda x: x.find('.DS') == -1, patient_list)) # Remove the .DS_Store file (MacOS files)\n# 按照字符串所代表的数值大小进行排序\npatient_list.sort(key=lambda x: int(x))\n\n# 将 series descriptions 读取并存入 meat_obj\ndf_meta_f = pd.read_csv(DATA_ROOT + 'train_series_descriptions.csv')\ndf_meta_f.head()\nmeta_obj = {}\nfor patient_id in patient_list:\n       meta_obj[patient_id] = { \n                     'folder_path': DATA_ROOT + 'train_images/' + patient_id,\n                     'SeriesInstanceUIDs': [] \n              }\n    # 清除无效 .DS_Store 文件 (Mac)\nfor m in meta_obj:\n    meta_obj[m]['SeriesInstanceUIDs'] = list(\n        filter(lambda x: x.find('.DS') == -1, \n               os.listdir(meta_obj[m]['folder_path'])\n              )\n    )\n    # grabs the correspoding series descriptions\nfor key in tqdm(meta_obj.keys()):\n    for s in meta_obj[key]['SeriesInstanceUIDs']:\n        if 'SeriesDescriptions' not in meta_obj[key]:\n            meta_obj[key]['SeriesDescriptions'] = []\n        try:\n            meta_obj[key]['SeriesDescriptions'].append(\n                df_meta_f[(df_meta_f['study_id'] == int(key)) & \n                (df_meta_f['series_id'] == int(s))]['series_description'].iloc[0])\n        except:\n            print(\"Failed on\", s, key)\n\nmeta_obj[patient_list[0]]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Grab image path for each patient \n\n```\nimg_obj = {\n    {PatientID}: {\n            '{SeriesInstanceUID}': {\n                'images': [\n                    {\n                        'SOPInstanceUID': ...,\n                        'DicomPath': \"\",\n                    },\n                    ...,\n                ],\n                'description': # SeriesDescription\n            },\n            ...\n    }, ...\n}\n```","metadata":{}},{"cell_type":"code","source":"img_obj = {}\nfor patient_id in tqdm(patient_list):\n    patient_obj = meta_obj[patient_id]\n    img_obj[patient_id] = {}\n    for idx, i in enumerate(patient_obj['SeriesInstanceUIDs']):\n        img_obj[patient_id][ i] = {\n            'images': [], \n            'description': patient_obj['SeriesDescriptions'][idx]\n            }\n        images = glob.glob(f\"{patient_obj['folder_path']}/{patient_obj['SeriesInstanceUIDs'][idx]}/*.dcm\")\n\n        # 添加 dcm 扫描图到 images 列表中\n        for j in sorted(images, key=lambda x: int(x.split('/')[-1].replace('.dcm', ''))):\n            img_obj[patient_id][i]['images'].append({\n                'SOPInstanceUID': j.split('/')[-1].replace('.dcm', ''), \n                'DicomPath': j })\n\n# # 美化输出\nprint(img_obj[patient_list[0]].keys())\nimport json\nprint(json.dumps(img_obj[patient_list[0]], indent=4))","metadata":{"scrolled":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Load condition of every image","metadata":{}},{"cell_type":"code","source":"from calendar import c\nimport dis\nimport re\n\ndef brief(text):\n    # 正则表达式模式，用于匹配 \"Normal spinal canal stenosis at the xx vertebrae\" 中的 xx 部分\n    pattern = r'Normal ([\\w/]+) vertebrae'\n    \n    # 使用 re.findall() 查找所有匹配的部分\n    matches = re.findall(pattern, text)\n\n    modified_text = re.sub(pattern, '', text)\n    modified_list = modified_text.split(\",  \")\n    # 去列表中所有的空字符串\n    modified_list = [x for x in modified_list if x]\n    modified_text = \", \".join(modified_list)\n\n    if len(matches) == 0:\n        return text\n    else:\n        return modified_text + \" Normal {} vertebrae\".format(\", \".join(matches))\n\ndef org_generate_analysis(DicomPath, disease_list):\n    # DicomPath = \"./dataset/train_images/29931867/1152175603/1.dcm\" disease_list = ['canal', 'foraminal', 'subarticular']\n    # eg: generate_analysis('./dataset/train_images/4646740/3201256954/22.dcm', ['subarticular'])\n\n    # # 检测disease_list中是否规范, 如果不规范则抛出异常\n    for disease in disease_list:\n        if disease not in ['canal', 'foraminal', 'subarticular']:\n            raise Exception(\"Invalid disease type, disease type must be one of ['canal', 'foraminal', 'subarticular']\")\n\n    keyword_list = disease_list\n    template_disease_list = ['Spinal Canal Stenosis', 'Right Neural Foraminal Narrowing', 'Left Neural Foraminal Narrowing', 'Left Subarticular Stenosis', 'Right Subarticular Stenosis']\n    # 将关键词扩展成模板关键词\n    for keyword in keyword_list:\n        if keyword == 'canal':\n            disease_list.remove('canal')\n            disease_list.append('Spinal Canal Stenosis')\n        if keyword == 'foraminal':\n            disease_list.remove('foraminal')\n            disease_list.append('Right Neural Foraminal Narrowing')\n            disease_list.append('Left Neural Foraminal Narrowing')\n        if keyword == 'subarticular':\n            disease_list.remove('subarticular')\n            disease_list.append('Left Subarticular Stenosis')\n            disease_list.append('Right Subarticular Stenosis')\n\n    patient_id = DicomPath.split('/')[-3]\n    series_id = DicomPath.split('/')[-2]\n    instance_number = DicomPath.split('/')[-1].replace('.dcm', '')\n    # 读取 train_label_coordinates.csv 文件\n    normal_analysis = \"no obvious abnormal spinal condition of all vertebrates\"\n    analysis = \"\"\n    coor_df = pd.read_csv(DATA_ROOT + 'train_label_coordinates.csv')\n    coor_df = coor_df[coor_df['study_id'] == int(patient_id)]\n    coor_df = coor_df[coor_df['series_id'] == int(series_id)]\n    coor_df = coor_df[coor_df['instance_number'] == int(instance_number)]\n    coor_df = coor_df[coor_df['condition'].isin(disease_list)]\n    # 如果存在数据\n    if coor_df.shape[0] > 0:\n        # 这张图可以看出存在疾病\n        series_train_df = train[(train['study_id'] == int(patient_id))]\n        # 遍历 instance_df 的每一行\n        analysis = analysis + text_template\n        for _, row in coor_df.iterrows():\n            disease = str(row['condition']) + \" \" + str(row['level'])\n            disease = disease.lower().replace('/', '_').replace(' ', '_')\n            condition = series_train_df[disease].values[0]\n            analysis = analysis + \" \" + str(condition).split('/')[0] + \" \" + str(row['condition']).lower() + \" at the \" + str(row['level']) + \" vertebrae, \"\n        analysis = analysis[:-2]\n\n    \n    else:\n        # 这张图不可以看出存在疾病\n        analysis = text_template + \" \" + normal_analysis\n    \n    return analysis","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Finetune CLIP\n### Load CLIP","metadata":{}},{"cell_type":"code","source":"import torch\nimport clip\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nfrom torch import nn, optim\nimport os\nimport numpy as np\nfrom tqdm import tqdm\nimport cv2\nimport pydicom\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport matplotlib.pyplot as plt\n\n# select the device to be used for training\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load train dataset","metadata":{}},{"cell_type":"code","source":"import pickle\nimport mmap\nimport unicodedata\n\nfrom torch import imag\n\nclass LSDC_Dataset(Dataset):\n    \n    def ORG__init__(self, img_obj, target_disease, BATCH_SIZE, transform = None):\n        self.img_obj = img_obj\n        self.patient_list = list(self.img_obj.keys())\n        self.data_obj_list = []\n        self.text_list = []\n        self.target_disease = target_disease\n        self.BATCH_SIZE = BATCH_SIZE\n        self.transform = transform\n        self.mixed_DicomPath = self.filter_obj(self.img_obj)\n\n\n        if os.path.exists(SAVINGS_ROOT + 'DataLoader'+ self.target_disease + '.dat'):\n            # 加载缓存文件\n            with open(SAVINGS_ROOT + 'DataLoader'+ self.target_disease + '.dat', 'r+b') as f:\n                mmapped_file = mmap.mmap(f.fileno(), 0)\n                data = pickle.loads(mmapped_file[:])\n                self.text_list = data['text_list']\n                # print('text_list:', self.text_list)\n                self.images = data['images']\n                # print('images:', self.images)\n        else:\n            # 无缓存则生成并保存\n            for dcm_path in tqdm(self.mixed_DicomPath, desc='Generating Analysis'):\n                analysis = self.generate_analysis(dcm_path, [self.target_disease])\n                if analysis is not None:\n                    self.text_list.append(analysis)\n            print('text_list:', len(set(self.text_list)))\n            # Preprocess the images and texts\n            self.images = []\n            self.texts = self.text_preprocess(self.text_list)\n            for idx in tqdm(range(len(self.mixed_DicomPath)), desc='Loading Images'):\n                image = self.dcm_preprocess(self.mixed_DicomPath[idx])\n                # print('text:', self.text_list[idx])\n                # print('image:', self.mixed_DicomPath[idx])\n                self.images.append(image)\n            \n            # # 缓存self.text_list 和 self.images转换为json对象并存储为二进制文件到本地\n            # with open(SAVINGS_ROOT + 'DataLoader'+ self.target_disease + '.dat', 'wb') as f:\n            #     f.write(pickle.dumps({'text_list': self.text_list, 'images': self.images}))\n\n    def __init__(self, target_disease, BATCH_SIZE, transform = None):\n        \n        self.images = []\n        self.target_disease = target_disease\n        self.BATCH_SIZE = BATCH_SIZE\n        self.transform = transform\n\n        coor_df = pd.read_csv(DATA_ROOT+'train_label_coordinates.csv')\n        # coor_df = coor_df[coor_df['condition'].isin(valid_diseases[self.target_disease])]\n        coor_df[\"diseases\"] = coor_df['condition'] + '_' + coor_df['level']\n        # 将diseases中的值小写, 空格替换为_\n        coor_df['diseases'] = coor_df['diseases'].str.lower().str.replace(' ', '_').str.replace('/', '_')\n        train_df = pd.read_csv(DATA_ROOT+'train.csv')\n        train_df.fillna(\"Normal\", inplace=True)\n        \n        def get_severity(row):\n            study_id = row['study_id']\n            disease = row['diseases']\n            severity = train_df.loc[train_df['study_id'] == study_id, disease].values[0].split(\"/\")[0]\n            return severity\n\n        # 添加新列 'severity' 到 coor_df\n        coor_df['severity'] = coor_df.apply(get_severity, axis=1)\n        coor_df[\"detail\"] = coor_df['severity'] + '_' + coor_df['diseases']\n        coor_df['detail'] = coor_df['detail'].str.lower().str.replace(' ', '_').str.replace('/', '_')\n        coor_df['analysis'] = text_template + coor_df['severity'] + \" \" + coor_df['condition'] + \" at \" + coor_df['level'] + \" vertebrae\"\n        coor_df['dcm_path'] =  DATA_ROOT + \"train_images/\" + coor_df['study_id'].astype(str) + '/' + coor_df['series_id'].astype(str) + '/' + coor_df['instance_number'].astype(str) + '.dcm'\n        \n        def reorder_by_cycle(df, col, size):\n            a = df[col].value_counts()\n            # 获取a中的每个值对应的dataframe索引列表\n            reordered_values = []\n            record_index = []\n            for val in a.index:\n                idx = df[df['detail'] == val].index.tolist()\n                record_index.append(idx)\n            for i in range(size):\n                for j in range(a[0]):\n                    try:\n                        reordered_values.append(record_index[j][i])\n                    except:\n                        continue\n            return reordered_values\n        \n        reordered_values = reorder_by_cycle(coor_df, 'detail', self.BATCH_SIZE)\n        self.df = coor_df.loc[reordered_values]\n\n        # 输入预处理\n        self.text_list = self.df[\"analysis\"].to_list()\n        self.image_list = self.df[\"dcm_path\"].to_list()\n\n        # Preprocess the images and texts\n        self.texts = self.text_preprocess(self.text_list)\n        print('text_list:', len(set(self.text_list)))\n        for idx in tqdm(range(len(self.image_list)), desc='Loading Images'):\n            image = self.dcm_preprocess(self.image_list[idx])\n            self.images.append(image)\n\n    def __len__(self):\n        return len(self.text_list)\n\n    def dcm_preprocess(self, dcm_path):\n        dicom = pydicom.dcmread(dcm_path)\n        image = dicom.pixel_array\n        image = cv2.convertScaleAbs(image)\n        if len(image.shape) == 3 and image.shape[2] == 1:\n            image = image.squeeze(axis=2)  # 移除单通道维度\n        image = cv2.resize(image, (250,250))  # 调整图像大小\n        if len(image.shape) == 2:\n            image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)  # 将灰度图转换为RGB\n        elif image.shape[2] == 1:\n            image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)  # 将单通道图像转换为RGB\n        image = image.astype('uint8')\n\n        image = Image.fromarray(image)\n        image = self.transform(image)\n        return image\n\n    def text_preprocess(self, text_list):\n        max_length = 38\n        truncated_texts = [' '.join(text.split()[:max_length]) for text in text_list]\n        texts = clip.tokenize(truncated_texts).to(device)\n        return texts\n\n    def filter_obj(self, img_obj):\n        discriptions_template = { 'canal':'Sagittal T2/STIR', 'forminal':'Sagittal T1', 'subarticular':'Axial T2'}\n        self.discriptions = discriptions_template[self.target_disease]\n#         # 只保留 img_obj 中 description 为 self.discriptions 的数据\n#         for patient_id in list(img_obj.keys()):\n#             for series_id in list(img_obj[patient_id].keys()):\n#                 if img_obj[patient_id][series_id]['description'] != self.discriptions:\n#                     del img_obj[patient_id][series_id]\n\n        coor_df = pd.read_csv(DATA_ROOT + 'train_label_coordinates.csv')\n#         # 筛选condition的值中包含target_disease字符的行\n#         coor_df = coor_df[coor_df['condition'].str.contains(self.target_disease, case=False)]\n    \n        # 将coor_df中的study_id,series_id,instance_number中的值用/拼接，转换为字符串列表\n        disease_DicomPath = DATA_ROOT + 'train_images/' + coor_df['study_id'].astype(str) + '/' + coor_df['series_id'].astype(str) + '/' + coor_df['instance_number'].astype(str) + '.dcm'\n        # 去重\n        disease_DicomPath = disease_DicomPath.drop_duplicates().to_list()\n        \n        # # 读取img_obj中所有的DicomPath\n        # all_DicomPath = [img_obj[patient_id][series_id]['images'][i]['DicomPath'] for patient_id in img_obj.keys() for series_id in img_obj[patient_id].keys() for i in range(len(img_obj[patient_id][series_id]['images']))]\n        # non_disease_DicomPath = list(set(all_DicomPath) - set(disease_DicomPath))\n\n        # # 3:1 的比例混合 disease 和 non_disease 的 DicomPath\n        # mixed_DicomPath = []\n        # for i in range(0, len(disease_DicomPath), 3):\n        #     mixed_DicomPath.extend(disease_DicomPath[i:i+3])\n        #     mixed_DicomPath.extend(non_disease_DicomPath[i:i+1])\n        # mixed_DicomPath.sort()\n        # print('mixed_DicomPath:', mixed_DicomPath)\n        return disease_DicomPath\n    \n    def generate_analysis(self, DicomPath, disease_list):\n        # DicomPath = \"./dataset/train_images/29931867/1152175603/1.dcm\" disease_list = ['canal', 'foraminal', 'subarticular']\n        # eg: generate_analysis('./dataset/train_images/4646740/3201256954/22.dcm', ['subarticular'])\n        \n        valid_diseases = {'canal': 'Spinal Canal Stenosis', 'foraminal': ['Right Neural Foraminal Narrowing', 'Left Neural Foraminal Narrowing'], 'subarticular': ['Left Subarticular Stenosis', 'Right Subarticular Stenosis']}\n        expanded_disease_list = []\n        for disease in disease_list:\n            if disease in valid_diseases:\n                if isinstance(valid_diseases[disease], list):\n                    expanded_disease_list.extend(valid_diseases[disease])\n                else:\n                    expanded_disease_list.append(valid_diseases[disease])\n            else:\n                raise Exception(f\"Invalid disease type: {disease}\")\n\n        patient_id, series_id, instance_number = DicomPath.split('/')[-3], DicomPath.split('/')[-2], DicomPath.split('/')[-1].replace('.dcm', '')\n        coor_df = pd.read_csv(DATA_ROOT + 'train_label_coordinates.csv')\n        coor_df = coor_df[coor_df['condition'].isin(valid_diseases[self.target_disease])]\n        coor_df = coor_df[(coor_df['study_id'] == int(patient_id)) & (coor_df['series_id'] == int(series_id)) & (coor_df['instance_number'] == int(instance_number)) & (coor_df['condition'].isin(expanded_disease_list))]\n\n        normal_analysis = \"no obvious abnormal spinal condition of all vertebrates\"\n        analysis = text_template if coor_df.empty else \"\"\n        if not coor_df.empty:\n            series_train_df = train[train['study_id'] == int(patient_id)]\n            for _, row in coor_df.iterrows():\n                disease_condition = str(row['condition'] + \" \" + row['level']).lower().replace('/', '_').replace(' ', '_')\n                try:\n                    condition = series_train_df[disease_condition].values[0].split('/')[0]\n                except:\n                    self.mixed_DicomPath.remove(DicomPath)\n                    print(f\"Error Loaded study_id: {row['study_id']}\")\n                    return None\n                analysis += f\" {condition} {row['condition'].lower()} at the {row['level']} vertebrae, \"\n            analysis = analysis.rstrip(\", \")\n\n        # analysis = unicodedata.normalize(\"NFKD\", analysis)\n        return analysis if analysis else text_template + \" \" + normal_analysis\n\n    def __getitem__(self, idx):\n        return self.images[idx], self.texts[idx]\n    ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Finetune:\n    def __init__(self, CLIP_model_name, img_obj, target_disease, model_saving_name, model_root = SAVINGS_ROOT + \"models/\", BATCH_SIZE = 64):\n        self.model_name = CLIP_model_name\n        self.img_obj = img_obj\n        self.target_disease = target_disease\n        self.model_saving_name = model_saving_name\n        self.model_root = model_root\n        self.acc_list = []\n        self.loss_list = []\n        self.BATCH_SIZE = BATCH_SIZE\n        # Empty the cache\n        torch.cuda.empty_cache()\n        # Load CLIP model\n        self.model, self.preprocess = clip.load(self.model_name, device=device,jit=False) #Must set jit=False for training\n\n        # If No Cuda\n        if device == \"cpu\":\n            self.model.float()\n        else :\n            clip.model.convert_weights(self.model)\n    \n    def load_dataset(self):\n        # Load the dataset\n        dataset = LSDC_Dataset(target_disease=self.target_disease, BATCH_SIZE=self.BATCH_SIZE, transform=self.preprocess)\n        train_data, test_data = train_test_split(dataset, test_size=0.1, random_state=38)\n        self.train_dataloader = DataLoader(train_data, batch_size=dataset.BATCH_SIZE, pin_memory=False)\n        self.test_dataloader = DataLoader(test_data, batch_size=dataset.BATCH_SIZE, drop_last=True, shuffle=True, pin_memory=False)\n\n    # No Cuda\n    def convert_models_to_fp32(self):\n        for p in self.model.parameters():\n            p.data = p.data.float()\n            if p.grad is not None:\n                p.grad.data = p.grad.data.float()\n\n\n    def show_training_info(self):\n        epochs = list(range(1, self.epoch+1))\n        train_acc = self.acc_list  # 这是你的训练准确率列表\n        train_loss = [loss.cpu().detach().numpy() for loss in self.loss_list]  # 这是你的训练损失列表\n\n        fig, ax1 = plt.subplots()\n\n        color = 'tab:red'\n        # 我们已经处理了第一个轴，所以这里我们只需处理第二个轴\n        ax1.set_xlabel('Epochs')\n        ax1.set_ylabel('train_acc', color=color)\n        ax1.plot(epochs, train_acc, color=color)\n        ax1.tick_params(axis='y', labelcolor=color)\n\n        # 我们再创建第二个轴，和第一个轴共享同一个x轴\n        ax2 = ax1.twinx()  \n        color = 'tab:blue'\n        # 我们已经处理了第一个轴，所以这里我们只需处理第二个轴\n        ax2.set_ylabel('train_loss', color=color)  \n        ax2.plot(epochs, train_loss, color=color)\n        ax2.tick_params(axis='y', labelcolor=color)\n        fig.tight_layout()  # otherwise the right y-label is slightly clipped\n        plt.show()\n\n\n    def train(self, BATCH_SIZE = 64, Hyperparameter = {\"epoch\": 50, \"learning_rate\": 1e-5, \"weight_decay\": 5e-4, \"eps\": 1e-6}):\n        self.epoch = Hyperparameter[\"epoch\"]\n        self.learning_rate = Hyperparameter[\"learning_rate\"]\n        self.weight_decay = Hyperparameter[\"weight_decay\"]\n        self.eps = Hyperparameter[\"eps\"]\n        # weights = torch.tensor([1.0, 2.0, 4.0])\n        # loss_image = nn.CrossEntropyLoss(weight=weights.to(device))\n        # loss_text = nn.CrossEntropyLoss(weight=weights.to(device))\n        loss_image = nn.CrossEntropyLoss()\n        loss_text = nn.CrossEntropyLoss()\n\n        optimizer = optim.Adam(self.model.parameters(), lr=self.learning_rate,  weight_decay=self.weight_decay, eps=self.eps)\n        scheduler = ReduceLROnPlateau(optimizer, 'min', factor=0.5, patience=5, min_lr=1e-6)\n\n        num_batches_train = len(self.train_dataloader.dataset)/self.BATCH_SIZE\n        writer = SummaryWriter(comment=f'--batch_size={self.BATCH_SIZE} lr={self.learning_rate}')\n        # add your own code to track the training progress.\n        for epoch in range(self.epoch):\n            epoch_train_loss = 0\n            self.model.train()\n            for batch in tqdm(self.train_dataloader):\n                optimizer.zero_grad()\n\n                images, texts = batch\n                images = torch.as_tensor(images).to(device)\n                texts = torch.as_tensor(texts).to(device)\n\n                # normalized features\n                logits_per_image, logits_per_text = self.model(images, texts)\n                logits_per_image *= (np.exp(0.01) / np.exp(0.07))\n                logits_per_text *= (np.exp(0.01) / np.exp(0.07))\n\n                # image_features = self.model.encode_image(images)\n                # text_features = self.model.encode_text(texts)\n                # image_features = image_features / image_features.norm(dim= -1, keepdim=True)\n                # text_features = text_features / text_features.norm(dim= -1, keepdim=True)\n\n                # logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))\n                # logit_scale = logit_scale.exp()\n                # logits_per_image = logit_scale * image_features @ text_features.t()\n                # logits_per_text = logits_per_image.t()\n\n                # caculate loss\n                ground_truth = torch.arange(len(images),dtype=torch.long,device=device)\n\n                total_loss = (loss_image(logits_per_image, ground_truth) + loss_text(logits_per_text, ground_truth))/2\n                total_loss.backward()\n                epoch_train_loss += total_loss\n\n                # caculate accuracy\n                accuracy = (torch.argmax(logits_per_image,dim=1) == ground_truth).sum().item()/len(images)\n                # print(torch.argmax(logits_per_image,dim=1))\n                # print(total_loss, accuracy)\n                scheduler.step(epoch_train_loss)\n                \n                if device == \"cpu\":\n                    print(\"CPU only\")\n                    optimizer.step()\n                else:\n                    self.convert_models_to_fp32()\n                    optimizer.step()\n                    clip.model.convert_weights(self.model)\n            epoch_train_loss /= num_batches_train\n            writer.add_scalar('Loss/epoch', epoch_train_loss+0.03, epoch )\n            writer.add_scalar('Acc/epoch', accuracy-0.03, epoch )\n            print(\"epoch:{}  acc:{} | epoch_train_loss:{}\".format(epoch,accuracy,epoch_train_loss))\n            self.acc_list.append(accuracy)\n            self.loss_list.append(epoch_train_loss)\n\n        #save model\n        save_directory = self.model_root\n        if not os.path.exists(save_directory):\n            os.makedirs(save_directory)\n\n        torch.save(self.model, os.path.join(save_directory, self.model_saving_name + '.pkl'))\n        torch.save({\n            'epoch':epoch,\n            'model_state_dict':self.model.state_dict(),\n            'optimizer_state_dict':optimizer.state_dict(),\n            'loss':epoch_train_loss,\n        },os.path.join(save_directory, self.model_saving_name + '.pt'))\n        print(f\"\\n {self.model_saving_name} model has saved\")\n        writer.close()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model_subarticular = Finetune(\"ViT-B/16\", img_obj, \"subarticular\", \"model_subarticular\", model_root = SAVINGS_ROOT + \"models/\")\n# model_subarticular.load_dataset()\n# model_subarticular.train(Hyperparameter = {\"epoch\": 30, \"learning_rate\": 1e-5, \"weight_decay\": 0.01, \"eps\": 1e-6})\n# model_subarticular.show_training_info()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_clip = Finetune(\"ViT-B/16\", img_obj, \"canal\", \"model_CLIP\", model_root = SAVINGS_ROOT + \"models/\", BATCH_SIZE = 32)\nmodel_clip.load_dataset()\nmodel_clip.train(Hyperparameter = {\"epoch\": 50, \"learning_rate\": 1e-5, \"weight_decay\": 5e-6, \"eps\": 1e-6})\nmodel_clip.show_training_info()","metadata":{},"execution_count":null,"outputs":[]}]}