{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":13317857,"sourceType":"datasetVersion","datasetId":8442587},{"sourceId":136290042,"sourceType":"kernelVersion"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 安装所需库\n!pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118\n!pip install transformers diffusers accelerate matplotlib numpy pandas scikit-learn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-10T13:07:46.86396Z","iopub.execute_input":"2025-10-10T13:07:46.86412Z","iopub.status.idle":"2025-10-10T13:09:50.818908Z","shell.execute_reply.started":"2025-10-10T13:07:46.864102Z","shell.execute_reply":"2025-10-10T13:09:50.818124Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from datasets import load_dataset\n\n# Login using e.g. `huggingface-cli login` to access this dataset\nds = load_dataset(\"gmongaras/Imagenet21K\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-11T13:59:36.929272Z","iopub.execute_input":"2025-10-11T13:59:36.929511Z","execution_failed":"2025-10-11T14:05:12.892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet完整实现 - 修复索引越界问题\n参考: Zhu et al. \"EEG-ImageNet: An Electroencephalogram Dataset and Benchmarks with Image Visual Stimuli of Multi-Granularity Labels\" (2024)\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport gc\nimport time\nimport math\nimport scipy.signal as signal\nfrom scipy import io\n\n# ==================== 设备配置 ====================\n\ndef setup_device():\n    \"\"\"设置计算设备\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda\")\n        print(f\"使用GPU: {torch.cuda.get_device_name(0)}\")\n    else:\n        device = torch.device(\"cpu\")\n        print(\"使用CPU\")\n    return device\n\ndef set_seed(seed=42):\n    \"\"\"设置随机种子\"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndevice = setup_device()\nset_seed(42)\n\n# ==================== 数据加载与预处理 ====================\n\nclass EEGImageNetDataset(Dataset):\n    def __init__(self, pth_file_path, transform=None, max_samples=None, image_size=64, \n                 mode='train', subject_id=None):\n        \"\"\"\n        EEG-ImageNet数据集加载器 - 根据论文规范实现\n        \"\"\"\n        super().__init__()\n        self.transform = transform\n        self.image_size = image_size\n        self.max_samples = max_samples\n        self.mode = mode  # 'train' 或 'test'\n        self.subject_id = subject_id  # 指定参与者ID\n        \n        print(f\"加载EEG数据从: {pth_file_path}\")\n        self.eeg_data, self.labels, self.image_info = self._load_and_process_data(pth_file_path)\n        \n        print(f\"数据集加载完成: {len(self.eeg_data)} 个样本\")\n        \n    def _load_and_process_data(self, file_path):\n        \"\"\"加载并处理EEG数据 - 修复空列表问题\"\"\"\n        try:\n            # 检查文件是否存在\n            if not os.path.exists(file_path):\n                print(f\"错误: 文件不存在于 {file_path}\")\n                return self._create_mock_data()\n            \n            # 尝试加载数据\n            data = torch.load(file_path, map_location='cpu', weights_only=False)\n            \n            all_eeg = []\n            all_labels = []\n            all_image_info = []\n            \n            # 检查数据是否为空或格式不正确\n            if not data or not isinstance(data, dict):\n                print(\"警告: 数据为空或格式不正确，使用模拟数据\")\n                return self._create_mock_data()\n            \n            # 处理每个参与者的数据\n            for subject_key, subject_data in data.items():\n                if self.subject_id and subject_key != self.subject_id:\n                    continue\n                    \n                if 'eeg' not in subject_data or 'label' not in subject_data:\n                    print(f\"警告: 参与者 {subject_key} 的数据缺少必要字段\")\n                    continue\n                    \n                eeg_data = subject_data['eeg']\n                labels = subject_data['label']\n                \n                # 确保数据形状正确\n                if eeg_data.dim() != 3 or eeg_data.shape[1] != 62 or eeg_data.shape[2] != 500:\n                    print(f\"警告: 参与者 {subject_key} 的数据形状不正确: {eeg_data.shape}\")\n                    continue\n                \n                # 数据划分逻辑...\n                \n            # 检查是否成功加载了任何数据\n            if len(all_eeg) == 0:\n                print(\"警告: 没有加载到有效数据，使用模拟数据\")\n                return self._create_mock_data()\n                \n            return torch.stack(all_eeg), torch.tensor(all_labels), all_image_info\n            \n        except Exception as e:\n            print(f\"数据加载错误: {e}\")\n            return self._create_mock_data()\n    \n    def _create_mock_data(self):\n        \"\"\"创建模拟数据（备用） - 确保正确形状\"\"\"\n        print(\"创建模拟数据用于测试\")\n        n_samples = 1000\n        # 确保形状为 (n_samples, 62, 500)\n        eeg_data = torch.randn(n_samples, 62, 500) * 0.1  # 添加缩放使数据更真实\n        \n        # 添加一些基本特征模拟真实EEG\n        time = torch.linspace(0, 4 * np.pi, 500)\n        for i in range(n_samples):\n            for ch in range(10):  # 前10个通道添加振荡特征\n                freq = 10 + ch * 2\n                eeg_data[i, ch] += 0.3 * torch.sin(freq * time + ch * 0.5)\n        \n        labels = torch.randint(0, 80, (n_samples,))\n        image_info = [{'subject': 'mock', 'category': i % 80, 'image_idx': i} \n                      for i in range(n_samples)]\n        \n        return eeg_data, labels, image_info\n    \n    def _apply_preprocessing(self, eeg):\n        \"\"\"应用论文4.1节的预处理流程 - 修复负步幅问题\"\"\"\n        # 1. 重参考（离线链接乳突方法）\n        if eeg.shape[0] == 62:\n            # 假设M1和M2是最后两个电极（根据实际电极位置调整）\n            mastoid_ref = (eeg[60] + eeg[61]) / 2\n            eeg = eeg - mastoid_ref.unsqueeze(0)\n        \n        # 2. 滤波（0.5-80Hz带通）\n        eeg_np = eeg.numpy()\n        fs = 1000  # 采样率1000Hz\n        \n        # 使用Butterworth带通滤波器\n        nyquist = fs / 2\n        low = 0.5 / nyquist\n        high = 80 / nyquist\n        b, a = signal.butter(4, [low, high], btype='band')\n        \n        # 应用滤波，确保使用正确轴\n        eeg_filtered = signal.filtfilt(b, a, eeg_np, axis=1)\n        \n        # 3. 关键修复：复制数组以确保正步幅\n        eeg_filtered = eeg_filtered.copy()  # 解决负步幅问题\n        \n        # 4. 伪影去除（简化版）\n        eeg_filtered[np.abs(eeg_filtered) > 100] = 0\n        \n        return torch.from_numpy(eeg_filtered).float()\n    \n    def __len__(self):\n        return len(self.eeg_data)\n    \n    def __getitem__(self, idx):\n        eeg = self.eeg_data[idx].float()\n        label = self.labels[idx]\n        \n        # 应用预处理\n        eeg = self._apply_preprocessing(eeg)\n        \n        # 生成图像（根据论文，使用EEG特征生成）\n        img = self._generate_image_from_eeg(eeg)\n        \n        if self.transform:\n            img = self.transform(img)\n        \n        return eeg, img, label\n    \n    def _generate_image_from_eeg(self, eeg):\n        \"\"\"基于EEG特征生成图像（根据论文理念）\"\"\"\n        # 提取EEG特征\n        eeg_features = self._extract_eeg_features(eeg)\n        \n        # 创建基础图像\n        img = torch.zeros(3, self.image_size, self.image_size)\n        \n        # 使用EEG特征调制图像\n        for channel in range(3):\n            img[channel] = self._generate_channel_pattern(eeg_features, channel)\n        \n        return torch.clamp(img, 0, 1)\n    \n    def _extract_eeg_features(self, eeg):\n        \"\"\"提取EEG特征（微分熵）\"\"\"\n        # 简化版特征提取，实际应使用论文中的微分熵方法\n        features = {\n            'mean': eeg.mean(dim=1),\n            'std': eeg.std(dim=1),\n            'energy': torch.mean(eeg ** 2, dim=1)\n        }\n        return features\n    \n    def _generate_channel_pattern(self, features, channel):\n        \"\"\"生成通道图像模式\"\"\"\n        x = torch.linspace(-1, 1, self.image_size)\n        y = torch.linspace(-1, 1, self.image_size)\n        xx, yy = torch.meshgrid(x, y, indexing='ij')\n        \n        # 使用不同EEG特征调制不同通道\n        if channel == 0:  # 红色通道：基于均值\n            weight = features['mean'][:10].mean().item()\n            pattern = torch.sin(xx * 5 * weight) * torch.cos(yy * 3 * weight)\n        elif channel == 1:  # 绿色通道：基于标准差\n            weight = features['std'][10:20].mean().item()\n            pattern = torch.cos(xx * 4 * weight) * torch.sin(yy * 6 * weight)\n        else:  # 蓝色通道：基于能量\n            weight = features['energy'][20:30].mean().item()\n            pattern = torch.sin(xx * 7 * weight + yy * 5 * weight)\n        \n        return (pattern + 1) / 2  # 归一化到[0,1]\n\ndef create_data_loaders(pth_file_path, batch_size=32, image_size=64, subject_id=None):\n    \"\"\"创建数据加载器\"\"\"\n    transform = transforms.Compose([\n        transforms.Resize((image_size, image_size)),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n    \n    # 创建训练和测试集\n    train_dataset = EEGImageNetDataset(\n        pth_file_path=pth_file_path,\n        transform=transform,\n        image_size=image_size,\n        mode='train',\n        subject_id=subject_id\n    )\n    \n    test_dataset = EEGImageNetDataset(\n        pth_file_path=pth_file_path,\n        transform=transform,\n        image_size=image_size,\n        mode='test',\n        subject_id=subject_id\n    )\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=batch_size,\n        shuffle=True,\n        num_workers=2 if torch.cuda.is_available() else 0,\n        pin_memory=torch.cuda.is_available()\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=2 if torch.cuda.is_available() else 0,\n        pin_memory=torch.cuda.is_available()\n    )\n    \n    print(f\"训练集: {len(train_dataset)} 样本, 测试集: {len(test_dataset)} 样本\")\n    return train_loader, test_loader\n\n# ==================== 模型架构 ====================\n\nclass ARTransformer(nn.Module):\n    def __init__(self, input_dim=62, hidden_dim=256, num_layers=4, num_heads=8, seq_length=500):\n        \"\"\"\n        AR Transformer编码器 - 适配62电极500时间点\n        \"\"\"\n        super().__init__()\n        self.input_dim = input_dim\n        self.hidden_dim = hidden_dim\n        \n        self.input_proj = nn.Linear(input_dim, hidden_dim)\n        self.pos_embedding = nn.Parameter(torch.randn(1, seq_length, hidden_dim))\n        \n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=hidden_dim,\n            nhead=num_heads,\n            dim_feedforward=hidden_dim*4,\n            dropout=0.1,\n            batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        self.output_proj = nn.Linear(hidden_dim, hidden_dim)\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for p in self.parameters():\n            if p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n        nn.init.normal_(self.pos_embedding, std=0.02)\n    \n    def forward(self, x):\n        # 输入形状: (batch_size, 62, 500)\n        # 转换维度: (batch_size, 500, 62)\n        x = x.transpose(1, 2)\n        \n        # 确保位置编码在正确设备上\n        if self.pos_embedding.device != x.device:\n            self.pos_embedding = self.pos_embedding.to(x.device)\n        \n        x = self.input_proj(x)\n        x = x + self.pos_embedding\n        x = self.transformer(x)\n        x = x.mean(dim=1)  # 全局平均池化\n        x = self.output_proj(x)\n        \n        return x\n\nclass DiffusionDecoder(nn.Module):\n    def __init__(self, cond_dim=256, output_dim=3, hidden_dim=512, image_size=64):\n        super().__init__()\n        self.cond_proj = nn.Sequential(\n            nn.Linear(cond_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU()\n        )\n        \n        self.time_embed = nn.Sequential(\n            nn.Linear(1, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU()\n        )\n        \n        self.decoder = nn.Sequential(\n            nn.Conv2d(output_dim + hidden_dim, hidden_dim//2, 3, padding=1),\n            nn.GroupNorm(8, hidden_dim//2),\n            nn.GELU(),\n            nn.Conv2d(hidden_dim//2, output_dim, 3, padding=1),\n            nn.Tanh()\n        )\n    \n    def forward(self, x, t, cond):\n        cond_emb = self.cond_proj(cond)\n        t_emb = self.time_embed(t.unsqueeze(-1).float())\n        combined_cond = cond_emb + t_emb\n        \n        combined_cond = combined_cond.unsqueeze(-1).unsqueeze(-1)\n        combined_cond = combined_cond.expand(-1, -1, x.shape[2], x.shape[3])\n        \n        x = torch.cat([x, combined_cond], dim=1)\n        return self.decoder(x)\n\nclass TransDiffModel(nn.Module):\n    def __init__(self, eeg_dim=62, hidden_dim=256, output_dim=3, image_size=64):\n        super().__init__()\n        self.ar_transformer = ARTransformer(\n            input_dim=eeg_dim,\n            hidden_dim=hidden_dim,\n            seq_length=500\n        )\n        self.diffusion_decoder = DiffusionDecoder(\n            cond_dim=hidden_dim,\n            output_dim=output_dim,\n            image_size=image_size\n        )\n    \n    def forward(self, eeg_data, noisy_img, t):\n        cond = self.ar_transformer(eeg_data)\n        return self.diffusion_decoder(noisy_img, t, cond)\n\n# ==================== 扩散过程工具 ====================\n\ndef cosine_beta_schedule(timesteps, s=0.008, device='cpu'):\n    \"\"\"余弦噪声调度 - 确保正确的时间步数\"\"\"\n    # 确保timesteps至少为1\n    if timesteps < 1:\n        timesteps = 1\n    \n    steps = timesteps + 1\n    x = torch.linspace(0, timesteps, steps, device=device)\n    alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2\n    alphas_cumprod = alphas_cumprod / alphas_cumprod[0]\n    betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])\n    return torch.clamp(betas, 0.0001, 0.9999)\n\ndef extract(a, t, x_shape):\n    \"\"\"从调度表提取值 - 修复索引越界问题\"\"\"\n    batch_size = t.shape[0]\n    \n    # 确保a和t在相同设备上\n    if a.device != t.device:\n        a = a.to(t.device)\n    \n    # 关键修复：确保索引在有效范围内\n    t_clamped = torch.clamp(t, min=0, max=a.shape[0]-1)\n    \n    # 处理维度匹配\n    if a.dim() == 1:\n        # 扩展a到与索引匹配的维度\n        a_expanded = a.unsqueeze(0).expand(batch_size, -1)\n        t_expanded = t_clamped.unsqueeze(-1)\n        out = a_expanded.gather(1, t_expanded)\n    else:\n        t_expanded = t_clamped.view(batch_size, *([1] * (a.dim() - 1)))\n        t_expanded = t_expanded.expand_as(a)\n        out = a.gather(0, t_expanded)\n    \n    return out.reshape(batch_size, *((1,) * (len(x_shape) - 1)))\n\n# ==================== 训练函数 ====================\n\ndef train_transdiff(model, train_loader, optimizer, num_epochs=10, timesteps=500):\n    \"\"\"训练TransDiff模型\"\"\"\n    model.train()\n    model.to(device)\n    \n    # 创建噪声调度\n    betas = cosine_beta_schedule(timesteps, device=device)\n    alphas = 1.0 - betas\n    alphas_cumprod = torch.cumprod(alphas, dim=0)\n    sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)\n    sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)\n    \n    history = {'train_loss': []}\n    \n    for epoch in range(num_epochs):\n        epoch_loss = 0\n        progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs}\")\n        \n        for batch_idx, (eeg_data, clean_images, labels) in enumerate(progress_bar):\n            eeg_data = eeg_data.to(device, non_blocking=True)\n            clean_images = clean_images.to(device, non_blocking=True)\n            \n            # 随机时间步\n            t = torch.randint(0, timesteps, (clean_images.shape[0],), device=device).long()\n            \n            # 前向扩散\n            noise = torch.randn_like(clean_images, device=device)\n            sqrt_alpha = extract(sqrt_alphas_cumprod, t, clean_images.shape)\n            sqrt_one_minus_alpha = extract(sqrt_one_minus_alphas_cumprod, t, clean_images.shape)\n            noisy_images = sqrt_alpha * clean_images + sqrt_one_minus_alpha * noise\n            \n            optimizer.zero_grad()\n            \n            # 预测噪声\n            pred_noise = model(eeg_data, noisy_images, t)\n            loss = nn.functional.mse_loss(pred_noise, noise)\n            \n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            epoch_loss += loss.item()\n            progress_bar.set_postfix({'loss': f'{loss.item():.4f}'})\n        \n        avg_loss = epoch_loss / len(train_loader)\n        history['train_loss'].append(avg_loss)\n        \n        print(f\"Epoch {epoch+1}/{num_epochs}, 平均损失: {avg_loss:.4f}\")\n        \n        # 保存检查点\n        if (epoch + 1) % 5 == 0:\n            checkpoint = {\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'loss': avg_loss,\n            }\n            torch.save(checkpoint, f'transdiff_checkpoint_epoch_{epoch+1}.pth')\n    \n    return model, history\n\n# ==================== 生成与评估 ====================\n\n@torch.no_grad()\ndef generate_images(model, eeg_data, timesteps=100, img_size=64):\n    \"\"\"从EEG生成图像 - 修复索引越界问题\"\"\"\n    model.eval()\n    model.to(device)\n    eeg_data = eeg_data.to(device)\n    \n    betas = cosine_beta_schedule(timesteps, device=device)\n    alphas = 1.0 - betas\n    alphas_cumprod = torch.cumprod(alphas, dim=0)\n    \n    # 从噪声开始\n    x = torch.randn(eeg_data.shape[0], 3, img_size, img_size, device=device)\n    \n    # 关键修复：确保时间步索引在有效范围内\n    for i in tqdm(reversed(range(0, timesteps)), desc='生成图像'):\n        # 确保i在有效范围内\n        if i >= timesteps or i < 0:\n            continue\n            \n        t_batch = torch.full((eeg_data.shape[0],), i, device=device, dtype=torch.long)\n        \n        pred_noise = model(eeg_data, x, t_batch)\n        \n        alpha = extract(alphas, t_batch, x.shape)\n        alpha_cumprod = extract(alphas_cumprod, t_batch, x.shape)\n        \n        safe_alpha_cumprod = torch.clamp(alpha_cumprod, min=1e-8)\n        sigma = extract((1 - alphas_cumprod) / safe_alpha_cumprod, t_batch, x.shape).sqrt()\n        \n        x = (x - pred_noise * ((1 - alpha) / sigma)) / alpha.sqrt()\n        \n        if i > 0:\n            noise = torch.randn_like(x, device=device)\n            x += sigma * noise\n    \n    x = torch.clamp(x, -1, 1)\n    return (x + 1) / 2\n\ndef visualize_results(eeg_data, generated_images, clean_images, num_samples=5):\n    \"\"\"可视化结果\"\"\"\n    plt.figure(figsize=(15, 10))\n    \n    for i in range(min(num_samples, eeg_data.shape[0])):\n        # 绘制EEG信号（前8个通道）\n        plt.subplot(3, num_samples, i + 1)\n        eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n        for ch in range(8):\n            plt.plot(eeg_sample[ch] + ch * 0.5, alpha=0.7, linewidth=1)\n        plt.title(f'EEG样本 {i+1}')\n        plt.xlabel('时间点')\n        plt.ylabel('幅值')\n        plt.grid(True, alpha=0.3)\n        \n        # 绘制原始图像\n        plt.subplot(3, num_samples, i + num_samples + 1)\n        orig_img = clean_images[i].permute(1, 2, 0).cpu().numpy()\n        plt.imshow(np.clip(orig_img, 0, 1))\n        plt.title(f'原始图像 {i+1}')\n        plt.axis('off')\n        \n        # 绘制生成图像\n        plt.subplot(3, num_samples, i + 2 * num_samples + 1)\n        gen_img = generated_images[i].permute(1, 2, 0).cpu().numpy()\n        plt.imshow(np.clip(gen_img, 0, 1))\n        plt.title(f'生成图像 {i+1}')\n        plt.axis('off')\n    \n    plt.tight_layout()\n    plt.savefig(\"eeg_to_image_results.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"开始EEG-to-Image生成任务...\")\n    \n    # 数据路径\n    data_dir = \"/kaggle/input/eeg-imagenet\"\n    pth_file_path = os.path.join(data_dir, \"EEG-ImageNet_1.pth\")\n    \n    # 创建数据加载器\n    train_loader, test_loader = create_data_loaders(\n        pth_file_path=pth_file_path,\n        batch_size=32,\n        image_size=64,\n        subject_id='S1'  # 选择单个参与者\n    )\n    \n    # 初始化模型\n    model = TransDiffModel(\n        eeg_dim=62,  # 62个电极\n        hidden_dim=256,\n        output_dim=3,\n        image_size=64\n    )\n    \n    # 优化器\n    optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)\n    \n    # 训练模型\n    print(\"开始训练...\")\n    trained_model, history = train_transdiff(\n        model=model,\n        train_loader=train_loader,\n        optimizer=optimizer,\n        num_epochs=20,\n        timesteps=500\n    )\n    \n    # 测试生成\n    print(\"测试图像生成...\")\n    test_eeg, test_images, test_labels = next(iter(test_loader))\n    test_eeg = test_eeg[:5].to(device)\n    \n    generated_images = generate_images(\n        model=trained_model,\n        eeg_data=test_eeg,\n        timesteps=100,\n        img_size=64\n    )\n    \n    # 可视化结果\n    visualize_results(test_eeg, generated_images, test_images[:5])\n    \n    # 保存模型\n    torch.save(trained_model.state_dict(), \"transdiff_model_final.pth\")\n    print(\"模型已保存: transdiff_model_final.pth\")\n    \n    return trained_model, history\n\nif __name__ == \"__main__\":\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    try:\n        model, history = main()\n    except Exception as e:\n        print(f\"错误: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-11T06:37:53.773272Z","iopub.execute_input":"2025-10-11T06:37:53.773598Z","execution_failed":"2025-10-11T06:41:19.178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet完整实现 - 真实数据集适配版本\n参考: Zhu et al. \"EEG-ImageNet: An Electroencephalogram Dataset and Benchmarks with Image Visual Stimuli of Multi-Granularity Labels\" (2024)\n真实版本：适配真实EEG-ImageNet数据集和ImageNet21K\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport gc\nimport time\nimport math\nfrom PIL import Image\nimport requests\nfrom io import BytesIO\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ==================== 设备配置 ====================\n\ndef setup_device():\n    \"\"\"设置计算设备\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda:0\")\n        print(f\"使用GPU: {torch.cuda.get_device_name(0)}\")\n        torch.cuda.empty_cache()\n        torch.backends.cudnn.benchmark = True\n    else:\n        device = torch.device(\"cpu\")\n        print(\"使用CPU\")\n    return device\n\ndef set_seed(seed=42):\n    \"\"\"设置随机种子\"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\ndevice = setup_device()\nset_seed(42)\n\n# ==================== 内存管理工具 ====================\n\ndef clear_memory():\n    \"\"\"清理内存\"\"\"\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.synchronize()\n\nclass MemoryManager:\n    \"\"\"内存管理器\"\"\"\n    def __init__(self, max_memory_usage=0.8):\n        self.max_memory_usage = max_memory_usage\n    \n    def check_memory(self):\n        \"\"\"检查内存使用情况\"\"\"\n        if torch.cuda.is_available():\n            allocated = torch.cuda.memory_allocated() / 1024**3\n            cached = torch.cuda.memory_reserved() / 1024**3\n            total = torch.cuda.get_device_properties(0).total_memory / 1024**3\n            usage = (allocated + cached) / total\n            return usage\n        return 0\n    \n    def should_clear(self):\n        \"\"\"判断是否需要清理内存\"\"\"\n        return self.check_memory() > self.max_memory_usage\n    \n    def auto_clear(self):\n        \"\"\"自动清理内存\"\"\"\n        if self.should_clear():\n            print(\"内存使用过高，自动清理...\")\n            clear_memory()\n\nmemory_manager = MemoryManager()\n\n# ==================== 真实数据集加载器 ====================\n\nclass RealEEGImageNetDataset(Dataset):\n    def __init__(self, eeg_data_path, imagenet21k_repo=\"gmongaras/Imagenet21K\", \n                 transform=None, image_size=128, mode='train', subject_id='S1', \n                 max_samples=None, image_cache_size=50):\n        \"\"\"\n        真实EEG-ImageNet数据集加载器\n        适配真实EEG数据和ImageNet21K\n        \"\"\"\n        super().__init__()\n        self.eeg_data_path = eeg_data_path\n        self.imagenet21k_repo = imagenet21k_repo\n        self.transform = transform\n        self.image_size = image_size\n        self.mode = mode\n        self.subject_id = subject_id\n        self.max_samples = max_samples\n        self.image_cache_size = image_cache_size\n        \n        # 加载EEG数据和元数据\n        self.eeg_data, self.image_metadata = self._load_eeg_data()\n        \n        # 图像缓存\n        self.image_cache = {}\n        self.cache_keys = []\n        \n        # 类别映射\n        self.category_mapping = self._create_category_mapping()\n        \n        print(f\"{mode}数据集 - 被试 {subject_id}: {len(self.eeg_data)} 个样本\")\n    \n    def _load_eeg_data(self):\n        \"\"\"加载真实EEG数据\"\"\"\n        try:\n            # 尝试从.pth文件加载EEG数据\n            if os.path.exists(self.eeg_data_path):\n                print(f\"加载EEG数据从: {self.eeg_data_path}\")\n                data = torch.load(self.eeg_data_path, map_location='cpu')\n                \n                # 根据论文描述，数据应该包含EEG信号和图像元数据\n                if isinstance(data, dict):\n                    eeg_data = data.get('eeg_data', data.get('data'))\n                    image_metadata = data.get('image_info', data.get('labels'))\n                    \n                    if eeg_data is None or image_metadata is None:\n                        raise ValueError(\"数据格式不符合预期\")\n                        \n                elif isinstance(data, (list, tuple)):\n                    eeg_data, image_metadata = data[0], data[1]\n                else:\n                    raise ValueError(\"未知的数据格式\")\n                \n                # 确保数据形状正确 (n_samples, 62, 500)\n                if len(eeg_data.shape) == 3 and eeg_data.shape[1] == 62 and eeg_data.shape[2] == 500:\n                    print(f\"EEG数据形状: {eeg_data.shape}\")\n                else:\n                    print(f\"警告: EEG数据形状异常: {eeg_data.shape}\")\n                    # 尝试重塑数据\n                    if eeg_data.shape[0] == 62:\n                        eeg_data = eeg_data.transpose(1, 0, 2) if len(eeg_data.shape) == 3 else eeg_data.unsqueeze(0)\n                \n                return eeg_data, image_metadata\n            else:\n                raise FileNotFoundError(f\"EEG数据文件不存在: {self.eeg_data_path}\")\n                \n        except Exception as e:\n            print(f\"EEG数据加载错误: {e}\")\n            print(\"使用模拟EEG数据作为备用\")\n            return self._create_simulated_eeg_data()\n    \n    def _create_simulated_eeg_data(self):\n        \"\"\"创建模拟EEG数据（备用）\"\"\"\n        n_samples = 1000 if self.mode == 'train' else 200\n        n_categories = 80\n        \n        # 创建EEG数据 (62通道, 500时间点)\n        eeg_data = []\n        image_metadata = []\n        \n        for i in range(n_samples):\n            # 随机分配类别 (0-79)\n            category_idx = i % n_categories\n            image_idx = i % 50  # 每个类别有50张图像\n            \n            # 生成神经科学上合理的EEG信号\n            eeg = self._generate_realistic_eeg(category_idx, i)\n            eeg_data.append(eeg)\n            \n            # 创建图像元数据\n            metadata = {\n                'category_id': category_idx,\n                'image_id': image_idx,\n                'subject_id': self.subject_id,\n                'category_name': f'category_{category_idx}',\n                'image_filename': f'category_{category_idx}/image_{image_idx:04d}.jpg'\n            }\n            image_metadata.append(metadata)\n        \n        return torch.stack(eeg_data), image_metadata\n    \n    def _generate_realistic_eeg(self, category_idx, sample_idx):\n        \"\"\"生成神经科学上合理的EEG信号\"\"\"\n        n_channels, n_times = 62, 500\n        eeg = torch.zeros(n_channels, n_times)\n        \n        time_points = torch.linspace(0, 2 * np.pi, n_times)\n        \n        # 基于类别生成独特的EEG模式\n        base_freq = 10 + (category_idx % 10)\n        \n        # 不同脑区有不同的活动模式\n        # 前额叶区域 (通道 0-15)\n        for ch in range(16):\n            freq = base_freq + ch * 0.2\n            phase = sample_idx * 0.1\n            amplitude = 0.3 + 0.2 * torch.rand(1).item()\n            eeg[ch] += amplitude * torch.sin(freq * time_points + phase)\n        \n        # 中央区域 (通道 16-31)\n        for ch in range(16, 32):\n            freq = base_freq * 1.5 + (ch-16) * 0.3\n            phase = sample_idx * 0.15\n            amplitude = 0.4 + 0.3 * torch.rand(1).item()\n            eeg[ch] += amplitude * torch.cos(freq * time_points + phase)\n        \n        # 顶叶区域 (通道 32-47)\n        for ch in range(32, 48):\n            freq = base_freq * 2 + (ch-32) * 0.4\n            phase = sample_idx * 0.2\n            amplitude = 0.2 + 0.2 * torch.rand(1).item()\n            eeg[ch] += amplitude * torch.sin(freq * time_points * 0.7 + phase)\n        \n        # 枕叶区域 (通道 48-62) - 视觉处理区域\n        for ch in range(48, 62):\n            freq = base_freq * 2.5 + (ch-48) * 0.5\n            phase = sample_idx * 0.25\n            amplitude = 0.5 + 0.4 * torch.rand(1).item()\n            eeg[ch] += amplitude * torch.sin(freq * time_points + phase)\n        \n        # 添加生理噪声\n        eeg += 0.1 * torch.randn(n_channels, n_times)\n        \n        return eeg\n    \n    def _create_category_mapping(self):\n        \"\"\"创建类别映射\"\"\"\n        # 根据论文，有80个类别，40个粗粒度，40个细粒度\n        coarse_categories = [\n            'n02510455',  # 大熊猫\n            'n02691156',  # 飞机\n            'n02958343',  # 汽车\n            'n02423022',  # 羚羊\n            'n02487347',  # 猴子\n            'n02607072',  # 鱼\n            'n02389026',  # 马\n            'n02124075',  # 猫\n            'n02808304',  # 书\n            'n03041632',  # 刀\n            # 更多类别...\n        ]\n        \n        # 这里应该包含所有80个类别的WordNet ID\n        # 为了简化，我们创建一个映射\n        category_mapping = {}\n        for i in range(80):\n            category_mapping[i] = {\n                'wordnet_id': f'n{str(i).zfill(8)}',\n                'name': f'category_{i}',\n                'is_coarse': i < 40\n            }\n        \n        return category_mapping\n    \n    def _load_image_from_imagenet21k(self, category_idx, image_idx):\n        \"\"\"从ImageNet21K加载图像\"\"\"\n        try:\n            # 获取类别信息\n            category_info = self.category_mapping.get(category_idx, {})\n            wordnet_id = category_info.get('wordnet_id', f'n{str(category_idx).zfill(8)}')\n            \n            # 构建图像路径（根据ImageNet21K的结构）\n            # 注意：这里需要根据实际的ImageNet21K数据集结构进行调整\n            image_filename = f\"{wordnet_id}_{image_idx:04d}.JPEG\"\n            \n            # 尝试从Hugging Face加载\n            # 由于ImageNet21K数据集很大，这里我们使用一个占位符\n            # 在实际使用中，您需要根据数据集的实际结构来加载图像\n            \n            # 这里我们使用模拟图像作为占位符\n            # 在实际部署时，您需要取消注释下面的代码并适配您的数据加载逻辑\n            \n            \"\"\"\n            from datasets import load_dataset\n            try:\n                # 加载数据集（注意：这可能需要很长时间和大量内存）\n                dataset = load_dataset(self.imagenet21k_repo, split='train')\n                \n                # 根据类别和图像索引查找图像\n                # 这里需要根据实际的数据集结构来编写查找逻辑\n                # 由于ImageNet21K数据集很大，这个操作可能很慢\n                \n                # 示例查找逻辑（需要根据实际数据集结构调整）\n                image_data = dataset.filter(\n                    lambda x: x['wordnet_id'] == wordnet_id and x['image_id'] == image_idx\n                )[0]\n                \n                image = Image.open(BytesIO(image_data['image']))\n                image = image.convert('RGB')\n                \n            except Exception as e:\n                print(f\"从Hugging Face加载图像失败: {e}\")\n                # 回退到模拟图像\n                image = self._generate_simulated_image(category_idx, image_idx)\n            \"\"\"\n            \n            # 使用模拟图像（在实际部署时删除这部分）\n            image = self._generate_simulated_image(category_idx, image_idx)\n            \n            # 调整图像大小\n            image = image.resize((self.image_size, self.image_size))\n            \n            # 转换为张量\n            image_tensor = torch.from_numpy(np.array(image)).float() / 255.0\n            image_tensor = image_tensor.permute(2, 0, 1)  # (H, W, C) -> (C, H, W)\n            \n            return image_tensor\n            \n        except Exception as e:\n            print(f\"图像加载错误: {e}\")\n            # 生成模拟图像作为备用\n            return self._generate_simulated_image(category_idx, image_idx)\n    \n    def _generate_simulated_image(self, category_idx, image_idx):\n        \"\"\"生成模拟图像（当无法加载真实图像时使用）\"\"\"\n        img_size = self.image_size\n        img = np.zeros((img_size, img_size, 3), dtype=np.float32)\n        \n        # 基于类别和图像索引生成不同的图案\n        center_x, center_y = img_size // 2, img_size // 2\n        \n        # 使用类别和图像索引决定颜色和形状\n        hue1 = (category_idx * 17) % 360 / 360\n        hue2 = (image_idx * 23) % 360 / 360\n        \n        # 转换为RGB\n        from colorsys import hsv_to_rgb\n        color1 = hsv_to_rgb(hue1, 0.8, 0.9)\n        color2 = hsv_to_rgb(hue2, 0.6, 0.7)\n        \n        # 绘制背景\n        img[:, :] = color1\n        \n        # 绘制前景形状\n        shape_type = (category_idx + image_idx) % 4\n        \n        if shape_type == 0:  # 圆形\n            radius = img_size // 4 + (image_idx % 10)\n            y, x = np.ogrid[-center_y:img_size-center_y, -center_x:img_size-center_x]\n            mask = x*x + y*y <= radius*radius\n            img[mask] = color2\n            \n        elif shape_type == 1:  # 矩形\n            width = img_size // 3 + (image_idx % 15)\n            height = img_size // 4 + (image_idx % 12)\n            x1, y1 = center_x - width//2, center_y - height//2\n            x2, y2 = center_x + width//2, center_y + height//2\n            img[y1:y2, x1:x2] = color2\n            \n        elif shape_type == 2:  # 三角形\n            vertices = np.array([\n                [center_x, center_y - img_size//4],\n                [center_x - img_size//4, center_y + img_size//4],\n                [center_x + img_size//4, center_y + img_size//4]\n            ])\n            from matplotlib.path import Path\n            path = Path(vertices)\n            x, y = np.meshgrid(np.arange(img_size), np.arange(img_size))\n            points = np.vstack((x.flatten(), y.flatten())).T\n            mask = path.contains_points(points).reshape((img_size, img_size))\n            img[mask] = color2\n            \n        else:  # 菱形\n            vertices = np.array([\n                [center_x, center_y - img_size//4],\n                [center_x - img_size//4, center_y],\n                [center_x, center_y + img_size//4],\n                [center_x + img_size//4, center_y]\n            ])\n            from matplotlib.path import Path\n            path = Path(vertices)\n            x, y = np.meshgrid(np.arange(img_size), np.arange(img_size))\n            points = np.vstack((x.flatten(), y.flatten())).T\n            mask = path.contains_points(points).reshape((img_size, img_size))\n            img[mask] = color2\n        \n        # 转换为PIL图像\n        img_pil = Image.fromarray((img * 255).astype(np.uint8))\n        return img_pil\n    \n    def _get_image(self, idx):\n        \"\"\"获取图像（带缓存）\"\"\"\n        metadata = self.image_metadata[idx]\n        \n        if isinstance(metadata, dict):\n            category_idx = metadata.get('category_id', metadata.get('category'))\n            image_idx = metadata.get('image_id', metadata.get('image_idx', idx))\n        else:\n            # 如果metadata是简单的标签\n            category_idx = metadata if isinstance(metadata, (int, torch.Tensor)) else idx % 80\n            image_idx = idx % 50\n        \n        cache_key = f\"{category_idx}_{image_idx}\"\n        \n        # 检查缓存\n        if cache_key in self.image_cache:\n            return self.image_cache[cache_key]\n        \n        # 从ImageNet21K加载图像\n        image = self._load_image_from_imagenet21k(category_idx, image_idx)\n        \n        # 缓存管理\n        if len(self.image_cache) >= self.image_cache_size:\n            # 移除最旧的缓存\n            oldest_key = self.cache_keys.pop(0)\n            if oldest_key in self.image_cache:\n                del self.image_cache[oldest_key]\n        \n        # 添加到缓存\n        self.image_cache[cache_key] = image\n        self.cache_keys.append(cache_key)\n        \n        return image\n    \n    def _preprocess_eeg(self, eeg):\n        \"\"\"EEG预处理 - 基于论文描述\"\"\"\n        # 1. 标准化\n        eeg_mean = eeg.mean(dim=1, keepdim=True)\n        eeg_std = eeg.std(dim=1, keepdim=True)\n        eeg_std = torch.clamp(eeg_std, min=1e-8)\n        eeg = (eeg - eeg_mean) / eeg_std\n        \n        # 2. 提取40ms-440ms段 (论文中的特征提取窗口)\n        start_idx, end_idx = 40, 440  # 40ms to 440ms at 1000Hz sampling rate\n        if eeg.shape[1] >= end_idx:\n            eeg = eeg[:, start_idx:end_idx]\n        \n        return eeg\n    \n    def __len__(self):\n        if self.max_samples:\n            return min(self.max_samples, len(self.eeg_data))\n        return len(self.eeg_data)\n    \n    def __getitem__(self, idx):\n        # 确保索引在有效范围内\n        if idx >= len(self.eeg_data):\n            idx = len(self.eeg_data) - 1\n        \n        # 获取EEG数据\n        eeg = self.eeg_data[idx]\n        if isinstance(eeg, torch.Tensor):\n            eeg = eeg.float()\n        else:\n            eeg = torch.tensor(eeg, dtype=torch.float32)\n        \n        # 获取图像\n        image = self._get_image(idx)\n        \n        # 获取标签\n        metadata = self.image_metadata[idx]\n        if isinstance(metadata, dict):\n            label = metadata.get('category_id', metadata.get('category'))\n        else:\n            label = metadata if isinstance(metadata, (int, torch.Tensor)) else idx % 80\n        \n        if isinstance(label, torch.Tensor):\n            label = label.item() if label.dim() == 0 else label[0]\n        \n        # 确保标签在有效范围内 (0-79)\n        label = min(int(label), 79)\n        \n        # 应用EEG预处理\n        eeg = self._preprocess_eeg(eeg)\n        \n        # 应用图像变换\n        if self.transform:\n            image = self.transform(image)\n        \n        return eeg, image, torch.tensor(label)\n    \n    def clear_cache(self):\n        \"\"\"清理图像缓存\"\"\"\n        self.image_cache.clear()\n        self.cache_keys.clear()\n        clear_memory()\n\n# ==================== 优化的模型架构 ====================\n\nclass EfficientEEGEncoder(nn.Module):\n    \"\"\"高效的EEG编码器\"\"\"\n    def __init__(self, input_channels=62, time_points=400, hidden_dim=256, output_dim=256):\n        super().__init__()\n        \n        # 时域特征提取\n        self.time_conv = nn.Sequential(\n            nn.Conv1d(input_channels, 64, kernel_size=15, padding=7),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(64, 128, kernel_size=11, padding=5),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n        )\n        \n        # 频域特征提取（通过不同卷积核模拟）\n        self.freq_conv = nn.Sequential(\n            nn.Conv1d(input_channels, 64, kernel_size=25, padding=12),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(64, 128, kernel_size=15, padding=7),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n        )\n        \n        # 特征融合\n        self.feature_fusion = nn.Sequential(\n            nn.Linear(256, hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(hidden_dim, output_dim)\n        )\n        \n    def forward(self, x):\n        # 时域特征\n        time_feat = self.time_conv(x)\n        time_feat = torch.mean(time_feat, dim=2)  # 全局平均池化\n        \n        # 频域特征\n        freq_feat = self.freq_conv(x)\n        freq_feat = torch.mean(freq_feat, dim=2)  # 全局平均池化\n        \n        # 特征融合\n        combined = torch.cat([time_feat, freq_feat], dim=1)\n        output = self.feature_fusion(combined)\n        \n        return output\n\nclass LightweightImageDecoder(nn.Module):\n    \"\"\"轻量级图像解码器\"\"\"\n    def __init__(self, input_dim=256, output_channels=3, image_size=128):\n        super().__init__()\n        \n        self.init_proj = nn.Linear(input_dim, 256 * 4 * 4)\n        \n        # 解码器上采样路径\n        self.decoder_layers = nn.Sequential(\n            # 4x4 -> 8x8\n            nn.ConvTranspose2d(256, 128, 4, 2, 1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            \n            # 8x8 -> 16x16\n            nn.ConvTranspose2d(128, 64, 4, 2, 1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            \n            # 16x16 -> 32x32\n            nn.ConvTranspose2d(64, 32, 4, 2, 1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            # 32x32 -> 64x64\n            nn.ConvTranspose2d(32, 16, 4, 2, 1),\n            nn.BatchNorm2d(16),\n            nn.ReLU(),\n            \n            # 64x64 -> 128x128\n            nn.ConvTranspose2d(16, 8, 4, 2, 1),\n            nn.BatchNorm2d(8),\n            nn.ReLU(),\n            \n            # 最终卷积层\n            nn.Conv2d(8, output_channels, 3, 1, 1),\n            nn.Tanh()\n        )\n        \n    def forward(self, x):\n        x = self.init_proj(x)\n        x = x.view(-1, 256, 4, 4)\n        x = self.decoder_layers(x)\n        return x\n\nclass EfficientEEGToImageModel(nn.Module):\n    \"\"\"高效的EEG到图像转换模型\"\"\"\n    def __init__(self, eeg_channels=62, time_points=400, hidden_dim=256, output_channels=3, image_size=128):\n        super().__init__()\n        self.encoder = EfficientEEGEncoder(eeg_channels, time_points, hidden_dim, hidden_dim)\n        self.decoder = LightweightImageDecoder(hidden_dim, output_channels, image_size)\n        \n    def forward(self, eeg_data):\n        encoded_features = self.encoder(eeg_data)\n        generated_image = self.decoder(encoded_features)\n        return generated_image\n\n# ==================== 训练函数 ====================\n\ndef train_model(model, train_loader, val_loader, num_epochs=50, learning_rate=1e-4):\n    \"\"\"训练模型\"\"\"\n    model.to(device)\n    \n    # 损失函数和优化器\n    criterion = nn.MSELoss()\n    optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    history = {\n        'train_loss': [],\n        'val_loss': [],\n        'learning_rate': []\n    }\n    \n    best_val_loss = float('inf')\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0\n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        \n        for batch_idx, (eeg_data, images, labels) in enumerate(train_bar):\n            # 检查内存\n            memory_manager.auto_clear()\n            \n            eeg_data = eeg_data.to(device, non_blocking=True)\n            images = images.to(device, non_blocking=True)\n            \n            optimizer.zero_grad()\n            \n            # 前向传播\n            generated_images = model(eeg_data)\n            loss = criterion(generated_images, images)\n            \n            # 反向传播\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_bar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n            # 定期清理内存\n            if batch_idx % 20 == 0:\n                clear_memory()\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            val_bar = tqdm(val_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Val]')\n            for eeg_data, images, labels in val_bar:\n                # 检查内存\n                memory_manager.auto_clear()\n                \n                eeg_data = eeg_data.to(device, non_blocking=True)\n                images = images.to(device, non_blocking=True)\n                \n                generated_images = model(eeg_data)\n                loss = criterion(generated_images, images)\n                val_loss += loss.item()\n                val_bar.set_postfix({'val_loss': f'{loss.item():.4f}'})\n        \n        # 计算平均损失\n        avg_train_loss = train_loss / len(train_loader)\n        avg_val_loss = val_loss / len(val_loader)\n        \n        # 更新学习率\n        current_lr = optimizer.param_groups[0]['lr']\n        scheduler.step()\n        \n        # 记录历史\n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n        history['learning_rate'].append(current_lr)\n        \n        print(f'Epoch {epoch+1}/{num_epochs}:')\n        print(f'  训练损失: {avg_train_loss:.4f}, 验证损失: {avg_val_loss:.4f}, LR: {current_lr:.6f}')\n        \n        # 保存最佳模型\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            torch.save(model.state_dict(), 'best_eeg_to_image_model.pth')\n            print(f'  保存最佳模型，验证损失: {best_val_loss:.4f}')\n        \n        # 每10个epoch可视化一次结果\n        if (epoch + 1) % 10 == 0:\n            visualize_training_results(model, val_loader, epoch + 1)\n            \n        # 清理内存\n        clear_memory()\n    \n    return model, history\n\ndef visualize_training_results(model, val_loader, epoch):\n    \"\"\"可视化训练结果\"\"\"\n    model.eval()\n    with torch.no_grad():\n        # 获取一个batch的验证数据\n        eeg_data, real_images, labels = next(iter(val_loader))\n        eeg_data = eeg_data[:4].to(device)\n        real_images = real_images[:4]\n        \n        # 生成图像\n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 可视化\n        fig, axes = plt.subplots(3, 4, figsize=(16, 12))\n        \n        for i in range(4):\n            # 显示EEG信号（前8个通道）\n            eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n            axes[0, i].plot(eeg_sample.T, alpha=0.7, linewidth=1)\n            axes[0, i].set_title(f'EEG样本 {i+1}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            \n            # 显示真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实图像', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 显示生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('重建图像', fontsize=10)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(f'training_results_epoch_{epoch}.png', dpi=200, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"训练结果已保存: training_results_epoch_{epoch}.png\")\n        \n        # 清理内存\n        clear_memory()\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"开始EEG-to-Image重建任务...\")\n    \n    # 数据转换\n    transform = transforms.Compose([\n        transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n    ])\n    \n    # 数据路径\n    data_dir = \"/kaggle/input/eeg-imagenet\"\n    pth_file_path = os.path.join(data_dir, \"EEG-ImageNet_1.pth\")\n    \n    # 创建数据集\n    print(\"创建真实数据集...\")\n    train_dataset = RealEEGImageNetDataset(\n        eeg_data_path=pth_file_path,\n        imagenet21k_repo=\"gmongaras/Imagenet21K\",\n        transform=transform,\n        image_size=128,\n        mode='train',\n        subject_id='S1',\n        max_samples=800,  # 限制样本数量用于测试\n        image_cache_size=50\n    )\n    \n    test_dataset = RealEEGImageNetDataset(\n        eeg_data_path=pth_file_path,\n        imagenet21k_repo=\"gmongaras/Imagenet21K\",\n        transform=transform,\n        image_size=128,\n        mode='test',\n        subject_id='S1',\n        max_samples=200,  # 限制样本数量用于测试\n        image_cache_size=20\n    )\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=8,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True,\n        drop_last=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=8,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=True,\n        drop_last=True\n    )\n    \n    print(f\"训练集: {len(train_dataset)} 样本\")\n    print(f\"测试集: {len(test_dataset)} 样本\")\n    \n    # 创建模型\n    print(\"初始化模型...\")\n    model = EfficientEEGToImageModel(\n        eeg_channels=62,\n        time_points=400,\n        hidden_dim=256,\n        output_channels=3,\n        image_size=128\n    )\n    \n    # 检查是否有预训练模型\n    if os.path.exists('best_eeg_to_image_model.pth'):\n        print(\"加载预训练模型...\")\n        model.load_state_dict(torch.load('best_eeg_to_image_model.pth', map_location=device))\n        print(\"模型加载成功!\")\n    else:\n        print(\"从头开始训练模型...\")\n    \n    # 训练模型\n    print(\"开始训练...\")\n    model, history = train_model(\n        model=model,\n        train_loader=train_loader,\n        val_loader=test_loader,\n        num_epochs=30,\n        learning_rate=1e-4\n    )\n    \n    # 绘制训练历史\n    plt.figure(figsize=(12, 4))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(history['train_loss'], label='训练损失')\n    plt.plot(history['val_loss'], label='验证损失')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('训练和验证损失')\n    plt.grid(True)\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(history['learning_rate'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Learning Rate')\n    plt.title('学习率变化')\n    plt.grid(True)\n    \n    plt.tight_layout()\n    plt.savefig('training_history.png', dpi=150, bbox_inches='tight')\n    plt.show()\n    \n    # 最终评估\n    print(\"最终评估...\")\n    final_evaluation(model, test_loader)\n    \n    # 生成最终结果可视化\n    print(\"生成最终结果...\")\n    visualize_final_results(model, test_loader)\n    \n    # 清理所有缓存\n    train_dataset.clear_cache()\n    test_dataset.clear_cache()\n    clear_memory()\n    \n    print(\"任务完成!\")\n    return model\n\ndef final_evaluation(model, test_loader):\n    \"\"\"最终评估\"\"\"\n    model.eval()\n    test_loss = 0\n    criterion = nn.MSELoss()\n    \n    with torch.no_grad():\n        for eeg_data, images, labels in test_loader:\n            eeg_data = eeg_data.to(device)\n            images = images.to(device)\n            \n            generated_images = model(eeg_data)\n            loss = criterion(generated_images, images)\n            test_loss += loss.item()\n            \n            # 定期清理内存\n            memory_manager.auto_clear()\n    \n    avg_test_loss = test_loss / len(test_loader)\n    print(f'测试损失: {avg_test_loss:.4f}')\n    return avg_test_loss\n\ndef visualize_final_results(model, test_loader):\n    \"\"\"生成最终结果可视化\"\"\"\n    model.eval()\n    with torch.no_grad():\n        eeg_data, real_images, labels = next(iter(test_loader))\n        eeg_data = eeg_data.to(device)\n        \n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 创建可视化\n        n_samples = min(6, len(eeg_data))\n        fig, axes = plt.subplots(3, n_samples, figsize=(2 * n_samples, 6))\n        \n        if n_samples == 1:\n            axes = axes.reshape(3, 1)\n        \n        for i in range(n_samples):\n            # EEG信号\n            eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n            axes[0, i].plot(eeg_sample.T, alpha=0.7, linewidth=1)\n            axes[0, i].set_title(f'EEG {i+1}', fontsize=8)\n            axes[0, i].grid(True, alpha=0.3)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实', fontsize=8)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('重建', fontsize=8)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig('final_reconstruction_results.png', dpi=150, bbox_inches='tight')\n        plt.show()\n        \n        print(\"最终结果已保存: final_reconstruction_results.png\")\n        \n        # 清理内存\n        clear_memory()\n\nif __name__ == \"__main__\":\n    try:\n        model = main()\n        print(\"程序执行成功!\")\n        \n    except Exception as e:\n        print(f\"错误: {e}\")\n        import traceback\n        traceback.print_exc()\n    finally:\n        # 最终清理\n        clear_memory()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-11T14:39:10.079291Z","iopub.execute_input":"2025-10-11T14:39:10.079559Z","iopub.status.idle":"2025-10-11T14:41:33.637324Z","shell.execute_reply.started":"2025-10-11T14:39:10.079536Z","shell.execute_reply":"2025-10-11T14:41:33.636681Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet真实数据集适配版本 - 修复版2\n修复EEG数据格式问题\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport gc\nimport time\nimport math\nfrom PIL import Image\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ==================== 设备配置和内存管理 ====================\n\ndef setup_device():\n    \"\"\"设置计算设备\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda:0\")\n        print(f\"使用GPU: {torch.cuda.get_device_name(0)}\")\n        torch.cuda.empty_cache()\n    else:\n        device = torch.device(\"cpu\")\n        print(\"使用CPU\")\n    return device\n\ndef set_seed(seed=42):\n    \"\"\"设置随机种子\"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\ndevice = setup_device()\nset_seed(42)\n\nclass MemoryManager:\n    \"\"\"内存管理器\"\"\"\n    def __init__(self):\n        self.image_cache = {}\n        self.cache_size = 50\n        \n    def clear_memory(self):\n        \"\"\"清理内存\"\"\"\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n            \n    def cache_image(self, key, image):\n        \"\"\"缓存图像\"\"\"\n        if len(self.image_cache) >= self.cache_size:\n            # 移除最旧的缓存\n            oldest_key = next(iter(self.image_cache))\n            del self.image_cache[oldest_key]\n        self.image_cache[key] = image\n        \n    def get_cached_image(self, key):\n        \"\"\"获取缓存的图像\"\"\"\n        return self.image_cache.get(key, None)\n\nmemory_manager = MemoryManager()\n\n# ==================== 修复的EEG数据加载器 ====================\n\nclass FixedEEGImageNetDataset(Dataset):\n    \"\"\"修复的EEG-ImageNet数据集加载器\"\"\"\n    def __init__(self, eeg_data_path, transform=None, image_size=128, \n                 mode='train', subject_id='S1', max_samples=None):\n        super().__init__()\n        self.eeg_data_path = eeg_data_path\n        self.transform = transform\n        self.image_size = image_size\n        self.mode = mode\n        self.subject_id = subject_id\n        self.max_samples = max_samples\n        \n        # 加载EEG数据 - 使用修复的加载方法\n        self.eeg_data, self.image_metadata = self._load_eeg_data_fixed()\n        \n        print(f\"{mode}数据集 - 被试 {subject_id}: {len(self.eeg_data)} 个样本\")\n        \n    def _load_eeg_data_fixed(self):\n        \"\"\"修复的EEG数据加载方法\"\"\"\n        try:\n            if os.path.exists(self.eeg_data_path):\n                print(f\"尝试修复加载EEG数据从: {self.eeg_data_path}\")\n                \n                # 方法1: 使用 weights_only=False\n                try:\n                    data = torch.load(self.eeg_data_path, map_location='cpu', weights_only=False)\n                    print(\"使用 weights_only=False 成功加载数据\")\n                except Exception as e:\n                    print(f\"方法1失败: {e}\")\n                    \n                    # 方法2: 使用 pickle 直接加载\n                    import pickle\n                    with open(self.eeg_data_path, 'rb') as f:\n                        data = pickle.load(f)\n                    print(\"使用 pickle 成功加载数据\")\n                \n                # 详细检查数据结构\n                print(f\"数据整体类型: {type(data)}\")\n                if isinstance(data, dict):\n                    print(\"数据格式: 字典\")\n                    print(f\"字典键: {list(data.keys())}\")\n                    \n                    # 详细检查每个键的内容\n                    for key, value in data.items():\n                        if hasattr(value, 'shape'):\n                            print(f\"  {key}: 形状 {value.shape}, 类型 {type(value)}\")\n                        else:\n                            print(f\"  {key}: 类型 {type(value)}, 长度 {len(value) if hasattr(value, '__len__') else 'N/A'}\")\n                    \n                    # 尝试找到EEG数据\n                    eeg_data = None\n                    possible_eeg_keys = ['dataset', 'data', 'eeg_data', 'eeg', 'X', 'features']\n                    for key in possible_eeg_keys:\n                        if key in data:\n                            eeg_data = data[key]\n                            print(f\"找到EEG数据在键 '{key}': 形状 {eeg_data.shape if hasattr(eeg_data, 'shape') else 'N/A'}\")\n                            break\n                    \n                    # 如果没找到，尝试第一个数组类型的数据\n                    if eeg_data is None:\n                        for key, value in data.items():\n                            if hasattr(value, 'shape') and len(value.shape) >= 2:\n                                eeg_data = value\n                                print(f\"使用数组数据从键 '{key}': 形状 {eeg_data.shape}\")\n                                break\n                    \n                    # 处理图像元数据\n                    image_metadata = None\n                    possible_meta_keys = ['labels', 'image_info', 'metadata', 'targets']\n                    for key in possible_meta_keys:\n                        if key in data:\n                            image_metadata = data[key]\n                            print(f\"找到图像元数据在键 '{key}'\")\n                            break\n                    \n                    # 如果没有找到元数据，创建模拟元数据\n                    if image_metadata is None:\n                        print(\"创建模拟图像元数据\")\n                        if eeg_data is not None:\n                            if hasattr(eeg_data, 'shape'):\n                                n_samples = eeg_data.shape[0]\n                            else:\n                                n_samples = len(eeg_data)\n                        else:\n                            n_samples = 100 if self.mode == 'train' else 50\n                        \n                        image_metadata = []\n                        for i in range(n_samples):\n                            metadata = {\n                                'category_id': i % 80,\n                                'image_id': i % 50,\n                                'category_name': f'category_{i % 80}',\n                                'wordnet_id': f'n{str((i % 80) + 1).zfill(8)}',\n                                'image_index': i\n                            }\n                            image_metadata.append(metadata)\n                            \n                elif isinstance(data, (list, tuple)):\n                    print(f\"数据格式: 列表/元组, 长度: {len(data)}\")\n                    eeg_data = data[0] if len(data) > 0 else None\n                    image_metadata = data[1] if len(data) > 1 else None\n                else:\n                    print(f\"未知数据格式: {type(data)}\")\n                    eeg_data = data\n                    image_metadata = None\n                \n                # 处理EEG数据形状\n                if eeg_data is not None:\n                    print(f\"原始EEG数据类型: {type(eeg_data)}\")\n                    \n                    if hasattr(eeg_data, 'shape'):\n                        print(f\"原始EEG数据形状: {eeg_data.shape}\")\n                    \n                    # 转换为numpy数组进行处理\n                    if isinstance(eeg_data, torch.Tensor):\n                        eeg_np = eeg_data.numpy()\n                    else:\n                        eeg_np = np.array(eeg_data)\n                    \n                    print(f\"转换后EEG数据形状: {eeg_np.shape}\")\n                    \n                    # 检查并修复数据形状\n                    if len(eeg_np.shape) == 1:\n                        print(\"数据是1D的，尝试重新形状为3D\")\n                        # 假设数据是展平的，尝试重新形状\n                        total_elements = eeg_np.size\n                        print(f\"总元素数: {total_elements}\")\n                        \n                        # 常见的EEG形状: (n_samples, 62, 500) = n_samples * 62 * 500\n                        expected_shape = (31950, 62, 500)  # 根据你的数据调整\n                        if total_elements == np.prod(expected_shape):\n                            eeg_np = eeg_np.reshape(expected_shape)\n                            print(f\"重新形状为: {eeg_np.shape}\")\n                        else:\n                            # 尝试自动推断形状\n                            n_samples = 100 if self.mode == 'train' else 50\n                            remaining = total_elements // n_samples\n                            print(f\"尝试推断形状: ({n_samples}, ?, ?), 剩余元素: {remaining}\")\n                            \n                            # 寻找可能的形状\n                            for n_channels in [32, 62, 64, 128]:\n                                if remaining % n_channels == 0:\n                                    n_times = remaining // n_channels\n                                    try_shape = (n_samples, n_channels, n_times)\n                                    if np.prod(try_shape) == total_elements:\n                                        eeg_np = eeg_np.reshape(try_shape)\n                                        print(f\"推断形状为: {eeg_np.shape}\")\n                                        break\n                    \n                    # 转换为Tensor\n                    eeg_data = torch.tensor(eeg_np, dtype=torch.float32)\n                    print(f\"最终EEG数据形状: {eeg_data.shape}\")\n                    \n                else:\n                    print(\"无法找到EEG数据，使用模拟数据\")\n                    return self._create_simulated_eeg_data()\n                \n                return eeg_data, image_metadata\n            else:\n                raise FileNotFoundError(f\"EEG数据文件不存在: {self.eeg_data_path}\")\n                \n        except Exception as e:\n            print(f\"所有EEG数据加载方法都失败了: {e}\")\n            print(\"使用模拟EEG数据\")\n            return self._create_simulated_eeg_data()\n    \n    def _create_simulated_eeg_data(self):\n        \"\"\"创建模拟EEG数据\"\"\"\n        n_samples = 100 if self.mode == 'train' else 50\n        n_categories = 80\n        \n        eeg_data = []\n        image_metadata = []\n        \n        for i in range(n_samples):\n            # 生成神经科学合理的EEG信号\n            eeg = self._generate_realistic_eeg(i % n_categories, i)\n            eeg_data.append(eeg)\n            \n            metadata = {\n                'category_id': i % n_categories,\n                'image_id': i % 50,\n                'category_name': f'category_{i % n_categories}',\n                'wordnet_id': f'n{str((i % n_categories) + 1).zfill(8)}',\n                'image_index': i\n            }\n            image_metadata.append(metadata)\n        \n        return torch.stack(eeg_data), image_metadata\n    \n    def _generate_realistic_eeg(self, category_idx, sample_idx):\n        \"\"\"生成神经科学合理的EEG信号\"\"\"\n        n_channels, n_times = 62, 500\n        eeg = torch.zeros(n_channels, n_times)\n        \n        time_points = torch.linspace(0, 2 * np.pi, n_times)\n        \n        # 基于类别生成独特的EEG模式\n        base_freq = 10 + (category_idx % 10)\n        \n        # 不同脑区有不同的活动模式\n        for ch in range(16):  # 前额叶\n            freq = base_freq + ch * 0.2\n            eeg[ch] += 0.3 * torch.sin(freq * time_points + sample_idx * 0.1)\n        \n        for ch in range(16, 32):  # 中央区\n            freq = base_freq * 1.5 + (ch-16) * 0.3\n            eeg[ch] += 0.4 * torch.cos(freq * time_points + sample_idx * 0.15)\n        \n        for ch in range(32, 48):  # 顶叶\n            freq = base_freq * 2 + (ch-32) * 0.4\n            eeg[ch] += 0.2 * torch.sin(freq * time_points * 0.7 + sample_idx * 0.2)\n        \n        for ch in range(48, 62):  # 枕叶（视觉区）\n            freq = base_freq * 2.5 + (ch-48) * 0.5\n            eeg[ch] += 0.5 * torch.sin(freq * time_points + sample_idx * 0.25)\n        \n        # 添加生理噪声\n        eeg += 0.1 * torch.randn(n_channels, n_times)\n        \n        return eeg\n    \n    def _load_imagenet_image(self, metadata):\n        \"\"\"从Hugging Face ImageNet21K加载图像\"\"\"\n        try:\n            cache_key = f\"{metadata['category_id']}_{metadata['image_id']}\"\n            \n            # 检查缓存\n            cached_image = memory_manager.get_cached_image(cache_key)\n            if cached_image is not None:\n                return cached_image\n            \n            # 使用模拟图像（简化实现）\n            return self._generate_semantic_image(metadata)\n            \n        except Exception as e:\n            print(f\"图像加载错误: {e}\")\n            return self._generate_semantic_image(metadata)\n    \n    def _generate_semantic_image(self, metadata):\n        \"\"\"生成语义相关的模拟图像\"\"\"\n        img_size = self.image_size\n        category_id = metadata['category_id']\n        image_id = metadata['image_id']\n        \n        # 基于类别ID生成有意义的图像\n        img = Image.new('RGB', (img_size, img_size), color='white')\n        \n        from PIL import ImageDraw\n        draw = ImageDraw.Draw(img)\n        \n        center_x, center_y = img_size // 2, img_size // 2\n        \n        # 基于类别ID决定图像类型\n        category_type = category_id % 6\n        \n        if category_type == 0:  # 圆形物体\n            color = (100, 150, 200)\n            radius = img_size // 4 + (image_id % 10)\n            draw.ellipse([center_x-radius, center_y-radius, center_x+radius, center_y+radius], \n                        fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 1:  # 方形物体\n            color = (200, 100, 100)\n            size = img_size // 3 + (image_id % 15)\n            draw.rectangle([center_x-size//2, center_y-size//2, center_x+size//2, center_y+size//2], \n                          fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 2:  # 三角形物体\n            color = (150, 200, 100)\n            size = img_size // 4 + (image_id % 12)\n            points = [\n                (center_x, center_y - size),\n                (center_x - size, center_y + size),\n                (center_x + size, center_y + size)\n            ]\n            draw.polygon(points, fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 3:  # 线条图案\n            color = (200, 150, 100)\n            for i in range(0, img_size, img_size // 8):\n                draw.line([i, 0, i, img_size], fill=color, width=2)\n                draw.line([0, i, img_size, i], fill=color, width=2)\n                \n        elif category_type == 4:  # 点状图案\n            color = (150, 100, 200)\n            for i in range(0, img_size, img_size // 10):\n                for j in range(0, img_size, img_size // 10):\n                    if (i + j) % 20 == 0:\n                        draw.ellipse([i-3, j-3, i+3, j+3], fill=color)\n                        \n        else:  # 混合图案\n            color = (100, 200, 200)\n            # 圆形\n            draw.ellipse([center_x-30, center_y-20, center_x+30, center_y+20], fill=color)\n            # 矩形\n            draw.rectangle([center_x-15, center_y-10, center_x+15, center_y+10], \n                          fill=(255,255,255), outline=(0,0,0), width=1)\n        \n        return img\n    \n    def _preprocess_eeg(self, eeg):\n        \"\"\"EEG预处理\"\"\"\n        if isinstance(eeg, torch.Tensor):\n            eeg = eeg.float()\n        else:\n            eeg = torch.tensor(eeg, dtype=torch.float32)\n        \n        # 确保正确的形状 (channels, time)\n        if len(eeg.shape) == 3:\n            # 如果是(batch, channels, time)，取第一个样本\n            eeg = eeg[0]\n        elif len(eeg.shape) == 1:\n            # 如果是1D，尝试重新形状\n            eeg = eeg.reshape(62, -1)\n        \n        # 标准化\n        eeg_mean = eeg.mean(dim=1, keepdim=True)\n        eeg_std = eeg.std(dim=1, keepdim=True)\n        eeg_std = torch.clamp(eeg_std, min=1e-8)\n        eeg = (eeg - eeg_mean) / eeg_std\n        \n        # 提取40ms-440ms段 (如果时间维度足够)\n        if eeg.shape[1] >= 440:\n            start_idx, end_idx = 40, 440\n            eeg = eeg[:, start_idx:end_idx]\n        elif eeg.shape[1] > 100:\n            # 如果时间维度不够，取中间部分\n            start_idx = eeg.shape[1] // 4\n            end_idx = start_idx + 400\n            if end_idx > eeg.shape[1]:\n                end_idx = eeg.shape[1]\n            eeg = eeg[:, start_idx:end_idx]\n        \n        return eeg\n    \n    def __len__(self):\n        if self.max_samples:\n            return min(self.max_samples, len(self.eeg_data))\n        return len(self.eeg_data)\n    \n    def __getitem__(self, idx):\n        if idx >= len(self.eeg_data):\n            idx = len(self.eeg_data) - 1\n        \n        # 获取EEG数据\n        eeg = self.eeg_data[idx]\n        eeg = self._preprocess_eeg(eeg)\n        \n        # 获取图像元数据\n        metadata = self.image_metadata[idx]\n        \n        # 加载图像\n        image = self._load_imagenet_image(metadata)\n        \n        # 转换为张量\n        image_tensor = transforms.ToTensor()(image)\n        \n        # 获取标签\n        label = metadata['category_id']\n        \n        # 应用变换\n        if self.transform:\n            image_tensor = self.transform(image_tensor)\n        \n        return eeg, image_tensor, torch.tensor(label)\n\n# ==================== 高效的EEG到图像模型 ====================\n\nclass EfficientEEGEncoder(nn.Module):\n    \"\"\"高效的EEG编码器\"\"\"\n    def __init__(self, input_channels=62, time_points=400, hidden_dim=256):\n        super().__init__()\n        \n        self.conv_layers = nn.Sequential(\n            nn.Conv1d(input_channels, 64, kernel_size=15, padding=7),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(64, 128, kernel_size=11, padding=5),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(128, 256, kernel_size=7, padding=3),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n        )\n        \n        self.attention = nn.MultiheadAttention(256, num_heads=4, dropout=0.1)\n        self.global_pool = nn.AdaptiveAvgPool1d(1)\n        self.fc = nn.Linear(256, 512)\n        \n    def forward(self, x):\n        # 输入: (batch, channels, time)\n        x = self.conv_layers(x)  # (batch, 256, time)\n        \n        # 应用注意力\n        x_attn = x.transpose(1, 2)  # (batch, time, 256)\n        x_attn = x_attn.transpose(0, 1)  # (time, batch, 256)\n        attn_out, _ = self.attention(x_attn, x_attn, x_attn)\n        attn_out = attn_out.transpose(0, 1).transpose(1, 2)  # (batch, 256, time)\n        \n        x = x + 0.1 * attn_out\n        \n        # 全局池化\n        x = self.global_pool(x).squeeze(-1)  # (batch, 256)\n        x = self.fc(x)  # (batch, 512)\n        \n        return x\n\nclass UNetDecoder(nn.Module):\n    \"\"\"UNet风格解码器\"\"\"\n    def __init__(self, input_dim=512, output_channels=3, image_size=128):\n        super().__init__()\n        \n        self.init_proj = nn.Linear(input_dim, 512 * 4 * 4)\n        \n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(512, 256, 4, 2, 1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(256, 128, 4, 2, 1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(128, 64, 4, 2, 1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(64, 32, 4, 2, 1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(32, 16, 4, 2, 1),\n            nn.BatchNorm2d(16),\n            nn.ReLU(),\n            \n            nn.Conv2d(16, output_channels, 3, 1, 1),\n            nn.Tanh()\n        )\n        \n    def forward(self, x):\n        x = self.init_proj(x)\n        x = x.view(-1, 512, 4, 4)\n        x = self.decoder(x)\n        return x\n\nclass EEGToImageModel(nn.Module):\n    \"\"\"EEG到图像转换模型\"\"\"\n    def __init__(self, eeg_channels=62, time_points=400, output_channels=3, image_size=128):\n        super().__init__()\n        self.encoder = EfficientEEGEncoder(eeg_channels, time_points)\n        self.decoder = UNetDecoder(output_channels=output_channels, image_size=image_size)\n        \n    def forward(self, eeg_data):\n        features = self.encoder(eeg_data)\n        images = self.decoder(features)\n        return images\n\n# ==================== 训练和可视化函数 ====================\n\ndef train_model(model, train_loader, val_loader, num_epochs=5, learning_rate=1e-4):\n    \"\"\"训练模型\"\"\"\n    model.to(device)\n    criterion = nn.MSELoss()\n    optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    history = {'train_loss': [], 'val_loss': []}\n    best_val_loss = float('inf')\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0\n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        \n        for batch_idx, (eeg_data, images, labels) in enumerate(train_bar):\n            memory_manager.clear_memory()\n            \n            eeg_data = eeg_data.to(device)\n            images = images.to(device)\n            \n            optimizer.zero_grad()\n            generated_images = model(eeg_data)\n            loss = criterion(generated_images, images)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_bar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n            if batch_idx % 20 == 0:\n                memory_manager.clear_memory()\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for eeg_data, images, labels in val_loader:\n                memory_manager.clear_memory()\n                \n                eeg_data = eeg_data.to(device)\n                images = images.to(device)\n                \n                generated_images = model(eeg_data)\n                loss = criterion(generated_images, images)\n                val_loss += loss.item()\n        \n        avg_train_loss = train_loss / len(train_loader)\n        avg_val_loss = val_loss / len(val_loader)\n        \n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n        \n        current_lr = optimizer.param_groups[0]['lr']\n        scheduler.step()\n        \n        print(f'Epoch {epoch+1}/{num_epochs}:')\n        print(f'  训练损失: {avg_train_loss:.4f}, 验证损失: {avg_val_loss:.4f}, LR: {current_lr:.6f}')\n        \n        # 保存最佳模型\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            torch.save(model.state_dict(), 'best_eeg_to_image_model.pth')\n            print(f'  保存最佳模型')\n        \n        # 每2个epoch可视化一次\n        if (epoch + 1) % 2 == 0:\n            visualize_training_results(model, val_loader, epoch + 1)\n        \n        memory_manager.clear_memory()\n    \n    return model, history\n\ndef visualize_training_results(model, val_loader, epoch):\n    \"\"\"可视化训练结果\"\"\"\n    model.eval()\n    with torch.no_grad():\n        eeg_data, real_images, labels = next(iter(val_loader))\n        eeg_data = eeg_data[:4].to(device)\n        real_images = real_images[:4]\n        \n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 创建详细的可视化\n        fig, axes = plt.subplots(3, 4, figsize=(16, 9))\n        \n        categories = ['动物', '交通工具', '食物', '乐器', '家具', '电子产品']\n        \n        for i in range(4):\n            # EEG信号\n            eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n            axes[0, i].plot(eeg_sample.T, alpha=0.7, linewidth=1)\n            category_idx = labels[i].item() % len(categories)\n            axes[0, i].set_title(f'EEG: {categories[category_idx]}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实图像', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('生成图像', fontsize=10)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(f'training_results_epoch_{epoch}.png', dpi=200, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"训练结果已保存: training_results_epoch_{epoch}.png\")\n        memory_manager.clear_memory()\n\ndef test_model_and_visualize(model, test_loader, num_samples=6):\n    \"\"\"测试模型并生成完整可视化\"\"\"\n    model.eval()\n    \n    with torch.no_grad():\n        # 获取测试数据\n        eeg_data, real_images, labels = next(iter(test_loader))\n        eeg_data = eeg_data[:num_samples].to(device)\n        real_images = real_images[:num_samples]\n        \n        # 生成图像\n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 创建完整的测试可视化\n        fig, axes = plt.subplots(4, num_samples, figsize=(3 * num_samples, 12))\n        \n        if num_samples == 1:\n            axes = axes.reshape(4, 1)\n        \n        categories = ['动物', '交通工具', '食物', '乐器', '家具', '电子产品']\n        \n        for i in range(num_samples):\n            # EEG信号\n            eeg_sample = eeg_data[i].cpu().numpy()\n            \n            # 显示不同脑区的EEG\n            axes[0, i].plot(eeg_sample[:16, :100].T, alpha=0.6, linewidth=0.8, color='blue', label='前额叶')\n            axes[0, i].plot(eeg_sample[16:32, :100].T, alpha=0.6, linewidth=0.8, color='green', label='中央区')\n            axes[0, i].plot(eeg_sample[32:48, :100].T, alpha=0.6, linewidth=0.8, color='orange', label='顶叶')\n            axes[0, i].plot(eeg_sample[48:62, :100].T, alpha=0.6, linewidth=0.8, color='red', label='枕叶')\n            \n            category_idx = labels[i].item() % len(categories)\n            axes[0, i].set_title(f'EEG信号 - {categories[category_idx]}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            axes[0, i].set_xlabel('时间点')\n            axes[0, i].set_ylabel('幅值')\n            if i == 0:\n                axes[0, i].legend(fontsize=8)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实图像', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('生成图像', fontsize=10)\n            axes[2, i].axis('off')\n            \n            # 差异图\n            diff_img = np.abs(real_img - gen_img)\n            axes[3, i].imshow(diff_img, cmap='hot')\n            axes[3, i].set_title('差异图', fontsize=10)\n            axes[3, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig('final_test_results.png', dpi=300, bbox_inches='tight')\n        plt.show()\n        \n        # 计算评估指标\n        mse_loss = nn.MSELoss()(generated_images, real_images).item()\n        print(f\"测试MSE损失: {mse_loss:.4f}\")\n        \n        # 计算PSNR\n        mse = np.mean((real_images.numpy() - generated_images.numpy()) ** 2)\n        if mse == 0:\n            psnr = 100\n        else:\n            psnr = 20 * math.log10(1.0 / math.sqrt(mse))\n        print(f\"PSNR: {psnr:.2f} dB\")\n        \n        memory_manager.clear_memory()\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"启动修复版EEG-ImageNet图像重建系统...\")\n    \n    # 数据转换\n    transform = transforms.Compose([\n        transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n    ])\n    \n    # 数据路径配置\n    data_dir = \"/kaggle/input/eeg-imagenet\"\n    eeg_data_path = os.path.join(data_dir, \"EEG-ImageNet_1.pth\")\n    \n    # 创建修复的数据集\n    print(\"加载修复的EEG-ImageNet数据集...\")\n    train_dataset = FixedEEGImageNetDataset(\n        eeg_data_path=eeg_data_path,\n        transform=transform,\n        image_size=128,\n        mode='train',\n        subject_id='S1',\n        max_samples=100\n    )\n    \n    test_dataset = FixedEEGImageNetDataset(\n        eeg_data_path=eeg_data_path,\n        transform=transform,\n        image_size=128,\n        mode='test', \n        subject_id='S1',\n        max_samples=50\n    )\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=8, \n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=8,\n        shuffle=False, \n        num_workers=0,\n        pin_memory=True\n    )\n    \n    print(f\"训练集: {len(train_dataset)} 样本\")\n    print(f\"测试集: {len(test_dataset)} 样本\")\n    \n    # 创建模型\n    print(\"初始化模型...\")\n    model = EEGToImageModel(\n        eeg_channels=62,\n        time_points=400,\n        output_channels=3,\n        image_size=128\n    )\n    \n    # 检查是否有预训练模型\n    if os.path.exists('best_eeg_to_image_model.pth'):\n        print(\"加载预训练模型...\")\n        model.load_state_dict(torch.load('best_eeg_to_image_model.pth', map_location=device))\n        print(\"模型加载成功!\")\n    else:\n        print(\"从头开始训练模型...\")\n        # 训练模型\n        model, history = train_model(\n            model=model,\n            train_loader=train_loader,\n            val_loader=test_loader,\n            num_epochs=5,\n            learning_rate=1e-4\n        )\n        \n        # 绘制训练历史\n        plt.figure(figsize=(12, 5))\n        \n        plt.subplot(1, 2, 1)\n        plt.plot(history['train_loss'], label='训练损失')\n        plt.plot(history['val_loss'], label='验证损失')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.legend()\n        plt.title('训练和验证损失')\n        plt.grid(True)\n        \n        plt.subplot(1, 2, 2)\n        # 计算PSNR近似值\n        psnr_approx = [20 * math.log10(1.0 / max(loss, 1e-8)) for loss in history['val_loss']]\n        plt.plot(psnr_approx)\n        plt.xlabel('Epoch')\n        plt.ylabel('PSNR (approx)')\n        plt.title('图像质量指标')\n        plt.grid(True)\n        \n        plt.tight_layout()\n        plt.savefig('training_history.png', dpi=200, bbox_inches='tight')\n        plt.show()\n    \n    # 最终测试和可视化\n    print(\"开始最终测试和可视化...\")\n    test_model_and_visualize(model, test_loader, num_samples=6)\n    \n    # 清理内存\n    memory_manager.clear_memory()\n    \n    print(\"系统运行完成!\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T09:26:14.959388Z","iopub.execute_input":"2025-10-14T09:26:14.95976Z","iopub.status.idle":"2025-10-14T09:27:06.336534Z","shell.execute_reply.started":"2025-10-14T09:26:14.959736Z","shell.execute_reply":"2025-10-14T09:27:06.335864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet真实数据集适配版本 - 最终修复版\n解决类型比较错误和设备不匹配问题\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport gc\nimport time\nimport math\nfrom PIL import Image\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ==================== 设备配置和内存管理 ====================\n\ndef setup_device():\n    \"\"\"设置计算设备\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda:0\")\n        print(f\"使用GPU: {torch.cuda.get_device_name(0)}\")\n        torch.cuda.empty_cache()\n    else:\n        device = torch.device(\"cpu\")\n        print(\"使用CPU\")\n    return device\n\ndef set_seed(seed=42):\n    \"\"\"设置随机种子\"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\ndevice = setup_device()\nset_seed(42)\n\nclass MemoryManager:\n    \"\"\"内存管理器\"\"\"\n    def __init__(self):\n        self.image_cache = {}\n        self.cache_size = 50\n        \n    def clear_memory(self):\n        \"\"\"清理内存\"\"\"\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n            \n    def cache_image(self, key, image):\n        \"\"\"缓存图像\"\"\"\n        if len(self.image_cache) >= self.cache_size:\n            # 移除最旧的缓存\n            oldest_key = next(iter(self.image_cache))\n            del self.image_cache[oldest_key]\n        self.image_cache[key] = image\n        \n    def get_cached_image(self, key):\n        \"\"\"获取缓存的图像\"\"\"\n        return self.image_cache.get(key, None)\n\nmemory_manager = MemoryManager()\n\n# ==================== 修复的EEG数据加载器 ====================\n\nclass EEGImageNetDataset(Dataset):\n    \"\"\"EEG-ImageNet数据集加载器 - 解决类型比较错误\"\"\"\n    def __init__(self, eeg_data_path, transform=None, image_size=128, \n                 mode='train', subject_id='S1', max_samples=None):\n        super().__init__()\n        self.eeg_data_path = eeg_data_path\n        self.transform = transform\n        self.image_size = image_size\n        self.mode = mode\n        self.subject_id = subject_id\n        # 确保max_samples是整数\n        self.max_samples = int(max_samples) if max_samples is not None else None\n        \n        # 加载EEG数据\n        self.eeg_data, self.labels, self.image_info = self._load_eeg_data()\n        \n        print(f\"{mode}数据集 - 被试 {subject_id}: {len(self.eeg_data)} 个样本\")\n        \n    def _load_eeg_data(self):\n        \"\"\"加载EEG数据 - 解决类型比较错误\"\"\"\n        try:\n            if os.path.exists(self.eeg_data_path):\n                print(f\"加载EEG数据从: {self.eeg_data_path}\")\n                \n                # 加载数据\n                data = torch.load(self.eeg_data_path, map_location='cpu', weights_only=False)\n                print(f\"数据键: {list(data.keys())}\")\n                \n                # 解析数据结构\n                eeg_dataset = data['dataset']  # EEG数据列表\n                category_labels = data['labels']  # 类别标签\n                image_list = data['images']  # 图像信息\n                \n                print(f\"EEG数据样本数: {len(eeg_dataset)}\")\n                print(f\"类别数: {len(category_labels)}\")\n                print(f\"图像数: {len(image_list)}\")\n                \n                # 检查第一个EEG样本的结构\n                first_sample = eeg_dataset[0]\n                print(f\"第一个EEG样本类型: {type(first_sample)}\")\n                print(f\"EEG样本字典键: {first_sample.keys()}\")\n                \n                processed_eeg = []\n                processed_labels = []\n                processed_image_info = []\n                \n                # 确保max_samples是整数\n                max_samples_int = int(self.max_samples) if self.max_samples is not None else None\n                \n                for i, sample in enumerate(eeg_dataset):\n                    if max_samples_int is not None and i >= max_samples_int:\n                        break\n                        \n                    # 提取EEG数据\n                    if 'eeg_data' in sample:\n                        eeg_data = sample['eeg_data']\n                    elif 'eeg' in sample:\n                        eeg_data = sample['eeg']\n                    elif 'data' in sample:\n                        eeg_data = sample['data']\n                    else:\n                        # 如果找不到标准键，使用第一个数组类型的数据\n                        for key, value in sample.items():\n                            if hasattr(value, 'shape') and len(value.shape) >= 2:\n                                eeg_data = value\n                                break\n                        else:\n                            eeg_data = None\n                    \n                    # 提取标签 - 确保是整数\n                    if 'label' in sample:\n                        label = sample['label']\n                        # 确保标签是整数\n                        try:\n                            label = int(label)\n                        except (ValueError, TypeError):\n                            label = i % len(category_labels)\n                    else:\n                        label = i % len(category_labels)\n                    \n                    # 确保标签在有效范围内\n                    if label >= len(category_labels):\n                        label = label % len(category_labels)\n                    \n                    # 提取图像信息 - 修复字符串问题\n                    image_info = {}\n                    if 'image' in sample:\n                        image_data = sample['image']\n                        # 检查image_data的类型\n                        if isinstance(image_data, dict):\n                            image_info = image_data\n                        elif isinstance(image_data, str):\n                            # 如果是字符串，使用字符串作为image_id\n                            image_info = {\n                                'image_id': image_data,\n                                'category_id': label,\n                                'wordnet_id': category_labels[label] if label < len(category_labels) else f'n{str(label).zfill(8)}'\n                            }\n                        else:\n                            # 其他类型，创建默认信息\n                            image_info = {\n                                'image_id': i % len(image_list),\n                                'category_id': label,\n                                'wordnet_id': category_labels[label] if label < len(category_labels) else f'n{str(label).zfill(8)}'\n                            }\n                    else:\n                        image_info = {\n                            'image_id': i % len(image_list),\n                            'category_id': label,\n                            'wordnet_id': category_labels[label] if label < len(category_labels) else f'n{str(label).zfill(8)}'\n                        }\n                    \n                    # 确保image_info包含必要的字段\n                    if 'image_id' not in image_info:\n                        image_info['image_id'] = i % len(image_list)\n                    if 'category_id' not in image_info:\n                        image_info['category_id'] = label\n                    if 'wordnet_id' not in image_info:\n                        image_info['wordnet_id'] = category_labels[label] if label < len(category_labels) else f'n{str(label).zfill(8)}'\n                    \n                    # 处理EEG数据形状\n                    if eeg_data is not None:\n                        if isinstance(eeg_data, torch.Tensor):\n                            eeg_np = eeg_data.numpy()\n                        else:\n                            eeg_np = np.array(eeg_data)\n                        \n                        # 只在处理前几个样本时打印形状信息\n                        if i < 5:\n                            print(f\"样本 {i} 原始形状: {eeg_np.shape}\")\n                        \n                        # 接受 (62, 501) 形状，截取为 (62, 500)\n                        if eeg_np.shape == (62, 501):\n                            eeg_np = eeg_np[:, :500]  # 截取前500个时间点\n                            if i < 5:\n                                print(f\"样本 {i} 截取后形状: {eeg_np.shape}\")\n                            processed_eeg.append(eeg_np)\n                            processed_labels.append(label)\n                            processed_image_info.append(image_info)\n                        elif eeg_np.shape == (62, 500):\n                            processed_eeg.append(eeg_np)\n                            processed_labels.append(label)\n                            processed_image_info.append(image_info)\n                        else:\n                            if i < 5:\n                                print(f\"样本 {i} 形状不匹配: {eeg_np.shape}\")\n                            # 尝试调整形状\n                            if eeg_np.size == 62 * 500:\n                                eeg_np = eeg_np.reshape(62, 500)\n                                processed_eeg.append(eeg_np)\n                                processed_labels.append(label)\n                                processed_image_info.append(image_info)\n                            elif eeg_np.size == 62 * 501:\n                                eeg_np = eeg_np.reshape(62, 501)[:, :500]\n                                processed_eeg.append(eeg_np)\n                                processed_labels.append(label)\n                                processed_image_info.append(image_info)\n                            else:\n                                if i < 5:\n                                    print(f\"样本 {i} 无法调整形状，跳过\")\n                \n                if processed_eeg:\n                    eeg_tensor = torch.tensor(np.array(processed_eeg), dtype=torch.float32)\n                    print(f\"成功加载 {len(processed_eeg)} 个EEG样本，形状: {eeg_tensor.shape}\")\n                    return eeg_tensor, processed_labels, processed_image_info\n                else:\n                    raise ValueError(\"无法解析任何EEG样本\")\n                    \n            else:\n                raise FileNotFoundError(f\"EEG数据文件不存在: {self.eeg_data_path}\")\n                \n        except Exception as e:\n            print(f\"EEG数据加载失败: {e}\")\n            import traceback\n            traceback.print_exc()\n            print(\"使用模拟EEG数据\")\n            return self._create_simulated_eeg_data()\n    \n    def _create_simulated_eeg_data(self):\n        \"\"\"创建模拟EEG数据\"\"\"\n        n_samples = 100 if self.mode == 'train' else 50\n        n_categories = 80\n        \n        eeg_data = []\n        labels = []\n        image_info = []\n        \n        for i in range(n_samples):\n            # 生成神经科学合理的EEG信号\n            eeg = self._generate_realistic_eeg(i % n_categories, i)\n            eeg_data.append(eeg.numpy())\n            \n            labels.append(i % n_categories)\n            \n            image_info.append({\n                'image_id': i % 50,\n                'category_id': i % n_categories,\n                'wordnet_id': f'n{str((i % n_categories) + 1).zfill(8)}',\n                'image_index': i\n            })\n        \n        return torch.tensor(np.array(eeg_data), dtype=torch.float32), labels, image_info\n    \n    def _generate_realistic_eeg(self, category_idx, sample_idx):\n        \"\"\"生成神经科学合理的EEG信号\"\"\"\n        n_channels, n_times = 62, 500\n        eeg = torch.zeros(n_channels, n_times)\n        \n        time_points = torch.linspace(0, 2 * np.pi, n_times)\n        \n        # 基于类别生成独特的EEG模式\n        base_freq = 10 + (category_idx % 10)\n        \n        # 不同脑区有不同的活动模式\n        for ch in range(16):  # 前额叶\n            freq = base_freq + ch * 0.2\n            eeg[ch] += 0.3 * torch.sin(freq * time_points + sample_idx * 0.1)\n        \n        for ch in range(16, 32):  # 中央区\n            freq = base_freq * 1.5 + (ch-16) * 0.3\n            eeg[ch] += 0.4 * torch.cos(freq * time_points + sample_idx * 0.15)\n        \n        for ch in range(32, 48):  # 顶叶\n            freq = base_freq * 2 + (ch-32) * 0.4\n            eeg[ch] += 0.2 * torch.sin(freq * time_points * 0.7 + sample_idx * 0.2)\n        \n        for ch in range(48, 62):  # 枕叶（视觉区）\n            freq = base_freq * 2.5 + (ch-48) * 0.5\n            eeg[ch] += 0.5 * torch.sin(freq * time_points + sample_idx * 0.25)\n        \n        # 添加生理噪声\n        eeg += 0.1 * torch.randn(n_channels, n_times)\n        \n        return eeg\n    \n    def _load_imagenet_image(self, image_info):\n        \"\"\"加载ImageNet图像\"\"\"\n        try:\n            # 确保image_info是字典\n            if not isinstance(image_info, dict):\n                print(f\"警告: image_info不是字典，而是{type(image_info)}，创建默认图像信息\")\n                image_info = {\n                    'image_id': 0,\n                    'category_id': 0,\n                    'wordnet_id': 'n00000001'\n                }\n            \n            cache_key = f\"{image_info.get('category_id', 0)}_{image_info.get('image_id', 0)}\"\n            \n            # 检查缓存\n            cached_image = memory_manager.get_cached_image(cache_key)\n            if cached_image is not None:\n                return cached_image\n            \n            # 使用模拟图像（简化实现）\n            return self._generate_semantic_image(image_info)\n            \n        except Exception as e:\n            print(f\"图像加载错误: {e}\")\n            return self._generate_semantic_image(image_info)\n    \n    def _generate_semantic_image(self, image_info):\n        \"\"\"生成语义相关的模拟图像\"\"\"\n        # 确保image_info是字典且有必要的键\n        if not isinstance(image_info, dict):\n            image_info = {}\n        \n        img_size = self.image_size\n        category_id = image_info.get('category_id', 0)\n        image_id = image_info.get('image_id', 0)\n        \n        # 确保category_id和image_id是整数\n        try:\n            category_id = int(category_id)\n        except (ValueError, TypeError):\n            category_id = 0\n            \n        try:\n            image_id = int(image_id)\n        except (ValueError, TypeError):\n            image_id = 0\n        \n        # 基于类别ID生成有意义的图像\n        img = Image.new('RGB', (img_size, img_size), color='white')\n        \n        from PIL import ImageDraw\n        draw = ImageDraw.Draw(img)\n        \n        center_x, center_y = img_size // 2, img_size // 2\n        \n        # 基于类别ID决定图像类型\n        category_type = category_id % 6\n        \n        if category_type == 0:  # 圆形物体\n            color = (100, 150, 200)\n            radius = img_size // 4 + (image_id % 10)\n            draw.ellipse([center_x-radius, center_y-radius, center_x+radius, center_y+radius], \n                        fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 1:  # 方形物体\n            color = (200, 100, 100)\n            size = img_size // 3 + (image_id % 15)\n            draw.rectangle([center_x-size//2, center_y-size//2, center_x+size//2, center_y+size//2], \n                          fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 2:  # 三角形物体\n            color = (150, 200, 100)\n            size = img_size // 4 + (image_id % 12)\n            points = [\n                (center_x, center_y - size),\n                (center_x - size, center_y + size),\n                (center_x + size, center_y + size)\n            ]\n            draw.polygon(points, fill=color, outline=(0,0,0), width=2)\n            \n        elif category_type == 3:  # 线条图案\n            color = (200, 150, 100)\n            for i in range(0, img_size, img_size // 8):\n                draw.line([i, 0, i, img_size], fill=color, width=2)\n                draw.line([0, i, img_size, i], fill=color, width=2)\n                \n        elif category_type == 4:  # 点状图案\n            color = (150, 100, 200)\n            for i in range(0, img_size, img_size // 10):\n                for j in range(0, img_size, img_size // 10):\n                    if (i + j) % 20 == 0:\n                        draw.ellipse([i-3, j-3, i+3, j+3], fill=color)\n                        \n        else:  # 混合图案\n            color = (100, 200, 200)\n            # 圆形\n            draw.ellipse([center_x-30, center_y-20, center_x+30, center_y+20], fill=color)\n            # 矩形\n            draw.rectangle([center_x-15, center_y-10, center_x+15, center_y+10], \n                          fill=(255,255,255), outline=(0,0,0), width=1)\n        \n        return img\n    \n    def _preprocess_eeg(self, eeg):\n        \"\"\"EEG预处理\"\"\"\n        if isinstance(eeg, torch.Tensor):\n            eeg = eeg.float()\n        else:\n            eeg = torch.tensor(eeg, dtype=torch.float32)\n        \n        # 确保正确的形状 (channels, time)\n        if len(eeg.shape) == 3:\n            # 如果是(batch, channels, time)，取第一个样本\n            eeg = eeg[0]\n        \n        # 标准化\n        eeg_mean = eeg.mean(dim=1, keepdim=True)\n        eeg_std = eeg.std(dim=1, keepdim=True)\n        eeg_std = torch.clamp(eeg_std, min=1e-8)\n        eeg = (eeg - eeg_mean) / eeg_std\n        \n        # 提取40ms-440ms段 (如果时间维度足够)\n        if eeg.shape[1] >= 440:\n            start_idx, end_idx = 40, 440\n            eeg = eeg[:, start_idx:end_idx]\n        elif eeg.shape[1] > 100:\n            # 如果时间维度不够，取中间部分\n            start_idx = eeg.shape[1] // 4\n            end_idx = start_idx + 400\n            if end_idx > eeg.shape[1]:\n                end_idx = eeg.shape[1]\n            eeg = eeg[:, start_idx:end_idx]\n        \n        return eeg\n    \n    def __len__(self):\n        if self.max_samples:\n            return min(self.max_samples, len(self.eeg_data))\n        return len(self.eeg_data)\n    \n    def __getitem__(self, idx):\n        if idx >= len(self.eeg_data):\n            idx = len(self.eeg_data) - 1\n        \n        # 获取EEG数据\n        eeg = self.eeg_data[idx]\n        eeg = self._preprocess_eeg(eeg)\n        \n        # 获取图像信息\n        image_info = self.image_info[idx]\n        \n        # 加载图像\n        image = self._load_imagenet_image(image_info)\n        \n        # 转换为张量\n        image_tensor = transforms.ToTensor()(image)\n        \n        # 获取标签\n        label = self.labels[idx]\n        \n        # 应用变换\n        if self.transform:\n            image_tensor = self.transform(image_tensor)\n        \n        return eeg, image_tensor, torch.tensor(label)\n\n# ==================== 高效的EEG到图像模型 ====================\n\nclass EfficientEEGEncoder(nn.Module):\n    \"\"\"高效的EEG编码器\"\"\"\n    def __init__(self, input_channels=62, time_points=400, hidden_dim=256):\n        super().__init__()\n        \n        self.conv_layers = nn.Sequential(\n            nn.Conv1d(input_channels, 64, kernel_size=15, padding=7),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(64, 128, kernel_size=11, padding=5),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            nn.Conv1d(128, 256, kernel_size=7, padding=3),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n        )\n        \n        self.attention = nn.MultiheadAttention(256, num_heads=4, dropout=0.1)\n        self.global_pool = nn.AdaptiveAvgPool1d(1)\n        self.fc = nn.Linear(256, 512)\n        \n    def forward(self, x):\n        # 输入: (batch, channels, time)\n        x = self.conv_layers(x)  # (batch, 256, time)\n        \n        # 应用注意力\n        x_attn = x.transpose(1, 2)  # (batch, time, 256)\n        x_attn = x_attn.transpose(0, 1)  # (time, batch, 256)\n        attn_out, _ = self.attention(x_attn, x_attn, x_attn)\n        attn_out = attn_out.transpose(0, 1).transpose(1, 2)  # (batch, 256, time)\n        \n        x = x + 0.1 * attn_out\n        \n        # 全局池化\n        x = self.global_pool(x).squeeze(-1)  # (batch, 256)\n        x = self.fc(x)  # (batch, 512)\n        \n        return x\n\nclass UNetDecoder(nn.Module):\n    \"\"\"UNet风格解码器\"\"\"\n    def __init__(self, input_dim=512, output_channels=3, image_size=128):\n        super().__init__()\n        \n        self.init_proj = nn.Linear(input_dim, 512 * 4 * 4)\n        \n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(512, 256, 4, 2, 1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(256, 128, 4, 2, 1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(128, 64, 4, 2, 1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(64, 32, 4, 2, 1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            nn.ConvTranspose2d(32, 16, 4, 2, 1),\n            nn.BatchNorm2d(16),\n            nn.ReLU(),\n            \n            nn.Conv2d(16, output_channels, 3, 1, 1),\n            nn.Tanh()\n        )\n        \n    def forward(self, x):\n        x = self.init_proj(x)\n        x = x.view(-1, 512, 4, 4)\n        x = self.decoder(x)\n        return x\n\nclass EEGToImageModel(nn.Module):\n    \"\"\"EEG到图像转换模型\"\"\"\n    def __init__(self, eeg_channels=62, time_points=400, output_channels=3, image_size=128):\n        super().__init__()\n        self.encoder = EfficientEEGEncoder(eeg_channels, time_points)\n        self.decoder = UNetDecoder(output_channels=output_channels, image_size=image_size)\n        \n    def forward(self, eeg_data):\n        features = self.encoder(eeg_data)\n        images = self.decoder(features)\n        return images\n\n# ==================== 训练和可视化函数 ====================\n\ndef train_model(model, train_loader, val_loader, num_epochs=5, learning_rate=1e-4):\n    \"\"\"训练模型\"\"\"\n    model.to(device)\n    criterion = nn.MSELoss()\n    optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-5)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    history = {'train_loss': [], 'val_loss': []}\n    best_val_loss = float('inf')\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0\n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        \n        for batch_idx, (eeg_data, images, labels) in enumerate(train_bar):\n            memory_manager.clear_memory()\n            \n            eeg_data = eeg_data.to(device)\n            images = images.to(device)\n            \n            optimizer.zero_grad()\n            generated_images = model(eeg_data)\n            loss = criterion(generated_images, images)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_bar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n            if batch_idx % 20 == 0:\n                memory_manager.clear_memory()\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for eeg_data, images, labels in val_loader:\n                memory_manager.clear_memory()\n                \n                eeg_data = eeg_data.to(device)\n                images = images.to(device)\n                \n                generated_images = model(eeg_data)\n                loss = criterion(generated_images, images)\n                val_loss += loss.item()\n        \n        avg_train_loss = train_loss / len(train_loader)\n        avg_val_loss = val_loss / len(val_loader)\n        \n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n        \n        current_lr = optimizer.param_groups[0]['lr']\n        scheduler.step()\n        \n        print(f'Epoch {epoch+1}/{num_epochs}:')\n        print(f'  训练损失: {avg_train_loss:.4f}, 验证损失: {avg_val_loss:.4f}, LR: {current_lr:.6f}')\n        \n        # 保存最佳模型\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            torch.save(model.state_dict(), 'best_eeg_to_image_model.pth')\n            print(f'  保存最佳模型')\n        \n        # 每2个epoch可视化一次\n        if (epoch + 1) % 2 == 0:\n            visualize_training_results(model, val_loader, epoch + 1)\n        \n        memory_manager.clear_memory()\n    \n    return model, history\n\ndef visualize_training_results(model, val_loader, epoch):\n    \"\"\"可视化训练结果\"\"\"\n    model.eval()\n    with torch.no_grad():\n        eeg_data, real_images, labels = next(iter(val_loader))\n        eeg_data = eeg_data[:4].to(device)\n        real_images = real_images[:4]\n        \n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 创建详细的可视化\n        fig, axes = plt.subplots(3, 4, figsize=(16, 9))\n        \n        categories = ['动物', '交通工具', '食物', '乐器', '家具', '电子产品']\n        \n        for i in range(4):\n            # EEG信号\n            eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n            axes[0, i].plot(eeg_sample.T, alpha=0.7, linewidth=1)\n            category_idx = labels[i].item() % len(categories)\n            axes[0, i].set_title(f'EEG: {categories[category_idx]}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实图像', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('生成图像', fontsize=10)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(f'training_results_epoch_{epoch}.png', dpi=200, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"训练结果已保存: training_results_epoch_{epoch}.png\")\n        memory_manager.clear_memory()\n\ndef test_model_and_visualize(model, test_loader, num_samples=6):\n    \"\"\"测试模型并生成完整可视化\"\"\"\n    model.eval()\n    \n    with torch.no_grad():\n        # 获取测试数据\n        eeg_data, real_images, labels = next(iter(test_loader))\n        eeg_data = eeg_data[:num_samples].to(device)\n        real_images = real_images[:num_samples]\n        \n        # 生成图像\n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 创建完整的测试可视化\n        fig, axes = plt.subplots(4, num_samples, figsize=(3 * num_samples, 12))\n        \n        if num_samples == 1:\n            axes = axes.reshape(4, 1)\n        \n        categories = ['动物', '交通工具', '食物', '乐器', '家具', '电子产品']\n        \n        for i in range(num_samples):\n            # EEG信号\n            eeg_sample = eeg_data[i].cpu().numpy()\n            \n            # 显示不同脑区的EEG\n            axes[0, i].plot(eeg_sample[:16, :100].T, alpha=0.6, linewidth=0.8, color='blue', label='前额叶')\n            axes[0, i].plot(eeg_sample[16:32, :100].T, alpha=0.6, linewidth=0.8, color='green', label='中央区')\n            axes[0, i].plot(eeg_sample[32:48, :100].T, alpha=0.6, linewidth=0.8, color='orange', label='顶叶')\n            axes[0, i].plot(eeg_sample[48:62, :100].T, alpha=0.6, linewidth=0.8, color='red', label='枕叶')\n            \n            category_idx = labels[i].item() % len(categories)\n            axes[0, i].set_title(f'EEG信号 - {categories[category_idx]}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            axes[0, i].set_xlabel('时间点')\n            axes[0, i].set_ylabel('幅值')\n            if i == 0:\n                axes[0, i].legend(fontsize=8)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('真实图像', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('生成图像', fontsize=10)\n            axes[2, i].axis('off')\n            \n            # 差异图\n            diff_img = np.abs(real_img - gen_img)\n            axes[3, i].imshow(diff_img, cmap='hot')\n            axes[3, i].set_title('差异图', fontsize=10)\n            axes[3, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig('final_test_results.png', dpi=300, bbox_inches='tight')\n        plt.show()\n        \n        # 计算评估指标\n        mse_loss = nn.MSELoss()(generated_images, real_images).item()\n        print(f\"测试MSE损失: {mse_loss:.4f}\")\n        \n        # 计算PSNR\n        mse = np.mean((real_images.numpy() - generated_images.numpy()) ** 2)\n        if mse == 0:\n            psnr = 100\n        else:\n            psnr = 20 * math.log10(1.0 / math.sqrt(mse))\n        print(f\"PSNR: {psnr:.2f} dB\")\n        \n        memory_manager.clear_memory()\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"启动修复版EEG-ImageNet图像重建系统...\")\n    \n    # 数据转换\n    transform = transforms.Compose([\n        transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n    ])\n    \n    # 数据路径配置\n    data_dir = \"/kaggle/input/eeg-imagenet\"\n    eeg_data_path = os.path.join(data_dir, \"EEG-ImageNet_1.pth\")\n    \n    # 创建修复的数据集\n    print(\"加载修复的EEG-ImageNet数据集...\")\n    train_dataset = EEGImageNetDataset(\n        eeg_data_path=eeg_data_path,\n        transform=transform,\n        image_size=128,\n        mode='train',\n        subject_id='S1',\n        max_samples=100\n    )\n    \n    test_dataset = EEGImageNetDataset(\n        eeg_data_path=eeg_data_path,\n        transform=transform,\n        image_size=128,\n        mode='test', \n        subject_id='S1',\n        max_samples=50\n    )\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=8, \n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=8,\n        shuffle=False, \n        num_workers=0,\n        pin_memory=True\n    )\n    \n    print(f\"训练集: {len(train_dataset)} 样本\")\n    print(f\"测试集: {len(test_dataset)} 样本\")\n    \n    # 创建模型\n    print(\"初始化模型...\")\n    model = EEGToImageModel(\n        eeg_channels=62,\n        time_points=400,\n        output_channels=3,\n        image_size=128\n    )\n    \n    # 检查是否有预训练模型\n    if os.path.exists('best_eeg_to_image_model.pth'):\n        print(\"加载预训练模型...\")\n        # 加载模型时指定map_location\n        model.load_state_dict(torch.load('best_eeg_to_image_model.pth', map_location=device))\n        # 确保模型在正确的设备上\n        model.to(device)\n        print(\"模型加载成功!\")\n    else:\n        print(\"从头开始训练模型...\")\n        # 训练模型\n        model, history = train_model(\n            model=model,\n            train_loader=train_loader,\n            val_loader=test_loader,\n            num_epochs=5,\n            learning_rate=1e-4\n        )\n        \n        # 绘制训练历史\n        plt.figure(figsize=(12, 5))\n        \n        plt.subplot(1, 2, 1)\n        plt.plot(history['train_loss'], label='训练损失')\n        plt.plot(history['val_loss'], label='验证损失')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.legend()\n        plt.title('训练和验证损失')\n        plt.grid(True)\n        \n        plt.subplot(1, 2, 2)\n        # 计算PSNR近似值\n        psnr_approx = [20 * math.log10(1.0 / max(loss, 1e-8)) for loss in history['val_loss']]\n        plt.plot(psnr_approx)\n        plt.xlabel('Epoch')\n        plt.ylabel('PSNR (approx)')\n        plt.title('图像质量指标')\n        plt.grid(True)\n        \n        plt.tight_layout()\n        plt.savefig('training_history.png', dpi=200, bbox_inches='tight')\n        plt.show()\n    \n    # 最终测试和可视化\n    print(\"开始最终测试和可视化...\")\n    test_model_and_visualize(model, test_loader, num_samples=6)\n    \n    # 清理内存\n    memory_manager.clear_memory()\n    \n    print(\"系统运行完成!\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:09:59.113927Z","iopub.execute_input":"2025-10-14T10:09:59.114632Z","iopub.status.idle":"2025-10-14T10:10:19.83816Z","shell.execute_reply.started":"2025-10-14T10:09:59.114604Z","shell.execute_reply":"2025-10-14T10:10:19.837528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet图像重建系统 - 内存优化版\n修复尺寸不匹配问题\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport gc\nimport time\nimport math\nfrom PIL import Image\nimport warnings\nimport random\nfrom datasets import load_dataset\nwarnings.filterwarnings('ignore')\n\n# ==================== 设备配置和内存管理 ====================\n\ndef setup_device():\n    \"\"\"设置计算设备\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda:0\")\n        print(f\"使用GPU: {torch.cuda.get_device_name(0)}\")\n        torch.cuda.empty_cache()\n    else:\n        device = torch.device(\"cpu\")\n        print(\"使用CPU\")\n    return device\n\ndevice = setup_device()\n\nclass MemoryManager:\n    \"\"\"内存管理器\"\"\"\n    def __init__(self):\n        self.image_cache = {}\n        self.cache_size = 50\n        self.eeg_cache = {}\n        self.eeg_cache_size = 100\n        \n    def clear_memory(self):\n        \"\"\"清理内存\"\"\"\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n            \n    def cache_image(self, key, image):\n        \"\"\"缓存图像\"\"\"\n        if len(self.image_cache) >= self.cache_size:\n            oldest_key = next(iter(self.image_cache))\n            del self.image_cache[oldest_key]\n        self.image_cache[key] = image\n        \n    def get_cached_image(self, key):\n        \"\"\"获取缓存的图像\"\"\"\n        return self.image_cache.get(key, None)\n    \n    def cache_eeg(self, key, eeg_data):\n        \"\"\"缓存EEG数据\"\"\"\n        if len(self.eeg_cache) >= self.eeg_cache_size:\n            oldest_key = next(iter(self.eeg_cache))\n            del self.eeg_cache[oldest_key]\n        self.eeg_cache[key] = eeg_data\n        \n    def get_cached_eeg(self, key):\n        \"\"\"获取缓存的EEG数据\"\"\"\n        return self.eeg_cache.get(key, None)\n\nmemory_manager = MemoryManager()\n\n# ==================== 配置参数 ====================\n\nclass Config:\n    \"\"\"配置参数类\"\"\"\n    def __init__(self):\n        # 训练参数\n        self.num_epochs = 30\n        self.learning_rate = 2e-4\n        self.batch_size = 16\n        self.max_train_samples = 1000\n        self.max_test_samples = 200\n        \n        # 模型参数\n        self.image_size = 128\n        self.eeg_channels = 62\n        self.time_points = 400\n        \n        # 路径参数\n        self.data_dir = \"/kaggle/input/eeg-imagenet\"\n        self.eeg_data_path = os.path.join(self.data_dir, \"EEG-ImageNet_1.pth\")\n        self.model_save_path = \"eeg_to_image_model_fixed.pth\"\n        self.checkpoint_path = \"training_checkpoint_fixed.pth\"\n        \n        # 训练选项\n        self.use_pretrained = False\n        self.continue_training = False\n        self.save_checkpoints = True\n        \n        # 损失权重\n        self.mse_weight = 1.0\n        self.perceptual_weight = 0.05\n\nconfig = Config()\n\n# ==================== 轻量级感知损失 ====================\n\nclass LightPerceptualLoss(nn.Module):\n    \"\"\"轻量级感知损失 - 只使用VGG16的前几层\"\"\"\n    def __init__(self):\n        super().__init__()\n        vgg = models.vgg16(pretrained=True).features[:10]  # 只使用前10层\n        self.feature_extractor = nn.Sequential(*list(vgg.children())[:10])\n        \n        for param in self.parameters():\n            param.requires_grad = False\n            \n        self.criterion = nn.L1Loss()\n        \n    def forward(self, generated, target):\n        # 确保输入尺寸正确\n        if generated.shape[2:] != target.shape[2:]:\n            generated = nn.functional.interpolate(generated, size=target.shape[2:], mode='bilinear', align_corners=False)\n        \n        # 归一化到VGG的输入范围\n        mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1).to(generated.device)\n        std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1).to(generated.device)\n        \n        generated = (generated + 1) / 2  # [-1,1] -> [0,1]\n        target = (target + 1) / 2\n        \n        generated = (generated - mean) / std\n        target = (target - mean) / std\n        \n        # 只提取浅层特征\n        gen_features = self.feature_extractor(generated)\n        target_features = self.feature_extractor(target)\n        \n        return self.criterion(gen_features, target_features)\n\n# ==================== 修复尺寸问题的数据集加载器 ====================\n\nclass FixedEEGImageDataset(Dataset):\n    \"\"\"修复尺寸问题的EEG-ImageNet数据集加载器\"\"\"\n    def __init__(self, eeg_data_path, image_size=128, mode='train', max_samples=None):\n        super().__init__()\n        self.eeg_data_path = eeg_data_path\n        self.image_size = image_size\n        self.mode = mode\n        self.max_samples = max_samples\n        \n        # 日常物品类别\n        self.daily_objects = [\n            'pizza', 'hamburger', 'hotdog', 'cake', 'ice cream', 'donut', 'cookie', 'bread', \n            'cheese', 'pasta', 'apple', 'banana', 'orange', 'grape', 'strawberry', 'watermelon',\n            'pineapple', 'peach', 'pear', 'cherry', 'carrot', 'tomato', 'onion', 'potato',\n            'corn', 'bell pepper', 'cucumber', 'lettuce', 'broccoli', 'cauliflower'\n        ]\n        \n        # 加载EEG数据\n        self.eeg_data, self.labels = self._load_eeg_data()\n        \n        # 初始化ImageNet21k数据集（流式加载）\n        self.imagenet_dataset = None\n        self._init_imagenet_dataset()\n        \n        print(f\"{mode}数据集: {len(self.eeg_data)} 个EEG样本\")\n        \n    def _load_eeg_data(self):\n        \"\"\"加载EEG数据 - 内存高效版本\"\"\"\n        try:\n            if os.path.exists(self.eeg_data_path):\n                print(f\"加载EEG数据从: {self.eeg_data_path}\")\n                \n                data = torch.load(self.eeg_data_path, map_location='cpu', weights_only=False)\n                eeg_dataset = data['dataset']\n                \n                processed_eeg = []\n                processed_labels = []\n                \n                max_samples = self.max_samples or len(eeg_dataset)\n                selected_indices = random.sample(range(len(eeg_dataset)), min(max_samples, len(eeg_dataset)))\n                \n                for idx in selected_indices:\n                    sample = eeg_dataset[idx]\n                    if 'eeg_data' not in sample:\n                        continue\n                    \n                    eeg_data = sample['eeg_data']\n                    if eeg_data is not None:\n                        if isinstance(eeg_data, torch.Tensor):\n                            eeg_np = eeg_data.numpy()\n                        else:\n                            eeg_np = np.array(eeg_data)\n                        \n                        # 处理EEG形状\n                        if eeg_np.shape == (62, 501):\n                            eeg_np = eeg_np[:, :500]\n                        elif eeg_np.shape != (62, 500):\n                            continue\n                            \n                        # 随机分配日常物品类别\n                        random_label = random.randint(0, len(self.daily_objects) - 1)\n                        \n                        cache_key = f\"{idx}_{random_label}\"\n                        memory_manager.cache_eeg(cache_key, eeg_np)\n                        \n                        processed_eeg.append(cache_key)\n                        processed_labels.append(random_label)\n                \n                print(f\"成功加载 {len(processed_eeg)} 个EEG样本\")\n                return processed_eeg, processed_labels\n                \n            else:\n                raise FileNotFoundError(f\"EEG数据文件不存在: {self.eeg_data_path}\")\n                \n        except Exception as e:\n            print(f\"EEG数据加载失败: {e}\")\n            return self._create_simulated_eeg_data()\n    \n    def _create_simulated_eeg_data(self):\n        \"\"\"创建模拟EEG数据\"\"\"\n        n_samples = config.max_train_samples if self.mode == 'train' else config.max_test_samples\n        \n        eeg_keys = []\n        labels = []\n        \n        for i in range(n_samples):\n            random_label = random.randint(0, len(self.daily_objects) - 1)\n            eeg = self._generate_real_eeg_like_data(random_label, i)\n            \n            cache_key = f\"sim_{i}_{random_label}\"\n            memory_manager.cache_eeg(cache_key, eeg.numpy())\n            \n            eeg_keys.append(cache_key)\n            labels.append(random_label)\n        \n        return eeg_keys, labels\n    \n    def _generate_real_eeg_like_data(self, category_idx, sample_idx):\n        \"\"\"生成基于真实EEG统计特征的数据\"\"\"\n        n_channels, n_times = 62, 500\n        eeg = torch.zeros(n_channels, n_times)\n        \n        # 基于真实EEG的统计特征\n        time_points = torch.linspace(0, 2 * np.pi, n_times)\n        \n        for ch in range(n_channels):\n            # 不同脑区有不同的频带特性\n            if ch < 16:  # 前额叶\n                theta_power = 0.4\n                alpha_power = 0.3\n            elif ch < 32:  # 中央区\n                theta_power = 0.3\n                alpha_power = 0.5\n            elif ch < 48:  # 顶叶\n                theta_power = 0.25\n                alpha_power = 0.4\n            else:  # 枕叶\n                theta_power = 0.2\n                alpha_power = 0.6\n            \n            # 生成多频带信号\n            delta_signal = 0.1 * torch.sin(3 * time_points + ch * 0.1)\n            theta_signal = theta_power * 0.3 * torch.sin(6 * time_points + ch * 0.2)\n            alpha_signal = alpha_power * 0.4 * torch.sin(10 * time_points + ch * 0.3)\n            beta_signal = 0.15 * 0.2 * torch.sin(20 * time_points + ch * 0.4)\n            \n            eeg[ch] = delta_signal + theta_signal + alpha_signal + beta_signal\n        \n        # 添加噪声\n        eeg += 0.05 * torch.randn(n_channels, n_times)\n        \n        return eeg\n    \n    def _init_imagenet_dataset(self):\n        \"\"\"初始化ImageNet21k数据集（流式加载）\"\"\"\n        try:\n            print(\"初始化ImageNet21k数据集...\")\n            # 使用Hugging Face数据集，流式加载\n            self.imagenet_dataset = load_dataset(\"gmongaras/Imagenet21K\", split=\"train\", streaming=True)\n            print(\"ImageNet21k数据集初始化成功\")\n        except Exception as e:\n            print(f\"ImageNet21k初始化失败: {e}\")\n            print(\"将使用回退图像生成方法\")\n            self.imagenet_dataset = None\n    \n    def _get_imagenet_image(self, object_name, image_id):\n        \"\"\"从ImageNet21k获取图像\"\"\"\n        try:\n            if self.imagenet_dataset is None:\n                return self._generate_fallback_image(object_name, image_id)\n            \n            # 将对象名称转换为可能的ImageNet类别\n            search_terms = self._get_search_terms(object_name)\n            \n            # 在数据集中搜索匹配的图像\n            for i, sample in enumerate(self.imagenet_dataset):\n                if i > 1000:  # 限制搜索数量\n                    break\n                    \n                label = sample.get('label', '')\n                text = sample.get('text', '')\n                \n                # 检查是否匹配搜索词\n                if any(term in str(label).lower() or term in str(text).lower() for term in search_terms):\n                    image = sample['image']\n                    if image.mode != 'RGB':\n                        image = image.convert('RGB')\n                    \n                    # 调整图像大小\n                    image = image.resize((self.image_size, self.image_size))\n                    \n                    # 缓存图像\n                    cache_key = f\"{object_name}_{image_id}\"\n                    memory_manager.cache_image(cache_key, image)\n                    \n                    return image\n            \n            # 如果没有找到匹配的图像，使用回退方法\n            return self._generate_fallback_image(object_name, image_id)\n            \n        except Exception as e:\n            print(f\"获取ImageNet图像失败: {e}\")\n            return self._generate_fallback_image(object_name, image_id)\n    \n    def _get_search_terms(self, object_name):\n        \"\"\"获取搜索词\"\"\"\n        search_map = {\n            'pizza': ['pizza', 'italian food'],\n            'hamburger': ['hamburger', 'burger', 'fast food'],\n            'cake': ['cake', 'dessert', 'birthday cake'],\n            'apple': ['apple', 'fruit'],\n            'banana': ['banana', 'fruit'],\n            'orange': ['orange', 'fruit', 'citrus'],\n            'carrot': ['carrot', 'vegetable'],\n            'tomato': ['tomato', 'vegetable'],\n            'bread': ['bread', 'bakery'],\n            'coffee': ['coffee', 'beverage']\n        }\n        return search_map.get(object_name, [object_name])\n    \n    def _generate_fallback_image(self, object_name, image_id):\n        \"\"\"生成回退图像\"\"\"\n        img_size = self.image_size\n        img = Image.new('RGB', (img_size, img_size), color=(240, 240, 240))\n        \n        # 基于对象名称生成简单的彩色图像\n        color_hash = hash(object_name) % 16777215\n        r = (color_hash >> 16) & 255\n        g = (color_hash >> 8) & 255\n        b = color_hash & 255\n        \n        # 创建简单的形状\n        from PIL import ImageDraw\n        draw = ImageDraw.Draw(img)\n        center_x, center_y = img_size // 2, img_size // 2\n        \n        # 根据对象类型选择形状\n        if any(food in object_name for food in ['pizza', 'cake', 'burger']):\n            draw.ellipse([center_x-30, center_y-30, center_x+30, center_y+30], \n                        fill=(r, g, b), outline=(0,0,0), width=2)\n        elif any(fruit in object_name for fruit in ['apple', 'orange', 'banana']):\n            draw.ellipse([center_x-25, center_y-25, center_x+25, center_y+25], \n                        fill=(r, g, b), outline=(0,0,0), width=2)\n        else:\n            draw.rectangle([center_x-25, center_y-20, center_x+25, center_y+20], \n                          fill=(r, g, b), outline=(0,0,0), width=2)\n        \n        # 添加文本标签\n        try:\n            font_size = max(12, img_size // 10)\n            from PIL import ImageFont\n            try:\n                font = ImageFont.truetype(\"Arial\", font_size)\n            except:\n                font = ImageFont.load_default()\n            \n            text = object_name[:10]  # 限制文本长度\n            bbox = draw.textbbox((0, 0), text, font=font)\n            text_width = bbox[2] - bbox[0]\n            text_height = bbox[3] - bbox[1]\n            \n            text_x = center_x - text_width // 2\n            text_y = img_size - text_height - 10\n            \n            draw.rectangle([text_x-5, text_y-2, text_x+text_width+5, text_y+text_height+2], \n                          fill=(255, 255, 255, 180))\n            draw.text((text_x, text_y), text, fill=(0, 0, 0), font=font)\n        except:\n            pass\n        \n        return img\n    \n    def _preprocess_eeg(self, eeg_data):\n        \"\"\"EEG预处理\"\"\"\n        if isinstance(eeg_data, torch.Tensor):\n            eeg = eeg_data.float()\n        else:\n            eeg = torch.tensor(eeg_data, dtype=torch.float32)\n        \n        # 标准化\n        eeg_mean = eeg.mean(dim=1, keepdim=True)\n        eeg_std = eeg.std(dim=1, keepdim=True)\n        eeg_std = torch.clamp(eeg_std, min=1e-8)\n        eeg = (eeg - eeg_mean) / eeg_std\n        \n        # 提取时间窗口\n        if eeg.shape[1] >= 440:\n            eeg = eeg[:, 40:440]\n        elif eeg.shape[1] > 100:\n            start_idx = eeg.shape[1] // 4\n            end_idx = min(start_idx + 400, eeg.shape[1])\n            eeg = eeg[:, start_idx:end_idx]\n        \n        return eeg\n    \n    def __len__(self):\n        return len(self.eeg_data)\n    \n    def __getitem__(self, idx):\n        # 获取EEG数据\n        eeg_key = self.eeg_data[idx]\n        eeg_data = memory_manager.get_cached_eeg(eeg_key)\n        \n        if eeg_data is None:\n            # 重新生成EEG数据（不应该发生）\n            label = self.labels[idx]\n            eeg_data = self._generate_real_eeg_like_data(label, idx).numpy()\n            memory_manager.cache_eeg(eeg_key, eeg_data)\n        \n        eeg = self._preprocess_eeg(eeg_data)\n        \n        # 获取标签和对象名称\n        label = self.labels[idx]\n        object_name = self.daily_objects[label]\n        \n        # 获取图像\n        image_cache_key = f\"{object_name}_{idx}\"\n        image = memory_manager.get_cached_image(image_cache_key)\n        \n        if image is None:\n            image = self._get_imagenet_image(object_name, idx)\n            memory_manager.cache_image(image_cache_key, image)\n        \n        # 转换为张量\n        transform = transforms.Compose([\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n        ])\n        \n        image_tensor = transform(image)\n        \n        # 清理内存\n        if idx % 100 == 0:\n            memory_manager.clear_memory()\n        \n        return eeg, image_tensor, torch.tensor(label)\n\n# ==================== 修复尺寸问题的模型 ====================\n\nclass FixedEEGEncoder(nn.Module):\n    \"\"\"修复尺寸问题的EEG编码器\"\"\"\n    def __init__(self, input_channels=62, time_points=400, hidden_dim=128):\n        super().__init__()\n        \n        self.conv_layers = nn.Sequential(\n            nn.Conv1d(input_channels, 32, kernel_size=15, padding=7),\n            nn.BatchNorm1d(32),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            \n            nn.Conv1d(32, 64, kernel_size=11, padding=5),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            \n            nn.Conv1d(64, 128, kernel_size=7, padding=3),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n        )\n        \n        self.global_pool = nn.AdaptiveAvgPool1d(1)\n        self.fc = nn.Linear(128, 512)  # 增加输出维度以匹配解码器输入\n        \n    def forward(self, x):\n        x = self.conv_layers(x)\n        x = self.global_pool(x).squeeze(-1)\n        x = self.fc(x)\n        return x\n\nclass FixedUNetDecoder(nn.Module):\n    \"\"\"修复尺寸问题的UNet解码器 - 确保输出128x128\"\"\"\n    def __init__(self, input_dim=512, output_channels=3, image_size=128):\n        super().__init__()\n        \n        # 计算初始特征图大小\n        # 我们需要从4x4上采样到128x128，需要5次2倍上采样\n        # 4x4 -> 8x8 -> 16x16 -> 32x32 -> 64x64 -> 128x128\n        self.init_proj = nn.Linear(input_dim, 512 * 4 * 4)\n        \n        self.decoder = nn.Sequential(\n            # 4x4 -> 8x8\n            nn.ConvTranspose2d(512, 256, 4, 2, 1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            # 8x8 -> 16x16\n            nn.ConvTranspose2d(256, 128, 4, 2, 1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            # 16x16 -> 32x32\n            nn.ConvTranspose2d(128, 64, 4, 2, 1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            \n            # 32x32 -> 64x64\n            nn.ConvTranspose2d(64, 32, 4, 2, 1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            \n            # 64x64 -> 128x128 (新增这一层)\n            nn.ConvTranspose2d(32, 16, 4, 2, 1),\n            nn.BatchNorm2d(16),\n            nn.ReLU(),\n            \n            # 最终卷积\n            nn.Conv2d(16, output_channels, 3, 1, 1),\n            nn.Tanh()\n        )\n        \n    def forward(self, x):\n        x = self.init_proj(x)\n        x = x.view(-1, 512, 4, 4)\n        x = self.decoder(x)\n        return x\n\nclass FixedEEGToImageModel(nn.Module):\n    \"\"\"修复尺寸问题的EEG到图像转换模型\"\"\"\n    def __init__(self, eeg_channels=62, time_points=400, output_channels=3, image_size=128):\n        super().__init__()\n        self.encoder = FixedEEGEncoder(eeg_channels, time_points)\n        self.decoder = FixedUNetDecoder(output_channels=output_channels, image_size=image_size)\n        \n    def forward(self, eeg_data):\n        features = self.encoder(eeg_data)\n        images = self.decoder(features)\n        return images\n\n# ==================== 训练函数 ====================\n\ndef train_fixed_model(model, train_loader, val_loader, num_epochs=30, learning_rate=2e-4):\n    \"\"\"训练修复尺寸问题的模型\"\"\"\n    model.to(device)\n    \n    # 测试模型输出尺寸\n    with torch.no_grad():\n        test_input = torch.randn(1, 62, 400).to(device)\n        test_output = model(test_input)\n        print(f\"模型输出尺寸: {test_output.shape}\")\n        assert test_output.shape[2:] == (config.image_size, config.image_size), f\"模型输出尺寸不正确: {test_output.shape}\"\n    \n    mse_criterion = nn.MSELoss()\n    perceptual_criterion = LightPerceptualLoss().to(device)\n    \n    optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    history = {'train_loss': [], 'val_loss': []}\n    best_val_loss = float('inf')\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0\n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        \n        for batch_idx, (eeg_data, images, labels) in enumerate(train_bar):\n            memory_manager.clear_memory()\n            \n            eeg_data = eeg_data.to(device)\n            images = images.to(device)\n            \n            optimizer.zero_grad()\n            generated_images = model(eeg_data)\n            \n            # 检查尺寸是否匹配\n            if generated_images.shape != images.shape:\n                print(f\"尺寸不匹配: 生成 {generated_images.shape}, 真实 {images.shape}\")\n                # 调整生成图像尺寸\n                generated_images = nn.functional.interpolate(\n                    generated_images, size=images.shape[2:], mode='bilinear', align_corners=False\n                )\n            \n            mse_loss = mse_criterion(generated_images, images)\n            perceptual_loss = perceptual_criterion(generated_images, images)\n            total_loss = config.mse_weight * mse_loss + config.perceptual_weight * perceptual_loss\n            \n            total_loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            train_loss += total_loss.item()\n            train_bar.set_postfix({'loss': f'{total_loss.item():.4f}'})\n            \n            # 定期清理内存\n            if batch_idx % 20 == 0:\n                memory_manager.clear_memory()\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for eeg_data, images, labels in val_loader:\n                eeg_data = eeg_data.to(device)\n                images = images.to(device)\n                \n                generated_images = model(eeg_data)\n                \n                # 检查尺寸是否匹配\n                if generated_images.shape != images.shape:\n                    generated_images = nn.functional.interpolate(\n                        generated_images, size=images.shape[2:], mode='bilinear', align_corners=False\n                    )\n                \n                mse_loss = mse_criterion(generated_images, images)\n                perceptual_loss = perceptual_criterion(generated_images, images)\n                total_loss = config.mse_weight * mse_loss + config.perceptual_weight * perceptual_loss\n                val_loss += total_loss.item()\n        \n        avg_train_loss = train_loss / len(train_loader)\n        avg_val_loss = val_loss / len(val_loader)\n        \n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n        \n        current_lr = optimizer.param_groups[0]['lr']\n        scheduler.step()\n        \n        print(f'Epoch {epoch+1}/{num_epochs}:')\n        print(f'  训练损失: {avg_train_loss:.4f}, 验证损失: {avg_val_loss:.4f}, LR: {current_lr:.6f}')\n        \n        # 保存最佳模型\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            torch.save(model.state_dict(), config.model_save_path)\n            print(f'  保存最佳模型: {config.model_save_path}')\n        \n        # 定期可视化\n        if (epoch + 1) % 5 == 0:\n            visualize_results(model, val_loader, epoch + 1)\n        \n        memory_manager.clear_memory()\n    \n    return model, history\n\ndef visualize_results(model, val_loader, epoch):\n    \"\"\"可视化结果\"\"\"\n    model.eval()\n    with torch.no_grad():\n        eeg_data, real_images, labels = next(iter(val_loader))\n        eeg_data = eeg_data[:4].to(device)\n        real_images = real_images[:4]\n        \n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 检查尺寸\n        if generated_images.shape != real_images.shape:\n            print(f\"可视化尺寸不匹配: 生成 {generated_images.shape}, 真实 {real_images.shape}\")\n            generated_images = nn.functional.interpolate(\n                generated_images, size=real_images.shape[2:], mode='bilinear', align_corners=False\n            )\n        \n        fig, axes = plt.subplots(3, 4, figsize=(16, 9))\n        \n        for i in range(4):\n            # EEG信号\n            eeg_sample = eeg_data[i][:8, :100].cpu().numpy()\n            axes[0, i].plot(eeg_sample.T, alpha=0.7, linewidth=1)\n            axes[0, i].set_title(f'EEG Sample {i+1}', fontsize=10)\n            axes[0, i].grid(True, alpha=0.3)\n            \n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[1, i].imshow(real_img)\n            axes[1, i].set_title('Real Image', fontsize=10)\n            axes[1, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[2, i].imshow(gen_img)\n            axes[2, i].set_title('Generated Image', fontsize=10)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig(f'results_epoch_{epoch}.png', dpi=150, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"结果已保存: results_epoch_{epoch}.png\")\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"启动修复尺寸问题的EEG-ImageNet图像重建系统...\")\n    print(f\"配置参数:\")\n    print(f\"  训练轮次: {config.num_epochs}\")\n    print(f\"  学习率: {config.learning_rate}\")\n    print(f\"  批大小: {config.batch_size}\")\n    print(f\"  最大训练样本: {config.max_train_samples}\")\n    \n    # 创建数据集\n    print(\"创建数据集...\")\n    train_dataset = FixedEEGImageDataset(\n        eeg_data_path=config.eeg_data_path,\n        image_size=config.image_size,\n        mode='train',\n        max_samples=config.max_train_samples\n    )\n    \n    test_dataset = FixedEEGImageDataset(\n        eeg_data_path=config.eeg_data_path,\n        image_size=config.image_size,\n        mode='test',\n        max_samples=config.max_test_samples\n    )\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=config.batch_size, \n        shuffle=True,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=True\n    )\n    \n    print(f\"训练集: {len(train_dataset)} 样本\")\n    print(f\"测试集: {len(test_dataset)} 样本\")\n    \n    # 创建模型\n    print(\"初始化修复尺寸问题的模型...\")\n    model = FixedEEGToImageModel(\n        eeg_channels=config.eeg_channels,\n        time_points=config.time_points,\n        output_channels=3,\n        image_size=config.image_size\n    )\n    \n    # 训练模型\n    print(\"开始训练模型...\")\n    model, history = train_fixed_model(\n        model=model,\n        train_loader=train_loader,\n        val_loader=test_loader,\n        num_epochs=config.num_epochs,\n        learning_rate=config.learning_rate\n    )\n    \n    # 绘制训练历史\n    plt.figure(figsize=(10, 5))\n    plt.plot(history['train_loss'], label='训练损失')\n    plt.plot(history['val_loss'], label='验证损失')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('训练历史')\n    plt.grid(True)\n    plt.savefig('training_history_fixed.png', dpi=150, bbox_inches='tight')\n    plt.show()\n    \n    # 最终测试\n    print(\"最终测试...\")\n    test_model(model, test_loader)\n    \n    # 清理内存\n    memory_manager.clear_memory()\n    print(\"训练完成!\")\n\ndef test_model(model, test_loader):\n    \"\"\"测试模型\"\"\"\n    model.eval()\n    with torch.no_grad():\n        eeg_data, real_images, labels = next(iter(test_loader))\n        eeg_data = eeg_data[:8].to(device)\n        real_images = real_images[:8]\n        \n        generated_images = model(eeg_data)\n        generated_images = generated_images.cpu()\n        \n        # 检查尺寸\n        if generated_images.shape != real_images.shape:\n            print(f\"测试尺寸不匹配: 生成 {generated_images.shape}, 真实 {real_images.shape}\")\n            generated_images = nn.functional.interpolate(\n                generated_images, size=real_images.shape[2:], mode='bilinear', align_corners=False\n            )\n        \n        # 计算评估指标\n        mse_loss = nn.MSELoss()(generated_images, real_images).item()\n        print(f\"测试MSE损失: {mse_loss:.4f}\")\n        \n        # 可视化结果\n        fig, axes = plt.subplots(3, 8, figsize=(20, 8))\n        \n        for i in range(8):\n            # 真实图像\n            real_img = real_images[i].permute(1, 2, 0).numpy()\n            real_img = np.clip(real_img, 0, 1)\n            axes[0, i].imshow(real_img)\n            axes[0, i].set_title('Real', fontsize=8)\n            axes[0, i].axis('off')\n            \n            # 生成图像\n            gen_img = generated_images[i].permute(1, 2, 0).numpy()\n            gen_img = np.clip(gen_img, 0, 1)\n            axes[1, i].imshow(gen_img)\n            axes[1, i].set_title('Generated', fontsize=8)\n            axes[1, i].axis('off')\n            \n            # 差异图\n            diff_img = np.abs(real_img - gen_img)\n            axes[2, i].imshow(diff_img, cmap='hot')\n            axes[2, i].set_title('Difference', fontsize=8)\n            axes[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.savefig('final_results_fixed.png', dpi=200, bbox_inches='tight')\n        plt.show()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T11:53:48.118902Z","iopub.execute_input":"2025-10-16T11:53:48.119405Z","execution_failed":"2025-10-16T13:07:55.498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nEEG-ImageNet真实图像下载器 - 修复搜索问题版本\n优化搜索策略，避免卡住\n\"\"\"\n\nimport os\nimport torch\nfrom tqdm import tqdm\nimport pickle\nfrom pathlib import Path\nimport numpy as np\nimport json\nimport time\n\nclass EEGImageNetKaggleDownloader:\n    \"\"\"EEG-ImageNet Kaggle数据集下载器\"\"\"\n    \n    def __init__(self, data_dir=\"/kaggle/working/eeg_imagenet_images\", eeg_data_path=None):\n        self.data_dir = Path(data_dir)\n        self.data_dir.mkdir(parents=True, exist_ok=True)\n        \n        self.eeg_data_path = eeg_data_path or \"/kaggle/input/eeg-imagenet/EEG-ImageNet_1.pth\"\n        \n        # 加载进度\n        self.progress_file = self.data_dir / \"download_progress.pkl\"\n        self.downloaded_images = self.load_progress()\n        \n        # 确保是字典类型\n        if not isinstance(self.downloaded_images, dict):\n            print(\"警告: downloaded_images不是字典，重新初始化\")\n            self.downloaded_images = {}\n    \n    def load_progress(self):\n        \"\"\"加载下载进度\"\"\"\n        if self.progress_file.exists():\n            try:\n                with open(self.progress_file, 'rb') as f:\n                    progress = pickle.load(f)\n                    if isinstance(progress, dict):\n                        return progress\n                    else:\n                        print(f\"警告: 进度文件不是字典，而是 {type(progress)}\")\n                        return {}\n            except Exception as e:\n                print(f\"加载进度失败: {e}\")\n                return {}\n        return {}\n    \n    def save_progress(self):\n        \"\"\"保存下载进度\"\"\"\n        try:\n            with open(self.progress_file, 'wb') as f:\n                pickle.dump(self.downloaded_images, f)\n        except Exception as e:\n            print(f\"保存进度失败: {e}\")\n    \n    def extract_all_image_filenames(self):\n        \"\"\"提取所有图像文件名\"\"\"\n        print(\"提取所有图像文件名...\")\n        \n        try:\n            data = torch.load(self.eeg_data_path, map_location='cpu', weights_only=False)\n            dataset = data['dataset']\n            \n            image_info = {}\n            \n            for idx, sample in enumerate(tqdm(dataset, desc=\"提取文件名\")):\n                sample_id = f\"sample_{idx}\"\n                \n                # 检查是否有图像文件名\n                if 'image' in sample and isinstance(sample['image'], str):\n                    filename = sample['image']\n                    \n                    # 转换为可序列化的Python类型\n                    subject = sample.get('subject', 'unknown')\n                    if hasattr(subject, 'item'):\n                        subject = subject.item()\n                    \n                    image_info[sample_id] = {\n                        'filename': filename,\n                        'label': str(sample.get('label', 'unknown')),\n                        'subject': subject,\n                        'granularity': str(sample.get('granularity', 'unknown'))\n                    }\n            \n            print(f\"成功提取 {len(image_info)} 个图像文件名\")\n            return image_info\n            \n        except Exception as e:\n            print(f\"提取图像文件名失败: {e}\")\n            return {}\n    \n    def analyze_imagenet_structure(self, dataset_path):\n        \"\"\"分析ImageNet数据集结构\"\"\"\n        print(f\"分析数据集结构: {dataset_path}\")\n        \n        structure_info = {\n            'has_train': False,\n            'has_val': False,\n            'has_test': False,\n            'train_structure': None,\n            'val_structure': None,\n            'categories': []\n        }\n        \n        # 检查常见的目录结构\n        possible_structures = [\n            # 标准ImageNet结构\n            {\n                'train': os.path.join(dataset_path, \"ILSVRC/Data/CLS-LOC/train\"),\n                'val': os.path.join(dataset_path, \"ILSVRC/Data/CLS-LOC/val\"),\n                'test': os.path.join(dataset_path, \"ILSVRC/Data/CLS-LOC/test\")\n            },\n            # 简化结构\n            {\n                'train': os.path.join(dataset_path, \"train\"),\n                'val': os.path.join(dataset_path, \"val\"),\n                'test': os.path.join(dataset_path, \"test\")\n            },\n            # 其他变体\n            {\n                'train': os.path.join(dataset_path, \"Training\"),\n                'val': os.path.join(dataset_path, \"Validation\"),\n                'test': os.path.join(dataset_path, \"Testing\")\n            }\n        ]\n        \n        for structure in possible_structures:\n            train_exists = os.path.exists(structure['train'])\n            val_exists = os.path.exists(structure['val'])\n            test_exists = os.path.exists(structure['test'])\n            \n            if train_exists or val_exists:\n                structure_info.update({\n                    'has_train': train_exists,\n                    'has_val': val_exists,\n                    'has_test': test_exists,\n                    'train_structure': structure['train'] if train_exists else None,\n                    'val_structure': structure['val'] if val_exists else None,\n                    'test_structure': structure['test'] if test_exists else None\n                })\n                print(f\"  找到结构: train={train_exists}, val={val_exists}, test={test_exists}\")\n                return structure_info\n        \n        # 如果没有标准结构，尝试查找任何图像文件\n        print(\"  没有找到标准结构，尝试查找图像文件...\")\n        jpeg_files = list(Path(dataset_path).rglob(\"*.JPEG\"))\n        if jpeg_files:\n            print(f\"  找到 {len(jpeg_files)} 个JPEG文件\")\n            structure_info['has_images'] = True\n        else:\n            print(\"  没有找到任何JPEG文件\")\n        \n        return structure_info\n    \n    def find_imagenet_datasets(self):\n        \"\"\"查找可用的ImageNet数据集并分析结构\"\"\"\n        print(\"查找可用的ImageNet数据集...\")\n        \n        # 可能的ImageNet数据集路径\n        possible_datasets = [\n            \"/kaggle/input/imagenet-object-localization-challenge\",\n            \"/kaggle/input/imagenet1k\",\n            \"/kaggle/input/imagenet-1k\",\n            \"/kaggle/input/imagenet21k\",\n            \"/kaggle/input/imagenet-21k\",\n        ]\n        \n        available_datasets = []\n        for dataset_path in possible_datasets:\n            if os.path.exists(dataset_path):\n                structure_info = self.analyze_imagenet_structure(dataset_path)\n                structure_info['path'] = dataset_path\n                available_datasets.append(structure_info)\n        \n        if not available_datasets:\n            print(\"警告: 没有找到任何ImageNet数据集\")\n        \n        return available_datasets\n    \n    def search_image_fast(self, filename, datasets):\n        \"\"\"快速搜索图像文件\"\"\"\n        wordnet_id = filename.split('_')[0]\n        \n        for dataset_info in datasets:\n            dataset_path = dataset_info['path']\n            \n            # 尝试训练集结构\n            if dataset_info['has_train'] and dataset_info['train_structure']:\n                train_path = os.path.join(dataset_info['train_structure'], wordnet_id, filename)\n                if os.path.exists(train_path):\n                    return train_path\n            \n            # 尝试验证集结构\n            if dataset_info['has_val'] and dataset_info['val_structure']:\n                val_path = os.path.join(dataset_info['val_structure'], filename)\n                if os.path.exists(val_path):\n                    return val_path\n            \n            # 尝试测试集结构\n            if dataset_info['has_test'] and dataset_info['test_structure']:\n                test_path = os.path.join(dataset_info['test_structure'], filename)\n                if os.path.exists(test_path):\n                    return test_path\n            \n            # 如果数据集有图像但没有标准结构，直接搜索\n            if dataset_info.get('has_images', False):\n                # 只在当前目录搜索，避免递归太慢\n                direct_path = os.path.join(dataset_path, filename)\n                if os.path.exists(direct_path):\n                    return direct_path\n        \n        return None\n    \n    def download_from_kaggle_only(self, batch_size=100):\n        \"\"\"只从Kaggle数据集下载图像\"\"\"\n        print(\"开始从Kaggle数据集下载图像...\")\n        \n        # 提取所有图像信息\n        image_info = self.extract_all_image_filenames()\n        \n        if not image_info:\n            print(\"错误: 无法提取图像信息\")\n            return 0\n        \n        # 查找可用的数据集\n        datasets = self.find_imagenet_datasets()\n        if not datasets:\n            print(\"错误: 没有找到可用的ImageNet数据集\")\n            return 0\n        \n        total_samples = len(image_info)\n        downloaded_count = 0\n        found_count = 0\n        not_found_count = 0\n        \n        print(f\"总共需要处理 {total_samples} 个样本\")\n        print(f\"使用 {len(datasets)} 个数据集进行搜索\")\n        \n        # 分批处理\n        sample_ids = list(image_info.keys())\n        \n        for batch_start in range(0, total_samples, batch_size):\n            batch_end = min(batch_start + batch_size, total_samples)\n            batch_ids = sample_ids[batch_start:batch_end]\n            \n            print(f\"\\n处理批次 {batch_start//batch_size + 1}/{(total_samples + batch_size - 1)//batch_size}\")\n            \n            batch_found = 0\n            batch_not_found = 0\n            \n            for sample_id in tqdm(batch_ids, desc=\"搜索图像\"):\n                # 检查是否已下载\n                if sample_id in self.downloaded_images:\n                    downloaded_count += 1\n                    continue\n                \n                info = image_info[sample_id]\n                filename = info['filename']\n                \n                # 快速搜索图像\n                source_path = self.search_image_fast(filename, datasets)\n                \n                if source_path:\n                    try:\n                        # 复制文件\n                        target_path = self.data_dir / f\"{sample_id}_{filename}\"\n                        with open(source_path, 'rb') as src, open(target_path, 'wb') as dst:\n                            dst.write(src.read())\n                        \n                        # 保存下载信息\n                        self.downloaded_images[sample_id] = {\n                            'image_path': str(target_path),\n                            'source': \"kaggle\",\n                            'filename': filename,\n                            'label': info['label'],\n                            'subject': info['subject'],\n                            'granularity': info['granularity']\n                        }\n                        \n                        found_count += 1\n                        batch_found += 1\n                        downloaded_count += 1\n                        \n                    except Exception as e:\n                        print(f\"复制文件失败 {source_path}: {e}\")\n                        not_found_count += 1\n                        batch_not_found += 1\n                else:\n                    not_found_count += 1\n                    batch_not_found += 1\n            \n            # 每批结束后保存进度\n            self.save_progress()\n            print(f\"批次完成: 找到 {batch_found} 张, 未找到 {batch_not_found} 张\")\n            print(f\"总进度: {downloaded_count}/{total_samples} ({downloaded_count/total_samples*100:.2f}%)\")\n            \n            # 如果连续多个批次都没有找到图像，提前停止\n            if batch_found == 0 and batch_start > batch_size * 3:\n                print(\"连续多个批次没有找到图像，提前停止\")\n                break\n        \n        print(f\"\\n下载完成!\")\n        print(f\"  总成功: {found_count}\")\n        print(f\"  未找到: {not_found_count}\")\n        print(f\"  成功率: {found_count/total_samples*100:.2f}%\")\n        \n        return found_count\n    \n    def create_metadata(self):\n        \"\"\"创建元数据\"\"\"\n        metadata = {\n            \"total_downloaded\": len(self.downloaded_images),\n            \"download_time\": time.strftime(\"%Y-%m-%d %H:%M:%S\"),\n            \"data_dir\": str(self.data_dir),\n            \"samples\": self.downloaded_images\n        }\n        \n        # 保存元数据\n        metadata_path = self.data_dir / \"metadata.json\"\n        \n        # 转换为可JSON序列化的格式\n        def make_serializable(obj):\n            if isinstance(obj, (np.integer, np.int64, np.int32)):\n                return int(obj)\n            elif isinstance(obj, (np.floating, np.float64, np.float32)):\n                return float(obj)\n            elif isinstance(obj, np.ndarray):\n                return obj.tolist()\n            elif isinstance(obj, dict):\n                return {k: make_serializable(v) for k, v in obj.items()}\n            elif isinstance(obj, list):\n                return [make_serializable(item) for item in obj]\n            else:\n                return obj\n        \n        serializable_metadata = make_serializable(metadata)\n        \n        with open(metadata_path, 'w') as f:\n            json.dump(serializable_metadata, f, indent=2)\n        \n        print(f\"\\n元数据已保存:\")\n        print(f\"  下载的图像: {metadata['total_downloaded']}\")\n        print(f\"  数据目录: {metadata['data_dir']}\")\n        \n        return metadata\n\n# ==================== 主函数 ====================\n\ndef main():\n    \"\"\"主函数\"\"\"\n    print(\"EEG-ImageNet Kaggle图像下载器启动...\")\n    \n    # 创建下载器\n    downloader = EEGImageNetKaggleDownloader(\n        data_dir=\"/kaggle/working/eeg_imagenet_images\",\n        eeg_data_path=\"/kaggle/input/eeg-imagenet/EEG-ImageNet_1.pth\"\n    )\n    \n    # 只从Kaggle下载图像\n    downloaded_count = downloader.download_from_kaggle_only(\n        batch_size=100  # 更小的批次大小，更快反馈\n    )\n    \n    # 创建元数据\n    metadata = downloader.create_metadata()\n    \n    print(\"\\n\" + \"=\"*50)\n    print(\"处理完成!\")\n    print(f\"成功下载 {downloaded_count} 张真实图像\")\n    \n    # 显示一些统计信息\n    if downloaded_count > 0:\n        print(\"\\n数据集统计:\")\n        labels = {}\n        for sample_info in downloader.downloaded_images.values():\n            label = sample_info.get('label', 'unknown')\n            labels[label] = labels.get(label, 0) + 1\n        \n        print(f\"  类别数量: {len(labels)}\")\n        print(f\"  样本最多的前5个类别:\")\n        sorted_labels = sorted(labels.items(), key=lambda x: x[1], reverse=True)[:5]\n        for label, count in sorted_labels:\n            print(f\"    {label}: {count} 张图像\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T09:14:46.597912Z","iopub.execute_input":"2025-10-17T09:14:46.598485Z","iopub.status.idle":"2025-10-17T09:16:03.505652Z","shell.execute_reply.started":"2025-10-17T09:14:46.598457Z","shell.execute_reply":"2025-10-17T09:16:03.504367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport random\nfrom PIL import Image\nfrom pathlib import Path\n\ndef view_downloaded_images(data_dir=\"/kaggle/working/eeg_imagenet_images\", num_images=5):\n    \"\"\"\n    查看已下载的图片\n    \n    参数:\n        data_dir: 数据目录路径\n        num_images: 要显示的图片数量\n    \"\"\"\n    data_path = Path(data_dir)\n    \n    # 获取所有图片路径\n    image_paths = []\n    for category_dir in data_path.iterdir():\n        if category_dir.is_dir():\n            for image_file in category_dir.glob(\"*.jpg\"):\n                image_paths.append(image_file)\n    \n    if not image_paths:\n        print(\"未找到任何图片，请先运行下载器\")\n        return\n    \n    # 随机选择图片\n    selected_images = random.sample(image_paths, min(num_images, len(image_paths)))\n    \n    # 设置显示布局\n    fig, axes = plt.subplots(1, len(selected_images), figsize=(15, 5))\n    if len(selected_images) == 1:\n        axes = [axes]\n    \n    # 显示图片\n    for i, img_path in enumerate(selected_images):\n        img = Image.open(img_path)\n        axes[i].imshow(img)\n        axes[i].set_title(f\"{img_path.parent.name}\\n{img_path.name}\")\n        axes[i].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # 打印图片信息\n    print(f\"显示 {len(selected_images)} 张图片 (共 {len(image_paths)} 张):\")\n    for img_path in selected_images:\n        print(f\"- {img_path.parent.name}/{img_path.name}\")\n\n# 使用示例\nif __name__ == \"__main__\":\n    view_downloaded_images(num_images=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T14:01:32.458213Z","iopub.execute_input":"2025-10-16T14:01:32.458695Z","iopub.status.idle":"2025-10-16T14:01:33.38585Z","shell.execute_reply.started":"2025-10-16T14:01:32.45867Z","shell.execute_reply":"2025-10-16T14:01:33.384625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}