{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71447,"databundleVersionId":8208918,"sourceType":"competition"},{"sourceId":438523,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":357771,"modelId":379106},{"sourceId":438644,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":357871,"modelId":379212},{"sourceId":449986,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":365245,"modelId":386127}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport random\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport sklearn\nfrom sklearn.metrics import mean_absolute_error\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport pydicom\nimport warnings\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\nfrom sklearn.preprocessing import StandardScaler\nimport seaborn as sns\n\n# 忽略警告\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nos.environ[\"NO_ALBUMENTATIONS_UPDATE\"] = \"1\"  # 关闭Albumentations更新检查\n\n# 随机种子设置函数\ndef seed_everything(seed):\n    \"\"\"设置所有随机数生成器的种子，确保实验可重复\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n# SuperDAM模型定义\nclass DGM(nn.Module):\n    def __init__(self, in_channels, reduction=4, groups=4):\n        super(DGM, self).__init__()\n        self.groups = groups\n        mid_channels = max(1, in_channels // reduction)\n        self.weight_layer = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, 1, groups=self.groups, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, in_channels, 1, groups=self.groups, bias=False),\n            nn.Sigmoid()\n        )\n        self.pointwise_groups = nn.Conv2d(in_channels, in_channels, 1, groups=self.groups, bias=False)\n\n    def forward(self, x):\n        weights = self.weight_layer(x)\n        x = x * weights\n        x = self.pointwise_groups(x)\n        return x\n\nclass GAM_Attention(nn.Module):\n    def __init__(self, in_channels):\n        super(GAM_Attention, self).__init__()\n        self.global_avgpool = nn.AdaptiveAvgPool2d(1)\n        self.channel_attention = nn.Sequential(\n            nn.Linear(in_channels, in_channels // 16),\n            nn.ReLU(inplace=True),\n            nn.Linear(in_channels // 16, in_channels),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, h, w = x.size()\n        x_global = self.global_avgpool(x).view(b, c)\n        x_channel_att = self.channel_attention(x_global).view(b, c, 1, 1)\n        x = x * x_channel_att\n        return x\n\nclass LinearBottleNeck_1(nn.Module):\n    def __init__(self, in_c, out_c, s, t):\n        super().__init__()\n        self.residual = nn.Sequential(\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=s, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        self.stride = s\n        self.in_channels = in_c\n        self.out_channels = out_c\n        self.attention = GAM_Attention(out_c)\n\n    def forward(self, x):\n        residual = self.residual(x)\n        if self.stride == 1 and self.in_channels == self.out_channels:\n            residual += x\n        residual = self.attention(residual)\n        return residual\n\nclass LinearBottleNeck_2(nn.Module):\n    def __init__(self, in_c, out_c, s, t):\n        super().__init__()\n        self.residual = nn.Sequential(\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=s, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        self.residual_1 = nn.Sequential(\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 5, stride=s, padding=2, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        self.residual_2 = nn.Sequential(\n            nn.Conv2d(in_c, out_c, 1, stride=2),\n            nn.BatchNorm2d(out_c)\n        )\n        self.stride = s\n        self.in_channels = in_c\n        self.out_channels = out_c\n\n    def forward(self, x):\n        residual = self.residual(x)\n        residual_1 = self.residual_1(x)\n        residual_2 = self.residual_2(x)\n        out_feature = residual_1 + residual + residual_2\n        return out_feature\n\nclass SuperDAM(nn.Module):\n    def __init__(self, class_num=1, use_multi_slice=True):\n        super().__init__()\n        self.use_multi_slice = use_multi_slice\n        self.pre = nn.Sequential(\n            nn.Conv2d(3, 32, 3, stride=2, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU6(inplace=True),\n        )\n        self.stage1 = nn.Sequential(\n            LinearBottleNeck_1(32, 16, 1, 1),\n            DGM(16)\n        )\n        self.stage2 = self.make_stage(2, 16, 24, 2, 6)\n        self.stage3 = self.make_stage(3, 24, 32, 2, 6)\n        self.stage4 = nn.Sequential(\n            self.make_stage(4, 32, 64, 2, 6),\n            DGM(64)\n        )\n        self.stage5 = self.make_stage(3, 64, 96, 1, 6)\n        self.stage6 = nn.Sequential(\n            self.make_stage(3, 96, 160, 2, 6),\n            DGM(160)\n        )\n        self.stage7 = LinearBottleNeck_1(160, 320, 1, 6)\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(320, 1280, 1),\n            nn.BatchNorm2d(1280),\n            nn.ReLU6(inplace=True),\n            DGM(1280)\n        )\n        self.conv2 = nn.Conv2d(1280, class_num, 1)\n\n    def forward(self, x):\n        if self.use_multi_slice:\n            batch_size, n_patches, channels, height, width = x.size()\n            x = x.view(batch_size * n_patches, channels, height, width)\n        x = self.pre(x)\n        x = self.stage1(x)\n        x = self.stage2(x)\n        x = self.stage3(x)\n        x = self.stage4(x)\n        x = self.stage5(x)\n        x = self.stage6(x)\n        x = self.stage7(x)\n        x = self.conv1(x)\n        x = F.adaptive_avg_pool2d(x, 1)\n        x = self.conv2(x)\n        if self.use_multi_slice:\n            x = x.view(batch_size, n_patches, -1)\n            x = x.mean(dim=1)\n        return x\n\n    def make_stage(self, repeat, in_c, out_c, s, t):\n        layers = []\n        if s == 1:\n            layers.append(LinearBottleNeck_1(in_c, out_c, s, t))\n        else:\n            layers.append(LinearBottleNeck_2(in_c, out_c, s, t))\n        while repeat - 1:\n            layers.append(LinearBottleNeck_1(out_c, out_c, 1, t))\n            repeat -= 1\n        return nn.Sequential(*layers)\n\n# 优化后的DICOM转图像函数（动态窗口+标准化）\ndef load_dicom_to_image(dicom_path):\n    try:\n        ds = pydicom.dcmread(dicom_path, force=True)\n        img = ds.pixel_array.astype(np.float32)\n        \n        # 自动计算窗宽窗位（5-95分位数）\n        valid_pixels = img[img != -2000]  # 假设-2000为无效值\n        if valid_pixels.size == 0:\n            return np.zeros((512, 512, 3), dtype=np.float32)\n        img_min = np.percentile(valid_pixels, 5)\n        img_max = np.percentile(valid_pixels, 95)\n        \n        img = (img - img_min) / (img_max - img_min)  # 归一化到[0,1]\n        img = (img * 255).astype(np.float32)  # 转换为浮点型\n        return cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n    \n    except Exception as e:\n        print(f\"Error loading DICOM: {dicom_path}, {e}\")\n        return np.zeros((512, 512, 3), dtype=np.float32)\n\n# 数据加载与预处理\nclass BrainAgeDataset(Dataset):\n    def __init__(self, df, cfg, transform=None, mode=\"train\"):\n        self.df = df\n        self.cfg = cfg\n        self.transform = transform\n        self.mode = mode\n        self.extensions = (\".dcm\",)\n        self.valid_study_ids = []\n        \n        if mode == \"train\":\n            self.study_ids = df['StudyID'].unique()\n            self.label_dict = df.set_index('StudyID')['Age'].to_dict()\n            self.study_to_folder = self._build_study_to_folder_map()\n            self.valid_study_ids = [sid for sid in self.study_ids if sid in self.study_to_folder]\n            print(f\"有效训练样本数: {len(self.valid_study_ids)}\")\n        else:\n            self.img_ids = df['StudyID'].values\n    \n    def __len__(self):\n        return len(self.valid_study_ids) if self.mode == \"train\" else len(self.img_ids)\n    \n    def _build_study_to_folder_map(self):\n        study_to_folder = {}\n        base_dir = os.path.join(self.cfg.root, \"dataset_jpr_train\", \"dataset_jpr_train\")\n        try:\n            for folder_name in os.listdir(base_dir):\n                folder_path = os.path.join(base_dir, folder_name)\n                if not os.path.isdir(folder_path):\n                    continue\n                for study_id in os.listdir(folder_path):\n                    study_to_folder[study_id] = folder_name\n        except Exception as e:\n            print(f\"构建映射失败: {e}\")\n        print(f\"成功映射 {len(study_to_folder)} 个StudyID\")\n        return study_to_folder\n    \n    def __getitem__(self, idx):\n        if self.mode == \"train\":\n            study_id = self.valid_study_ids[idx]\n            label = self.label_dict[study_id]\n            folder_name = self.study_to_folder[study_id]\n            base_path = os.path.join(\n                self.cfg.root, \"dataset_jpr_train\", \"dataset_jpr_train\", folder_name, study_id\n            )\n            \n            files = [os.path.join(base_path, f) for f in os.listdir(base_path) if f.lower().endswith(self.extensions)]\n            files.sort()  # 排序后取中间切片\n            n_patches = self.cfg.n_patches if self.cfg.use_multi_slice else 1\n            start_idx = max(0, len(files) // 2 - n_patches // 2)\n            selected_files = files[start_idx:start_idx + n_patches]\n            \n            # 补全切片\n            selected_files += [None] * (n_patches - len(selected_files))\n            \n            imgs = []\n            for path in selected_files:\n                img = np.zeros((self.cfg.img_size, self.cfg.img_size, 3)) if path is None else load_dicom_to_image(path)\n                if self.transform:\n                    img = self.transform(image=img)[\"image\"]\n                imgs.append(img)\n            \n            return torch.stack(imgs, dim=0), torch.tensor(label, dtype=torch.float32)\n        \n        else:\n            study_id = self.img_ids[idx]\n            base_path = os.path.join(\n                self.cfg.test_root, \"dataset_jpr_test2\", \"dataset_jpr_test2\", study_id\n            )\n            \n            files = [os.path.join(base_path, f) for f in os.listdir(base_path) if f.lower().endswith(self.extensions)]\n            img_path = files[0] if files else None\n            img = np.zeros((self.cfg.img_size, self.cfg.img_size, 3)) if img_path is None else load_dicom_to_image(img_path)\n            \n            if self.transform:\n                img = self.transform(image=img)[\"image\"]\n            \n            return img.unsqueeze(0), torch.tensor(0.0, dtype=torch.float32)\n\n# 训练函数\ndef train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device, cfg):\n    model.train()\n    scaler = GradScaler(enabled=cfg.amp and device.type == \"cuda\")\n    losses = []\n    \n    for step, (images, labels) in enumerate(train_loader):\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        with autocast(enabled=cfg.amp and device.type == \"cuda\"):\n            y_preds = model(images)\n            loss = criterion(y_preds.view(-1), labels)\n        \n        losses.append(loss.item())\n        \n        if scaler is not None:\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ga_accum == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n        else:\n            loss.backward()\n            if (step + 1) % cfg.ga_accum == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n        \n        if scheduler is not None and isinstance(scheduler, lr_scheduler.LRScheduler):\n            scheduler.step()\n    \n    return np.mean(losses)\n\n# 验证函数\ndef valid_fn(valid_loader, model, criterion, device, cfg):\n    model.eval()\n    preds = []\n    losses = []\n    \n    with torch.no_grad():\n        for images, labels in valid_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            \n            with autocast(enabled=cfg.amp and device.type == \"cuda\"):\n                y_preds = model(images)\n                loss = criterion(y_preds.view(-1), labels)\n            \n            losses.append(loss.item())\n            preds.append(y_preds.detach().cpu().numpy())\n    \n    preds = np.concatenate(preds)\n    return np.mean(losses), preds\n\n# 特征提取模型包装器\nclass FeatureExtractor(nn.Module):\n    def __init__(self, model, layer_name='conv1'):\n        super().__init__()\n        self.features = nn.Sequential(\n            *list(model.children())[:list(model.children()).index(getattr(model, layer_name))+1]\n        )\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        \n    def forward(self, x):\n        x = self.features(x)\n        x = self.pooling(x)\n        x = x.view(x.size(0), -1)\n        return x\n\n# 特征提取函数\ndef extract_features(data_loader, model, device):\n    model.eval()\n    features = []\n    labels = []\n    \n    with torch.no_grad():\n        for images, target in data_loader:\n            images = images.to(device)\n            batch_size, n_patches, channels, height, width = images.size()\n            images = images.view(batch_size * n_patches, channels, height, width)\n            \n            # 提取特征\n            batch_features = model(images)\n            batch_features = batch_features.view(batch_size, n_patches, -1).mean(dim=1)  # 多切片特征平均\n            \n            features.append(batch_features.cpu().numpy())\n            labels.append(target.cpu().numpy())\n    \n    features = np.vstack(features)\n    labels = np.hstack(labels)\n    return features, labels\n\n# TSNE可视化函数\ndef visualize_tsne(features, labels, save_path='tsne_visualization.png'):\n    # 数据标准化\n    scaler = StandardScaler()\n    scaled_features = scaler.fit_transform(features)\n    \n    # 使用TSNE降维\n    tsne = TSNE(n_components=2, perplexity=30, n_iter=3000, random_state=42)\n    tsne_results = tsne.fit_transform(scaled_features)\n    \n    # 创建DataFrame用于绘图\n    df = pd.DataFrame({\n        'tsne-2d-one': tsne_results[:,0],\n        'tsne-2d-two': tsne_results[:,1],\n        'Age': labels\n    })\n    \n    # 创建年龄分组\n    age_bins = [0, 20, 40, 60, 80, 100]\n    age_labels = ['0-20', '21-40', '41-60', '61-80', '81+']\n    df['Age Group'] = pd.cut(df['Age'], bins=age_bins, labels=age_labels, include_lowest=True)\n    \n    # 设置中文字体\n    plt.rcParams[\"font.family\"] = [\"SimHei\", \"WenQuanYi Micro Hei\", \"Heiti TC\"]\n    \n    # 绘制TSNE图\n    plt.figure(figsize=(12, 10))\n    sns.scatterplot(\n        x=\"tsne-2d-one\", y=\"tsne-2d-two\",\n        hue=\"Age Group\",\n        palette=sns.color_palette(\"hls\", len(age_labels)),\n        data=df,\n        legend=\"full\",\n        alpha=0.7\n    )\n    \n    plt.title('t-SNE可视化：基于脑部CT的年龄特征分布')\n    plt.xlabel('t-SNE维度1')\n    plt.ylabel('t-SNE维度2')\n    \n    # 添加年龄分布直方图\n    plt.figure(figsize=(10, 6))\n    sns.histplot(df['Age'], bins=20, kde=True)\n    plt.title('数据集中的年龄分布')\n    plt.xlabel('年龄')\n    plt.ylabel('样本数量')\n    \n    # 保存图像\n    plt.tight_layout()\n    plt.savefig(save_path)\n    print(f\"TSNE可视化结果已保存至 {save_path}\")\n    \n    return df\n\ndef main():\n    class CFG:\n        seed = 42\n        n_fold = 1          # 设置为1折\n        use_folds = [0]     # 仅使用第0折\n        img_size = 224\n        batch_size = 8\n        epochs = 50         # 训练50轮\n        lr = 5e-5\n        min_lr = 1e-6\n        amp = True\n        ga_accum = 2\n        num_workers = 4\n        root = \"/kaggle/input/spr-head-ct-age-prediction-challenge\"\n        test_root = \"/kaggle/input/spr-head-ct-age-prediction-challenge\"\n        fold_csv = \"/kaggle/input/spr-head-ct-age-prediction-challenge/train.csv\"\n        test_csv = \"/kaggle/input/spr-head-ct-age-prediction-challenge/test.csv\"\n        pretrained_weights = \"/kaggle/input/1111/pytorch/default/1/newmodel_weights.pth\"\n        n_patches = 10\n        use_multi_slice = True\n        target_column = 'Age'\n        weight_decay = 1e-5\n        warmup_epochs = 5    # 保持热身轮数\n        scheduler_name = \"onecycle\"\n        train_mode = \"train\"\n        drop_last = True\n        val_ratio = 0.2      # 验证集比例20%\n        visualize_tsne = True  # 是否执行TSNE可视化\n    \n    cfg = CFG()\n    seed_everything(cfg.seed)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    print(f\"PyTorch version: {torch.__version__}\")\n    print(f\"CUDA available: {torch.cuda.is_available()}\")\n    print(f\"Device: {device}\")\n    \n    train_df = pd.read_csv(cfg.fold_csv)\n    test_df = pd.read_csv(cfg.test_csv)\n    \n    # 直接划分训练集和验证集，不再使用交叉验证\n    train_df, valid_df = train_test_split(\n        train_df, \n        test_size=cfg.val_ratio, \n        random_state=cfg.seed,\n        shuffle=True\n    )\n    print(f\"训练集样本数: {len(train_df)}, 验证集样本数: {len(valid_df)}\")\n    \n    # 数据增强配置\n    train_transform = A.Compose([\n        A.Resize(cfg.img_size, cfg.img_size),\n        A.HorizontalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=20, p=0.7),\n        A.ElasticTransform(p=0.3, alpha=120, sigma=120*0.05, alpha_affine=120*0.03),\n        A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n        ToTensorV2(),\n    ])\n    \n    valid_transform = A.Compose([\n        A.Resize(cfg.img_size, cfg.img_size),\n        A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n        ToTensorV2(),\n    ])\n    \n    # 创建数据集和数据加载器\n    train_dataset = BrainAgeDataset(train_df, cfg, train_transform, mode=\"train\")\n    valid_dataset = BrainAgeDataset(valid_df, cfg, valid_transform, mode=\"train\")\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=cfg.batch_size,\n        shuffle=True,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n        drop_last=cfg.drop_last,\n    )\n    \n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=cfg.batch_size,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n        drop_last=False,\n    )\n    \n    # 获取验证集有效标签\n    valid_indices = valid_dataset.valid_study_ids\n    valid_labels = np.array([valid_dataset.label_dict[sid] for sid in valid_indices])\n    print(f\"原始验证样本数: {len(valid_df)}, 实际验证样本数: {len(valid_labels)}\")\n    \n    # 初始化模型\n    model = SuperDAM(class_num=1, use_multi_slice=cfg.use_multi_slice)\n    if cfg.pretrained_weights and os.path.exists(cfg.pretrained_weights):\n        try:\n            model.load_state_dict(torch.load(cfg.pretrained_weights, map_location=device), strict=False)\n            print(\"加载预训练权重成功\")\n        except Exception as e:\n            print(f\"加载预训练权重失败: {e}\")\n    model.to(device)\n    \n    # 定义损失函数和优化器\n    criterion = nn.L1Loss()\n    optimizer = optim.AdamW(\n        model.parameters(), \n        lr=cfg.lr, \n        weight_decay=cfg.weight_decay\n    )\n    \n    # 学习率调度器（OneCycleLR）\n    if cfg.scheduler_name == \"onecycle\":\n        steps_per_epoch = len(train_loader)\n        total_steps = ((len(train_dataset) + cfg.batch_size - 1) // cfg.batch_size) * cfg.epochs\n        pct_start = cfg.warmup_epochs / cfg.epochs\n        \n        scheduler = lr_scheduler.OneCycleLR(\n            optimizer,\n            max_lr=cfg.lr * 5,\n            steps_per_epoch=steps_per_epoch,\n            epochs=cfg.epochs,\n            total_steps=total_steps,\n            pct_start=pct_start,\n            anneal_strategy='cos'\n        )\n    else:\n        scheduler = None\n    \n    # 初始化最佳MAE\n    best_mae = float('inf')\n    \n    for epoch in range(cfg.epochs):\n        start_time = time.time()\n        \n        # 训练一轮\n        avg_train_loss = train_fn(\n            train_loader=train_loader,\n            model=model,\n            criterion=criterion,\n            optimizer=optimizer,\n            epoch=epoch,\n            scheduler=scheduler,\n            device=device,\n            cfg=cfg\n        )\n        \n        # 验证一轮\n        avg_val_loss, preds = valid_fn(\n            valid_loader=valid_loader,\n            model=model,\n            criterion=criterion,\n            device=device,\n            cfg=cfg\n        )\n        \n        # 计算MAE\n        mae = mean_absolute_error(valid_labels, preds)\n        \n        # 更新学习率调度器\n        if scheduler:\n            scheduler.step()\n        \n        # 保存最佳模型\n        if mae < best_mae:\n            best_mae = mae\n            torch.save(model.state_dict(), f\"best_model.pth\")\n            print(f\"在第 {epoch+1} 轮保存最佳模型，MAE: {best_mae:.4f}\")\n        \n        # 打印训练信息\n        elapsed = time.time() - start_time\n        print(f\"Epoch {epoch+1}/{cfg.epochs} | 训练损失: {avg_train_loss:.4f} | 验证损失: {avg_val_loss:.4f} | MAE: {mae:.4f} | 耗时: {elapsed:.2f}s\")\n    \n    # 测试阶段\n    if cfg.train_mode == \"test\":\n        test_transform = A.Compose([\n            A.Resize(cfg.img_size, cfg.img_size),\n            A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n            ToTensorV2(),\n        ])\n        \n        test_dataset = BrainAgeDataset(test_df, cfg, test_transform, mode=\"test\")\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=cfg.batch_size,\n            shuffle=False,\n            num_workers=cfg.num_workers,\n            pin_memory=True,\n        )\n        \n        model = SuperDAM(class_num=1, use_multi_slice=cfg.use_multi_slice)\n        model.to(device)\n        model.load_state_dict(torch.load(\"best_model.pth\"))\n        model.eval()\n        \n        all_preds = []\n        with torch.no_grad():\n            for images, _ in test_loader:\n                images = images.to(device)\n                y_preds = model(images)\n                all_preds.append(y_preds.detach().cpu().numpy())\n        \n        final_preds = np.concatenate(all_preds).reshape(-1)\n        \n        submission = pd.DataFrame({\n            'ID': test_df['StudyID___SeriesID'],\n            'Age': final_preds\n        })\n        \n        submission.to_csv('submission.csv', index=False)\n        print(\"预测结果已保存至submission.csv\")\n    \n    # 执行TSNE可视化\n    if cfg.visualize_tsne:\n        print(\"\\n开始TSNE可视化...\")\n        # 加载最佳模型\n        model.load_state_dict(torch.load(\"best_model.pth\"))\n        \n        # 创建特征提取器\n        feature_extractor = FeatureExtractor(model).to(device)\n        \n        # 提取特征\n        print(\"正在提取特征...\")\n        features, labels = extract_features(valid_loader, feature_extractor, device)\n        print(f\"提取完成，特征维度: {features.shape}\")\n        \n        # 可视化\n        print(\"正在生成TSNE可视化...\")\n        visualize_tsne(features, labels)\n\nif __name__ == \"__main__\":\n    main()    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T06:18:44.691972Z","iopub.execute_input":"2025-06-26T06:18:44.692787Z","iopub.status.idle":"2025-06-26T08:55:15.160652Z","shell.execute_reply.started":"2025-06-26T06:18:44.692748Z","shell.execute_reply":"2025-06-26T08:55:15.159745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport random\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport sklearn\nfrom sklearn.metrics import mean_absolute_error\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport pydicom\nimport warnings\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\nfrom sklearn.preprocessing import StandardScaler\nimport seaborn as sns\n\n# Ignore warnings\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nos.environ[\"NO_ALBUMENTATIONS_UPDATE\"] = \"1\"\n\n# Set default font to English\nplt.rcParams[\"font.family\"] = [\"DejaVu Sans\", \"Arial\", \"sans-serif\"]\n\n# Random seed setting function\ndef seed_everything(seed):\n    \"\"\"Set seeds for all random number generators to ensure reproducibility\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n# Enhanced DGM attention module\nclass EnhancedDGM(nn.Module):\n    def __init__(self, in_channels, reduction=4, groups=4):\n        super(EnhancedDGM, self).__init__()\n        self.groups = groups\n        mid_channels = max(1, in_channels // reduction)\n        \n        # Spatial attention branch\n        self.spatial_attention = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, 1, groups=self.groups, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, in_channels, 1, groups=self.groups, bias=False),\n            nn.Sigmoid()\n        )\n        \n        # Channel attention branch\n        self.channel_attention = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channels, mid_channels, 1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, in_channels, 1, bias=False),\n            nn.Sigmoid()\n        )\n        \n        self.pointwise_groups = nn.Conv2d(in_channels, in_channels, 1, groups=self.groups, bias=False)\n\n    def forward(self, x):\n        # Apply spatial and channel attention\n        spatial_weights = self.spatial_attention(x)\n        channel_weights = self.channel_attention(x)\n        weights = spatial_weights * channel_weights\n        \n        x = x * weights\n        x = self.pointwise_groups(x)\n        return x\n\n# Enhanced GAM attention\nclass EnhancedGAM_Attention(nn.Module):\n    def __init__(self, in_channels):\n        super(EnhancedGAM_Attention, self).__init__()\n        self.global_avgpool = nn.AdaptiveAvgPool2d(1)\n        \n        # Multi-layer channel attention\n        self.channel_attention = nn.Sequential(\n            nn.Linear(in_channels, in_channels // 16),\n            nn.BatchNorm1d(in_channels // 16),\n            nn.ReLU(inplace=True),\n            nn.Linear(in_channels // 16, in_channels),\n            nn.Sigmoid()\n        )\n        \n        # Spatial attention\n        self.spatial_attention = nn.Sequential(\n            nn.Conv2d(in_channels, in_channels // 4, kernel_size=7, padding=3),\n            nn.BatchNorm2d(in_channels // 4),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(in_channels // 4, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, h, w = x.size()\n        \n        # Channel attention\n        x_global = self.global_avgpool(x).view(b, c)\n        x_channel_att = self.channel_attention(x_global).view(b, c, 1, 1)\n        \n        # Spatial attention\n        x_spatial_att = self.spatial_attention(x)\n        \n        # Combined attention\n        x = x * x_channel_att * x_spatial_att.expand_as(x)\n        return x\n\n# Improved linear bottleneck module\nclass ImprovedLinearBottleNeck_1(nn.Module):\n    def __init__(self, in_c, out_c, s, t):\n        super().__init__()\n        self.residual = nn.Sequential(\n            # First 1x1 convolution to expand channels\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            \n            # Depthwise separable convolution\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=s, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            \n            # Second 1x1 convolution to compress channels\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            \n            # Final 1x1 convolution to adjust output channels\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        \n        self.stride = s\n        self.in_channels = in_c\n        self.out_channels = out_c\n        self.attention = EnhancedGAM_Attention(out_c)  # Use enhanced attention\n\n    def forward(self, x):\n        residual = self.residual(x)\n        if self.stride == 1 and self.in_channels == self.out_channels:\n            residual += x\n        residual = self.attention(residual)\n        return residual\n\n# Improved multi-scale linear bottleneck module\nclass ImprovedLinearBottleNeck_2(nn.Module):\n    def __init__(self, in_c, out_c, s, t):\n        super().__init__()\n        \n        # Branch 1: 3x3 convolution\n        self.branch1 = nn.Sequential(\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=s, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        \n        # Branch 2: 5x5 convolution (using two 3x3 convolutions instead)\n        self.branch2 = nn.Sequential(\n            nn.Conv2d(in_c, in_c * t, 1),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=s, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 3, stride=1, padding=1, groups=in_c * t),\n            nn.BatchNorm2d(in_c * t),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(in_c * t, in_c * t, 1, stride=1, padding=0, groups=1),\n            nn.BatchNorm2d(in_c * t),\n            nn.Conv2d(in_c * t, out_c, 1),\n            nn.BatchNorm2d(out_c)\n        )\n        \n        # Shortcut connection\n        self.shortcut = nn.Sequential()\n        if s != 1 or in_c != out_c:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_c, out_c, 1, stride=s),\n                nn.BatchNorm2d(out_c)\n            )\n\n    def forward(self, x):\n        branch1 = self.branch1(x)\n        branch2 = self.branch2(x)\n        shortcut = self.shortcut(x)\n        out = branch1 + branch2 + shortcut\n        return out\n\n# Age-sensitive attention module\nclass AgeSensitiveAttention(nn.Module):\n    def __init__(self, in_channels, num_age_groups=5):\n        super().__init__()\n        self.num_age_groups = num_age_groups\n        \n        # Global feature extraction\n        self.global_pool = nn.AdaptiveAvgPool2d(1)\n        \n        # Age classifier\n        self.age_classifier = nn.Sequential(\n            nn.Conv2d(in_channels, 256, 1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, num_age_groups, 1)\n        )\n        \n        # Age-conditioned attention weights\n        self.attention_gates = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(in_channels, in_channels // 4, 1),\n                nn.BatchNorm2d(in_channels // 4),\n                nn.ReLU(inplace=True),\n                nn.Conv2d(in_channels // 4, in_channels, 1),\n                nn.Sigmoid()\n            ) for _ in range(num_age_groups)\n        ])\n    \n    def forward(self, x):\n        batch_size, channels, height, width = x.size()\n        \n        # Extract global features and predict age group\n        global_features = self.global_pool(x)  # [B, C, 1, 1]\n        age_logits = self.age_classifier(global_features).view(batch_size, self.num_age_groups)  # [B, num_groups]\n        age_probs = F.softmax(age_logits, dim=1)  # [B, num_groups]\n        \n        # Generate attention weights for each age group\n        attention_weights = torch.zeros(batch_size, channels, 1, 1, device=x.device)\n        for i in range(self.num_age_groups):\n            attention_weights += age_probs[:, i].view(batch_size, 1, 1, 1) * self.attention_gates[i](global_features)\n        \n        # Apply attention weights\n        x = x * attention_weights.expand_as(x)\n        \n        return x, age_logits\n\n# Enhanced SuperDAM model\nclass EnhancedSuperDAM(nn.Module):\n    def __init__(self, class_num=1, use_multi_slice=True, num_age_groups=5):\n        super().__init__()\n        self.use_multi_slice = use_multi_slice\n        \n        # Enhanced feature extraction front-end\n        self.pre = nn.Sequential(\n            nn.Conv2d(3, 64, 3, stride=2, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(64, 64, 3, stride=1, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU6(inplace=True),\n            EnhancedDGM(64)  # Use enhanced DGM attention\n        )\n        \n        # Stage 1\n        self.stage1 = nn.Sequential(\n            ImprovedLinearBottleNeck_1(64, 32, 1, 1),\n            EnhancedDGM(32)\n        )\n        \n        # Stage 2\n        self.stage2 = self.make_stage(2, 32, 48, 2, 6)\n        \n        # Stage 3\n        self.stage3 = self.make_stage(3, 48, 64, 2, 6)\n        \n        # Stage 4\n        self.stage4 = nn.Sequential(\n            self.make_stage(4, 64, 128, 2, 6),\n            EnhancedDGM(128)\n        )\n        \n        # Stage 5\n        self.stage5 = self.make_stage(3, 128, 192, 1, 6)\n        \n        # Stage 6\n        self.stage6 = nn.Sequential(\n            self.make_stage(3, 192, 320, 2, 6),\n            EnhancedDGM(320)\n        )\n        \n        # Stage 7\n        self.stage7 = ImprovedLinearBottleNeck_1(320, 512, 1, 6)\n        \n        # Age-sensitive attention\n        self.age_attention = AgeSensitiveAttention(512, num_age_groups)\n        \n        # Feature fusion\n        self.feature_fusion = nn.Sequential(\n            nn.Conv2d(512, 1024, 3, padding=1),\n            nn.BatchNorm2d(1024),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(1024, 1024, 3, padding=1),\n            nn.BatchNorm2d(1024),\n            nn.ReLU6(inplace=True),\n            EnhancedDGM(1024)\n        )\n        \n        # Regression head\n        self.regression_head = nn.Sequential(\n            nn.Conv2d(1024, 1280, 1),\n            nn.BatchNorm2d(1280),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(1280, class_num, 1)\n        )\n        \n        # Age group prediction head (auxiliary task)\n        self.age_group_head = nn.Sequential(\n            nn.Conv2d(1024, 512, 1),\n            nn.BatchNorm2d(512),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(512, num_age_groups, 1)\n        )\n\n    def forward(self, x):\n        if self.use_multi_slice:\n            batch_size, n_patches, channels, height, width = x.size()\n            x = x.view(batch_size * n_patches, channels, height, width)\n        \n        # Feature extraction\n        x = self.pre(x)\n        x = self.stage1(x)\n        x = self.stage2(x)\n        x = self.stage3(x)\n        x = self.stage4(x)\n        x = self.stage5(x)\n        x = self.stage6(x)\n        x = self.stage7(x)\n        \n        # Apply age-sensitive attention\n        x, age_logits = self.age_attention(x)\n        \n        # Feature fusion\n        x = self.feature_fusion(x)\n        \n        # Regression prediction\n        x_reg = self.regression_head(x)\n        x_reg = F.adaptive_avg_pool2d(x_reg, 1)\n        \n        # Age group prediction (auxiliary task)\n        x_age = self.age_group_head(x)\n        x_age = F.adaptive_avg_pool2d(x_age, 1)\n        \n        if self.use_multi_slice:\n            x_reg = x_reg.view(batch_size, n_patches, -1)\n            x_reg = x_reg.mean(dim=1)\n            \n            x_age = x_age.view(batch_size, n_patches, -1)\n            x_age = x_age.mean(dim=1)\n        \n        return x_reg, x_age, age_logits  # Return age prediction, age group prediction, and attention age prediction\n\n    def make_stage(self, repeat, in_c, out_c, s, t):\n        layers = []\n        if s == 1:\n            layers.append(ImprovedLinearBottleNeck_1(in_c, out_c, s, t))\n        else:\n            layers.append(ImprovedLinearBottleNeck_2(in_c, out_c, s, t))\n        while repeat - 1:\n            layers.append(ImprovedLinearBottleNeck_1(out_c, out_c, 1, t))\n            repeat -= 1\n        return nn.Sequential(*layers)\n\n# Optimized DICOM to image function (dynamic window + normalization)\ndef load_dicom_to_image(dicom_path):\n    try:\n        ds = pydicom.dcmread(dicom_path, force=True)\n        img = ds.pixel_array.astype(np.float32)\n        \n        # Automatically calculate window width and level (5-95 percentile)\n        valid_pixels = img[img != -2000]  # Assume -2000 is invalid value\n        if valid_pixels.size == 0:\n            return np.zeros((512, 512, 3), dtype=np.float32)\n        img_min = np.percentile(valid_pixels, 5)\n        img_max = np.percentile(valid_pixels, 95)\n        \n        img = (img - img_min) / (img_max - img_min)  # Normalize to [0,1]\n        img = np.clip(img, 0, 1)  # Clip to [0,1] range\n        img = (img * 255).astype(np.float32)  # Convert to float\n        \n        # Convert to RGB\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        \n        # Enhance contrast (optional)\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n        img_rgb[:, :, 0] = clahe.apply(img_rgb[:, :, 0].astype(np.uint8)).astype(np.float32)\n        img_rgb[:, :, 1] = clahe.apply(img_rgb[:, :, 1].astype(np.uint8)).astype(np.float32)\n        img_rgb[:, :, 2] = clahe.apply(img_rgb[:, :, 2].astype(np.uint8)).astype(np.float32)\n        \n        return img_rgb\n    \n    except Exception as e:\n        print(f\"Error loading DICOM: {dicom_path}, {e}\")\n        return np.zeros((512, 512, 3), dtype=np.float32)\n\n# Data loading and preprocessing\nclass BrainAgeDataset(Dataset):\n    def __init__(self, df, cfg, transform=None, mode=\"train\"):\n        self.df = df\n        self.cfg = cfg\n        self.transform = transform\n        self.mode = mode\n        self.extensions = (\".dcm\",)\n        self.valid_study_ids = []\n        \n        if mode == \"train\":\n            self.study_ids = df['StudyID'].unique()\n            self.label_dict = df.set_index('StudyID')['Age'].to_dict()\n            self.study_to_folder = self._build_study_to_folder_map()\n            self.valid_study_ids = [sid for sid in self.study_ids if sid in self.study_to_folder]\n            print(f\"Number of valid training samples: {len(self.valid_study_ids)}\")\n        else:\n            self.img_ids = df['StudyID'].values\n    \n    def __len__(self):\n        return len(self.valid_study_ids) if self.mode == \"train\" else len(self.img_ids)\n    \n    def _build_study_to_folder_map(self):\n        study_to_folder = {}\n        base_dir = os.path.join(self.cfg.root, \"dataset_jpr_train\", \"dataset_jpr_train\")\n        try:\n            for folder_name in os.listdir(base_dir):\n                folder_path = os.path.join(base_dir, folder_name)\n                if not os.path.isdir(folder_path):\n                    continue\n                for study_id in os.listdir(folder_path):\n                    study_to_folder[study_id] = folder_name\n        except Exception as e:\n            print(f\"Failed to build mapping: {e}\")\n        print(f\"Successfully mapped {len(study_to_folder)} StudyIDs\")\n        return study_to_folder\n    \n    def __getitem__(self, idx):\n        if self.mode == \"train\":\n            study_id = self.valid_study_ids[idx]\n            label = self.label_dict[study_id]\n            folder_name = self.study_to_folder[study_id]\n            base_path = os.path.join(\n                self.cfg.root, \"dataset_jpr_train\", \"dataset_jpr_train\", folder_name, study_id\n            )\n            \n            files = [os.path.join(base_path, f) for f in os.listdir(base_path) if f.lower().endswith(self.extensions)]\n            files.sort()  # Sort and select middle slices\n            n_patches = self.cfg.n_patches if self.cfg.use_multi_slice else 1\n            start_idx = max(0, len(files) // 2 - n_patches // 2)\n            selected_files = files[start_idx:start_idx + n_patches]\n            \n            # Pad with empty slices if necessary\n            selected_files += [None] * (n_patches - len(selected_files))\n            \n            imgs = []\n            for path in selected_files:\n                img = np.zeros((self.cfg.img_size, self.cfg.img_size, 3)) if path is None else load_dicom_to_image(path)\n                if self.transform:\n                    img = self.transform(image=img)[\"image\"]\n                imgs.append(img)\n            \n            return torch.stack(imgs, dim=0), torch.tensor(label, dtype=torch.float32)\n        \n        else:\n            study_id = self.img_ids[idx]\n            base_path = os.path.join(\n                self.cfg.test_root, \"dataset_jpr_test2\", \"dataset_jpr_test2\", study_id\n            )\n            \n            files = [os.path.join(base_path, f) for f in os.listdir(base_path) if f.lower().endswith(self.extensions)]\n            img_path = files[0] if files else None\n            img = np.zeros((self.cfg.img_size, self.cfg.img_size, 3)) if img_path is None else load_dicom_to_image(img_path)\n            \n            if self.transform:\n                img = self.transform(image=img)[\"image\"]\n            \n            return img.unsqueeze(0), torch.tensor(0.0, dtype=torch.float32)\n\n# Multi-task loss function\ndef multi_task_loss(y_preds, age_labels, age_group_preds=None, attention_age_logits=None, alpha=0.3, beta=0.2):\n    \"\"\"\n    Multi-task loss function: age regression loss + age group classification loss + attention age classification loss\n    y_preds: predicted age values\n    age_labels: ground truth ages\n    age_group_preds: predicted age groups\n    attention_age_logits: age predictions from attention module\n    alpha: weight for age group classification loss\n    beta: weight for attention age classification loss\n    \"\"\"\n    # Regression loss (MAE)\n    regression_loss = F.l1_loss(y_preds.view(-1), age_labels)\n    \n    total_loss = regression_loss\n    losses_dict = {\"regression_loss\": regression_loss.item()}\n    \n    # Create age group labels (for classification task)\n    age_bins = [0, 20, 40, 60, 80, 100]\n    age_groups = torch.zeros_like(age_labels, dtype=torch.long)\n    \n    for i in range(1, len(age_bins)):\n        age_groups[(age_labels >= age_bins[i-1]) & (age_labels < age_bins[i])] = i-1\n    \n    # Age group classification loss\n    if age_group_preds is not None:\n        age_group_loss = F.cross_entropy(age_group_preds, age_groups.to(age_group_preds.device))\n        total_loss += alpha * age_group_loss\n        losses_dict[\"age_group_loss\"] = age_group_loss.item()\n    \n    # Attention age classification loss\n    if attention_age_logits is not None:\n        attention_loss = F.cross_entropy(attention_age_logits, age_groups.to(attention_age_logits.device))\n        total_loss += beta * attention_loss\n        losses_dict[\"attention_loss\"] = attention_loss.item()\n    \n    return total_loss, losses_dict\n\n# Training function\ndef train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device, cfg):\n    model.train()\n    scaler = GradScaler(enabled=cfg.amp and device.type == \"cuda\")\n    total_losses = []\n    losses_history = {\"regression_loss\": [], \"age_group_loss\": [], \"attention_loss\": []}\n    \n    for step, (images, labels) in enumerate(train_loader):\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        with autocast(enabled=cfg.amp and device.type == \"cuda\"):\n            y_preds, age_group_preds, attention_age_logits = model(images)\n            loss, losses_dict = criterion(y_preds, labels, age_group_preds, attention_age_logits)\n            \n            for key, value in losses_dict.items():\n                losses_history[key].append(value)\n            \n            total_losses.append(loss.item())\n        \n        # Backpropagation and optimization\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            if (step + 1) % cfg.ga_accum == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n        else:\n            loss.backward()\n            if (step + 1) % cfg.ga_accum == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n        \n        if scheduler is not None and isinstance(scheduler, lr_scheduler.LRScheduler):\n            scheduler.step()\n    \n    # Calculate average losses\n    avg_losses = {key: np.mean(values) for key, values in losses_history.items()}\n    avg_losses[\"total_loss\"] = np.mean(total_losses)\n    \n    return avg_losses\n\n# Validation function\ndef valid_fn(valid_loader, model, criterion, device, cfg):\n    model.eval()\n    total_losses = []\n    losses_history = {\"regression_loss\": [], \"age_group_loss\": [], \"attention_loss\": []}\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for images, labels in valid_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            all_labels.append(labels.cpu().numpy())\n            \n            y_preds, age_group_preds, attention_age_logits = model(images)\n            loss, losses_dict = criterion(y_preds, labels, age_group_preds, attention_age_logits)\n            \n            for key, value in losses_dict.items():\n                losses_history[key].append(value)\n            \n            total_losses.append(loss.item())\n            all_preds.append(y_preds.detach().cpu().numpy())\n    \n    # Calculate average losses\n    avg_losses = {key: np.mean(values) for key, values in losses_history.items()}\n    avg_losses[\"total_loss\"] = np.mean(total_losses)\n    \n    # Calculate MAE\n    all_preds = np.concatenate(all_preds).reshape(-1)\n    all_labels = np.concatenate(all_labels)\n    mae = mean_absolute_error(all_labels, all_preds)\n    \n    return avg_losses, mae, all_preds\n\n# Feature extractor model wrapper\nclass FeatureExtractor(nn.Module):\n    def __init__(self, model):\n        super().__init__()\n        # Extract features up to the feature fusion layer\n        self.features = nn.Sequential(\n            model.pre,\n            model.stage1,\n            model.stage2,\n            model.stage3,\n            model.stage4,\n            model.stage5,\n            model.stage6,\n            model.stage7,\n            model.age_attention,  # Returns features and age predictions\n        )\n        self.feature_fusion = model.feature_fusion\n        \n    def forward(self, x):\n        if len(x.shape) == 5:  # Multi-slice input\n            batch_size, n_patches, channels, height, width = x.size()\n            x = x.view(batch_size * n_patches, channels, height, width)\n        \n        # Extract features\n        x, _ = self.features(x)\n        x = self.feature_fusion(x)\n        x = F.adaptive_avg_pool2d(x, 1).view(x.size(0), -1)\n        \n        if len(x.shape) == 2 and batch_size is not None:  # Restore multi-slice dimension\n            x = x.view(batch_size, n_patches, -1).mean(dim=1)  # Average multi-slice features\n        \n        return x\n\n# Feature extraction function\ndef extract_features(data_loader, model, device):\n    model.eval()\n    features = []\n    labels = []\n    \n    with torch.no_grad():\n        for images, target in data_loader:\n            images = images.to(device)\n            batch_features = model(images)\n            features.append(batch_features.cpu().numpy())\n            labels.append(target.cpu().numpy())\n    \n    features = np.vstack(features)\n    labels = np.hstack(labels)\n    return features, labels\n\n# Improved TSNE visualization function\ndef visualize_tsne(features, labels, save_path='tsne_visualization.png'):\n    \"\"\"\n    Visualize t-SNE dimensionality reduction results of features, grouped by age\n    \"\"\"\n    # Standardize data\n    scaler = StandardScaler()\n    scaled_features = scaler.fit_transform(features)\n    \n    # Perform t-SNE dimensionality reduction\n    print(\"Performing t-SNE dimensionality reduction...\")\n    tsne = TSNE(n_components=2, perplexity=30, n_iter=3000, random_state=42)\n    tsne_results = tsne.fit_transform(scaled_features)\n    \n    # Create DataFrame for plotting\n    df = pd.DataFrame({\n        'tsne-2d-one': tsne_results[:,0],\n        'tsne-2d-two': tsne_results[:,1],\n        'Age': labels\n    })\n    \n    # Create age groups\n    age_bins = [0, 20, 40, 60, 80, 100]\n    age_labels = ['0-20 years', '21-40 years', '41-60 years', '61-80 years', '81+ years']\n    df['Age Group'] = pd.cut(df['Age'], bins=age_bins, labels=age_labels, include_lowest=True)\n    \n    # Plot t-SNE results\n    plt.figure(figsize=(14, 10))\n    scatter = sns.scatterplot(\n        x=\"tsne-2d-one\", y=\"tsne-2d-two\",\n        hue=\"Age Group\",\n        palette=sns.color_palette(\"hls\", len(age_labels)),\n        data=df,\n        legend=\"full\",\n        alpha=0.7,\n        s=80  # Increase point size\n    )\n    \n    # Add titles and labels\n    plt.title('t-SNE Visualization of Feature Space for Brain CT Age Prediction Model', fontsize=16)\n    plt.xlabel('t-SNE Dimension 1', fontsize=14)\n    plt.ylabel('t-SNE Dimension 2', fontsize=14)\n    plt.legend(title='Age Groups', fontsize=12, title_fontsize=13)\n    \n    # Add age distribution histogram\n    plt.figure(figsize=(12, 6))\n    sns.histplot(df['Age'], bins=20, kde=True)\n    plt.title('Age Distribution in the Dataset', fontsize=16)\n    plt.xlabel('Age', fontsize=14)\n    plt.ylabel('Number of Samples', fontsize=14)\n    \n    # Save figures\n    plt.tight_layout()\n    plt.savefig(save_path)\n    print(f\"t-SNE visualization saved to {save_path}\")\n    \n    # Analyze clustering performance\n    from sklearn.metrics import silhouette_score\n    if len(df['Age Group'].dropna().unique()) > 1:\n        silhouette_avg = silhouette_score(scaled_features, df['Age Group'].dropna())\n        print(f\"Silhouette score for feature clustering: {silhouette_avg:.4f} (closer to 1 indicates better clustering)\")\n    \n    return df\n\ndef main():\n    class CFG:\n        seed = 42\n        n_fold = 1          # Set to 1 fold\n        use_folds = [0]     # Use only fold 0\n        img_size = 224\n        batch_size = 8\n        epochs = 50         # Train for 50 epochs\n        lr = 5e-5\n        min_lr = 1e-6\n        amp = True\n        ga_accum = 2\n        num_workers = 4\n        root = \"/kaggle/input/spr-head-ct-age-prediction-challenge\"\n        test_root = \"/kaggle/input/spr-head-ct-age-prediction-challenge\"\n        fold_csv = \"/kaggle/input/spr-head-ct-age-prediction-challenge/train.csv\"\n        test_csv = \"/kaggle/input/spr-head-ct-age-prediction-challenge/test.csv\"\n        pretrained_weights = \"/kaggle/input/1111/pytorch/default/1/newmodel_weights.pth\"\n        n_patches = 10\n        use_multi_slice = True\n        target_column = 'Age'\n        weight_decay = 1e-5\n        warmup_epochs = 5    # Warmup epochs\n        scheduler_name = \"onecycle\"\n        train_mode = \"train\"\n        drop_last = True\n        val_ratio = 0.2      # Validation set ratio 20%\n        visualize_tsne = True  # Whether to perform t-SNE visualization\n        # Multi-task loss weights\n        alpha = 0.3  # Weight for age group classification loss\n        beta = 0.2   # Weight for attention age classification loss\n    \n    cfg = CFG()\n    seed_everything(cfg.seed)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    print(f\"PyTorch version: {torch.__version__}\")\n    print(f\"CUDA available: {torch.cuda.is_available()}\")\n    print(f\"Device: {device}\")\n    \n    # Load data\n    train_df = pd.read_csv(cfg.fold_csv)\n    test_df = pd.read_csv(cfg.test_csv)\n    \n    # Split train and validation sets\n    train_df, valid_df = train_test_split(\n        train_df, \n        test_size=cfg.val_ratio, \n        random_state=cfg.seed,\n        shuffle=True\n    )\n    print(f\"Number of training samples: {len(train_df)}, validation samples: {len(valid_df)}\")\n    \n    # Data augmentation configuration\n    train_transform = A.Compose([\n        A.Resize(cfg.img_size, cfg.img_size),\n        A.HorizontalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=20, p=0.7),\n        A.ElasticTransform(p=0.3, alpha=120, sigma=120*0.05, alpha_affine=120*0.03),\n        A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n        ToTensorV2(),\n    ])\n    \n    valid_transform = A.Compose([\n        A.Resize(cfg.img_size, cfg.img_size),\n        A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n        ToTensorV2(),\n    ])\n    \n    # Create datasets and dataloaders\n    train_dataset = BrainAgeDataset(train_df, cfg, train_transform, mode=\"train\")\n    valid_dataset = BrainAgeDataset(valid_df, cfg, valid_transform, mode=\"train\")\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=cfg.batch_size,\n        shuffle=True,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n        drop_last=cfg.drop_last,\n    )\n    \n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=cfg.batch_size,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n        drop_last=False,\n    )\n    \n    # Get valid validation labels\n    valid_indices = valid_dataset.valid_study_ids\n    valid_labels = np.array([valid_dataset.label_dict[sid] for sid in valid_indices])\n    print(f\"Original validation samples: {len(valid_df)}, actual validation samples: {len(valid_labels)}\")\n    \n    # Initialize model\n    model = EnhancedSuperDAM(class_num=1, use_multi_slice=cfg.use_multi_slice)\n    \n    # Load pretrained weights with detailed logging\n    if cfg.pretrained_weights and os.path.exists(cfg.pretrained_weights):\n        try:\n            pretrained_dict = torch.load(cfg.pretrained_weights, map_location=device)\n            model_dict = model.state_dict()\n            \n            # Filter out unnecessary keys and print matching information\n            pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}\n            matched_layers = set(pretrained_dict.keys())\n            all_layers = set(model_dict.keys())\n            unmatched_layers = all_layers - matched_layers\n            \n            print(f\"Loading pretrained weights: {len(matched_layers)} layers matched, {len(unmatched_layers)} layers unmatched\")\n            print(\"Matched layers:\", sorted(matched_layers))\n            print(\"Unmatched layers:\", sorted(unmatched_layers))\n            \n            # Load the filtered weights\n            model_dict.update(pretrained_dict)\n            model.load_state_dict(model_dict, strict=False)\n            print(\"Successfully loaded pretrained weights\")\n        except Exception as e:\n            print(f\"Failed to load pretrained weights: {e}\")\n            print(\"Training from scratch...\")\n    else:\n        print(\"No pretrained weights found. Training from scratch...\")\n    \n    model.to(device)\n    \n    # Define loss function and optimizer\n    criterion = lambda y_preds, labels, age_group_preds, attention_age_logits: multi_task_loss(\n        y_preds, labels, age_group_preds, attention_age_logits, cfg.alpha, cfg.beta\n    )\n    \n    optimizer = optim.AdamW(\n        model.parameters(), \n        lr=cfg.lr, \n        weight_decay=cfg.weight_decay\n    )\n    \n    # Learning rate scheduler (OneCycleLR)\n    if cfg.scheduler_name == \"onecycle\":\n        steps_per_epoch = len(train_loader)\n        total_steps = ((len(train_dataset) + cfg.batch_size - 1) // cfg.batch_size) * cfg.epochs\n        pct_start = cfg.warmup_epochs / cfg.epochs\n        \n        scheduler = lr_scheduler.OneCycleLR(\n            optimizer,\n            max_lr=cfg.lr * 5,\n            steps_per_epoch=steps_per_epoch,\n            epochs=cfg.epochs,\n            total_steps=total_steps,\n            pct_start=pct_start,\n            anneal_strategy='cos'\n        )\n    else:\n        scheduler = None\n    \n    # Initialize best MAE\n    best_mae = float('inf')\n    best_epoch = 0\n    \n    # Training loop\n    for epoch in range(cfg.epochs):\n        start_time = time.time()\n        \n        # Train for one epoch\n        train_losses = train_fn(\n            train_loader=train_loader,\n            model=model,\n            criterion=criterion,\n            optimizer=optimizer,\n            epoch=epoch,\n            scheduler=scheduler,\n            device=device,\n            cfg=cfg\n        )\n        \n        # Validate for one epoch\n        valid_losses, mae, preds = valid_fn(\n            valid_loader=valid_loader,\n            model=model,\n            criterion=criterion,\n            device=device,\n            cfg=cfg\n        )\n        \n        # Save best model\n        if mae < best_mae:\n            best_mae = mae\n            best_epoch = epoch + 1\n            torch.save(model.state_dict(), f\"best_model.pth\")\n            print(f\"Saved best model at epoch {epoch+1} with MAE: {best_mae:.4f}\")\n        \n        # Print training info\n        elapsed = time.time() - start_time\n        print(f\"Epoch {epoch+1}/{cfg.epochs} | Time elapsed: {elapsed:.2f}s\")\n        print(f\"Training losses: Total={train_losses['total_loss']:.4f}, Regression={train_losses['regression_loss']:.4f}, \"\n              f\"Age Group={train_losses.get('age_group_loss', 0):.4f}, Attention={train_losses.get('attention_loss', 0):.4f}\")\n        print(f\"Validation losses: Total={valid_losses['total_loss']:.4f}, Regression={valid_losses['regression_loss']:.4f}, \"\n              f\"MAE={mae:.4f} | Best MAE: {best_mae:.4f} (Epoch {best_epoch})\")\n        print(\"-\" * 80)\n    \n    # Testing phase\n    if cfg.train_mode == \"test\":\n        test_transform = A.Compose([\n            A.Resize(cfg.img_size, cfg.img_size),\n            A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),\n            ToTensorV2(),\n        ])\n        \n        test_dataset = BrainAgeDataset(test_df, cfg, test_transform, mode=\"test\")\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=cfg.batch_size,\n            shuffle=False,\n            num_workers=cfg.num_workers,\n            pin_memory=True,\n        )\n        \n        model = EnhancedSuperDAM(class_num=1, use_multi_slice=cfg.use_multi_slice)\n        model.to(device)\n        model.load_state_dict(torch.load(\"best_model.pth\"))\n        model.eval()\n        \n        all_preds = []\n        with torch.no_grad():\n            for images, _ in test_loader:\n                images = images.to(device)\n                y_preds, _, _ = model(images)\n                all_preds.append(y_preds.detach().cpu().numpy())\n        \n        # Prepare submission\n        all_preds = np.concatenate(all_preds).reshape(-1)\n        submission = pd.DataFrame({\n            'ID': test_df['StudyID'],\n            'Age': all_preds\n        })\n        submission.to_csv('submission.csv', index=False)\n        print(\"Submission file created.\")\n    \n    # Visualize features with t-SNE\n    if cfg.visualize_tsne:\n        feature_extractor = FeatureExtractor(model)\n        feature_extractor.to(device)\n        \n        print(\"Extracting features from validation set...\")\n        features, labels = extract_features(valid_loader, feature_extractor, device)\n        \n        print(\"Visualizing features with t-SNE...\")\n        visualize_tsne(features, labels)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-26T09:44:13.596069Z","iopub.execute_input":"2025-06-26T09:44:13.596501Z","iopub.status.idle":"2025-06-26T09:44:20.90724Z","shell.execute_reply.started":"2025-06-26T09:44:13.59647Z","shell.execute_reply":"2025-06-26T09:44:20.905618Z"}},"outputs":[],"execution_count":null}]}