OSCR

Concept2Brain: an AI model for predicting neurophysiological responses to text and pictures.

Code ↔ Paper

8 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 8 matches
  1. [1] § Methods › The conditioned variational autoencoder model for compressing and predicting electrophysiological data ↔ Concept2Brain_image_datarecovery.py, lines 24–139 · score 0.86 · reparameterization trick, bottleneck layer, standard deviation, latent vector, variational, decoded
  2. [2] § Methods › Training of the cVAE model to abstract electrophysiological features ↔ Concept2Brain_training_script.py, lines 342–374 · score 0.80 · Kullback Leibler, Squared Error, reconstruction loss, divergence, MSE, VAE
  3. [3] § Methods › Training of the cVAE model to abstract electrophysiological features ↔ Concept2Brain_training_script.py, lines 213–340 · score 0.79 · convolutional layers, RELU, kernel, padding, stride, flattened
  4. [4] § Methods › The conditioned variational autoencoder model for compressing and predicting electrophysiological data ↔ Concept2Brain_training_script.py, lines 213–340 · score 0.79 · reparameterization trick, bottleneck layer, latent vector, variational, decoded, encodes
  5. [5] § Methods › Training of the cVAE model to abstract electrophysiological features ↔ Concept2Brain_image_datarecovery.py, lines 24–139 · score 0.69 · RELU, kernel, padding, stride, flattened, concatenated
  6. [6] § Methods › Training of the cVAE model to abstract electrophysiological features ↔ Concept2Brain_training_script.py, lines 342–374 · score 0.67 · KL loss, reconstruction loss, summed, annealing, sigmoid, MSE
  7. [7] § Methods › Inference of the conceptual space and training of the cross-domain network ↔ Concept2Brain_training_script.py, lines 877–915 · score 0.53 · neural network, ELU, latent space, activation, matching, linear
  8. [8] § Methods › The cross-domain bridge between the concept-level and electrophysiological latent spaces ↔ Concept2Brain_training_script.py, lines 877–915 · score 0.51 · latent vector, neural network, linear, layer, mapping, space

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,260 lines · 56 KB · GPL-3.0 · 6 matches

  1. # -*- coding: utf-8 -*-
  2. """
  3. Concept2Brain script
  4. Code for processing EEG data, extracting image features with CLIP,
  5. training a Conditional Variational Autoencoder (CVAE) on EEG data,
  6. and training a neural network to map CLIP features to the CVAE's latent space.
  7. """
  8. # =============================================================================
  9. # Section 1: Imports and Helper Functions
  10. # =============================================================================
  11. # Description: Imports necessary libraries and defines utility functions.
  12. import matplotlib.pyplot as plt
  13. import torch
  14. import torch.nn as nn
  15. import torch.optim as optim
  16. from torch.utils.data import DataLoader, TensorDataset
  17. import clip
  18. from PIL import Image
  19. import numpy as np
  20. import pandas as pd
  21. import random
  22. import scipy.io as sio
  23. from scipy.stats import shapiro
  24. import gc # Garbage Collector to explicitly free memory if needed
  25. import os # To check if files exist (e.g., for loading best model)
  26. # Function to set a seed for reproducibility
  27. def set_seed(seed):
  28. """Sets the seed for PyTorch (CPU and CUDA), NumPy, and random."""
  29. torch.manual_seed(seed)
  30. if torch.cuda.is_available():
  31. torch.cuda.manual_seed(seed)
  32. torch.cuda.manual_seed_all(seed) # if using multi-GPU.
  33. np.random.seed(seed)
  34. random.seed(seed)
  35. # These settings help reproducibility on CUDA
  36. torch.backends.cudnn.benchmark = False
  37. torch.backends.cudnn.deterministic = True
  38. # =============================================================================
  39. # Section 2: Initial Configuration and Environment Setup
  40. # =============================================================================
  41. # Description: Sets the reproducibility seed and configures the device (CPU/GPU).
  42. # Set seed for reproducibility
  43. set_seed(42)
  44. # Configure device to use GPU if available, otherwise CPU
  45. device = "cuda" if torch.cuda.is_available() else "cpu"
  46. print(f"Using device: {device}")
  47. # =============================================================================
  48. # Section 3: Image Feature Extraction with CLIP
  49. # =============================================================================
  50. # Description: Loads the CLIP model, processes images from a specific directory,
  51. # extracts their features (embeddings), and saves them to a .npy file.
  52. print("\n--- Starting CLIP feature extraction ---")
  53. # Load Excel file with image information
  54. # Ensure this path is correct for your system
  55. file_path = '../PictureLabels.xlsx'
  56. excel = pd.read_excel(file_path)
  57. # Load the CLIP model and image preprocessing function
  58. # Using ViT-B/32 model and moving it to the configured device (GPU/CPU)
  59. model_clip, preprocess_clip = clip.load("ViT-B/32", device=device)
  60. # Path to the folder containing the images
  61. # Ensure this path is correct for your system
  62. folder_path = '../Pictures_Version12'
  63. # Initialize array to store features (360 images, 512 features per image)
  64. image_features = np.zeros((360, 512))
  65. # Iterate over the images specified in the Excel file
  66. for i in range(360):
  67. pic_filename = f"{excel.iloc[i, 1]}.jpg"
  68. pic2analyze = f"{folder_path}/{pic_filename}"
  69. # print(f"Processing: {pic2analyze}") # Uncomment if debugging needed
  70. try:
  71. # Open, preprocess the image, and move it to the device
  72. image = preprocess_clip(Image.open(pic2analyze)).unsqueeze(0).to(device)
  73. # Extract image features using CLIP model without gradient calculation
  74. with torch.no_grad():
  75. features = model_clip.encode_image(image)
  76. image_features[i, :] = features.cpu().numpy() # Move features to CPU and store
  77. except FileNotFoundError:
  78. print(f"Warning: File not found {pic2analyze}. Skipping.")
  79. except Exception as e:
  80. print(f"Error processing {pic2analyze}: {e}")
  81. # Save the extracted features to a NumPy file
  82. np.save('image_features.npy', image_features)
  83. print("Image features saved to image_features.npy")
  84. # Free CLIP model memory if possible (optional)
  85. del model_clip
  86. del preprocess_clip
  87. if torch.cuda.is_available():
  88. torch.cuda.empty_cache()
  89. # =============================================================================
  90. # Section 4: EEG Data Loading and Preprocessing
  91. # =============================================================================
  92. # Description: Loads EEG data, labels, and subject information from .mat files.
  93. # Performs normalization on the EEG data.
  94. print("\n--- Loading and preprocessing EEG data ---")
  95. # Base path where scripts and data are located
  96. # Ensure this path is correct for your system
  97. path = './' # Example path, adjust as needed
  98. # Load Clean voltage EEG data (raw) from .mat files and concatenate them
  99. # All trials (> 10k) were divided for python loading reasons into 3 parts
  100. # Ensure these file paths are correct
  101. mat_contents = sio.loadmat(path + 'DATA_EEG_PARTS.ma.mat')
  102. DATA_EEG = np.concatenate((
  103. np.asarray(mat_contents['DATA_RAW_ALL_PART1']),
  104. np.asarray(mat_contents['DATA_RAW_ALL_PART2']),
  105. np.asarray(mat_contents['DATA_RAW_ALL_PART3'])
  106. ), axis=2)
  107. print(f"Shape of loaded EEG data: {DATA_EEG.shape}") # E.g., (Electrodes, Timepoints, Trials)
  108. # Load image labels (numbers corresponding to rows in Excel/CLIP features)
  109. # Ensure this file path is correct
  110. tmp = sio.loadmat(path + 'LABEL_ALL_MyAPS12.mat')
  111. LABELS = tmp['PICTURE_XCEL_NUM'].T
  112. print(f"Shape of loaded labels: {LABELS.shape}") # E.g., (Trials, 1)
  113. # Load subject identifiers
  114. # Ensure this file path is correct
  115. tmp = sio.loadmat(path + 'SUBJECTS_ALL_MyAPS12.mat')
  116. SUBJECTS = tmp['SUBJECT_NUM'].T
  117. SUBJECTS_IDS = np.unique(SUBJECTS)
  118. print(f"Shape of loaded subject IDs: {SUBJECTS.shape}") # E.g., (Trials, 1)
  119. print(f"Number of unique subjects: {len(SUBJECTS_IDS)}")
  120. # Create condition vector (one-hot encoding for subjects)
  121. ConditionedVector = np.zeros((len(SUBJECTS), len(SUBJECTS_IDS)))
  122. for i in range(len(SUBJECTS)):
  123. ConditionedVector[i, SUBJECTS_IDS == SUBJECTS[i]] = 1
  124. print(f"Shape of condition vector (subjects one-hot): {ConditionedVector.shape}") # E.g., (Trials, NumSubjects)
  125. # Z-score normalization of EEG data based on the baseline (first 50 timepoints)
  126. print("Normalizing EEG data...")
  127. # Calculate mean and standard deviation of the baseline period.
  128. # It corresponds with first 50 time points, i.e 100 ms, for each channel and trial
  129. meanData = DATA_EEG[:, :50, :].mean(axis=1, keepdims=True)
  130. stdData = DATA_EEG[:, :50, :].std(axis=1, keepdims=True)
  131. # Avoid division by zero if standard deviation is very small
  132. stdData[stdData < 1e-6] = 1e-6
  133. # Apply Z-score normalization
  134. DATA_EEG_NORM = (DATA_EEG - meanData) / stdData
  135. print("Normalization completed.")
  136. # Free up memory
  137. del mat_contents, tmp, meanData, stdData, DATA_EEG
  138. gc.collect()
  139. # =============================================================================
  140. # Section 5: Data Splitting into Training and Testing Sets
  141. # =============================================================================
  142. # Description: Randomly permutes the data and splits it into training (80%)
  143. # and testing (20%) sets.
  144. print("\n--- Splitting data into training and testing sets ---")
  145. # Generate randomly permuted indices
  146. num_trials = DATA_EEG_NORM.shape[2]
  147. newInds = np.random.permutation(num_trials)
  148. # Calculate the split index for 80/20 division
  149. th_traintest = int(num_trials * 0.8)
  150. # Split the normalized EEG data
  151. # Shape: (Electrodes, Timepoints, Trials) -> (E, T, N_train/N_test)
  152. X_train = DATA_EEG_NORM[:, :, newInds[:th_traintest]]
  153. X_test = DATA_EEG_NORM[:, :, newInds[th_traintest:]]
  154. # Split labels, condition vectors, and subject IDs
  155. # Shape: (Trials, Features) -> (N_train/N_test, Features)
  156. y_train = LABELS[newInds[:th_traintest]]
  157. y_test = LABELS[newInds[th_traintest:]]
  158. c_train = ConditionedVector[newInds[:th_traintest]]
  159. c_test = ConditionedVector[newInds[th_traintest:]]
  160. s_train = SUBJECTS[newInds[:th_traintest]]
  161. s_test = SUBJECTS[newInds[th_traintest:]]
  162. print(f"Training set size: {X_train.shape[2]} trials")
  163. print(f"Testing set size: {X_test.shape[2]} trials")
  164. # Free up memory
  165. del DATA_EEG_NORM, LABELS, ConditionedVector, SUBJECTS, newInds
  166. gc.collect()
  167. # =============================================================================
  168. # Section 6: Conditional Variational Autoencoder (CVAE) Definition
  169. # =============================================================================
  170. # Description: Defines the CVAE architecture using PyTorch, including the
  171. # encoder, decoder, and loss function. The model is designed for multi-GPU use.
  172. print("\n--- Defining the CVAE model ---")
  173. # Function to free GPU memory (if needed)
  174. def free_gpu_memory():
  175. """Clears the GPU memory cache."""
  176. if torch.cuda.is_available():
  177. # print("Clearing GPU cache...") # Can be verbose, uncomment if needed
  178. torch.cuda.empty_cache()
  179. gc.collect()
  180. # Definition of the conditional CVAE structure
  181. class ConditionalVAE(nn.Module):
  182. """
  183. Conditional Variational Autoencoder (CVAE) with convolutional layers.
  184. Designed to operate across multiple GPUs:
  185. - Encoder on GPU 0
  186. - Bottleneck layers (mu, logvar) on GPU 1
  187. - Decoder on GPU 2
  188. Adjusts to the main device if fewer than 3 GPUs are available.
  189. """
  190. def __init__(self, input_channels, condition_dim, latent_dim):
  191. super(ConditionalVAE, self).__init__()
  192. self.condition_dim = condition_dim
  193. self.latent_dim = latent_dim
  194. # Check availability of sufficient GPUs
  195. num_gpus = torch.cuda.device_count()
  196. if num_gpus < 3:
  197. print(f"Warning: Found {num_gpus} GPUs, but the model is designed for 3. Adjusting to the main device: {device}")
  198. self.encoder_device = torch.device(device)
  199. self.bottleneck_device = torch.device(device)
  200. self.decoder_device = torch.device(device)
  201. else:
  202. self.encoder_device = torch.device("cuda:0")
  203. self.bottleneck_device = torch.device("cuda:1")
  204. self.decoder_device = torch.device("cuda:2")
  205. # --- Encoder (GPU 0) ---
  206. # Input shape: (N, input_channels, 129, 750) [N, C, H, W]
  207. self.conv1 = nn.Conv2d(input_channels, 32, kernel_size=3, stride=2, padding=1).to(self.encoder_device) # Output: (N, 32, 65, 375)
  208. self.conv2 = nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1).to(self.encoder_device) # Output: (N, 64, 33, 188)
  209. self.conv3 = nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1).to(self.encoder_device) # Output: (N, 128, 17, 94)
  210. # Flattened size after conv3: 128 * 17 * 94 = 204224
  211. # --- Bottleneck Layers (GPU 1) ---
  212. self.flattened_size = 128 * 17 * 94
  213. self.fc_mu = nn.Linear(self.flattened_size + condition_dim, latent_dim).to(self.bottleneck_device)
  214. self.fc_logvar = nn.Linear(self.flattened_size + condition_dim, latent_dim).to(self.bottleneck_device)
  215. # --- Decoder (GPU 2) ---
  216. self.fc_decode = nn.Linear(latent_dim + condition_dim, self.flattened_size).to(self.decoder_device)
  217. # Input to deconv1: (N, 128, 17, 94)
  218. self.deconv1 = nn.ConvTranspose2d(128, 64, kernel_size=3, stride=2, padding=1, output_padding=1).to(self.decoder_device) # Output: (N, 64, 33, 188) approx
  219. self.deconv2 = nn.ConvTranspose2d(64, 32, kernel_size=3, stride=2, padding=1, output_padding=1).to(self.decoder_device) # Output: (N, 32, 65, 375) approx
  220. self.deconv3 = nn.ConvTranspose2d(32, input_channels, kernel_size=3, stride=2, padding=1, output_padding=1).to(self.decoder_device) # Output: (N, C, 129, 750) approx
  221. # The output size might slightly differ due to padding/stride, will be cropped.
  222. def encode(self, x, c):
  223. """Encodes input x conditioned by c into mu and logvar."""
  224. x = x.to(self.encoder_device)
  225. # No need to move c here if it's already on the correct device before calling encode
  226. # c = c.to(self.encoder_device) # Will be moved in the main forward pass if needed
  227. h1 = torch.relu(self.conv1(x))
  228. h2 = torch.relu(self.conv2(h1))
  229. h3 = torch.relu(self.conv3(h2)) # Using relu here is also common
  230. # Flatten and move to the bottleneck GPU
  231. h3_flat = h3.view(h3.size(0), -1).to(self.bottleneck_device)
  232. # Move condition c to the bottleneck GPU as well
  233. c_bottleneck = c.to(self.bottleneck_device)
  234. # Concatenate the flattened representation with the condition
  235. combined_bottleneck = torch.cat([h3_flat, c_bottleneck], dim=1)
  236. # Calculate mu and logvar
  237. mu = self.fc_mu(combined_bottleneck)
  238. logvar = self.fc_logvar(combined_bottleneck)
  239. return mu, logvar
  240. def reparameterize(self, mu, logvar):
  241. """Performs the reparameterization trick."""
  242. # Ensure tensors are on the correct device (bottleneck)
  243. mu = mu.to(self.bottleneck_device)
  244. logvar = logvar.to(self.bottleneck_device)
  245. std = torch.exp(0.5 * logvar)
  246. eps = torch.randn_like(std) # Creates tensor on the same device as std
  247. # Optional clamping for numerical stability
  248. # std = torch.clamp(std, min=1e-6, max=10) # Example
  249. return mu + eps * std
  250. def decode(self, z, c):
  251. """Decodes latent vector z conditioned by c into reconstruction."""
  252. # Move z and c to the decoder GPU
  253. z = z.to(self.decoder_device)
  254. c_decoder = c.to(self.decoder_device)
  255. # Concatenate latent vector and condition
  256. combined_decoder_input = torch.cat([z, c_decoder], dim=1)
  257. # Pass through the dense layer and reshape
  258. h6 = torch.relu(self.fc_decode(combined_decoder_input))
  259. # Reshape to the expected shape for the first deconvolutional layer (N, C, H, W)
  260. h6_reshaped = h6.view(h6.size(0), 128, 17, 94) # Use post-conv3 dimensions
  261. # Pass through the deconvolutional layers
  262. h9 = torch.relu(self.deconv1(h6_reshaped))
  263. h10 = torch.relu(self.deconv2(h9))
  264. # Final layer might not have activation (depends on data domain, e.g., sigmoid for [0,1])
  265. x_recon_raw = self.deconv3(h10)
  266. # Crop the output to the original input size (129 electrodes, 750 time points)
  267. # Original input dimensions were (N, 1, 129, 750)
  268. original_height, original_width = 129, 750
  269. x_recon = x_recon_raw[:, :, :original_height, :original_width]
  270. return x_recon
  271. def forward(self, x, c):
  272. """Complete forward pass of the CVAE."""
  273. # Move initial condition c to the encoder device
  274. c_enc = c.to(self.encoder_device)
  275. x_enc = x.to(self.encoder_device)
  276. # Encode
  277. mu, logvar = self.encode(x_enc, c_enc)
  278. # Reparameterize (occurs on bottleneck_device)
  279. z = self.reparameterize(mu, logvar)
  280. # Decode (requires z and c on decoder_device)
  281. c_dec = c.to(self.decoder_device)
  282. x_recon = self.decode(z, c_dec)
  283. # Ensure mu and logvar are on the correct device for loss calculation (bottleneck device)
  284. mu_final = mu.to(self.bottleneck_device)
  285. logvar_final = logvar.to(self.bottleneck_device)
  286. return x_recon, mu_final, logvar_final
  287. # VAE Loss Function
  288. def vae_loss(x_recon, x, mu, logvar, epoch, beta_factor=0.00001, annealing_midpoint=100, annealing_steepness=0.1):
  289. """
  290. Calculates the VAE loss (Reconstruction + KL Divergence).
  291. Includes a beta factor with sigmoidal annealing for the KL divergence.
  292. Assumes inputs (mu, logvar) determine the device for calculation.
  293. """
  294. # Move all tensors to the same device for loss calculation
  295. # Assuming loss is calculated on the bottleneck GPU (where mu/logvar reside)
  296. loss_device = mu.device
  297. x_recon = x_recon.to(loss_device)
  298. x = x.to(loss_device)
  299. # mu and logvar should already be on loss_device
  300. # 1. Reconstruction Loss (Mean Squared Error)
  301. # Sum over pixels/timepoints, then average over the batch
  302. recon_loss = nn.MSELoss(reduction='sum')(x_recon, x) / x.size(0)
  303. # 2. Kullback-Leibler (KL) Divergence
  304. # KL divergence between the latent distribution q(z|x,c) and the prior p(z)=N(0,I)
  305. # Summed over latent dimensions, averaged over the batch
  306. kl_loss = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)
  307. kl_loss = torch.mean(kl_loss) # Average over the batch
  308. # 3. Beta Factor with Sigmoidal Annealing
  309. # Beta starts near 0 and grows towards beta_factor
  310. beta = beta_factor / (1 + np.exp(-annealing_steepness * (epoch - annealing_midpoint)))
  311. # 4. Total Loss
  312. total_loss = recon_loss + beta * kl_loss
  313. # Return total loss and KL component (for monitoring)
  314. return total_loss, kl_loss
  315. # =============================================================================
  316. # Section 7: CVAE Training Preparation
  317. # =============================================================================
  318. # Description: Converts data to PyTorch tensors, creates DataLoaders for batch
  319. # handling, instantiates the CVAE model, and defines the optimizer.
  320. print("\n--- Preparing CVAE training ---")
  321. # Convert NumPy data to PyTorch tensors
  322. # Permute dimensions to be (N, C, H, W) where N=trials, C=channels(1), H=electrodes, W=time
  323. # Add channel dimension (unsqueeze(1))
  324. X_train_torch = torch.tensor(X_train, dtype=torch.float32).permute(2, 0, 1).unsqueeze(1)
  325. X_test_torch = torch.tensor(X_test, dtype=torch.float32).permute(2, 0, 1).unsqueeze(1)
  326. c_train_torch = torch.tensor(c_train, dtype=torch.float32)
  327. c_test_torch = torch.tensor(c_test, dtype=torch.float32)
  328. print(f"Shape of X_train_torch: {X_train_torch.shape}") # E.g., (N_train, 1, E, T)
  329. print(f"Shape of c_train_torch: {c_train_torch.shape}") # E.g., (N_train, NumSubjects)
  330. # Create Datasets and DataLoaders
  331. batch_size_cvae = 64
  332. train_dataset = TensorDataset(X_train_torch, c_train_torch)
  333. # Use multiple workers and pin memory for potentially faster data loading
  334. train_loader = DataLoader(train_dataset, batch_size=batch_size_cvae, shuffle=True, num_workers=4, pin_memory=True)
  335. test_dataset = TensorDataset(X_test_torch, c_test_torch)
  336. test_loader = DataLoader(test_dataset, batch_size=batch_size_cvae, shuffle=False, num_workers=4, pin_memory=True)
  337. print(f"CVAE Batch size: {batch_size_cvae}")
  338. print(f"Number of training batches: {len(train_loader)}")
  339. print(f"Number of testing batches: {len(test_loader)}")
  340. # Instantiate the CVAE model
  341. input_channels = 1 # EEG is single-channel (in this context)
  342. latent_dim = 256 # Latent space dimension
  343. condition_dim = c_train.shape[1] # Number of subjects (dimension of one-hot vector)
  344. # Instantiate the model. Internal modules will be moved to designated GPUs.
  345. model_cvae = ConditionalVAE(input_channels, condition_dim, latent_dim)
  346. # The model object itself resides on the CPU, but its parameters are registered correctly across GPUs.
  347. # Define the optimizer
  348. # Adam is a common choice
  349. learning_rate_cvae = 1e-6 # Very low learning rate for precise tunin
  350. optimizer_cvae = optim.Adam(model_cvae.parameters(), lr=learning_rate_cvae)
  351. print(f"CVAE Optimizer: Adam, Learning Rate: {learning_rate_cvae}")
  352. # Lists to store loss history
  353. train_loss_all_cvae = []
  354. val_loss_all_cvae = []
  355. kl_loss_all_cvae = [] # To monitor KL divergence
  356. # =============================================================================
  357. # Section 8: CVAE Training Loop
  358. # =============================================================================
  359. # Description: Performs the training and validation cycle for the CVAE over a
  360. # defined number of epochs. Logs training progress to a file.
  361. print("\n--- Starting CVAE training ---")
  362. epochs_cvae = 1300 # Number of epochs for training the cVAE
  363. # Annealing parameters for VAE loss (consistent with original settings)
  364. beta_factor_cvae = 0.00001 # Final beta value for KL term weight
  365. annealing_midpoint_cvae = 100 # Epoch where beta reaches half its final value
  366. annealing_steepness_cvae = 0.1 # Controls how quickly beta increases
  367. # Load state if exists (e.g., if training was interrupted)
  368. # This would require saving/loading model state_dict, optimizer state_dict, and loss lists.
  369. # For simplicity, starting from scratch here as in the original script.
  370. start_epoch = 0 # Start from epoch 0
  371. # Open log file to append results
  372. log_filename = 'training_log_cvae.txt'
  373. with open(log_filename, 'a') as f:
  374. f.write("\n--- New CVAE Training Session ---\n")
  375. f.write(f"Epochs: {epochs_cvae}, Batch Size: {batch_size_cvae}, LR: {learning_rate_cvae}, Latent Dim: {latent_dim}\n")
  376. f.write(f"Beta Factor: {beta_factor_cvae}, Annealing Midpoint: {annealing_midpoint_cvae}, Steepness: {annealing_steepness_cvae}\n")
  377. for epoch in range(start_epoch, epochs_cvae):
  378. # --- Training Phase ---
  379. model_cvae.train() # Set model to training mode
  380. train_loss_epoch = 0
  381. kl_loss_epoch_train = 0
  382. for batch_idx, (data, condition) in enumerate(train_loader):
  383. # Data and conditions will be moved to the correct devices inside the model
  384. # Move initial input to the encoder device
  385. data = data.to(model_cvae.encoder_device, non_blocking=True)
  386. condition = condition.to(model_cvae.encoder_device, non_blocking=True)
  387. # Reset optimizer gradients
  388. optimizer_cvae.zero_grad()
  389. # Forward pass: get reconstruction, mu, and logvar
  390. x_recon, mu, logvar = model_cvae(data, condition)
  391. # Calculate loss (occurs on bottleneck_device)
  392. # Move original 'data' to that device for comparison
  393. data_loss_device = data.to(mu.device, non_blocking=True)
  394. loss, kl_component = vae_loss(x_recon, data_loss_device, mu, logvar, epoch,
  395. beta_factor=beta_factor_cvae,
  396. annealing_midpoint=annealing_midpoint_cvae,
  397. annealing_steepness=annealing_steepness_cvae)
  398. # Backward pass: compute gradients
  399. loss.backward()
  400. # Update model parameters
  401. optimizer_cvae.step()
  402. # Accumulate epoch losses
  403. train_loss_epoch += loss.item()
  404. kl_loss_epoch_train += kl_component.item()
  405. # Calculate average training loss for the epoch (per sample)
  406. train_loss_avg = train_loss_epoch / len(train_loader.dataset)
  407. kl_loss_avg_train = kl_loss_epoch_train / len(train_loader) # KL average per batch
  408. train_loss_all_cvae.append(train_loss_avg)
  409. # --- Validation Phase ---
  410. model_cvae.eval() # Set model to evaluation mode
  411. val_loss_epoch = 0
  412. kl_loss_epoch_val = 0
  413. mu_all_val = [] # To collect mu values for Shapiro test
  414. with torch.no_grad(): # Disable gradient calculations during validation
  415. for batch_idx, (data, condition) in enumerate(test_loader):
  416. data = data.to(model_cvae.encoder_device, non_blocking=True)
  417. condition = condition.to(model_cvae.encoder_device, non_blocking=True)
  418. x_recon, mu, logvar = model_cvae(data, condition)
  419. data_loss_device = data.to(mu.device, non_blocking=True)
  420. loss, kl_component = vae_loss(x_recon, data_loss_device, mu, logvar, epoch,
  421. beta_factor=beta_factor_cvae,
  422. annealing_midpoint=annealing_midpoint_cvae,
  423. annealing_steepness=annealing_steepness_cvae)
  424. val_loss_epoch += loss.item()
  425. kl_loss_epoch_val += kl_component.item()
  426. mu_all_val.append(mu.cpu().numpy()) # Store mu for Shapiro test (move to CPU)
  427. # Calculate average validation loss for the epoch (per sample)
  428. val_loss_avg = val_loss_epoch / len(test_loader.dataset)
  429. kl_loss_avg_val = kl_loss_epoch_val / len(test_loader) # KL average per batch
  430. val_loss_all_cvae.append(val_loss_avg)
  431. kl_loss_all_cvae.append(kl_loss_avg_val) # Store average validation KL
  432. # Shapiro-Wilk normality test on 'mu' values from the validation set
  433. # Concatenate mus from all batches and flatten for Shapiro test
  434. mu_val_np = np.concatenate(mu_all_val, axis=0).flatten()
  435. shapiro_stat, shapiro_p = -1, -1 # Default values
  436. if len(mu_val_np) >= 3: # Shapiro requires at least 3 samples
  437. try:
  438. shapiro_stat, shapiro_p = shapiro(mu_val_np)
  439. except ValueError:
  440. print(f"Warning: Could not compute Shapiro test for epoch {epoch+1} (possibly zero variance).")
  441. # Print progress and save to log
  442. log_message = (
  443. f"Epoch [{epoch + 1}/{epochs_cvae}], "
  444. f"Train Loss: {train_loss_avg:.4f}, Val Loss: {val_loss_avg:.4f}, "
  445. f"KL Val: {kl_loss_avg_val:.3f}, Shapiro p: {shapiro_p:.3g}, "
  446. # Use last batch mu/logvar for min/max/std stats (representative)
  447. f"Mu Val (min/max/std): {mu.min():.3f}/{mu.max():.3f}/{mu.std():.3f}, "
  448. f"LogVar Val (min/max/std): {logvar.min():.3f}/{logvar.max():.3f}/{logvar.std():.3f}"
  449. )
  450. print(log_message)
  451. with open(log_filename, 'a') as f:
  452. f.write(log_message + "\n")
  453. # Free GPU memory at the end of each epoch (optional, but can help)
  454. # free_gpu_memory() # Can slow down training if called too often
  455. print("CVAE training completed.")
  456. # Save the trained model (optional but recommended)
  457. # Save state_dict to CPU for easier loading in different environments
  458. model_cvae_save_path = "MODEL_CCVAE_AnnelingBeta00001_01-100sigmoid_lr1e-6_1300epc_dict.pth"
  459. # Ensure the model's parameters are moved to CPU before saving state_dict
  460. # Or save directly, but specify map_location='cpu' when loading if needed
  461. # torch.save(model_cvae.cpu().state_dict(), model_cvae_save_path)
  462. # print(f"CVAE model state dictionary saved to {model_cvae_save_path}")
  463. # Note: Multi-GPU models might require careful handling for saving/loading. Saving state_dict is generally safer.
  464. # =============================================================================
  465. # Section 9: CVAE Post-Training Visualization and Analysis
  466. # =============================================================================
  467. # Description: Plots loss curves, generates reconstruction examples, and performs
  468. # visual analysis to evaluate the trained CVAE's performance.
  469. print("\n--- CVAE Post-Training Analysis ---")
  470. # 1. Plot training and validation loss curves
  471. plt.figure(figsize=(10, 5))
  472. plt.plot(train_loss_all_cvae, label='Train Loss (Total)')
  473. plt.plot(val_loss_all_cvae, label='Validation Loss (Total)')
  474. # Could also plot KL loss if desired
  475. # plt.plot(kl_loss_all_cvae, label='Validation KL Loss (Avg per Batch)')
  476. plt.xlabel('Epoch')
  477. plt.ylabel('Loss (Per Sample)')
  478. plt.title('CVAE Loss Curves during Training')
  479. plt.legend()
  480. plt.grid(True)
  481. plt.savefig('cvae_loss_curve.png')
  482. # plt.show() # Uncomment if running interactively
  483. # 2. Generate and visualize a reconstruction example from the test set
  484. model_cvae.eval() # Ensure the model is in evaluation mode
  485. with torch.no_grad():
  486. # Take a sample from the test set
  487. idx_sample = 0 # Use the first sample
  488. x_sample = X_test_torch[idx_sample].unsqueeze(0) # Add batch dimension
  489. c_sample = c_test_torch[idx_sample].unsqueeze(0)
  490. # Move sample to the appropriate initial device (encoder)
  491. x_sample_dev = x_sample.to(model_cvae.encoder_device)
  492. c_sample_dev = c_sample.to(model_cvae.encoder_device)
  493. # Generate reconstruction
  494. x_recon_sample, mu_sample, logvar_sample = model_cvae(x_sample_dev, c_sample_dev)
  495. # Move results to CPU for visualization
  496. x_original_np = x_sample.squeeze().cpu().numpy() # Remove batch/channel, move to CPU
  497. x_reconstructed_np = x_recon_sample.squeeze().cpu().numpy() # Remove batch/channel, move to CPU
  498. mu_sample_np = mu_sample.cpu().numpy()
  499. logvar_sample_np = logvar_sample.cpu().numpy()
  500. # Visualize original vs. reconstructed (as image and time series)
  501. plt.figure(figsize=(15, 10))
  502. # Original image
  503. plt.subplot(2, 3, 1)
  504. plt.imshow(x_original_np, aspect='auto', cmap='viridis')
  505. plt.title(f'Original EEG (Test Sample {idx_sample})')
  506. plt.xlabel('Time (points)')
  507. plt.ylabel('Electrode')
  508. plt.colorbar()
  509. # Reconstructed image
  510. plt.subplot(2, 3, 2)
  511. plt.imshow(x_reconstructed_np, aspect='auto', cmap='viridis')
  512. plt.title('CVAE Reconstruction')
  513. plt.xlabel('Time (points)')
  514. plt.ylabel('Electrode')
  515. plt.colorbar()
  516. # Difference image
  517. plt.subplot(2, 3, 3)
  518. diff = x_original_np - x_reconstructed_np
  519. vmax = np.max(np.abs(diff))
  520. plt.imshow(diff, aspect='auto', cmap='coolwarm', vmin=-vmax, vmax=vmax)
  521. plt.title('Difference (Original - Reconstructed)')
  522. plt.xlabel('Time (points)')
  523. plt.ylabel('Electrode')
  524. plt.colorbar()
  525. # Original time series (all electrodes overlaid)
  526. plt.subplot(2, 3, 4)
  527. plt.plot(x_original_np.T) # Transpose so time is the x-axis
  528. plt.title('Original EEG (Time Series)')
  529. plt.xlabel('Time (points)')
  530. plt.ylabel('Normalized Amplitude')
  531. # Reconstructed time series
  532. plt.subplot(2, 3, 5)
  533. plt.plot(x_reconstructed_np.T)
  534. plt.title('CVAE Reconstruction (Time Series)')
  535. plt.xlabel('Time (points)')
  536. plt.ylabel('Normalized Amplitude')
  537. # Latent space (mu and logvar) for this sample
  538. plt.subplot(2, 3, 6)
  539. plt.plot(mu_sample_np.flatten(), label=f'Mu (std={mu_sample_np.std():.2f})')
  540. plt.plot(logvar_sample_np.flatten(), label=f'LogVar (std={logvar_sample_np.std():.2f})')
  541. plt.title('Latent Space (mu, logvar)')
  542. plt.xlabel('Latent Dimension')
  543. plt.ylabel('Value')
  544. plt.legend()
  545. plt.tight_layout()
  546. plt.savefig('cvae_reconstruction_sample.png')
  547. # plt.show()
  548. # 3. Analysis of Subject Effects (requires generating multiple reconstructions)
  549. # Original code generated reconstructions for the first 3 subject conditions (c_test[0], c_test[1], c_test[2])
  550. # applied to *all* X_test inputs. This checks if the model can impose a subject's "style".
  551. print("Generating reconstructions conditioned by different subjects...")
  552. model_cvae.eval()
  553. num_test_samples = X_test_torch.shape[0]
  554. num_subjects_to_test = min(3, c_test.shape[1]) # Test with the first 3 subjects
  555. reconstructions_by_subject = {} # Dictionary: {subject_idx: [reconstructions]}
  556. with torch.no_grad():
  557. for subj_idx in range(num_subjects_to_test):
  558. # Create condition vector for this subject (one-hot)
  559. c_subject = torch.zeros(1, condition_dim)
  560. c_subject[0, subj_idx] = 1
  561. # Move the *target* condition to the decoder device
  562. c_subject_target = c_subject.to(model_cvae.decoder_device)
  563. subject_recons = []
  564. # Iterate over a subset of test samples for efficiency
  565. samples_to_process = min(num_test_samples, 100) # Limit to 100 samples
  566. print(f" Processing {samples_to_process} samples for subject {subj_idx+1} condition...")
  567. for i in range(samples_to_process):
  568. x_input = X_test_torch[i].unsqueeze(0).to(model_cvae.encoder_device)
  569. # Use the *original* condition associated with x_input for encoding
  570. c_input_original = c_test_torch[i].unsqueeze(0).to(model_cvae.encoder_device)
  571. # Encode to get mu, logvar using the original input and condition
  572. mu, logvar = model_cvae.encode(x_input, c_input_original)
  573. # Reparameterize to get z (on bottleneck device)
  574. z = model_cvae.reparameterize(mu, logvar)
  575. # Decode using the obtained z but with the *new target* subject condition
  576. x_recon_subj = model_cvae.decode(z, c_subject_target)
  577. subject_recons.append(x_recon_subj.squeeze().cpu().numpy()) # Collect reconstructions
  578. # Store all reconstructions for this target subject condition
  579. reconstructions_by_subject[subj_idx] = np.stack(subject_recons, axis=-1) # Shape: (E, T, N_samples)
  580. # Visualize the average ERP for each simulated subject condition
  581. plt.figure(figsize=(15, 5))
  582. plt.suptitle('Average Simulated ERP Conditioned by Subject (Using Test Set Inputs)')
  583. for subj_idx in range(num_subjects_to_test):
  584. # Average across the samples processed for this condition
  585. erp_mean = reconstructions_by_subject[subj_idx].mean(axis=2)
  586. plt.subplot(1, num_subjects_to_test, subj_idx + 1)
  587. plt.plot(erp_mean.T) # Plot time series
  588. plt.title(f'Simulated Subject {subj_idx + 1}')
  589. plt.xlabel('Time (points)')
  590. if subj_idx == 0:
  591. plt.ylabel('Normalized Amplitude')
  592. plt.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust layout for suptitle
  593. plt.savefig('cvae_subject_conditioning_erp.png')
  594. # plt.show()
  595. # 4. Analysis of Emotional Condition Effects (requires mapping labels to conditions)
  596. print("Analyzing average reconstructions by emotional condition (Pleasant, Neutral, Unpleasant)...")
  597. # Define label ranges for each condition (based on original code comments)
  598. pics_pleasant = list(range(322, 362))
  599. pics_neutral = list(range(282, 322))
  600. pics_unpleasant = list(range(242, 282))
  601. # Need the *actual* reconstructions from the test set (not simulated ones)
  602. all_test_recons = []
  603. all_test_mu = [] # Also collect mu if needed later
  604. print(" Generating reconstructions for the entire test set...")
  605. with torch.no_grad():
  606. for data, condition in test_loader:
  607. data_dev = data.to(model_cvae.encoder_device)
  608. cond_dev = condition.to(model_cvae.encoder_device)
  609. x_recon, mu, _ = model_cvae(data_dev, cond_dev)
  610. all_test_recons.append(x_recon.cpu().numpy())
  611. all_test_mu.append(mu.cpu().numpy())
  612. all_test_recons_np = np.concatenate(all_test_recons, axis=0).squeeze() # Shape: (N_test, E, T)
  613. # y_test has shape (N_test, 1), need (N_test,) for isin
  614. y_test_flat = y_test.squeeze()
  615. # Find indices in the test set corresponding to each emotional condition
  616. indices_p = np.where(np.isin(y_test_flat, pics_pleasant))[0]
  617. indices_n = np.where(np.isin(y_test_flat, pics_neutral))[0]
  618. indices_u = np.where(np.isin(y_test_flat, pics_unpleasant))[0]
  619. # Calculate average ERP for each condition based on the reconstructions
  620. erp_p_recon = all_test_recons_np[indices_p].mean(axis=0) # Shape: (E, T)
  621. erp_n_recon = all_test_recons_np[indices_n].mean(axis=0)
  622. erp_u_recon = all_test_recons_np[indices_u].mean(axis=0)
  623. # Visualize average reconstructed ERPs by condition
  624. plt.figure(figsize=(18, 6))
  625. plt.suptitle('Average Reconstructed ERP by Emotional Condition')
  626. # Time series plot (all electrodes)
  627. plt.subplot(1, 2, 1)
  628. # Use different linestyles or colors if plotting all electrodes is too crowded
  629. plt.plot(erp_p_recon.T, label='Pleasant Recon', alpha=0.1, color='green')
  630. plt.plot(erp_n_recon.T, label='Neutral Recon', alpha=0.1, color='blue')
  631. plt.plot(erp_u_recon.T, label='Unpleasant Recon', alpha=0.1, color='red')
  632. # Create a custom legend for clarity
  633. from matplotlib.lines import Line2D
  634. custom_lines = [Line2D([0], [0], color='green', lw=2),
  635. Line2D([0], [0], color='blue', lw=2),
  636. Line2D([0], [0], color='red', lw=2)]
  637. plt.legend(custom_lines, ['Pleasant', 'Neutral', 'Unpleasant'])
  638. plt.title('All Reconstructed Signals')
  639. plt.xlabel('Time (points)')
  640. plt.ylabel('Normalized Amplitude')
  641. # Plot for a specific electrode (e.g., electrode 61 corresponds to Pz)
  642. plt.subplot(1, 2, 2)
  643. electrode_idx = 61 # Check if this index is valid (0 to E-1)
  644. if 0 <= electrode_idx < erp_p_recon.shape[0]:
  645. plt.plot(erp_p_recon[electrode_idx, :], label='Pleasant Recon', color='green')
  646. plt.plot(erp_n_recon[electrode_idx, :], label='Neutral Recon', color='blue')
  647. plt.plot(erp_u_recon[electrode_idx, :], label='Unpleasant Recon', color='red')
  648. plt.title(f'Reconstructed ERP - Electrode {electrode_idx+1}')
  649. plt.xlabel('Time (points)')
  650. plt.ylabel('Normalized Amplitude')
  651. plt.legend()
  652. plt.grid(True)
  653. else:
  654. plt.text(0.5, 0.5, f'Electrode index {electrode_idx} out of range [0, {erp_p_recon.shape[0]-1}]',
  655. horizontalalignment='center', verticalalignment='center')
  656. plt.tight_layout(rect=[0, 0.03, 1, 0.95])
  657. plt.savefig('cvae_condition_analysis_reconstructed.png')
  658. # plt.show()
  659. # Free up memory
  660. del X_train_torch, X_test_torch, c_train_torch, c_test_torch, train_loader, test_loader
  661. del all_test_recons, all_test_recons_np, all_test_mu
  662. if 'reconstructions_by_subject' in locals(): del reconstructions_by_subject
  663. free_gpu_memory()
  664. # =============================================================================
  665. # Section 10: Data Preparation for Latent Neural Network (CLIP -> CVAE)
  666. # =============================================================================
  667. # Description: Loads the saved CLIP features and creates corresponding datasets
  668. # to train a network mapping from CLIP's latent space to the CVAE's latent space.
  669. print("\n--- Preparing data for the Latent Neural Network (CLIP -> CVAE) ---")
  670. # Load previously saved CLIP image features
  671. try:
  672. # Ensure the path matches where features were saved
  673. image_features_loaded = np.load('image_features.npy')
  674. print(f"CLIP features loaded from image_features.npy, shape: {image_features_loaded.shape}")
  675. except FileNotFoundError:
  676. print("Error: File image_features.npy not found. Run Section 3 first.")
  677. exit() # Exit if features cannot be loaded
  678. # Create input datasets (CLIP features) for training and testing
  679. # Use y_train, y_test labels to find the correct CLIP features
  680. # Labels (y_train, y_test) are 1-based image numbers, subtract 1 for 0-based indexing
  681. CLIP_latent_train = image_features_loaded[y_train.squeeze() - 1]
  682. CLIP_latent_test = image_features_loaded[y_test.squeeze() - 1]
  683. print(f"Shape of CLIP_latent_train: {CLIP_latent_train.shape}") # (N_train, 512)
  684. print(f"Shape of CLIP_latent_test: {CLIP_latent_test.shape}") # (N_test, 512)
  685. # The 'targets' for this network are not directly the CVAE's mu/logvar.
  686. # Instead, the loss is calculated by comparing the *reconstructed EEG* generated
  687. # from the 'z' predicted by this network, with the original EEG.
  688. # Therefore, we need the original EEG data (X) and conditions (c) again.
  689. # Convert data to PyTorch tensors for the latent network
  690. input_clip_train = torch.tensor(CLIP_latent_train, dtype=torch.float32)
  691. input_clip_test = torch.tensor(CLIP_latent_test, dtype=torch.float32)
  692. # Need original EEG data (X) and conditions (c) again
  693. # Reload or reuse tensors if still in memory (reloading for clarity is safer)
  694. X_train_torch_latent = torch.tensor(X_train, dtype=torch.float32).permute(2, 0, 1).unsqueeze(1)
  695. X_test_torch_latent = torch.tensor(X_test, dtype=torch.float32).permute(2, 0, 1).unsqueeze(1)
  696. c_train_torch_latent = torch.tensor(c_train, dtype=torch.float32)
  697. c_test_torch_latent = torch.tensor(c_test, dtype=torch.float32)
  698. # Create DataLoaders for the latent network
  699. batch_size_latent = 512 # Based on original code
  700. # Dataset contains (CLIP_input, Original_EEG_output, Condition_input)
  701. latent_train_dataset = TensorDataset(input_clip_train, X_train_torch_latent, c_train_torch_latent)
  702. latent_train_loader = DataLoader(latent_train_dataset, batch_size=batch_size_latent, shuffle=True, num_workers=4, pin_memory=True)
  703. latent_test_dataset = TensorDataset(input_clip_test, X_test_torch_latent, c_test_torch_latent)
  704. latent_test_loader = DataLoader(latent_test_dataset, batch_size=batch_size_latent, shuffle=False, num_workers=4, pin_memory=True)
  705. print(f"DataLoader for LatentNN created. Batch size: {batch_size_latent}")
  706. print(f"Number of LatentNN training batches: {len(latent_train_loader)}")
  707. print(f"Number of LatentNN testing batches: {len(latent_test_loader)}")
  708. # Free memory (original X_train, X_test arrays no longer needed if tensors created)
  709. del X_train, X_test, y_train, c_train, s_train, y_test, c_test, s_test
  710. gc.collect()
  711. # =============================================================================
  712. # Section 11: Latent Neural Network (LatentNN) Definition
  713. # =============================================================================
  714. # Description: Defines the architecture of the neural network mapping CLIP
  715. # features and subject condition to the CVAE's latent space.
  716. print("\n--- Defining the Latent Neural Network (LatentNN) ---")
  717. # Network parameters (based on original code)
  718. input_features_clip = CLIP_latent_train.shape[1] # 512
  719. condition_features_subj = c_train_torch_latent.shape[1] # NumSubjects (e.g., 88)
  720. output_features_latent = latent_dim # 256 (CVAE's latent dimension)
  721. hidden_units_1 = 256
  722. hidden_units_2 = 512
  723. # Original code had commented-out layers (BatchNorm, Dropout), omitted here to match active version.
  724. class LatentNN(nn.Module):
  725. """
  726. Neural Network to map (CLIP features + Subject Condition) -> CVAE Latent Space (z).
  727. Designed to operate on a specific GPU (GPU 1 according to CVAE setup).
  728. """
  729. def __init__(self):
  730. super(LatentNN, self).__init__()
  731. # Device where this network will operate (GPU 1 if available, matching CVAE bottleneck)
  732. num_gpus = torch.cuda.device_count()
  733. if num_gpus >= 2:
  734. # Assumes GPU 1 corresponds to the CVAE bottleneck device
  735. self.device_latent = torch.device("cuda:1")
  736. else:
  737. print("Warning: Not enough GPUs for original assignment (LatentNN on GPU 1). Using main device.")
  738. self.device_latent = torch.device(device)
  739. print(f"LatentNN will operate on: {self.device_latent}")
  740. # Define layers and move them to the designated device
  741. self.fc1 = nn.Linear(input_features_clip + condition_features_subj, hidden_units_1).to(self.device_latent)
  742. self.fc2 = nn.Linear(hidden_units_1, hidden_units_2).to(self.device_latent)
  743. self.fc3 = nn.Linear(hidden_units_2, output_features_latent).to(self.device_latent)
  744. self.elu = nn.ELU(alpha=1.0) # ELU activation used in original code
  745. def forward(self, x_clip, c_subj):
  746. """Forward pass."""
  747. # Move inputs to this network's device
  748. x_clip = x_clip.to(self.device_latent)
  749. c_subj = c_subj.to(self.device_latent)
  750. # Concatenate CLIP features and subject condition
  751. xc_combined = torch.cat([x_clip, c_subj], dim=1)
  752. # Pass through layers
  753. hidden1 = self.elu(self.fc1(xc_combined))
  754. hidden2 = self.elu(self.fc2(hidden1))
  755. z_predicted = self.fc3(hidden2) # Linear output (predicting the latent vector z)
  756. return z_predicted
  757. # Instantiate the LatentNN model
  758. model_latent = LatentNN()
  759. # The model already moves its layers to the correct device in __init__
  760. # =============================================================================
  761. # Section 12: LatentNN Training Preparation
  762. # =============================================================================
  763. # Description: Defines the loss function (MSE on reconstructed EEG), the
  764. # optimizer, and sets the CVAE to evaluation mode (freezing its weights).
  765. print("\n--- Preparing LatentNN training ---")
  766. # Loss function: Mean Squared Error (MSE) between original EEG and EEG reconstructed
  767. # from the z predicted by LatentNN.
  768. criterion_latent = nn.MSELoss()
  769. # Optimizer for LatentNN
  770. learning_rate_latent = 1e-3 # Based on original code
  771. optimizer_latent = optim.Adam(model_latent.parameters(), lr=learning_rate_latent)
  772. print(f"LatentNN Optimizer: Adam, LR: {learning_rate_latent}")
  773. # Freeze CVAE weights: we don't want to train it further
  774. print("Freezing CVAE weights...")
  775. for param in model_cvae.parameters():
  776. param.requires_grad = False
  777. model_cvae.eval() # Set CVAE to evaluation mode (important if it uses Dropout/BatchNorm)
  778. # Lists to store LatentNN loss history
  779. train_loss_all_latent = []
  780. val_loss_all_latent = []
  781. # Parameters for Early Stopping (based on original code)
  782. patience = 30 # Number of epochs to wait for improvement before stopping
  783. epochs_no_improve = 0
  784. best_val_loss = float('inf')
  785. best_model_latent_path = 'MODEL_Latent_10epc_512batch_dict.pth' # File to save the best model
  786. # In original script a model fit in 10 epochs was the best found option.
  787. # =============================================================================
  788. # Section 13: LatentNN Training Loop
  789. # =============================================================================
  790. # Description: Performs the training and validation cycle for LatentNN, using
  791. # the frozen CVAE to decode the predicted 'z' and calculate loss on the EEG
  792. # reconstruction. Implements Early Stopping.
  793. print("\n--- Starting LatentNN training ---")
  794. epochs_latent = 1000 # Maximum number of epochs (may stop earlier due to Early Stopping)
  795. #Original script stopped at 10 epochs when evaluation data was not better.
  796. for epoch in range(epochs_latent):
  797. # --- Training Phase ---
  798. model_latent.train() # Set LatentNN to training mode
  799. epoch_loss_train = 0.0
  800. for batch_clip, batch_eeg_original, batch_condition in latent_train_loader:
  801. # batch_clip & batch_condition go to LatentNN device
  802. # batch_eeg_original goes to CVAE Decoder/Loss device
  803. # batch_condition also needed on Decoder device
  804. # Move condition to LatentNN device (and later to decoder device)
  805. batch_condition_latent = batch_condition.to(model_latent.device_latent, non_blocking=True)
  806. # Move clip features to LatentNN device
  807. batch_clip_latent = batch_clip.to(model_latent.device_latent, non_blocking=True)
  808. # Reset LatentNN gradients
  809. optimizer_latent.zero_grad()
  810. # 1. Predict z using LatentNN
  811. z_hat = model_latent(batch_clip_latent, batch_condition_latent)
  812. # 2. Decode z_hat using the frozen CVAE to get reconstructed EEG
  813. # Move predicted z_hat and condition to the CVAE decoder device
  814. z_hat_decoder = z_hat.to(model_cvae.decoder_device, non_blocking=True)
  815. batch_condition_decoder = batch_condition.to(model_cvae.decoder_device, non_blocking=True)
  816. X_hat_reconstructed = model_cvae.decode(z_hat_decoder, batch_condition_decoder)
  817. # 3. Calculate MSE loss between reconstruction and original EEG
  818. # Move original EEG to the device where loss is calculated (decoder device)
  819. batch_eeg_original_loss = batch_eeg_original.to(model_cvae.decoder_device, non_blocking=True)
  820. loss = criterion_latent(X_hat_reconstructed, batch_eeg_original_loss)
  821. # Backward pass (computes gradients only for LatentNN)
  822. loss.backward()
  823. # Update LatentNN weights
  824. optimizer_latent.step()
  825. epoch_loss_train += loss.item()
  826. # Calculate average training loss for the epoch (per sample)
  827. avg_epoch_loss_train = epoch_loss_train / len(latent_train_loader.dataset)
  828. train_loss_all_latent.append(avg_epoch_loss_train)
  829. # --- Validation Phase ---
  830. model_latent.eval() # Set LatentNN to evaluation mode
  831. epoch_loss_val = 0.0
  832. z_hat_last_batch = None # To print min/max/std stats
  833. with torch.no_grad():
  834. for batch_clip, batch_eeg_original, batch_condition in latent_test_loader:
  835. batch_condition_latent = batch_condition.to(model_latent.device_latent, non_blocking=True)
  836. batch_clip_latent = batch_clip.to(model_latent.device_latent, non_blocking=True)
  837. # 1. Predict z
  838. z_hat = model_latent(batch_clip_latent, batch_condition_latent)
  839. z_hat_last_batch = z_hat # Store for stats
  840. # 2. Decode z_hat
  841. z_hat_decoder = z_hat.to(model_cvae.decoder_device, non_blocking=True)
  842. batch_condition_decoder = batch_condition.to(model_cvae.decoder_device, non_blocking=True)
  843. X_hat_reconstructed = model_cvae.decode(z_hat_decoder, batch_condition_decoder)
  844. # 3. Calculate loss
  845. batch_eeg_original_loss = batch_eeg_original.to(model_cvae.decoder_device, non_blocking=True)
  846. loss = criterion_latent(X_hat_reconstructed, batch_eeg_original_loss)
  847. epoch_loss_val += loss.item()
  848. # Calculate average validation loss for the epoch (per sample)
  849. avg_epoch_loss_val = epoch_loss_val / len(latent_test_loader.dataset)
  850. val_loss_all_latent.append(avg_epoch_loss_val)
  851. # Print progress
  852. z_min, z_max, z_std = -1, -1, -1 # Default values
  853. if z_hat_last_batch is not None:
  854. z_min = z_hat_last_batch.min().item()
  855. z_max = z_hat_last_batch.max().item()
  856. z_std = z_hat_last_batch.std().item()
  857. print(f"Epoch {epoch+1}/{epochs_latent}, Train Loss: {avg_epoch_loss_train:.4f}, Val Loss: {avg_epoch_loss_val:.4f}, "
  858. f"z_hat Val (min/max/std): {z_min:.3f}/{z_max:.3f}/{z_std:.3f}")
  859. # Early Stopping logic
  860. if avg_epoch_loss_val < best_val_loss:
  861. best_val_loss = avg_epoch_loss_val
  862. epochs_no_improve = 0
  863. # Save the best model found so far
  864. torch.save(model_latent.state_dict(), best_model_latent_path)
  865. print(f" Validation Loss improved. Saving model to {best_model_latent_path}")
  866. else:
  867. epochs_no_improve += 1
  868. print(f" Validation Loss did not improve for {epochs_no_improve} epochs.")
  869. if epochs_no_improve >= patience:
  870. print(f"Early stopping triggered after epoch {epoch + 1}.")
  871. break
  872. # Optional memory clearing
  873. # free_gpu_memory()
  874. print("LatentNN training completed.")
  875. # Load the best model saved by Early Stopping
  876. if os.path.exists(best_model_latent_path):
  877. print(f"Loading best model state from {best_model_latent_path}")
  878. # Ensure map_location matches the device LatentNN is intended for
  879. model_latent.load_state_dict(torch.load(best_model_latent_path, map_location=model_latent.device_latent))
  880. else:
  881. print(f"Warning: Best model file {best_model_latent_path} not found. Using the model's last state.")
  882. # =============================================================================
  883. # Section 14: LatentNN Post-Training Visualization and Analysis
  884. # =============================================================================
  885. # Description: Plots LatentNN loss curves, generates EEG from test CLIP features,
  886. # and analyzes the results by emotional condition.
  887. print("\n--- LatentNN Post-Training Analysis ---")
  888. # 1. Plot LatentNN loss curves
  889. plt.figure(figsize=(10, 5))
  890. plt.plot(train_loss_all_latent, label='Train Loss (LatentNN)')
  891. plt.plot(val_loss_all_latent, label='Validation Loss (LatentNN)')
  892. plt.xlabel('Epoch')
  893. plt.ylabel('MSE Loss (on EEG reconstruction, per sample)')
  894. plt.title('LatentNN Loss Curves during Training')
  895. plt.legend()
  896. plt.grid(True)
  897. plt.yscale('log') # Loss might decrease significantly, log scale can be useful
  898. plt.savefig('latentnn_loss_curve.png')
  899. # plt.show()
  900. # 2. Generate EEG from test CLIP features using LatentNN + CVAE(decoder)
  901. print("Generating EEG from test CLIP features...")
  902. model_latent.eval()
  903. model_cvae.eval()
  904. # Process the test set (or a large subset as in original code)
  905. num_samples_to_generate = min(2000, len(latent_test_dataset)) # Limit to 2000
  906. generated_eegs = []
  907. generated_latents = [] # Store the predicted 'z' vectors
  908. # Use DataLoader or iterate manually. Manual iteration for clarity:
  909. # Need the CLIP inputs and conditions for the subset
  910. input_clip_test_subset = input_clip_test[:num_samples_to_generate]
  911. c_test_torch_subset = c_test_torch_latent[:num_samples_to_generate]
  912. # Also need the corresponding labels for analysis
  913. y_test_subset = y_test[:num_samples_to_generate] # Assuming y_test numpy array is still available
  914. print(f" Processing {num_samples_to_generate} test samples...")
  915. with torch.no_grad():
  916. # Process in batches if the subset is too large for GPU memory
  917. batch_size_inference = 256 # Adjust based on GPU memory
  918. for i in range(0, num_samples_to_generate, batch_size_inference):
  919. # Prepare batch inputs
  920. batch_clip = input_clip_test_subset[i:i+batch_size_inference].to(model_latent.device_latent)
  921. batch_cond = c_test_torch_subset[i:i+batch_size_inference].to(model_latent.device_latent)
  922. # 1. Predict z using LatentNN
  923. z_hat_batch = model_latent(batch_clip, batch_cond)
  924. generated_latents.append(z_hat_batch.cpu().numpy()) # Store predicted z
  925. # 2. Decode z using CVAE
  926. z_hat_dec = z_hat_batch.to(model_cvae.decoder_device)
  927. batch_cond_dec = batch_cond.to(model_cvae.decoder_device) # Condition also needed for decoder
  928. X_hat_batch = model_cvae.decode(z_hat_dec, batch_cond_dec)
  929. generated_eegs.append(X_hat_batch.cpu().numpy()) # Store generated EEG
  930. # Concatenate results from all batches
  931. generated_eegs_np = np.concatenate(generated_eegs, axis=0).squeeze() # Shape: (N_gen, E, T)
  932. generated_latents_np = np.concatenate(generated_latents, axis=0) # Shape: (N_gen, latent_dim)
  933. print(f"Shape of generated EEGs: {generated_eegs_np.shape}")
  934. print(f"Shape of generated latents (z_hat): {generated_latents_np.shape}")
  935. # 3. Visualize the generated latent space (z_hat) by emotional condition
  936. # Use the labels corresponding to the generated subset
  937. y_test_subset_flat = y_test_subset.squeeze()
  938. # Find indices within the subset for each condition
  939. indices_p_gen = np.where(np.isin(y_test_subset_flat, pics_pleasant))[0]
  940. indices_n_gen = np.where(np.isin(y_test_subset_flat, pics_neutral))[0]
  941. indices_u_gen = np.where(np.isin(y_test_subset_flat, pics_unpleasant))[0]
  942. # Group predicted latents by condition
  943. latents_p = generated_latents_np[indices_p_gen]
  944. latents_n = generated_latents_np[indices_n_gen]
  945. latents_u = generated_latents_np[indices_u_gen]
  946. # Concatenate for visualization (ensure order P, N, U for consistency)
  947. concatenated_latents = np.concatenate((latents_p, latents_n, latents_u), axis=0)
  948. plt.figure(figsize=(12, 6))
  949. plt.subplot(1, 2, 1)
  950. plt.imshow(concatenated_latents, aspect='auto', cmap='viridis')
  951. plt.title('Generated Latent Space Z_hat (P, N, U ordered)')
  952. plt.xlabel('Latent Dimension')
  953. plt.ylabel('Samples (Ordered Pleasant-Neutral-Unpleasant)')
  954. plt.colorbar()
  955. plt.subplot(1, 2, 2)
  956. plt.plot(latents_p.mean(axis=0), label='Pleasant (Mean Z_hat)')
  957. plt.plot(latents_n.mean(axis=0), label='Neutral (Mean Z_hat)')
  958. plt.plot(latents_u.mean(axis=0), label='Unpleasant (Mean Z_hat)')
  959. plt.title('Average Predicted Latent Vector Z_hat by Condition')
  960. plt.xlabel('Latent Dimension')
  961. plt.ylabel('Average Value')
  962. plt.legend()
  963. plt.grid(True)
  964. plt.tight_layout()
  965. plt.savefig('latentnn_latent_space_analysis.png')
  966. # plt.show()
  967. # 4. Analysis of the generated EEGs by emotional condition
  968. # Calculate average generated ERP for each condition
  969. erp_p_gen = generated_eegs_np[indices_p_gen].mean(axis=0) # Shape: (E, T)
  970. erp_n_gen = generated_eegs_np[indices_n_gen].mean(axis=0)
  971. erp_u_gen = generated_eegs_np[indices_u_gen].mean(axis=0)
  972. plt.figure(figsize=(18, 12))
  973. plt.suptitle('Average Generated ERP (CLIP -> LatentNN -> CVAE) by Emotional Condition')
  974. # Time series plot (all electrodes overlaid - might be messy)
  975. plt.subplot(2, 2, 1)
  976. plt.plot(erp_p_gen.T, alpha=0.1, color='green') # Pleasant generated
  977. plt.plot(erp_n_gen.T, alpha=0.1, color='blue') # Neutral generated
  978. plt.plot(erp_u_gen.T, alpha=0.1, color='red') # Unpleasant generated
  979. # Custom legend for clarity
  980. custom_lines_gen = [Line2D([0], [0], color='green', lw=2),
  981. Line2D([0], [0], color='blue', lw=2),
  982. Line2D([0], [0], color='red', lw=2)]
  983. plt.legend(custom_lines_gen, ['Pleasant (Gen)', 'Neutral (Gen)', 'Unpleasant (Gen)'])
  984. plt.title('Generated Signals (All Electrodes)')
  985. plt.xlabel('Time (points)')
  986. plt.ylabel('Normalized Amplitude')
  987. # Image plots (topographical maps over time)
  988. # Determine symmetric color limits based on max absolute amplitude
  989. clim_max_gen = np.max(np.abs(np.stack([erp_p_gen, erp_n_gen, erp_u_gen])))
  990. clim_gen = (-clim_max_gen, clim_max_gen)
  991. plt.subplot(2, 2, 2)
  992. plt.imshow(erp_p_gen, aspect='auto', cmap='viridis', vmin=clim_gen[0], vmax=clim_gen[1])
  993. plt.title('Pleasant (Generated)')
  994. plt.xlabel('Time (points)'), plt.ylabel('Electrode'), plt.colorbar()
  995. plt.subplot(2, 2, 3)
  996. plt.imshow(erp_n_gen, aspect='auto', cmap='viridis', vmin=clim_gen[0], vmax=clim_gen[1])
  997. plt.title('Neutral (Generated)')
  998. plt.xlabel('Time (points)'), plt.ylabel('Electrode'), plt.colorbar()
  999. plt.subplot(2, 2, 4)
  1000. plt.imshow(erp_u_gen, aspect='auto', cmap='viridis', vmin=clim_gen[0], vmax=clim_gen[1])
  1001. plt.title('Unpleasant (Generated)')
  1002. plt.xlabel('Time (points)'), plt.ylabel('Electrode'), plt.colorbar()
  1003. plt.tight_layout(rect=[0, 0.03, 1, 0.95]) # Adjust layout for suptitle
  1004. plt.savefig('latentnn_condition_analysis_generated_eeg_maps.png')
  1005. # plt.show()
  1006. # Plot for the specific electrode (61) as in the original code
  1007. plt.figure(figsize=(8, 5))
  1008. electrode_idx = 61 # Ensure this is a valid index (0 to E-1)
  1009. if 0 <= electrode_idx < erp_p_gen.shape[0]:
  1010. plt.plot(erp_p_gen[electrode_idx, :], label='Pleasant (Generated)', color='green')
  1011. plt.plot(erp_n_gen[electrode_idx, :], label='Neutral (Generated)', color='blue')
  1012. plt.plot(erp_u_gen[electrode_idx, :], label='Unpleasant (Generated)', color='red')
  1013. plt.title(f'Generated EEG - Electrode {electrode_idx+1}')
  1014. plt.xlabel('Time (points)')
  1015. plt.ylabel('Normalized Amplitude')
  1016. plt.legend()
  1017. plt.grid(True)
  1018. else:
  1019. plt.text(0.5, 0.5, f'Electrode index {electrode_idx} out of range [0, {erp_p_gen.shape[0]-1}]',
  1020. horizontalalignment='center', verticalalignment='center')
  1021. plt.savefig('latentnn_condition_analysis_generated_eeg_electrode_61.png')
  1022. # plt.show()
  1023. print("\n--- Analysis completed ---")
  1024. # Optional: Explicitly clear memory at the end
  1025. del model_cvae, model_latent
  1026. del latent_train_dataset, latent_test_dataset, latent_train_loader, latent_test_loader
  1027. del generated_eegs_np, generated_latents_np
  1028. free_gpu_memory()

Concept2Brain_training_script.py at commit c7c8ae2, under GPL-3.0 · at the source

Overview

Authors: Alejandro Santos-Mayo1,2,3, Faith Gilbert1,2, Arash Mirifar1,2, Anna-Lena Tebbe1,2, Ruogu Fang4, Mingzhou Ding4, Andreas Keil1,2
  1. Laboratory of Brain, Body, and Behavior, University of Florida, Gainesville, FL USA
  2. Department of Psychology, University of Florida, Gainesville, FL USA
  3. Center for Cognitive and Computational Neuroscience, Complutense University of Madrid, Madrid, Spain
  4. J. Crayton Pruitt Family Department of Biomedical Engineering, University of Florida, Gainesville, FL USA
Institutions: Universidad Complutense de Madrid (Spain); University of Florida (United States)
Journal: Nature communications, volume 17, issue 1, article 8961
Dates: received 16 September 2025; accepted 7 July 2026; published online 23 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41467-026-75653-x · PMID 42637720 · PMCID PMC13503920 · OpenAlex W7170083615
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism), cognitive (subfield)
Methods: Spectral & time-frequency, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Evoked potentials, Connectivity
Keywords: Perception, Neural encoding
MeSH: Artificial Intelligence*, Brain*, Models, Neurological*, Electroencephalography, Humans, Semantics (* major topic)
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: NIMH NIH HHS (R01 MH112558)
Citations: not cited yet (Europe PMC); 88 references in the paper

Abstract

Evolving methods rooted in artificial intelligence (AI) offer new opportunities for linking human behavior and experience to brain function. Here, we introduce the Concept2Brain model, a deep network designed to generate synthetic electrophysiological responses to semantic/emotional information conveyed through pictures or text. Leveraging AI solutions like CLIP from OpenAI, the model generates a representation of a stimulus and maps it into an electrophysiological latent space. We demonstrate that this openly available resource generates synthetic neural responses that closely resemble those observed empirically. The Concept2Brain model is provided as a web service tool for creating open and reproducible EEG datasets by predicting brain responses to any semantic concept or picture. Beyond its practical applications, it also paves the way for AI-driven brain activity modeling, offering new possibilities for studying how the brain represents the world.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 8 matches between paragraphs and lines of code.

ASantosMayo/Concept2Brain

License: GPL-3.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: c7c8ae2e82d7b21275b35718efa76dc9f72fd347, 5 May 2025
Languages: Python (3)
Size: 5 files, 3 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (3 files), pandas (3 files), Pillow (3 files), PyTorch (3 files), SciPy (3 files), Matplotlib (2 files), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
5 files

Code availability

The complete source code for the Concept2Brain model architecture, training pipelines, and evaluation scripts is publicly available in a GitHub repository: https://github.com/ASantosMayo/Concept2Brain.

Reproduced under the paper's license (CC BY), from the paper cited above.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 3 scripts, each with its path and the digest of its content;
  • 8 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

  • osf:7dpgm, at OSF; found in the text, “Evaluating the Concept2Brain model”

Data availability

The empirical EEG datasets used for training and evaluating the Concept2Brain model, along with the generated synthetic neurophysiological responses underlying the main figures, have been deposited in the OSF repository: https://osf.io/7dpgm.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 2 keywords, 6 MeSH terms, 1 funder, 76 references.

Cite

This paper

Santos-Mayo, A., Gilbert, F., Mirifar, A., Tebbe, A.-L., Fang, R., Ding, M., & Keil, A. (2026). Concept2Brain: an AI model for predicting neurophysiological responses to text and pictures. Nature communications, 17(1), 8961. https://doi.org/10.1038/s41467-026-75653-x

BibTeX

@article{santosmayo2026concept2brain,
author = {Santos-Mayo, Alejandro and Gilbert, Faith and Mirifar, Arash and Tebbe, Anna-Lena and Fang, Ruogu and Ding, Mingzhou and Keil, Andreas},
title = {{Concept2Brain: an AI model for predicting neurophysiological responses to text and pictures}},
journal = {Nature communications},
year = {2026},
month = jul,
volume = {17},
number = {1},
pages = {8961},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-75653-x},
url = {https://doi.org/10.1038/s41467-026-75653-x},
pmid = {42637720},
pmcid = {PMC13503920}
}

RIS

TY - JOUR
AU - Santos-Mayo, Alejandro
AU - Gilbert, Faith
AU - Mirifar, Arash
AU - Tebbe, Anna-Lena
AU - Fang, Ruogu
AU - Ding, Mingzhou
AU - Keil, Andreas
TI - Concept2Brain: an AI model for predicting neurophysiological responses to text and pictures
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/07/23
VL - 17
IS - 1
SP - 8961
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-75653-x
UR - https://doi.org/10.1038/s41467-026-75653-x
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-75653-x",
"type": "article-journal",
"title": "Concept2Brain: an AI model for predicting neurophysiological responses to text and pictures",
"container-title": "Nature communications",
"author": [
{
"family": "Santos-Mayo",
"given": "Alejandro"
},
{
"family": "Gilbert",
"given": "Faith"
},
{
"family": "Mirifar",
"given": "Arash"
},
{
"family": "Tebbe",
"given": "Anna-Lena"
},
{
"family": "Fang",
"given": "Ruogu"
},
{
"family": "Ding",
"given": "Mingzhou"
},
{
"family": "Keil",
"given": "Andreas"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "8961",
"DOI": "10.1038/s41467-026-75653-x",
"PMID": "42637720",
"PMCID": "PMC13503920",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-75653-x",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
23
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.3758/s13415-026-01446-w
People, places, and things: The impact of scene animacy on emotional modulation of the early posterior negativity.
Journal: Cognitive, affective & behavioral neuroscience
In common: EEG, cognitive, 6 references
[2] doi:10.1371/journal.pone.0351872 [code]
Decoding visual object recognition from EEG signals.
Journal: PloS one
In common: Pillow, PyTorch, scikit-learn, 4 other tools, EEG, 3 references
[3] doi:10.1162/imag.a.1207 [code]
Investigating the temporal dynamics and modeling of mid-level feature representations in humans.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Pillow, PyTorch, scikit-learn, 4 other tools, EEG, cognitive, 2 references
[4] doi:10.1038/s41467-026-71267-5 [code]
Human-like cognitive generalization for large models via mental representation-guided supervision.
Journal: Nature communications
In common: Pillow, PyTorch, scikit-learn, 4 other tools, cognitive, 2 references
[5] doi:10.7554/elife.105081 [code]
Movie reconstruction from mouse visual cortex activity.
Journal: eLife
In common: Pillow, PyTorch, pandas, 3 other tools, 3 references
[6] doi:10.1038/s41597-025-05174-7 [code]
A large-scale MEG and EEG dataset for object recognition in naturalistic scenes
Journal: n/a
In common: Pillow, PyTorch, scikit-learn, 4 other tools, EEG, 2 references
[7] doi:10.1038/s41593-026-02285-1 [code]
Fixation duration on natural scenes is explained by memory encoding not processing demand.
Journal: Nature neuroscience
In common: Pillow, PyTorch, scikit-learn, 4 other tools, cognitive, 2 references
[8] doi:10.1002/hbm.70528 [code]
Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates.
Journal: Human brain mapping
In common: Pillow, PyTorch, scikit-learn, 4 other tools, EEG, cognitive, 1 reference
[9] doi:10.1167/jov.26.8.2 [code]
Disentangling objects' contextual associations from perceptual and conceptual attributes using time-resolved neural decoding.
Journal: Journal of vision
In common: pandas, SciPy, Matplotlib, 1 other tool, EEG, cognitive, 3 references
[10] doi:10.1167/jov.26.8.1 [code]
MAME: Multidimensional adaptive metamer exploration with human perceptual feedback.
Journal: Journal of vision
In common: Pillow, PyTorch, scikit-learn, 4 other tools, cognitive, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.