{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553},{"sourceType":"datasetVersion","sourceId":15987623,"datasetId":10253099,"databundleVersionId":16949612},{"sourceType":"datasetVersion","sourceId":15669048,"datasetId":10033134,"databundleVersionId":16606230},{"sourceType":"datasetVersion","sourceId":15648661,"datasetId":10017903,"databundleVersionId":16584492},{"sourceType":"datasetVersion","sourceId":15529269,"datasetId":9935314,"databundleVersionId":16456830},{"sourceType":"datasetVersion","sourceId":15653647,"datasetId":10021709,"databundleVersionId":16589768}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# SYNAPSE: EEG-to-Image Generation\n\nEnd-to-end inference pipeline for the SYNAPSE framework (arXiv:2511.17547).\nGenerates images from EEG brain signals using a CLIP-aligned encoder + Stable Diffusion v2.1.\n\n**Datasets required:**\n- `kavenio/eeg-based-visual-classification-dataset` — EEG signal files\n- `jul503/stable-diffusion-2-1` — SD 2.1 model weights\n- `jul503/blip2-opt-2-7b` — BLIP2 model\n- `heyitsrup/stable-diffusion-checkpoints` — SYNAPSE Stage 2 checkpoints\n- `imagenet-object-localization-challenge` (competition) — ImageNet images","metadata":{}},{"cell_type":"markdown","source":"## Step 1: Clone Repo & Install Dependencies","metadata":{}},{"cell_type":"code","source":"import os\n\n# (Skips clone if repo already exists)\nif not os.path.exists('/kaggle/working/synapse'):\n    !git clone https://github.com/CVMILab-CUK/synapse.git\n    print('Cloned fresh')\nelse:\n    print('Repo already exists, skipping clone')\n\nos.chdir('/kaggle/working/synapse')\nprint('Working dir:', os.getcwd())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:40:30.410766Z","iopub.execute_input":"2026-04-29T10:40:30.411703Z","iopub.status.idle":"2026-04-29T10:40:30.422625Z","shell.execute_reply.started":"2026-04-29T10:40:30.411659Z","shell.execute_reply":"2026-04-29T10:40:30.42198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install open_clip_torch accelerate einops omegaconf wandb \\\n    clean-fid torch-fidelity transformers diffusers \\\n    peft xformers safetensors pytorch-lightning \\\n    torchmetrics scipy scikit-image gdown \\\n    natsort torcheval lightning \\\n    k-diffusion \"numpy>=2.0\" -q\nprint('Core dependencies ready')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:40:42.672628Z","iopub.execute_input":"2026-04-29T10:40:42.67324Z","iopub.status.idle":"2026-04-29T10:40:46.795246Z","shell.execute_reply.started":"2026-04-29T10:40:42.673209Z","shell.execute_reply":"2026-04-29T10:40:46.794153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install taming-transformers (required by dc_ldm)\n!pip install git+https://github.com/CompVis/taming-transformers.git -q\nprint('taming-transformers ready')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:40:50.045651Z","iopub.execute_input":"2026-04-29T10:40:50.046132Z","iopub.status.idle":"2026-04-29T10:40:59.557331Z","shell.execute_reply.started":"2026-04-29T10:40:50.046095Z","shell.execute_reply":"2026-04-29T10:40:59.556495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Copy dc_ldm, sc_mbm, eval_metrics and config from DreamDiffusion baseline\nimport shutil, os\n\nif not os.path.exists('/tmp/dreamdiffusion'):\n    !git clone https://github.com/bbaaii/DreamDiffusion.git /tmp/dreamdiffusion\n\nos.chdir('/kaggle/working/synapse/src')\n\nfor item in ['dc_ldm', 'sc_mbm']:\n    dst = f'/kaggle/working/synapse/src/{item}'\n    src = f'/tmp/dreamdiffusion/code/{item}'\n    if not os.path.exists(dst) and os.path.exists(src):\n        shutil.copytree(src, dst)\n        print(f'Copied {item}')\n    elif os.path.exists(dst):\n        print(f'{item} already exists')\n    else:\n        print(f'WARNING: {item} not found in DreamDiffusion')\n\nfor fname in ['eval_metrics.py', 'config.py']:\n    src = f'/tmp/dreamdiffusion/code/{fname}'\n    dst = f'/kaggle/working/synapse/src/{fname}'\n    if os.path.exists(src) and not os.path.exists(dst):\n        shutil.copy2(src, dst)\n        print(f'Copied {fname}')\n    elif os.path.exists(dst):\n        print(f'{fname} already exists')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:14.667824Z","iopub.execute_input":"2026-04-29T10:41:14.668563Z","iopub.status.idle":"2026-04-29T10:41:14.677211Z","shell.execute_reply.started":"2026-04-29T10:41:14.668526Z","shell.execute_reply":"2026-04-29T10:41:14.676559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 2: Verify Dataset Paths","metadata":{}},{"cell_type":"code","source":"paths_to_check = {\n    'EEG Data':               '/kaggle/input/datasets/heyitsrup/eef-5-95-ad-simulated-pth/eeg_5_95_std_AD_simulated.pth',\n    'EEG Splits':             '/kaggle/input/datasets/kavenio/eeg-based-visual-classification-dataset/block_splits_by_image_all.pth',\n    'SD 2.1':                 '/kaggle/input/datasets/jul503/stable-diffusion-2-1/model_index.json',\n    'BLIP2':                  '/kaggle/input/datasets/jul503/blip2-opt-2-7b/blip2-opt-2-7b',\n    'Checkpoint ip_adaption': '/kaggle/input/datasets/heyitsrup/stable-diffusion-checkpoints/checkpoint-23000/ip_adaption/diffusion_pytorch_model.safetensors',\n    'Checkpoint ema_unet':    '/kaggle/input/datasets/heyitsrup/stable-diffusion-checkpoints/checkpoint-23000/ema_unet/diffusion_pytorch_model.safetensors',\n    'ImageNet train':         '/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train',\n}\n\nall_ok = True\nfor name, path in paths_to_check.items():\n    status = 'OK' if os.path.exists(path) else 'MISSING'\n    if status == 'MISSING': all_ok = False\n    print(f'[{status}]  {name}')\n\nprint('\\nAll paths OK' if all_ok else '\\nFix missing paths before continuing')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:26.34296Z","iopub.execute_input":"2026-04-29T10:41:26.343384Z","iopub.status.idle":"2026-04-29T10:41:26.364582Z","shell.execute_reply.started":"2026-04-29T10:41:26.343352Z","shell.execute_reply":"2026-04-29T10:41:26.36398Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 3: Download Stage 1 Encoder Weights","metadata":{}},{"cell_type":"code","source":"os.makedirs('/kaggle/working/pretrain_models', exist_ok=True)\n\nENCODER_PATH = '/kaggle/working/pretrain_models/pretrain_models/EEGEncoder_BLIP2_69000.pth'\n\n# (Skips download if already present)\nif not os.path.exists(ENCODER_PATH):\n    !gdown --folder \"https://drive.google.com/drive/folders/1aQe2bxIijPKFn0fYnp-TnmjTVZCfi_T-\" \\\n        -O /kaggle/working/pretrain_models/\n    print('Download complete')\nelse:\n    print('Encoder already downloaded, skipping')\n\nfor root, dirs, files in os.walk('/kaggle/working/pretrain_models'):\n    for f in files:\n        print(os.path.join(root, f))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:28.63973Z","iopub.execute_input":"2026-04-29T10:41:28.640127Z","iopub.status.idle":"2026-04-29T10:41:28.64637Z","shell.execute_reply.started":"2026-04-29T10:41:28.640101Z","shell.execute_reply":"2026-04-29T10:41:28.645771Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 4: Preprocess EEG Data","metadata":{}},{"cell_type":"code","source":"import sys, os, glob, random, shutil\nimport torch\nfrom torch.utils.data import DataLoader\nimport torch.backends.cudnn as cudnn\ncudnn.benchmark = True\nfrom scipy.fftpack import fft, rfft, fftfreq, irfft, ifft, rfftfreq\nfrom scipy import signal\nimport numpy as np\nimport cv2\nfrom tqdm.notebook import tqdm\nprint('Imports OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:30.85059Z","iopub.execute_input":"2026-04-29T10:41:30.851035Z","iopub.status.idle":"2026-04-29T10:41:33.758525Z","shell.execute_reply.started":"2026-04-29T10:41:30.851005Z","shell.execute_reply":"2026-04-29T10:41:33.757775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EEGDataset:\n    def __init__(self, eeg_signals_path, eeg_data_path):\n        print('Loading EEG data...')\n        loaded = torch.load(eeg_signals_path)\n        self.data = loaded['dataset']\n        self.labels = loaded['labels']\n        self.images = loaded['images']\n        self.image_path = eeg_data_path\n        self.size = len(self.data)\n        print(f'Loaded {self.size} EEG samples')\n    def __len__(self): return self.size\n    def __getitem__(self, i):\n        eeg = self.data[i]['eeg'].float().t()[20:460, :]\n        label = self.data[i]['label']\n        image = self.images[self.data[i]['image']]\n        subject = self.data[i]['subject']\n        return eeg, image, label, subject\n\nclass Splitter:\n    def __init__(self, dataset, split_path, split_num=0, split_name='train'):\n        self.dataset = dataset\n        loaded = torch.load(split_path)\n        self.split_idx = loaded['splits'][split_num][split_name]\n        self.split_idx = [i for i in self.split_idx if 450 <= self.dataset.data[i]['eeg'].size(1) <= 600]\n        self.size = len(self.split_idx)\n        print(f'Split [{split_name}]: {self.size} samples')\n    def __len__(self): return self.size\n    def __getitem__(self, i): return self.dataset[self.split_idx[i]]\n\nprint('Classes defined')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:38.33152Z","iopub.execute_input":"2026-04-29T10:41:38.332084Z","iopub.status.idle":"2026-04-29T10:41:38.340509Z","shell.execute_reply.started":"2026-04-29T10:41:38.332052Z","shell.execute_reply":"2026-04-29T10:41:38.339837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EEG_SIGNALS_PATH = '/kaggle/input/datasets/heyitsrup/eef-5-95-ad-simulated-pth/eeg_5_95_std_AD_simulated.pth'\nEEG_SPLITS_PATH  = '/kaggle/input/datasets/kavenio/eeg-based-visual-classification-dataset/block_splits_by_image_all.pth'\nIMG_PATH         = '/kaggle/input/datasets/kavenio/eeg-based-visual-classification-dataset'\nPREPROC_OUT_PATH = '/kaggle/working/preprocessing_data'\n\nfor split in ['train', 'val', 'test']:\n    os.makedirs(os.path.join(PREPROC_OUT_PATH, split), exist_ok=True)\n\ntrain_count = len(os.listdir(os.path.join(PREPROC_OUT_PATH, 'train')))\nPREPROC_DONE = train_count > 0\nif PREPROC_DONE:\n    print(f'Preprocessing already done ({train_count} train files). Skipping.')\nelse:\n    print('Preprocessing needed, continuing...')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:44.912033Z","iopub.execute_input":"2026-04-29T10:41:44.912818Z","iopub.status.idle":"2026-04-29T10:41:44.918866Z","shell.execute_reply.started":"2026-04-29T10:41:44.912785Z","shell.execute_reply":"2026-04-29T10:41:44.917977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not PREPROC_DONE:\n    dataset = EEGDataset(eeg_signals_path=EEG_SIGNALS_PATH, eeg_data_path=IMG_PATH)\n    loaders = {\n        split: DataLoader(\n            Splitter(dataset, split_path=EEG_SPLITS_PATH, split_num=0, split_name=split),\n            batch_size=1, drop_last=False, shuffle=False\n        ) for split in ['train', 'val', 'test']\n    }\n    file_name = EEG_SIGNALS_PATH.split('/')[-1].replace('.pth', '')\n    for split in ['train', 'val', 'test']:\n        for idx, data in tqdm(enumerate(loaders[split]), total=len(loaders[split]), desc=f'{split}'):\n            subject = data[3].item()\n            data_out = {'eeg': data[0].numpy().squeeze(), 'image': data[1][0],\n                        'label': data[2].item(), 'subject': data[3].item()}\n            class_dir = os.path.join(PREPROC_OUT_PATH, f'class_{subject}')\n            os.makedirs(class_dir, exist_ok=True)\n            torch.save(data_out, os.path.join(class_dir, f'{file_name}_{split}_{idx}.pth'))\n    frac = 0.8\n    class_paths = glob.glob(os.path.join(PREPROC_OUT_PATH, 'class_*'))\n    tr_idx = te_idx = va_idx = 0\n    for p in tqdm(class_paths, desc='Splitting'):\n        ep_lst = os.listdir(p)\n        name = p.split('/')[-1]\n        random.shuffle(ep_lst)\n        length    = len(ep_lst)\n        train_len = int(length * frac)\n        valid_len = int(length * (1 - frac) // 2)\n        for t in ep_lst[:train_len]:\n            shutil.copy2(os.path.join(p, t), os.path.join(PREPROC_OUT_PATH, 'train', f'{name}_{tr_idx}.pth'))\n            tr_idx += 1\n        for t in ep_lst[train_len:train_len + valid_len]:\n            shutil.copy2(os.path.join(p, t), os.path.join(PREPROC_OUT_PATH, 'val', f'{name}_{va_idx}.pth'))\n            va_idx += 1\n        for t in ep_lst[train_len + valid_len:]:\n            shutil.copy2(os.path.join(p, t), os.path.join(PREPROC_OUT_PATH, 'test', f'{name}_{te_idx}.pth'))\n            te_idx += 1\n    print(f'Done — Train: {tr_idx} | Val: {va_idx} | Test: {te_idx}')\n\nprint('Train:', len(os.listdir(os.path.join(PREPROC_OUT_PATH, 'train'))))\nprint('Val:  ', len(os.listdir(os.path.join(PREPROC_OUT_PATH, 'val'))))\nprint('Test: ', len(os.listdir(os.path.join(PREPROC_OUT_PATH, 'test'))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:41:59.914309Z","iopub.execute_input":"2026-04-29T10:41:59.915171Z","iopub.status.idle":"2026-04-29T10:42:48.835153Z","shell.execute_reply.started":"2026-04-29T10:41:59.915127Z","shell.execute_reply":"2026-04-29T10:42:48.83417Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 5: Copy ImageNet Images\n\n> **Manual step required before running this cell.**\n>\n> You must attach the ImageNet competition dataset to this notebook:\n>\n> 1. Click Add Input\n> 2. Search \"imagenet-object-localization-challenge\"\n> 3. Add \"imagenet-object-localization-challenge\"","metadata":{}},{"cell_type":"code","source":"# Copy the 40 ImageNet synsets used in EEGCVPR40 from the competition dataset\nsynsets = ['n02106662', 'n02124075', 'n02281787', 'n02389026', 'n02492035',\n           'n02504458', 'n02510455', 'n02607072', 'n02690373', 'n02906734',\n           'n02951358', 'n02992529', 'n03063599', 'n03100240', 'n03180011',\n           'n03197337', 'n03272010', 'n03272562', 'n03297495', 'n03376595',\n           'n03445777', 'n03452741', 'n03584829', 'n03590841', 'n03709823',\n           'n03773504', 'n03775071', 'n03792782', 'n03792972', 'n03877472',\n           'n03888257', 'n03982430', 'n04044716', 'n04069434', 'n04086273',\n           'n04120489', 'n07753592', 'n07873807', 'n11939491', 'n13054560']\n\nIMAGENET_TRAIN = '/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train'\nIMG_OUT_PATH   = '/kaggle/working/imagenet_images'\nos.makedirs(IMG_OUT_PATH, exist_ok=True)\n\nfor synset in synsets:\n    src = os.path.join(IMAGENET_TRAIN, synset)\n    dst = os.path.join(IMG_OUT_PATH, synset)\n    if os.path.exists(dst):\n        print(f'Already exists: {synset}')\n    elif os.path.exists(src):\n        shutil.copytree(src, dst)\n        print(f'Copied: {synset}')\n    else:\n        print(f'MISSING: {synset} — accept ImageNet competition rules at kaggle.com/c/imagenet-object-localization-challenge')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:43:18.066441Z","iopub.execute_input":"2026-04-29T10:43:18.067368Z","iopub.status.idle":"2026-04-29T10:49:05.060312Z","shell.execute_reply.started":"2026-04-29T10:43:18.067335Z","shell.execute_reply":"2026-04-29T10:49:05.059294Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 6: Configure & Patch","metadata":{}},{"cell_type":"code","source":"import json\n\nos.chdir('/kaggle/working/synapse/src')\nconfig_path = 'config/Test_LDM.json'\n\nwith open(config_path, 'r') as f:\n    config = json.load(f)\n\n# Point all paths to Kaggle locations\nconfig['eeg_train_path']    = '/kaggle/working/preprocessing_data/train'\nconfig['eeg_val_path']      = '/kaggle/working/preprocessing_data/val'\nconfig['eeg_test_path']     = '/kaggle/working/preprocessing_data/test'\nconfig['img_path']          = '/kaggle/working/imagenet_images'\nconfig['eeg_pretrian_path'] = '/kaggle/working/pretrain_models/pretrain_models/EEGEncoder_BLIP2_69000.pth'\nconfig['sd_path']           = '/kaggle/input/datasets/jul503/stable-diffusion-2-1'\nconfig['blip2_path']        = '/kaggle/input/datasets/jul503/blip2-opt-2-7b/blip2-opt-2-7b'\nconfig['ldm_pretrain_path'] = '/kaggle/input/datasets/heyitsrup/stable-diffusion-checkpoints/checkpoint-23000'\n\n# Performance settings\nconfig['num_workers'] = 2      # Kaggle works best with 2 workers\nconfig['ddim_steps']  = 50     # 50 steps: ~5x faster than 250, still good quality\n\nwith open(config_path, 'w') as f:\n    json.dump(config, f, indent=2)\n\nprint('Config updated')\nprint(json.dumps({k: v for k, v in config.items() if 'path' in k or k in ['ddim_steps', 'num_workers', 'cfg_scale']}, indent=2))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:50:12.287231Z","iopub.execute_input":"2026-04-29T10:50:12.287948Z","iopub.status.idle":"2026-04-29T10:50:12.300654Z","shell.execute_reply.started":"2026-04-29T10:50:12.287867Z","shell.execute_reply":"2026-04-29T10:50:12.299846Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Patch eeg_LDM2.py: replace hardcoded HuggingFace ID with local SD path\nwith open('models/eeg_LDM2.py', 'r') as f:\n    content = f.read()\n\ncontent = content.replace(\n    '\"stabilityai/stable-diffusion-2-1\"',\n    '\"/kaggle/input/datasets/jul503/stable-diffusion-2-1\"'\n)\n\nwith open('models/eeg_LDM2.py', 'w') as f:\n    f.write(content)\n\nprint('eeg_LDM2.py patched — SD path now points to local dataset')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:51:02.193605Z","iopub.execute_input":"2026-04-29T10:51:02.194513Z","iopub.status.idle":"2026-04-29T10:51:02.205852Z","shell.execute_reply.started":"2026-04-29T10:51:02.19448Z","shell.execute_reply":"2026-04-29T10:51:02.205224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Patch trainer: add MAX_BATCHES env var support to limit test run size\nwith open('trainer/eeg_ldm2_ddp_trainer.py', 'r') as f:\n    content = f.read()\n\nif 'MAX_BATCHES' not in content:\n    content = content.replace(\n        '        for idx, val_data in enumerate(self.loader_test):\\n            img = val_data[\"image\"].to(self.device)',\n        '        max_batches = int(os.environ.get(\"MAX_BATCHES\", len(self.loader_test)))\\n        for idx, val_data in enumerate(self.loader_test):\\n            if idx >= max_batches:\\n                break\\n            img = val_data[\"image\"].to(self.device)'\n    )\n    with open('trainer/eeg_ldm2_ddp_trainer.py', 'w') as f:\n        f.write(content)\n    print('Trainer patched — MAX_BATCHES support added')\nelse:\n    print('Trainer already patched')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:51:04.249534Z","iopub.execute_input":"2026-04-29T10:51:04.250383Z","iopub.status.idle":"2026-04-29T10:51:04.258313Z","shell.execute_reply.started":"2026-04-29T10:51:04.250352Z","shell.execute_reply":"2026-04-29T10:51:04.257428Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 7: Generate Images\n\nSet `MAX_BATCHES` to control how many EEG samples to generate images for:\n- `20` = 20 samples × 5 images = 100 images (~15-20 min, good for testing)\n- `1199` = full test set (~1.5 hours with 50 DDIM steps)\n\nEach batch produces 6 images: 1 ground truth + 5 SYNAPSE-generated.","metadata":{}},{"cell_type":"code","source":"# Verify GPU is available\nimport torch\nprint('CUDA available:', torch.cuda.is_available())\nprint('GPU:', torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'None')\nprint('CWD:', os.getcwd())\nassert os.path.exists('gen_images.py'), 'gen_images.py not found'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:51:08.338663Z","iopub.execute_input":"2026-04-29T10:51:08.339235Z","iopub.status.idle":"2026-04-29T10:51:08.660916Z","shell.execute_reply.started":"2026-04-29T10:51:08.339205Z","shell.execute_reply":"2026-04-29T10:51:08.660184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run generation\n# Change MAX_BATCHES=20 to MAX_BATCHES=1199 for the full test set\n!MAX_BATCHES=20 python gen_images.py \\\n  --config ./config/Test_LDM.json \\\n  --pretrained_path /kaggle/input/datasets/heyitsrup/stable-diffusion-checkpoints/checkpoint-23000 \\\n  --output_path /kaggle/working/output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:51:11.101215Z","iopub.execute_input":"2026-04-29T10:51:11.101643Z","iopub.status.idle":"2026-04-29T13:49:48.759209Z","shell.execute_reply.started":"2026-04-29T10:51:11.10161Z","shell.execute_reply":"2026-04-29T13:49:48.75835Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 8: View Results","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom PIL import Image\nimport glob\n\nimage_files = sorted(glob.glob('/kaggle/working/output/test/*.png'))\nprint(f'Generated {len(image_files)} images')\n\n# Display in groups of 6 (1 ground truth + 5 generated per EEG sample)\nnum_rows = min(4, len(image_files) // 6)\nfig, axes = plt.subplots(num_rows, 6, figsize=(24, num_rows * 4))\n\nfor i, ax in enumerate(axes.flatten()):\n    if i < len(image_files):\n        img = Image.open(image_files[i])\n        ax.imshow(img)\n        label = os.path.basename(image_files[i])[:12]\n        # Mark ground truth (index ending in -0)\n        if image_files[i].endswith('-0.png'):\n            ax.set_title(f'GT: {label}', fontsize=7, color='green', fontweight='bold')\n        else:\n            ax.set_title(label, fontsize=7)\n    ax.axis('off')\n\nplt.suptitle('SYNAPSE EEG-to-Image Results\\nGreen = Ground Truth | Others = Generated from EEG', fontsize=12)\nplt.tight_layout()\nplt.savefig('/kaggle/working/synapse_results.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint('Saved to /kaggle/working/synapse_results.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T14:01:12.649923Z","iopub.execute_input":"2026-04-29T14:01:12.650478Z","iopub.status.idle":"2026-04-29T14:01:21.601864Z","shell.execute_reply.started":"2026-04-29T14:01:12.650447Z","shell.execute_reply":"2026-04-29T14:01:21.600963Z"}},"outputs":[],"execution_count":null}]}