{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"colab":{"provenance":[],"collapsed_sections":["MXtVUZnHiC5F","3YnUqwF6wLQu","_3zYdiwqxsx_","7B7iSXrPazbt","UUeXu0ZSTv-5","8dPv92Ra9t-T"],"gpuType":"T4"},"accelerator":"GPU","kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":23870,"databundleVersionId":1781260,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 0. Prepare the data","metadata":{"id":"OcLmGk2KNIxD"}},{"cell_type":"code","source":"!pip install kaggle","metadata":{"id":"J3xg-x8qLrPs","executionInfo":{"status":"ok","timestamp":1742262594777,"user_tz":-660,"elapsed":2420,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"outputId":"0f110df5-68c6-4b6c-f958-25909dbd86bb","jupyter":{"source_hidden":true},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:25:58.402587Z","iopub.status.idle":"2025-05-16T02:25:58.402889Z","shell.execute_reply.started":"2025-05-16T02:25:58.402718Z","shell.execute_reply":"2025-05-16T02:25:58.402729Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Check GPU and empty the cache, but may caused by the difference between kaggle and local GPU, Kaggle can not empty GPU usage by this code.","metadata":{}},{"cell_type":"code","source":"import torch\nprint(f\"Available device: {'GPU ✅' if torch.cuda.is_available() else 'CPU ⚠️'}\")","metadata":{"id":"SHmy9NfLL0PM","outputId":"0e3ef4c8-f276-4fe8-ee82-b9c8fa0944a7","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:25:58.404163Z","iopub.status.idle":"2025-05-16T02:25:58.404514Z","shell.execute_reply.started":"2025-05-16T02:25:58.404315Z","shell.execute_reply":"2025-05-16T02:25:58.404326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()# Clean up useless cache(although this is generally of limited help for large-scale memory allocation)","metadata":{"id":"SDsdP_SkeVQM","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:25:58.405596Z","iopub.status.idle":"2025-05-16T02:25:58.405852Z","shell.execute_reply.started":"2025-05-16T02:25:58.405744Z","shell.execute_reply":"2025-05-16T02:25:58.405754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# here is the first time I try to running it on the colab, so you can directly ignore them.\nfrom google.colab import drive\ndrive.mount('/content/drive')\nfilepath = \"/content/drive/MyDrive/checkpoint.pth\"\n\nimport os\nos.environ['KAGGLE_CONFIG_DIR'] = '/content/drive/MyDrive/kaggle'\n\n!kaggle competitions download -c ranzcr-clip-catheter-line-classification\n\nimport os\n\n# 检查指定目录中的文件\ndirectory = '/content/drive/MyDrive/kaggle'\nprint(os.listdir(directory))\n\nimport zipfile\nimport os\n\n# 定义 ZIP 文件的路径\nzip_file_path = '/content/ranzcr-clip-catheter-line-classification.zip'\n\n# 解压 ZIP 文件到指定目录\nwith zipfile.ZipFile(zip_file_path, 'r') as zip_ref:\n    zip_ref.extractall('/content/ranzcr_dataset')  # 指定解压目标文件夹\n\n# 列出解压后的文件\nextracted_files = os.listdir('/content/ranzcr_dataset')\nprint(extracted_files)","metadata":{"id":"w0r3xGpw_eHD","outputId":"4037fc52-daa1-4906-aca9-e5d024fec490","executionInfo":{"status":"ok","timestamp":1742262807130,"user_tz":-660,"elapsed":2671,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"jupyter":{"source_hidden":true},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:25:58.407147Z","iopub.status.idle":"2025-05-16T02:25:58.407489Z","shell.execute_reply.started":"2025-05-16T02:25:58.407323Z","shell.execute_reply":"2025-05-16T02:25:58.407338Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## EDA","metadata":{"id":"tjsIq1KALpzE"}},{"cell_type":"code","source":"from PIL import Image\nimport matplotlib.pyplot as plt\n\nimage_path = '/kaggle/input/ranzcr-clip-catheter-line-classification/train/1.2.826.0.1.3680043.8.498.10013648386639413183439376787966642105.jpg'  # 替换为实际图像路径\nimage = Image.open(image_path)\nplt.imshow(image)\nplt.axis('off')  # do not show the axis\nplt.show()","metadata":{"id":"7XliKMC0eo8w","outputId":"7f512504-5786-4394-ce79-65c83789ec21","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:25:58.408292Z","iopub.status.idle":"2025-05-16T02:25:58.408622Z","shell.execute_reply.started":"2025-05-16T02:25:58.408453Z","shell.execute_reply":"2025-05-16T02:25:58.408467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dividing the validation set seems to be repetitive\n\nimport os\nimport random\nimport shutil\nimport pandas as pd\n\n# Dataset Path\ndata_root = r\"/kaggle/input/ranzcr-clip-catheter-line-classification/\"\ntrain_path = os.path.join(data_root, \"train\")\ntest_path = os.path.join(data_root, \"test\")\ncsv_path = os.path.join(data_root, \"train.csv\")\n\n# Create the validation set and training set folders in the writable path\nval_path = \"/kaggle/working/val\"\ntrain_new_path = \"/kaggle/working/train\"\nos.makedirs(val_path, exist_ok=True)\nos.makedirs(train_new_path, exist_ok=True)\n\n# Load CSV file\ndf = pd.read_csv(csv_path)\n\n# Get all the pictures in the training set\ntrain_images = [f for f in os.listdir(train_path) if f.endswith('.jpg')]\nrandom.shuffle(train_images)\n\n# Divide the validation set (20%)\nsplit_idx = int(0.2 * len(train_images))\nval_images = train_images[:split_idx]\ntrain_remaining_images = train_images[split_idx:]  # Remaining 80%\n\n# Copy the validation set images to the validation set folder\nfor img in val_images:\n    src = os.path.join(train_path, img)\n    dst = os.path.join(val_path, img)\n    shutil.copy(src, dst)\n\n# Copy the remaining training set images to the new training set folder\nfor img in train_remaining_images:\n    src = os.path.join(train_path, img)\n    dst = os.path.join(train_new_path, img)\n    shutil.copy(src, dst)\n\n# Split the CSV file\nval_df = df[df['StudyInstanceUID'].isin([img.replace('.jpg', '') for img in val_images])]\ntrain_df = df[df['StudyInstanceUID'].isin([img.replace('.jpg', '') for img in train_remaining_images])]\n\n# Save the new CSV files\ntrain_df.to_csv(os.path.join(train_new_path, \"train.csv\"), index=False)\nval_df.to_csv(os.path.join(val_path, \"val.csv\"), index=False)\n\nprint(f\"Verification set partitioning completed! Training set: {len(train_remaining_images)} images, Validation set: {len(val_images)} images\")","metadata":{"id":"nppe4t0KRt6R","outputId":"ef114a7c-e218-4e15-cb8e-5561789bf69a","executionInfo":{"status":"ok","timestamp":1742263180398,"user_tz":-660,"elapsed":201,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:29:09.70559Z","iopub.execute_input":"2025-05-16T02:29:09.706412Z","iopub.status.idle":"2025-05-16T02:33:11.709144Z","shell.execute_reply.started":"2025-05-16T02:29:09.706388Z","shell.execute_reply":"2025-05-16T02:33:11.708366Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"First of all, what I must clarify is that I should adjust my mindset first. Too many problems are encountered now. Then comes the plan I envisioned for tomorrow.\nIt is necessary to separate different modalities or different sets in the dataset.\n2. Is it necessary to use dinov2 separately for individual modalities or can all modalities be used together? My plan is to run the 3-mode dinov2 smoothly tomorrow first, and then add parts such as optimizing the video memory.\n3. The current code seems to be running on the cpu. I guess the problem might lie within it. Fortunately, the dataset that was initially thought to work here can be run smoothly because there is about 100G of online memory.","metadata":{"id":"RobiDdDAZR8X"}},{"cell_type":"code","source":"# Install necessary libraries (DINOv2 dependencies)\n!pip install torchvision pytorch-lightning psutil","metadata":{"collapsed":true,"id":"JgiBHk9vgA5v","outputId":"1e5ba93e-4b83-4652-aa8e-6580de74e1df","executionInfo":{"status":"ok","timestamp":1742263292357,"user_tz":-660,"elapsed":108410,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"jupyter":{"outputs_hidden":true,"source_hidden":true},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:26:41.978528Z","iopub.execute_input":"2025-05-16T02:26:41.978858Z","iopub.status.idle":"2025-05-16T02:28:07.234129Z","shell.execute_reply.started":"2025-05-16T02:26:41.978828Z","shell.execute_reply":"2025-05-16T02:28:07.23317Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport time\nimport matplotlib.pyplot as plt","metadata":{"id":"-0VrYEuLklHl","executionInfo":{"status":"ok","timestamp":1742263384279,"user_tz":-660,"elapsed":9364,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:33:49.998737Z","iopub.execute_input":"2025-05-16T02:33:49.999038Z","iopub.status.idle":"2025-05-16T02:33:50.003815Z","shell.execute_reply.started":"2025-05-16T02:33:49.99902Z","shell.execute_reply":"2025-05-16T02:33:50.003054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    # Dataset parameters\n    data_root = r\"/kaggle/input/ranzcr-clip-catheter-line-classification/\"  # Replace with your path\n    num_classes = 11                       # Modify according to the dataset\n    input_size = 112                       # Input size (recommended: 112, 224, or 160)\n\n    # Training parameters\n    train_mode = \"feature_extract\"         # Options: feature_extract / fine_tune\n    batch_size = 64                       # Reduce to 16 or 8 if memory is insufficient; started with 32, but got an error\n    # Can gradually increase until memory is full\n    # Note: Increase batch_size, adjust learning rate as needed (e.g., linear scaling rule)\n\n    num_epochs = 10\n    lr = 1e-4\n\n    # Lightweight model selection\n    model_name = \"dinov2_vitl14\"           # Smallest model: dinov2_vits14 s b l g\n    freeze_backbone = True                  # Whether to freeze the backbone\n\n    # Device configuration\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    train_mode = \"fine_tune\"               # Training mode\n\n    '''The following 3 lines are new'''\n    use_clip = True                         # Whether to use CLIP\n    use_monai = True                        # Whether to use MONAI preprocessing\n    save_activation_maps = True             # Whether to save activation maps","metadata":{"id":"txj3OU2tkuLO","executionInfo":{"status":"ok","timestamp":1742263388387,"user_tz":-660,"elapsed":21,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:41:14.474297Z","iopub.execute_input":"2025-05-16T02:41:14.475099Z","iopub.status.idle":"2025-05-16T02:41:14.480166Z","shell.execute_reply.started":"2025-05-16T02:41:14.475069Z","shell.execute_reply":"2025-05-16T02:41:14.479271Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Here, the training set needs to be divided into a new training set and a validation set.","metadata":{"id":"Vg7TTbVEPN15"}},{"cell_type":"markdown","source":"### 2.1 CLIP Tool Functions","metadata":{"id":"MXtVUZnHiC5F"}},{"cell_type":"code","source":"!pip install git+https://github.com/openai/CLIP.git","metadata":{"id":"23cr5ewYzkl4","outputId":"4d68944e-b1fe-410c-fabb-2b8df17b288d","collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:33:57.432868Z","iopub.execute_input":"2025-05-16T02:33:57.433175Z","iopub.status.idle":"2025-05-16T02:34:03.003787Z","shell.execute_reply.started":"2025-05-16T02:33:57.433155Z","shell.execute_reply":"2025-05-16T02:34:03.002703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import clip\n\nclass CLIPFeatureExtractor:\n    def __init__(self, model_name=\"ViT-B/32\"):\n        self.device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        self.model, self.preprocess = clip.load(model_name, device=self.device)\n\n    def extract_features(self, image_path):\n        image = Image.open(image_path)\n        image_input = self.preprocess(image).unsqueeze(0).to(self.device)\n        with torch.no_grad():\n            features = self.model.encode_image(image_input)\n        return features.cpu()","metadata":{"id":"Mpm7wXgEiHaC","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:08.907874Z","iopub.execute_input":"2025-05-16T02:34:08.909116Z","iopub.status.idle":"2025-05-16T02:34:08.915457Z","shell.execute_reply.started":"2025-05-16T02:34:08.909053Z","shell.execute_reply":"2025-05-16T02:34:08.914546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2.2 Add an activation graph generation function\ndef generate_activation_map(model, image, output_path=\"activation_map.png\"):\n    model.eval()\n    image = image.unsqueeze(0).to(config.device)\n    image.requires_grad_()\n    \n    # Forward pass\n    output = model(image)\n    class_idx = output.argmax().item()\n    \n    # Backward pass\n    model.zero_grad()\n    output[0, class_idx].backward()\n    \n    # Obtain the Gradient\n    gradients = image.grad.data.abs().squeeze()\n    \n    # Visualization\n    plt.imshow(gradients.cpu().numpy(), cmap='hot')\n    plt.savefig(output_path)\n    plt.close()","metadata":{"id":"DyVp7Q1ZwOC-","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:11.65188Z","iopub.execute_input":"2025-05-16T02:34:11.652602Z","iopub.status.idle":"2025-05-16T02:34:11.657453Z","shell.execute_reply.started":"2025-05-16T02:34:11.652579Z","shell.execute_reply":"2025-05-16T02:34:11.656683Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2.3 New Data enhancements Added","metadata":{"id":"_3zYdiwqxsx_"}},{"cell_type":"code","source":"!pip install monai\n","metadata":{"id":"U_Ezijgs0ZqP","outputId":"eefae1b9-9ff8-40a7-e6aa-cddc0de37afd","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:14.701264Z","iopub.execute_input":"2025-05-16T02:34:14.701954Z","iopub.status.idle":"2025-05-16T02:34:18.26993Z","shell.execute_reply.started":"2025-05-16T02:34:14.701929Z","shell.execute_reply":"2025-05-16T02:34:18.268773Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.transforms import (\n    Compose, LoadImage, EnsureChannelFirst, ScaleIntensity,\n    RandRotate, RandFlip, RandZoom, EnsureType\n)\n\nclass MedicalTransform:\n    def __init__(self, input_size):\n        self.transform = Compose([\n            LoadImage(image_only=True),\n            EnsureChannelFirst(),  # 替换 AddChannel\n            ScaleIntensity(),\n            RandRotate(range_x=15, prob=0.5),\n            RandFlip(prob=0.5),\n            RandZoom(min_zoom=0.9, max_zoom=1.1, prob=0.5),\n            EnsureType(),\n            transforms.Resize((input_size, input_size)),\n        ])\n\n    def __call__(self, image_path):\n        return self.transform(image_path)","metadata":{"id":"LNehBMKIxvxz","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:24.457532Z","iopub.execute_input":"2025-05-16T02:34:24.457889Z","iopub.status.idle":"2025-05-16T02:34:24.464959Z","shell.execute_reply.started":"2025-05-16T02:34:24.457857Z","shell.execute_reply":"2025-05-16T02:34:24.464166Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Custom dataset loading","metadata":{"id":"W4_wW54Vb1vr"}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, root, csv_path, transform=None, use_clip=False, use_monai=False):\n        # Check if the path exists\n        if not os.path.exists(root):\n            raise FileNotFoundError(f\"Dataset path does not exist: {root}\")\n\n        # Get all image files\n        self.images = [os.path.join(root, f) for f in os.listdir(root) if f.endswith('.jpg')]\n\n        # Load label file\n        self.labels = self._load_labels(csv_path)  # Use CSV file as label file\n\n        # Light data augmentation (adjust based on memory)\n        self.transform = transform or transforms.Compose([\n            transforms.Resize((Config.input_size, Config.input_size)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n\n        # The following part is for CLIP\n        self.use_clip = use_clip\n        if self.use_clip:\n            self.clip_extractor = CLIPFeatureExtractor()\n\n        # The following part is for preprocessing\n        if use_monai:\n            self.transform = MedicalTransform(Config.input_size)\n        else:\n            self.transform = transform or transforms.Compose([\n                transforms.Resize((Config.input_size, Config.input_size)),\n                transforms.ToTensor(),\n                transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n            ])\n\n    def _load_labels(self, csv_path):\n        # Load label information from CSV file\n        import pandas as pd\n        df = pd.read_csv(csv_path)\n\n        # Extract all label column names (excluding 'StudyInstanceUID' and 'PatientID')\n        label_columns = [\n            'ETT - Abnormal', 'ETT - Borderline', 'ETT - Normal',\n            'NGT - Abnormal', 'NGT - Borderline', 'NGT - Incompletely Imaged', 'NGT - Normal',\n            'CVC - Abnormal', 'CVC - Borderline', 'CVC - Normal',\n            'Swan Ganz Catheter Present'\n        ]\n\n        labels = {}\n        for _, row in df.iterrows():\n            image_name = row['StudyInstanceUID']  # Use 'StudyInstanceUID' column as image name\n            # Convert each row's label columns to a label vector (e.g., [1, 0, 1, ...])\n            label_vector = [row[col] for col in label_columns]\n            labels[image_name] = label_vector\n\n        # Save label column names for later use\n        self.label_columns = label_columns\n        return labels\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img_path = self.images[idx]\n        image = Image.open(img_path).convert('RGB')\n\n        # Get label from filename\n        img_name = os.path.basename(img_path).replace('.jpg', '')  # Remove file extension\n        label = self.labels.get(img_name, [0] * len(self.label_columns))  # If label not found, return all 0 vector\n\n        # Modify to use CLIP for the return\n\n        '''   The following are new additions   '''\n        if self.use_clip:\n            clip_feats = self.clip_extractor.extract_features(img_path)\n            return self.transform(image), torch.tensor(label, dtype=torch.float32), clip_feats\n        else:\n            return self.transform(image), torch.tensor(label, dtype=torch.float32)","metadata":{"id":"JLKXcqZeSTnk","executionInfo":{"status":"ok","timestamp":1742263398618,"user_tz":-660,"elapsed":25,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:29.14102Z","iopub.execute_input":"2025-05-16T02:34:29.141623Z","iopub.status.idle":"2025-05-16T02:34:29.155508Z","shell.execute_reply.started":"2025-05-16T02:34:29.141599Z","shell.execute_reply":"2025-05-16T02:34:29.154396Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check the dataset path\ntrain_path = \"/kaggle/working/train\"\nval_path = \"/kaggle/working/val\"\n\n# Instantiate the dataset\ntrain_csv_path = os.path.join(train_new_path, \"train.csv\") # CSV path of the new training set\nval_csv_path = os.path.join(val_path, \"val.csv\") # The new validation set CSV path\n\ntrain_dataset = CustomDataset(train_path, train_csv_path)\nval_dataset = CustomDataset(val_path, val_csv_path)","metadata":{"id":"xmr7lTbhWECA","outputId":"81146ab1-a48a-44af-e53d-d7bf62ddc78b","executionInfo":{"status":"ok","timestamp":1742263409020,"user_tz":-660,"elapsed":6724,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:33.419344Z","iopub.execute_input":"2025-05-16T02:34:33.419689Z","iopub.status.idle":"2025-05-16T02:34:35.424458Z","shell.execute_reply.started":"2025-05-16T02:34:33.419665Z","shell.execute_reply":"2025-05-16T02:34:35.423757Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Build a lightweight DINOv2 model","metadata":{"id":"fc4vrI-3PfSg"}},{"cell_type":"code","source":"def build_lightweight_dinov2(config):\n    \"\"\"\n    Build a lightweight DINOv2 model.\n    - Load the pre-trained model.\n    - Freeze the backbone network (optional).\n    - Replace the classification head.\n    - Fine-tune partially (unfreeze the last few layers).\n    \"\"\"\n    # Load the pre-trained model\n    model = torch.hub.load('facebookresearch/dinov2', config.model_name, pretrained=True)\n\n    # Freeze the backbone network\n    if config.freeze_backbone:\n        for param in model.parameters():\n            param.requires_grad = False\n\n    # Dynamically get the feature dimension\n    if \"vits14\" in config.model_name:\n        feature_dim = 384\n    elif \"vitb14\" in config.model_name:\n        feature_dim = 768\n    elif \"vitl14\" in config.model_name:\n        feature_dim = 1024  # The feature dimension of vitl14 is 1024\n    elif \"vitg14\" in config.model_name:\n        feature_dim = 1536  # The feature dimension of vitg14 is 1536\n    else:\n        raise ValueError(f\"Unsupported model name: {config.model_name}\")\n\n    # Replace the classification head (with a lighter structure)\n    model.head = nn.Sequential(\n        nn.Linear(feature_dim, config.num_classes)  # Directly output classification results\n    )\n\n    # Partial fine-tuning: unfreeze the last 2 Transformer blocks\n    if config.train_mode == \"fine_tune\":\n        total_layers = len(model.blocks)\n        for i in range(total_layers - 2, total_layers):  # Unfreeze the last 2 layers\n            for param in model.blocks[i].parameters():\n                param.requires_grad = True\n\n    return model.to(config.device)\n\n    # Additional code for CLIP\n\n    if config.use_clip:\n        # Concatenate DINOv2 and CLIP features\n        model.head = nn.Sequential(\n            nn.Linear(feature_dim + 128, 64),  # CLIP feature dimension is 512, 256\n            nn.ReLU(),\n            nn.Linear(256, config.num_classes)\n        )\n    return model.to(config.device)\n    '''\n    The model's forward pass needs to support CLIP feature input.\n\n    Confirm the CLIP feature dimension (suggest changing from 512 to the actual value).\n    '''","metadata":{"id":"fpVQjXRvk02A","executionInfo":{"status":"ok","timestamp":1742263417806,"user_tz":-660,"elapsed":27,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:34:38.587726Z","iopub.execute_input":"2025-05-16T02:34:38.588358Z","iopub.status.idle":"2025-05-16T02:34:38.596067Z","shell.execute_reply.started":"2025-05-16T02:34:38.588332Z","shell.execute_reply":"2025-05-16T02:34:38.595235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create an instance of Config\nconfig = Config()\n\n# Instantiate the model\nmodel = build_lightweight_dinov2(config)  # Note: Passing the config instance, not the Config class\nprint(f\"Model parameter count: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M\")\n\n# Check model output\ndummy_input = torch.randn(2, 3, 224, 224).to(config.device)  # Move input data to GPU\noutput = model(dummy_input)\nprint(output.shape)  # Should output torch.Size([2, 11])","metadata":{"id":"NqPSFR8vuzFG","outputId":"62db080e-60a7-4e29-c846-249df5f7c707","executionInfo":{"status":"ok","timestamp":1742263434647,"user_tz":-660,"elapsed":14522,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:41:19.092793Z","iopub.execute_input":"2025-05-16T02:41:19.093106Z","iopub.status.idle":"2025-05-16T02:41:25.92611Z","shell.execute_reply.started":"2025-05-16T02:41:19.093086Z","shell.execute_reply":"2025-05-16T02:41:25.92535Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training Cycle (including resource monitoring)","metadata":{"id":"jjEBQaqcTTMF"}},{"cell_type":"markdown","source":"### 1 epoch","metadata":{"id":"oZkbLTqxTsYL"}},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch.amp  # Modern version of AMP uses torch.amp instead of torch.cuda.amp\n\nimport torch.backends.cudnn as cudnn\ncudnn.benchmark = True  # Automatically optimize convolution implementation\n\ndef train_lightweight(model, config):\n    # Initialize data loaders\n    train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=config.batch_size, pin_memory=True)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=config.lr)\n\n    # Properly initialize GradScaler without parameters\n    scaler = torch.amp.GradScaler()\n\n    # Training history\n    history = {'train_loss': [], 'val_acc': []}\n\n    # Run for config.num_epochs epochs\n    for epoch in range(config.num_epochs):\n        start_time = time.time()\n        model.train()\n        train_loss = 0.0\n        batch_count = 0\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Training Phase\")\n        for i, (inputs, labels) in enumerate(tqdm(train_loader, desc=\"Training\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            optimizer.zero_grad()\n\n            # Properly use autocast\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n            # Scale loss and backpropagate\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            train_loss += loss.item()\n            batch_count += 1\n\n            # Only run for 1 batch\n            #from tqdm import tqdm\nimport torch.amp  # Modern version of AMP uses torch.amp instead of torch.cuda.amp\n\nimport torch.backends.cudnn as cudnn\ncudnn.benchmark = True  # Automatically optimize convolution implementation\n\ndef train_lightweight(model, config):\n    # Initialize data loaders\n    train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=config.batch_size, pin_memory=True)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=config.lr)\n\n    # Properly initialize GradScaler without parameters\n    scaler = torch.amp.GradScaler()\n\n    # Training history\n    history = {'train_loss': [], 'val_acc': []}\n\n    # Run for config.num_epochs epochs\n    for epoch in range(config.num_epochs):\n        start_time = time.time()\n        model.train()\n        train_loss = 0.0\n        batch_count = 0\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Training Phase\")\n        for i, (inputs, labels) in enumerate(tqdm(train_loader, desc=\"Training\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            optimizer.zero_grad()\n\n            # Properly use autocast\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n            # Scale loss and backpropagate\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            train_loss += loss.item()\n            batch_count += 1\n\n            # Only run for 1 batch\n            if i == 40:\n                break\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Validation Phase\")\n        model.eval()\n\n        # Disable gradient computation to improve performance\n        correct = 0\n        total = 0\n        for i, (inputs, labels) in enumerate(tqdm(val_loader, desc=\"Validation\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n            _, predicted = torch.max(outputs, 1)\n\n            if labels.dim() > 1:\n                labels = torch.argmax(labels, dim=1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n            # Only run for 1 batch\n            if i == 4:\n                break\n\n        epoch_time = time.time() - start_time\n        avg_loss = train_loss / batch_count\n        val_acc = 100 * correct / total\n        history['train_loss'].append(avg_loss)\n        history['val_acc'].append(val_acc)\n\n        mem = psutil.virtual_memory()\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} | Time: {epoch_time:.1f}s | \"\n              f\"Loss: {avg_loss:.4f} | Val Acc: {val_acc:.2f}% | \"\n              f\"Memory: {mem.percent}% | CPU: {psutil.cpu_percent()}%\")\n\n    return history\n\n# Start training\nhistory = train_lightweight(model, config)  # Use the config instance\n            \n                break\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Validation Phase\")\n        model.eval()\n\n        # Disable gradient computation to improve performance\n        correct = 0\n        total = 0\n        for i, (inputs, labels) in enumerate(tqdm(val_loader, desc=\"Validation\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n            _, predicted = torch.max(outputs, 1)\n\n            if labels.dim() > 1:\n                labels = torch.argmax(labels, dim=1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n            # Only run for 1 batch\n            if i == 4:\n                break\n\n        epoch_time = time.time() - start_time\n        avg_loss = train_loss / batch_count\n        val_acc = 100 * correct / total\n        history['train_loss'].append(avg_loss)\n        history['val_acc'].append(val_acc)\n\n        mem = psutil.virtual_memory()\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} | Time: {epoch_time:.1f}s | \"\n              f\"Loss: {avg_loss:.4f} | Val Acc: {val_acc:.2f}% | \"\n              f\"Memory: {mem.percent}% | CPU: {psutil.cpu_percent()}%\")\n\n    return history\n\n# Start training\nhistory = train_lightweight(model, config)  # Use the config instance","metadata":{"id":"bsdQF8X_7D3k","outputId":"691782ec-d028-4e1a-d06b-3b8808bd00a8","executionInfo":{"status":"error","timestamp":1742267804627,"user_tz":-660,"elapsed":4364335,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:35:06.654305Z","iopub.execute_input":"2025-05-16T02:35:06.654604Z","iopub.status.idle":"2025-05-16T02:35:06.666083Z","shell.execute_reply.started":"2025-05-16T02:35:06.654583Z","shell.execute_reply":"2025-05-16T02:35:06.66501Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch.amp  # Modern version of AMP uses torch.amp instead of torch.cuda.amp\nimport torch.backends.cudnn as cudnn\nimport psutil\nimport time\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nimport torch.nn as nn\n\ncudnn.benchmark = True  # Automatically optimize convolution implementation\n\ndef train_lightweight(model, config):\n    # Initialize data loaders\n    train_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=config.batch_size, pin_memory=True)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=config.lr)\n\n    # Properly initialize GradScaler without parameters\n    scaler = torch.amp.GradScaler()\n\n    # Training history\n    history = {'train_loss': [], 'val_acc': []}\n\n    # Run for config.num_epochs epochs\n    for epoch in range(config.num_epochs):\n        start_time = time.time()\n        model.train()\n        train_loss = 0.0\n        batch_count = 0\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Training Phase\")\n        for i, (inputs, labels) in enumerate(tqdm(train_loader, desc=\"Training\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            optimizer.zero_grad()\n\n            # Properly use autocast\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n            # Scale loss and backpropagate\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            train_loss += loss.item()\n            batch_count += 1\n\n            # Only run for 1 batch\n            if i == 19:\n                break\n\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} - Validation Phase\")\n        model.eval()\n\n        # Disable gradient computation to improve performance\n        correct = 0\n        total = 0\n        for i, (inputs, labels) in enumerate(tqdm(val_loader, desc=\"Validation\")):\n            inputs, labels = inputs.to(config.device), labels.to(config.device)\n            with torch.amp.autocast(device_type=\"cuda\"):\n                outputs = model(inputs)\n            _, predicted = torch.max(outputs, 1)\n\n            if labels.dim() > 1:\n                labels = torch.argmax(labels, dim=1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n            # Only run for 1 batch\n            if i == 4:\n                break\n\n        epoch_time = time.time() - start_time\n        avg_loss = train_loss / batch_count\n        val_acc = 100 * correct / total\n        history['train_loss'].append(avg_loss)\n        history['val_acc'].append(val_acc)\n\n        mem = psutil.virtual_memory()\n        print(f\"Epoch {epoch + 1}/{config.num_epochs} | Time: {epoch_time:.1f}s | \"\n              f\"Loss: {avg_loss:.4f} | Val Acc: {val_acc:.2f}% | \"\n              f\"Memory: {mem.percent}% | CPU: {psutil.cpu_percent()}%\")\n\n    return history\n\n# Start training\nhistory = train_lightweight(model, config)  # Use the config instance","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:41:27.643715Z","iopub.execute_input":"2025-05-16T02:41:27.644379Z","iopub.status.idle":"2025-05-16T02:55:58.535061Z","shell.execute_reply.started":"2025-05-16T02:41:27.644353Z","shell.execute_reply":"2025-05-16T02:55:58.534321Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5.3 test","metadata":{"id":"F9824DIc5R6y"}},{"cell_type":"markdown","source":"Separate test dataset:","metadata":{"id":"mp9Ok-9p5W5Y"}},{"cell_type":"code","source":"dataset = CustomDataset(train_path, train_csv_path, use_clip=True)\nimage, label, clip_feats = dataset[0]\nprint(clip_feats.shape) # Check the CLIP feature dimension","metadata":{"id":"4n0wAmgV5VJB","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:58:23.491762Z","iopub.execute_input":"2025-05-16T02:58:23.492087Z","iopub.status.idle":"2025-05-16T02:58:32.241827Z","shell.execute_reply.started":"2025-05-16T02:58:23.492065Z","shell.execute_reply":"2025-05-16T02:58:32.241094Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Test the data loader separately:","metadata":{"id":"sdR3Y_oi5XyW"}},{"cell_type":"code","source":"loader = DataLoader(dataset, batch_size=2)\nfor data in loader:\n    if config.use_clip:\n        inputs, labels, clip_feats = data\n        print(inputs.shape, labels.shape, clip_feats.shape)\n    else:\n        inputs, labels = data\n        print(inputs.shape, labels.shape)\n    break","metadata":{"id":"meVLFKqk5V8x","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:58:36.104888Z","iopub.execute_input":"2025-05-16T02:58:36.10572Z","iopub.status.idle":"2025-05-16T02:58:36.32045Z","shell.execute_reply.started":"2025-05-16T02:58:36.105693Z","shell.execute_reply":"2025-05-16T02:58:36.319759Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Visualize the training results","metadata":{"id":"PG-Nowa_TWHr"}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"# Draw the loss and accuracy curves\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nplt.plot(history['train_loss'], label='Training Loss')\nplt.title(\"Training Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\n\nplt.subplot(1, 2, 2)\nplt.plot(history['val_acc'], label='Validation Accuracy')\nplt.title(\"Validation Accuracy\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy (%)\")\nplt.show()","metadata":{"id":"XyHd9npG9stE","outputId":"3d8bb09e-6c3c-4d6d-9e83-6a5616194dd9","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:58:38.122577Z","iopub.execute_input":"2025-05-16T02:58:38.123325Z","iopub.status.idle":"2025-05-16T02:58:38.449971Z","shell.execute_reply.started":"2025-05-16T02:58:38.123299Z","shell.execute_reply":"2025-05-16T02:58:38.449134Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Feature Visualization (t-SNE)","metadata":{"id":"8dPv92Ra9t-T"}},{"cell_type":"code","source":"# Install visual dependencies\n! pip install scikit-learn matplotlib\n\n# Import Library\nfrom sklearn.manifold import TSNE\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch.utils.data import DataLoader","metadata":{"id":"YDjQrbYu9xYK","outputId":"3cdf06a1-248c-4a8c-e758-5ad87bdfa53e","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:58:40.703532Z","iopub.execute_input":"2025-05-16T02:58:40.703845Z","iopub.status.idle":"2025-05-16T02:58:44.384244Z","shell.execute_reply.started":"2025-05-16T02:58:40.703822Z","shell.execute_reply":"2025-05-16T02:58:44.383306Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.manifold import TSNE\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\n\ndef visualize_features(model, dataloader):\n    features = []\n    labels = []\n    model.eval()\n    with torch.no_grad():\n        for inputs, targets in dataloader:\n            inputs = inputs.to(Config.device)  # Move data to GPU\n            feats = model(inputs).cpu().numpy()  # Move features back to CPU\n            features.append(feats)\n            labels.append(targets.numpy())\n\n    features = np.concatenate(features)\n    labels = np.concatenate(labels)\n\n    # Use a subset of data (e.g., the first 1000 samples)\n    if len(features) > 1000:\n        features = features[:1000]\n        labels = labels[:1000]\n\n    # If labels are one-hot encoded, convert to class indices\n    if labels.ndim > 1:\n        labels = np.argmax(labels, axis=1)\n\n    # Debug print\n    print(\"Features shape:\", features.shape)  # Expected: (1000, feature_dim)\n    print(\"Labels shape:\", labels.shape)  # Expected: (1000,)\n\n    # t-SNE dimensionality reduction\n    tsne = TSNE(n_components=2, random_state=42, perplexity=30, n_iter=300)  # Use n_iter instead of max_iter\n    reduced = tsne.fit_transform(features)\n\n    # Plot\n    plt.figure(figsize=(10, 8))\n    scatter = plt.scatter(reduced[:, 0], reduced[:, 1], c=labels, cmap='tab10', alpha=0.6)\n    plt.legend(*scatter.legend_elements(), title=\"Classes\")\n    plt.title(\"t-SNE Feature Visualization\")\n    plt.show()\n\n# Execute visualization (using a subset of data)\nsmall_loader = DataLoader(val_dataset, batch_size=64, shuffle=True)\nsmall_loader = list(small_loader)[:16]  # Only take the first 16 batches (about 1000 samples)\nvisualize_features(model, small_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T03:05:39.604048Z","iopub.execute_input":"2025-05-16T03:05:39.60467Z","iopub.status.idle":"2025-05-16T03:10:16.309515Z","shell.execute_reply.started":"2025-05-16T03:05:39.604629Z","shell.execute_reply":"2025-05-16T03:10:16.308706Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"I have a total of 11 categories. Why does this graph only generate the distribution of categories 0 to 9?\n\nPossible reason 1: For some categories, the data volume is relatively small. Only the first 1000 samples were selected and truncated - replaced by 2000 or 3000\n\nPossible reason 2: After PCA dimensionality reduction, some categories overlap - switch to t-SNE or UMAP for nonlinear dimensionality reduction\n\nPossible reason 3: The category index is not consecutive (skipped 10) - Switch to a colormap that supports more categories:","metadata":{"id":"ml9TG5xKuTE1"}},{"cell_type":"markdown","source":"# try","metadata":{"id":"iFS9N4BYezQt"}},{"cell_type":"code","source":"import os\n\n# Check if the file exists\nif os.path.exists('/content/checkpoint.pth'):\n    print(\"Checkpoint exists in /content/!\")\nelse:\n    print(\"Checkpoint missing in /content/.\")","metadata":{"id":"N0SYqxfee6Ht","executionInfo":{"status":"ok","timestamp":1742272123839,"user_tz":-660,"elapsed":60,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"outputId":"2e48f204-7a74-463b-b9ca-e97db8754efc","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:28:54.029404Z","iopub.status.idle":"2025-05-16T02:28:54.029812Z","shell.execute_reply.started":"2025-05-16T02:28:54.029596Z","shell.execute_reply":"2025-05-16T02:28:54.029613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\n\n# Check if checkpoint.pth exists\nif os.path.exists('/kaggle/working/checkpoint.pth'):\n    print(\"Checkpoint exists!\")\n    # Try to load the checkpoint\n    checkpoint = torch.load('/kaggle/working/checkpoint.pth')\n    print(\"Checkpoint keys:\", checkpoint.keys())\nelse:\n    print(\"Checkpoint missing.\")\n\n# Check if test.txt exists\nif os.path.exists('/kaggle/working/test.txt'):\n    print(\"Test file exists!\")\n    with open('/kaggle/working/test.txt', 'r') as f:\n        print(\"Test file content:\", f.read())\nelse:\n    print(\"Test file missing.\")","metadata":{"id":"JDNR2vRjfKY8","executionInfo":{"status":"ok","timestamp":1742271801174,"user_tz":-660,"elapsed":3186,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"outputId":"00485f17-672a-4d3f-f43e-7ab6ef78f3ef","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:28:54.03116Z","iopub.status.idle":"2025-05-16T02:28:54.031501Z","shell.execute_reply.started":"2025-05-16T02:28:54.031305Z","shell.execute_reply":"2025-05-16T02:28:54.031321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nshutil.make_archive('/kaggle/working/checkpoint', 'zip', '/kaggle/working')","metadata":{"id":"x8FTl3eWe2VI","executionInfo":{"status":"ok","timestamp":1742271895703,"user_tz":-660,"elapsed":81486,"user":{"displayName":"Yi Yang","userId":"13590681087948497577"}},"outputId":"c9b981af-a806-4e3e-ab28-a3894a96ee8c","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:28:54.032257Z","iopub.status.idle":"2025-05-16T02:28:54.03251Z","shell.execute_reply.started":"2025-05-16T02:28:54.032395Z","shell.execute_reply":"2025-05-16T02:28:54.032406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_checkpoint(model, optimizer, scaler, epoch, history, filepath=\"/kaggle/working/checkpoint.pth\"):\n    try:\n        # Ensure the parent directory exists\n        os.makedirs(os.path.dirname(filepath), exist_ok=True)\n\n        checkpoint = {\n            \"epoch\": epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"scaler_state_dict\": scaler.state_dict(),\n            \"history\": history\n        }\n        torch.save(checkpoint, filepath)\n        print(f\"Checkpoint saved at epoch {epoch + 1} to {os.path.abspath(filepath)}\")\n    except Exception as e:\n        print(f\"Error saving checkpoint: {e}\")","metadata":{"id":"7otsq03oftIy","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T02:28:54.034022Z","iopub.status.idle":"2025-05-16T02:28:54.034312Z","shell.execute_reply.started":"2025-05-16T02:28:54.034176Z","shell.execute_reply":"2025-05-16T02:28:54.034187Z"}},"outputs":[],"execution_count":null}]}