OSCR

Contrastive learning to fine-tune feature extraction models for the visual cortex.

Code ↔ Paper

5 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 5 matches
  1. [1] § Results ↔ code/results_utils.py, lines 739–824 · score 0.66 · mfs words, VWFA, FBA, FFA, OFA, OWFA
  2. [2] § Results ↔ code/utils.py, lines 314–355 · score 0.65 · mfs words, VWFA, FBA, FFA, OFA, OWFA
  3. [3] § Methodology › Implementation details ↔ code/models.py, lines 28–77 · score 0.59 · MLP projection head, ReLU, batch, layers, voxels
  4. [4] § Methodology › Baseline feature extraction models ↔ code/results_utils.py, lines 22–84 · score 0.55 · ImageNet, channels, resized, AlexNet
  5. [5] § Methodology › External dataset validation ↔ code/fit_encoding_models_nod.py, lines 212–271 · score 0.51 · AlexNet model, pooled model, Encoding Model, CL tuned, V1, NOD

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 · 890 lines · 43 KB · no license · 2 matches

  1. from models import CLR_model, fmri_reg, get_pooled_CL_model
  2. from utils import get_dataloaders
  3. from sklearn.metrics.pairwise import cosine_similarity
  4. import torchextractor as tx
  5. from tqdm import tqdm
  6. from sklearn.metrics import accuracy_score
  7. from sklearn.preprocessing import StandardScaler
  8. from sklearn.linear_model import LogisticRegression
  9. from torchvision.models.feature_extraction import create_feature_extractor
  10. from torchvision.models import AlexNet_Weights
  11. import torchvision
  12. import joblib
  13. from torchvision.models import AlexNet_Weights
  14. from torchvision import transforms
  15. import torch
  16. import numpy as np
  17. import os
  18. os.environ["OMP_NUM_THREADS"] = '1'
  19. # Get results for image classification task
  20. def image_classification_results(project_dir, subj_num, hemisphere, rois, device, tuning_method='cl', dataset_name='caltech256', save=False,
  21. pooled=False, pooled_h_method='const', save_probs=False):
  22. hemisphere_abbr = 'l' if hemisphere == 'left' else 'r'
  23. if pooled and pooled_h_method == 'avg':
  24. save_path = os.path.join(project_dir, "results", hemisphere_abbr + "h_" +
  25. dataset_name + "_" + tuning_method + "_results_pooled_havg.joblib")
  26. if save_probs:
  27. probs_save_folder = os.path.join(
  28. project_dir, "classification_preds", "Pooled")
  29. elif pooled and pooled_h_method == 'const':
  30. save_path = os.path.join(project_dir, "results", hemisphere_abbr + "h_" +
  31. dataset_name + "_" + tuning_method + "_results_pooled_hconst.joblib")
  32. if save_probs:
  33. probs_save_folder = os.path.join(
  34. project_dir, "classification_preds", "Pooled")
  35. else:
  36. save_path = project_dir + "/results/Subj" + str(subj_num) + "/subj" + str(
  37. subj_num) + "_" + hemisphere_abbr + "h_" + dataset_name + "_" + tuning_method + "_results.joblib"
  38. probs_save_folder = os.path.join(
  39. project_dir, "classification_preds", "Subj" + str(subj_num))
  40. # Seed RNG, define image transforms for alexnet
  41. torch.manual_seed(0)
  42. alex_transform = transforms.Compose([
  43. transforms.Resize(256),
  44. transforms.CenterCrop(224),
  45. transforms.ToTensor(), # convert the images to a PyTorch tensor
  46. # normalize the images color channels
  47. transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
  48. ])
  49. # Load the Data
  50. if (dataset_name == 'caltech256'):
  51. image_dir = os.path.join(project_dir, "caltech256")
  52. dataset = torchvision.datasets.ImageFolder(
  53. root=image_dir, transform=alex_transform)
  54. elif (dataset_name == 'places365'):
  55. image_dir = os.path.join(project_dir, "places365")
  56. dataset = torchvision.datasets.Places365(
  57. root=image_dir, split='val', small=True, transform=alex_transform)
  58. elif (dataset_name == 'sun397'):
  59. image_dir = os.path.join(project_dir, "sun397")
  60. dataset = torchvision.datasets.SUN397(
  61. root=image_dir, transform=alex_transform, download=True)
  62. elif (dataset_name == 'imagenet'):
  63. image_dir = os.path.join(project_dir, "imagenet")
  64. dataset = torchvision.datasets.ImageNet(
  65. root=image_dir, split='val', transform=alex_transform)
  66. total_num_images = len(dataset)
  67. generator = torch.Generator()
  68. generator.manual_seed(0)
  69. shuffled_idxs = torch.randperm(total_num_images, generator=generator)
  70. train_size = int(0.85 * total_num_images)
  71. test_size = total_num_images - train_size
  72. print(train_size, test_size)
  73. train_idxs = shuffled_idxs[:train_size]
  74. test_idxs = shuffled_idxs[train_size:train_size+test_size]
  75. # Create train and test dataloaders
  76. train_dataset = torch.utils.data.Subset(dataset, train_idxs)
  77. test_dataset = torch.utils.data.Subset(dataset, test_idxs)
  78. train_dataloader = torch.utils.data.DataLoader(
  79. train_dataset, batch_size=256, shuffle=False)
  80. test_dataloader = torch.utils.data.DataLoader(
  81. test_dataset, batch_size=256, shuffle=False)
  82. # Get image features for untuned AlexNet
  83. alex = torch.hub.load('pytorch/vision:v0.10.0', 'alexnet',
  84. weights=AlexNet_Weights.IMAGENET1K_V1)
  85. alex.to(device)
  86. alex.eval()
  87. feature_extractor = create_feature_extractor(
  88. alex, return_nodes=['classifier.5']).to(device)
  89. del alex
  90. # Get untuned alexnet features
  91. train_features_untuned = np.zeros((train_size, 4096))
  92. train_labels = np.zeros(train_size)
  93. for batch_index, data in tqdm(enumerate(train_dataloader), total=len(train_dataloader)):
  94. batch_size = data[0].shape[0]
  95. if batch_index == 0:
  96. low_idx = 0
  97. high_idx = batch_size
  98. else:
  99. low_idx = high_idx
  100. high_idx += batch_size
  101. # Extract features
  102. with torch.no_grad():
  103. ft = feature_extractor(data[0].to(device))
  104. # Flatten the features
  105. ft = torch.hstack([torch.flatten(l, start_dim=1)
  106. for l in ft.values()]).cpu().detach().numpy()
  107. train_features_untuned[low_idx:high_idx] = ft
  108. train_labels[low_idx:high_idx] = data[1]
  109. del ft
  110. test_features_untuned = np.zeros((test_size, 4096))
  111. test_labels = np.zeros(test_size)
  112. for batch_index, data in tqdm(enumerate(test_dataloader), total=len(test_dataloader)):
  113. batch_size = data[0].shape[0]
  114. if batch_index == 0:
  115. low_idx = 0
  116. high_idx = batch_size
  117. else:
  118. low_idx = high_idx
  119. high_idx += batch_size
  120. # Extract features
  121. with torch.no_grad():
  122. ft = feature_extractor(data[0].to(device))
  123. # Flatten the features
  124. ft = torch.hstack([torch.flatten(l, start_dim=1)
  125. for l in ft.values()]).cpu().detach().numpy()
  126. test_features_untuned[low_idx:high_idx] = ft
  127. del ft
  128. test_labels[low_idx:high_idx] = data[1]
  129. del feature_extractor
  130. # Save labels
  131. if save_probs:
  132. np.save(os.path.join(project_dir, "classification_preds",
  133. dataset_name + "_test_labels.npy"), test_labels)
  134. scaler = StandardScaler()
  135. fit_scaler = scaler.fit(train_features_untuned)
  136. train_features_untuned = fit_scaler.transform(train_features_untuned)
  137. test_features_untuned = fit_scaler.transform(test_features_untuned)
  138. print("Fitting linear classifier...")
  139. classifier = LogisticRegression(max_iter=5000).fit(
  140. train_features_untuned, train_labels)
  141. preds = classifier.predict(test_features_untuned)
  142. if save_probs:
  143. untuned_pred_probs = classifier.predict_proba(test_features_untuned)
  144. untuned_pred_probs_save_path = os.path.join(
  145. project_dir, "classification_preds", "Untuned", "untuned_" + dataset_name + "_test_probs.npy")
  146. np.save(untuned_pred_probs_save_path, untuned_pred_probs)
  147. acc = accuracy_score(test_labels, preds) * 100
  148. print("Untuned", acc)
  149. if rois[0] == 'all' and hemisphere == 'left':
  150. rois = ["V1v", "V1d", "V2v", "V2d", "V3v", "V3d", "hV4", "EBA", "FBA-1", "FBA-2",
  151. "OFA", "FFA-1", "FFA-2", "OPA", "PPA", "RSC", "OWFA", "VWFA-1", "VWFA-2", "mfs-words"]
  152. elif rois[0] == 'all' and hemisphere == 'right':
  153. rois = ["V1v", "V1d", "V2v", "V2d", "V3v", "V3d", "hV4", "EBA", "FBA-1", "FBA-2",
  154. "mTL-bodies", "OFA", "FFA-1", "FFA-2", "OPA",
  155. "PPA", "RSC", "OWFA", "VWFA-1", "VWFA-2", "mfs-words", "mTL-words"]
  156. # Go through list of rois
  157. for roi in rois:
  158. print(roi)
  159. found_subj = False
  160. # Get number of voxels for this model
  161. _, _, _, _, num_voxels = get_dataloaders(
  162. project_dir, device, subj_num, hemisphere, roi, 1024, shuffle=False)
  163. if num_voxels < 20 and not pooled:
  164. print(roi, "is too small or empty")
  165. elif num_voxels >= 20 and not pooled:
  166. found_subj = True
  167. # If using pooled models, may need dummy fmri input data matching shape of some other subject's acitvations
  168. elif num_voxels >= 20 and pooled:
  169. found_subj = True
  170. elif num_voxels < 20 and pooled:
  171. for subj_idx in range(2, 9):
  172. _, _, _, _, num_voxels = get_dataloaders(
  173. project_dir, device, subj_idx, hemisphere, roi, 1024, shuffle=False)
  174. if num_voxels > 0:
  175. found_subj = True
  176. break
  177. if found_subj:
  178. if tuning_method == 'cl':
  179. if pooled == True:
  180. # Get voxel counts for ROI across all subjects with the ROI from previously created dictionary
  181. voxel_counts_file = os.path.join(
  182. project_dir, hemisphere_abbr + "h_voxel_counts_rois.joblib")
  183. voxel_counts = joblib.load(voxel_counts_file)
  184. num_voxels_subjs = np.array(voxel_counts[roi])
  185. num_voxels_subjs = np.where(num_voxels_subjs == 0, np.nan, num_voxels_subjs)
  186. avg_voxel_dim = int(np.nanmean(num_voxels_subjs))
  187. if pooled_h_method == 'avg':
  188. # Load pooled CL model
  189. model, _ = get_pooled_CL_model(
  190. num_voxels_subjs, device, avg_voxel_dim)
  191. model_path = os.path.join(
  192. project_dir, "cl_models", hemisphere_abbr + "h_" + roi + "_pooled_model_e30_havg.pt")
  193. elif pooled_h_method == 'const':
  194. h_max_dim = 5741
  195. # Load pooled model
  196. model, _ = get_pooled_CL_model(
  197. num_voxels_subjs, device, h_max_dim)
  198. model_path = os.path.join(
  199. project_dir, "cl_models", hemisphere_abbr + "h_" + roi + "_pooled_model_e30_hconst.pt")
  200. else:
  201. model_dir = os.path.join(project_dir, "cl_models", "Subj" + str(subj_num))
  202. model_path = os.path.join(model_dir, "subj" + \
  203. str(subj_num) + "_" + hemisphere_abbr + \
  204. "h_" + roi + "_model_e30.pt")
  205. h_dim = int(num_voxels*0.8)
  206. z_dim = int(num_voxels*0.2)
  207. model = CLR_model(num_voxels, h_dim, z_dim)
  208. elif tuning_method == 'reg':
  209. model_dir = os.path.join(project_dir, "baseline_models", "nn_reg", "Subj" + str(subj_num))
  210. model_path = os.path.join(model_dir, "subj" + \
  211. str(subj_num) + "_" + hemisphere_abbr + \
  212. "h_" + roi + "_reg_model_e75.pt")
  213. model = fmri_reg(num_voxels)
  214. # Some models are saved differently
  215. try:
  216. model.load_state_dict(torch.load(
  217. model_path, map_location=torch.device('cpu'))[0].state_dict())
  218. except:
  219. try:
  220. model.load_state_dict(torch.load(
  221. model_path, map_location=torch.device('cpu')).state_dict())
  222. except:
  223. model.load_state_dict(torch.load(
  224. model_path, map_location=torch.device('cpu')))
  225. model.to(device)
  226. model.eval()
  227. feature_extractor = tx.Extractor(
  228. model, ["alex.classifier.5"]).to(device)
  229. train_features_tuned = np.zeros((train_size, 4096))
  230. train_labels = np.zeros(train_size)
  231. for batch_index, data in tqdm(enumerate(train_dataloader), total=len(train_dataloader)):
  232. batch_size = data[0].shape[0]
  233. if batch_index == 0:
  234. low_idx = 0
  235. high_idx = batch_size
  236. else:
  237. low_idx = high_idx
  238. high_idx += batch_size
  239. # Extract features
  240. with torch.no_grad():
  241. if tuning_method == 'cl':
  242. fmri_dummy = torch.zeros(
  243. (batch_size, num_voxels)).to(device)
  244. if pooled:
  245. _, alex_out_dict = feature_extractor(
  246. fmri_dummy, data[0], subj_num)
  247. else:
  248. _, alex_out_dict = feature_extractor(
  249. fmri_dummy, data[0])
  250. # _, alex_out_dict = feature_extractor(fmri_dummy, data[0].to(device))
  251. elif tuning_method == 'reg':
  252. _, alex_out_dict = feature_extractor(
  253. data[0].to(device))
  254. ft = alex_out_dict['alex.classifier.5'].detach().cpu().numpy()
  255. train_features_tuned[low_idx:high_idx] = ft
  256. train_labels[low_idx:high_idx] = data[1]
  257. del ft
  258. test_features_tuned = np.zeros((test_size, 4096))
  259. test_labels = np.zeros(test_size)
  260. for batch_index, data in tqdm(enumerate(test_dataloader), total=len(test_dataloader)):
  261. batch_size = data[0].shape[0]
  262. if batch_index == 0:
  263. low_idx = 0
  264. high_idx = batch_size
  265. else:
  266. low_idx = high_idx
  267. high_idx += batch_size
  268. # Extract features
  269. with torch.no_grad():
  270. if tuning_method == 'cl':
  271. fmri_dummy = torch.zeros(
  272. (batch_size, num_voxels)).to(device)
  273. if pooled:
  274. _, alex_out_dict = feature_extractor(
  275. fmri_dummy, data[0].to(torch.float), subj_num)
  276. else:
  277. _, alex_out_dict = feature_extractor(
  278. fmri_dummy, data[0])
  279. # _, alex_out_dict = feature_extractor(fmri_dummy, data[0].to(device))
  280. elif tuning_method == 'reg':
  281. _, alex_out_dict = feature_extractor(
  282. data[0].to(device))
  283. ft = alex_out_dict['alex.classifier.5'].detach().cpu().numpy()
  284. test_features_tuned[low_idx:high_idx] = ft
  285. test_labels[low_idx:high_idx] = data[1]
  286. del ft
  287. scaler = StandardScaler()
  288. fit_scaler = scaler.fit(train_features_tuned)
  289. train_features_tuned = fit_scaler.transform(train_features_tuned)
  290. test_features_tuned = fit_scaler.transform(test_features_tuned)
  291. print("Fitting linear classifier...")
  292. classifier = LogisticRegression(max_iter=5000).fit(
  293. train_features_tuned, train_labels)
  294. preds = classifier.predict(test_features_tuned)
  295. if save_probs:
  296. tuned_pred_probs = classifier.predict_proba(
  297. test_features_tuned)
  298. if pooled:
  299. if pooled_h_method == 'avg':
  300. tuned_pred_probs_save_path = os.path.join(
  301. probs_save_folder, hemisphere_abbr + "h_" + roi + "_pooled_havg_" + dataset_name + "_test_probs.npy")
  302. elif pooled_h_method == 'const':
  303. tuned_pred_probs_save_path = os.path.join(
  304. probs_save_folder, hemisphere_abbr + "h_" + roi + "_pooled_hconst_" + dataset_name + "_test_probs.npy")
  305. else:
  306. if tuning_method == 'cl':
  307. tuned_pred_probs_save_path = os.path.join(probs_save_folder, "subj" + str(
  308. subj_num) + "_" + hemisphere_abbr + "h_" + roi + "_" + dataset_name + "_test_probs.npy")
  309. else:
  310. tuned_pred_probs_save_path = os.path.join(probs_save_folder, "subj" + str(
  311. subj_num) + "_" + hemisphere_abbr + "h_" + roi + "_" + dataset_name + "_reg_test_probs.npy")
  312. np.save(tuned_pred_probs_save_path, tuned_pred_probs)
  313. acc = accuracy_score(test_labels, preds) * 100
  314. # results[roi] = acc
  315. print(roi, acc)
  316. if save:
  317. try:
  318. # Load existing results, add result for roi if not already in the existing results
  319. existing_results = joblib.load(save_path)
  320. if roi not in existing_results.keys():
  321. existing_results[roi] = acc
  322. joblib.dump(existing_results, save_path)
  323. except:
  324. results = {}
  325. results[roi] = acc
  326. joblib.dump(results, save_path)
  327. # Compute lower bound on mutual information between CNN features and ROI response using the testing data
  328. def compute_mi_lower_bound(project_dir, device, subj_num, roi, hemisphere, pooled=False, h_method='const'):
  329. # Temperature tau used to tune CL models
  330. tau = 0.3
  331. hemisphere_abbr = 'l' if hemisphere == 'left' else 'r'
  332. print(roi, hemisphere)
  333. # Get test dataloader
  334. _, test_dataloader, _, test_size, num_voxels = get_dataloaders(
  335. project_dir, device, subj_num, hemisphere, roi, 1024, shuffle=False)
  336. if (num_voxels == 0):
  337. print("Empty ROI")
  338. return -1, -1, -1, -1, -1, -1, -1
  339. elif (num_voxels < 20):
  340. print("Too few voxels")
  341. return -1, -1, -1, -1, -1, -1, -1
  342. if pooled:
  343. # Get list of number of voxels for each subj
  344. present_subjs = []
  345. num_voxels_subjs = []
  346. for subj_idx in range(1, 9):
  347. _, _, _, _, num_voxels = get_dataloaders(
  348. project_dir, device, subj_idx, hemisphere, roi, batch_size=1024)
  349. if num_voxels != 0:
  350. present_subjs.append(subj_idx)
  351. num_voxels_subjs.append(num_voxels)
  352. # Get voxel counts for ROI across all subjects with the ROI from previously created dictionary
  353. voxel_counts_file = os.path.join(
  354. project_dir, hemisphere_abbr + "h_voxel_counts_rois.joblib")
  355. voxel_counts = joblib.load(voxel_counts_file)
  356. num_voxels_subjs = np.array(voxel_counts[roi])
  357. num_voxels_subjs = np.where(num_voxels_subjs == 0, np.nan, num_voxels_subjs)
  358. avg_voxel_dim = int(np.nanmean(num_voxels_subjs))
  359. # Get pooled model
  360. print("Getting CL predictions...")
  361. cl_model_dir = os.path.join(project_dir, "cl_models")
  362. if h_method == 'avg':
  363. h_dim = avg_voxel_dim
  364. # Load pooled CL model
  365. cl_model, _ = get_pooled_CL_model(num_voxels_subjs, device, h_dim)
  366. cl_model_path = os.path.join(
  367. cl_model_dir, hemisphere_abbr + "h_" + roi + "_pooled_model_e30_havg.pt")
  368. elif h_method == 'const':
  369. h_dim = 5741
  370. # Load pooled model
  371. cl_model, _ = get_pooled_CL_model(num_voxels_subjs, device, h_dim)
  372. cl_model_path = os.path.join(
  373. cl_model_dir, hemisphere_abbr + "h_" + roi + "_pooled_model_e30_hconst.pt")
  374. z_dim = int(h_dim * 0.25)
  375. # Some models are saved differently
  376. try:
  377. cl_model.load_state_dict(torch.load(
  378. cl_model_path, map_location=torch.device('cpu'))[0].state_dict())
  379. except:
  380. try:
  381. cl_model.load_state_dict(torch.load(
  382. cl_model_path, map_location=torch.device('cpu')).state_dict())
  383. except:
  384. cl_model.load_state_dict(torch.load(
  385. cl_model_path, map_location=torch.device('cpu')))
  386. else:
  387. # Load CL-tuned model
  388. cl_model_dir = os.path.join(project_dir, "Subj" + str(subj_num))
  389. cl_model_path = os.path.join(cl_model_dir, "subj" + str(subj_num) + "_" + hemisphere_abbr +
  390. "h_" + roi + "_model_e30.pt")
  391. h_dim = int(num_voxels*0.8)
  392. z_dim = int(num_voxels*0.2)
  393. cl_model = CLR_model(num_voxels, h_dim, z_dim)
  394. # Some models are saved differently
  395. try:
  396. cl_model.load_state_dict(torch.load(
  397. cl_model_path, map_location=torch.device('cpu'))[0].state_dict())
  398. except:
  399. try:
  400. cl_model.load_state_dict(torch.load(
  401. cl_model_path, map_location=torch.device('cpu')).state_dict())
  402. except:
  403. cl_model.load_state_dict(torch.load(
  404. cl_model_path, map_location=torch.device('cpu')))
  405. cl_model.to(device)
  406. cl_model.eval()
  407. # Use just 1 batch
  408. K = test_size
  409. # Get outputs (after projection head) for CL model and fMRI data
  410. nn_z = np.zeros((K, z_dim))
  411. fmri_z = np.zeros((K, z_dim))
  412. for batch_index, data in tqdm(enumerate(test_dataloader), total=len(test_dataloader)):
  413. batch_size = data[0].shape[0]
  414. if batch_index == 0:
  415. low_idx = 0
  416. high_idx = batch_size
  417. else:
  418. low_idx = high_idx
  419. high_idx += batch_size
  420. # Extract features
  421. with torch.no_grad():
  422. if pooled:
  423. # Forward function for pooled CL models expects subj num = (1-based) position in list of subjects with roi present (exclusing subjects without roi),
  424. # so need to correct indexing
  425. subj_num_adjusted = present_subjs.index(subj_num) + 1
  426. # print(subj_num_adjusted)
  427. ft_fmri, ft_nn = cl_model(data[0], data[1], subj_num_adjusted)
  428. else:
  429. ft_fmri, ft_nn = cl_model(data[0], data[1])
  430. # Flatten the features, collect them
  431. ft_fmri = ft_fmri.cpu().detach().numpy()
  432. ft_nn = ft_nn.cpu().detach().numpy()
  433. fmri_z[low_idx:high_idx] = ft_fmri
  434. nn_z[low_idx:high_idx] = ft_nn
  435. critic_out = (1 / tau) * cosine_similarity(nn_z, fmri_z)
  436. exp_critic = np.exp(critic_out)
  437. lower_bound_mi = 0
  438. for i in range(K):
  439. lower_bound_mi += np.log(exp_critic[i, i] /
  440. ((1 / K) * np.sum(exp_critic[i, :])))
  441. lower_bound_mi_unscaled = (1 / K) * lower_bound_mi
  442. # Do post-hoc scaling for temp
  443. betas = np.logspace(-2, 2, 1000)
  444. best_beta = 1
  445. best_lower_bound = lower_bound_mi_unscaled
  446. for beta in betas:
  447. tau_new = tau * beta
  448. critic_out = (1 / tau_new) * cosine_similarity(nn_z, fmri_z)
  449. exp_critic = np.exp(critic_out)
  450. lower_bound_mi = 0
  451. for i in range(K):
  452. lower_bound_mi += np.log(exp_critic[i, i] /
  453. ((1 / K) * np.sum(exp_critic[i, :])))
  454. lower_bound_mi = (1 / K) * lower_bound_mi
  455. if lower_bound_mi > best_lower_bound:
  456. best_lower_bound = lower_bound_mi
  457. best_beta = beta
  458. print(best_beta, best_lower_bound)
  459. # Save results
  460. if pooled:
  461. save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(
  462. subj_num) + "_" + hemisphere_abbr + "h_mi_lower_bound_pooled_h" + h_method + ".joblib")
  463. else:
  464. save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  465. hemisphere_abbr + "h_mi_lower_bound.joblib")
  466. try:
  467. results = joblib.load(save_path)
  468. except:
  469. results = {}
  470. results[roi] = lower_bound_mi_unscaled, np.log(
  471. K), best_beta, best_lower_bound
  472. joblib.dump(results, save_path)
  473. # Load test cv results for single subject untuned for all layers and subjects to create figure 2 (matrix of encoding scores)
  474. def load_test_cv_single_subj_untuned_all_layers(project_dir):
  475. num_rows = 46
  476. results_matrix = np.zeros((num_rows, 8)) # Average across subjects
  477. # Keep track of how many subjs have each ROI
  478. row_counters = np.zeros((num_rows))
  479. for subj_num in range(1, 9):
  480. results_folder_path = os.path.join(
  481. project_dir, "best_alex_out_layers_test_cv", "Subj" + str(subj_num))
  482. results_file = os.path.join(
  483. results_folder_path, "best_alex_layers_mat_untuned.npy")
  484. subj_results = np.load(results_file)
  485. for row_idx in range(num_rows):
  486. if subj_results[row_idx, 0] > 0:
  487. row_counters[row_idx] += 1
  488. results_matrix[row_idx, :] += subj_results[row_idx, :]
  489. for row_idx in range(num_rows):
  490. if results_matrix[row_idx, 0] > 0:
  491. results_matrix[row_idx, :] /= row_counters[row_idx]
  492. return results_matrix
  493. # Load test cv results for single subject (untuned, cl-tuned, reg-tuned, pooled-avg, or pooled-specific)
  494. # split options are all, early, higher
  495. def load_test_cv_single_subj_results(project_dir, subj_num, split='all', return_pooled_results=False, roi=None):
  496. results_folder_path = os.path.join(
  497. project_dir, "best_alex_out_layers_test_cv", "Subj" + str(subj_num))
  498. untuned_results_path = os.path.join(
  499. results_folder_path, "best_alphas_voxel_accs_dict_untuned.joblib")
  500. cl_tuned_results_path = os.path.join(
  501. results_folder_path, "best_alphas_voxel_accs_dict_cl_tuned.joblib")
  502. reg_tuned_results_path = os.path.join(
  503. results_folder_path, "best_alphas_voxel_accs_dict_reg_tuned.joblib")
  504. pooled_avg_results_path = os.path.join(
  505. results_folder_path, "best_alphas_voxel_accs_dict_pooled_avg.joblib")
  506. pooled_const_results_path = os.path.join(
  507. results_folder_path, "best_alphas_voxel_accs_dict_pooled_const.joblib")
  508. untuned_results = joblib.load(untuned_results_path)
  509. cl_tuned_results = joblib.load(cl_tuned_results_path)
  510. reg_tuned_results = joblib.load(reg_tuned_results_path)
  511. pooled_avg_results = joblib.load(pooled_avg_results_path)
  512. pooled_const_results = joblib.load(pooled_const_results_path)
  513. if roi is not None:
  514. try:
  515. return untuned_results[roi], reg_tuned_results[roi], cl_tuned_results[roi], pooled_avg_results[roi], pooled_const_results[roi]
  516. except:
  517. return 0, 0, 0, 0, 0
  518. else:
  519. if split == 'all':
  520. rois = ["V1v", "V1d", "V2v", "V2d", "V3v", "V3d", "hV4", "EBA", "FBA-1", "FBA-2",
  521. "mTL-bodies", "OFA", "FFA-1", "FFA-2", "mTL-faces", "aTL-faces", "OPA",
  522. "PPA", "RSC", "OWFA", "VWFA-1", "VWFA-2", "mfs-words", "mTL-words"]
  523. elif split == 'early':
  524. rois = ["V1v", "V1d", "V2v", "V2d", "V3v", "V3d", "hV4"]
  525. elif split == 'higher':
  526. rois = ["EBA", "FBA-1", "FBA-2", "mTL-bodies", "OFA", "FFA-1", "FFA-2", "mTL-faces", "aTL-faces", "OPA",
  527. "PPA", "RSC", "OWFA", "VWFA-1", "VWFA-2", "mfs-words", "mTL-words"]
  528. else:
  529. print("Unsupported Split!")
  530. return
  531. untuned_mean_acc = 0
  532. reg_tuned_mean_acc = 0
  533. cl_tuned_mean_acc = 0
  534. pooled_avg_mean_acc = 0
  535. pooled_const_mean_acc = 0
  536. cl_over_untuned_voxel_mean_percentage = 0
  537. cl_over_reg_tuned_voxel_mean_percentage = 0
  538. pooled_avg_over_cl_voxel_mean_percentage = 0
  539. pooled_const_over_cl_voxel_mean_percentage = 0
  540. num_rois_present = 0
  541. for hemi in ["lh", "rh"]:
  542. for roi in rois:
  543. try:
  544. untuned_hemi_roi_results = untuned_results[hemi + '_' + roi]
  545. reg_tuned_hemi_roi_results = reg_tuned_results[hemi + '_' + roi]
  546. cl_tuned_hemi_roi_results = cl_tuned_results[hemi + '_' + roi]
  547. pooled_avg_hemi_roi_results = pooled_avg_results[hemi + '_' + roi]
  548. pooled_const_hemi_roi_results = pooled_const_results[hemi + '_' + roi]
  549. untuned_mean_acc += untuned_hemi_roi_results.mean()
  550. reg_tuned_mean_acc += reg_tuned_hemi_roi_results.mean()
  551. cl_tuned_mean_acc += cl_tuned_hemi_roi_results.mean()
  552. pooled_avg_mean_acc += pooled_avg_hemi_roi_results.mean()
  553. pooled_const_mean_acc += pooled_const_hemi_roi_results.mean()
  554. cl_over_untuned_voxel_mean_percentage += np.count_nonzero(
  555. cl_tuned_hemi_roi_results - untuned_hemi_roi_results > 0) / cl_tuned_hemi_roi_results.shape[0]
  556. cl_over_reg_tuned_voxel_mean_percentage += np.count_nonzero(
  557. cl_tuned_hemi_roi_results - reg_tuned_hemi_roi_results > 0) / cl_tuned_hemi_roi_results.shape[0]
  558. pooled_avg_over_cl_voxel_mean_percentage += np.count_nonzero(
  559. pooled_avg_hemi_roi_results - cl_tuned_hemi_roi_results > 0) / cl_tuned_hemi_roi_results.shape[0]
  560. pooled_const_over_cl_voxel_mean_percentage += np.count_nonzero(
  561. pooled_const_hemi_roi_results - cl_tuned_hemi_roi_results > 0) / cl_tuned_hemi_roi_results.shape[0]
  562. num_rois_present += 1
  563. except:
  564. pass # Skip missing ROIs
  565. untuned_mean_acc /= num_rois_present
  566. reg_tuned_mean_acc /= num_rois_present
  567. cl_tuned_mean_acc /= num_rois_present
  568. pooled_avg_mean_acc /= num_rois_present
  569. pooled_const_mean_acc /= num_rois_present
  570. cl_over_untuned_voxel_mean_percentage /= num_rois_present
  571. cl_over_reg_tuned_voxel_mean_percentage /= num_rois_present
  572. pooled_avg_over_cl_voxel_mean_percentage /= num_rois_present
  573. pooled_const_over_cl_voxel_mean_percentage /= num_rois_present
  574. if return_pooled_results:
  575. return pooled_avg_mean_acc, pooled_const_mean_acc, pooled_avg_over_cl_voxel_mean_percentage, pooled_const_over_cl_voxel_mean_percentage
  576. else:
  577. return untuned_mean_acc, reg_tuned_mean_acc, cl_tuned_mean_acc, cl_over_untuned_voxel_mean_percentage, cl_over_reg_tuned_voxel_mean_percentage
  578. # Return matrix of results for each layer/roi
  579. def load_test_cv_single_subj_results_all_layers(project_dir, subj_num):
  580. results_folder_path = os.path.join(
  581. project_dir, "best_alex_out_layers_test_cv", "Subj" + str(subj_num))
  582. untuned_results_path = os.path.join(
  583. results_folder_path, "best_alex_layers_mat_untuned.npy")
  584. cl_tuned_results_path = os.path.join(
  585. results_folder_path, "best_alex_layers_mat_cl_tuned.npy")
  586. reg_tuned_results_path = os.path.join(
  587. results_folder_path, "best_alex_layers_mat_reg_tuned.npy")
  588. pooled_avg_results_path = os.path.join(
  589. results_folder_path, "best_alex_layers_mat_pooled_avg.npy")
  590. pooled_const_results_path = os.path.join(
  591. results_folder_path, "best_alex_layers_mat_pooled_const.npy")
  592. untuned_results = np.load(untuned_results_path)
  593. cl_tuned_results = np.load(cl_tuned_results_path)
  594. reg_tuned_results = np.load(reg_tuned_results_path)
  595. pooled_avg_results = np.load(pooled_avg_results_path)
  596. pooled_const_results = np.load(pooled_const_results_path)
  597. return untuned_results, cl_tuned_results, reg_tuned_results, pooled_avg_results, pooled_const_results
  598. # Generate embeddings for test images from untuned, CL, or regression-tuned models. Save corresponding NSD IDs of images.
  599. # Options for tuning_method are 'untuned', 'CL' or 'reg'
  600. # Optionally get embeddings from best layer for encoding instead of final layer
  601. def save_embeddings(project_dir, subj_num, hemisphere, roi, device, tuning_method='CL', use_best_intermediate_layer=False, use_other_subj_images=False, cross_subj_num=1):
  602. # Strip whitespace from roi to handle cases where it comes from files with trailing spaces
  603. roi = roi.strip()
  604. hemisphere_abbr = 'l' if hemisphere == 'left' else 'r'
  605. if tuning_method == 'CL':
  606. if use_best_intermediate_layer:
  607. if use_other_subj_images:
  608. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  609. hemisphere_abbr + "h_" + roi + "_cl_embeddings_best_encoding_layer_cross_subj" + str(cross_subj_num) + ".npy")
  610. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  611. hemisphere_abbr + "h_" + roi + "_cl_embeddings_best_encoding_layer_cross_subj" + str(cross_subj_num) + "_img_ids.npy")
  612. else:
  613. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  614. hemisphere_abbr + "h_" + roi + "_cl_embeddings_best_encoding_layer.npy")
  615. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  616. hemisphere_abbr + "h_" + roi + "_cl_embeddings_best_encoding_layer_img_ids.npy")
  617. else:
  618. if use_other_subj_images:
  619. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  620. hemisphere_abbr + "h_" + roi + "_cl_embeddings_cross_subj" + str(cross_subj_num) + ".npy")
  621. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  622. hemisphere_abbr + "h_" + roi + "_cl_embeddings_cross_subj" + str(cross_subj_num) + "_img_ids.npy")
  623. else:
  624. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  625. hemisphere_abbr + "h_" + roi + "_cl_embeddings.npy")
  626. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  627. hemisphere_abbr + "h_" + roi + "_cl_embeddings_img_ids.npy")
  628. elif tuning_method == 'reg':
  629. if use_best_intermediate_layer:
  630. if use_other_subj_images:
  631. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  632. hemisphere_abbr + "h_" + roi + "_reg_embeddings_best_encoding_layer_cross_subj" + str(cross_subj_num) + ".npy")
  633. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  634. hemisphere_abbr + "h_" + roi + "_reg_embeddings_best_encoding_layer_cross_subj" + str(cross_subj_num) + "_img_ids.npy")
  635. else:
  636. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  637. hemisphere_abbr + "h_" + roi + "_reg_embeddings_best_encoding_layer.npy")
  638. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  639. hemisphere_abbr + "h_" + roi + "_reg_embeddings_best_encoding_layer_img_ids.npy")
  640. else:
  641. if use_other_subj_images:
  642. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  643. hemisphere_abbr + "h_" + roi + "_reg_embeddings_cross_subj" + str(cross_subj_num) + ".npy")
  644. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  645. hemisphere_abbr + "h_" + roi + "_reg_embeddings_cross_subj" + str(cross_subj_num) + "_img_ids.npy")
  646. else:
  647. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  648. hemisphere_abbr + "h_" + roi + "_reg_embeddings.npy")
  649. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  650. hemisphere_abbr + "h_" + roi + "_reg_embeddings_img_ids.npy")
  651. elif tuning_method == 'untuned':
  652. features_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_untuned_embeddings.npy")
  653. ids_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_untuned_embeddings_img_ids.npy")
  654. if use_other_subj_images:
  655. _, _, _, _, num_voxels = get_dataloaders(project_dir,
  656. device, subj_num, hemisphere, roi, batch_size=1024, use_all_data=False, shuffle=False, return_nsd_id=True)
  657. _, test_dataloader, _, test_size, _ = get_dataloaders(project_dir,
  658. device, cross_subj_num, hemisphere, roi, batch_size=1024, use_all_data=False, shuffle=False, return_nsd_id=True)
  659. else:
  660. _, test_dataloader, _, test_size, num_voxels = get_dataloaders(project_dir,
  661. device, subj_num, hemisphere, roi, batch_size=1024, use_all_data=False, shuffle=False, return_nsd_id=True)
  662. # Load tuned model
  663. if tuning_method == 'CL':
  664. model_dir = os.path.join(project_dir, "cl_models", "Subj" + str(subj_num))
  665. model_path = os.path.join(model_dir, "subj" + str(subj_num) + "_" + hemisphere_abbr + "h_" + roi + "_model_e30.pt")
  666. h_dim = int(num_voxels*0.8)
  667. z_dim = int(num_voxels*0.2)
  668. model = CLR_model(num_voxels, h_dim, z_dim)
  669. elif tuning_method == 'reg':
  670. model_dir = os.path.join(project_dir, "baseline_models", "nn_reg", "Subj" + str(subj_num))
  671. model_path = os.path.join(model_dir, "subj" + \
  672. str(subj_num) + "_" + hemisphere_abbr + \
  673. "h_" + roi + "_reg_model_e75.pt")
  674. model = fmri_reg(num_voxels)
  675. elif tuning_method == 'untuned':
  676. model = torch.hub.load('pytorch/vision:v0.10.0', 'alexnet',
  677. weights=AlexNet_Weights.IMAGENET1K_V1)
  678. # Some tuned models are saved differently
  679. if tuning_method == 'CL' or tuning_method == 'reg':
  680. try:
  681. model.load_state_dict(torch.load(
  682. model_path, map_location=torch.device('cpu'), weights_only=False)[0].state_dict())
  683. except:
  684. try:
  685. model.load_state_dict(torch.load(
  686. model_path, map_location=torch.device('cpu'), weights_only=False).state_dict())
  687. except:
  688. model.load_state_dict(torch.load(
  689. model_path, map_location=torch.device('cpu'), weights_only=False))
  690. model.to(device)
  691. model.eval()
  692. roi_names = ["V1v", "V1d", "V2v", "V2d", "V3v", "V3d", "hV4", "EBA", "FBA-1", "FBA-2",
  693. "mTL-bodies", "OFA", "FFA-1", "FFA-2", "mTL-faces", "OPA",
  694. "PPA", "RSC", "OWFA", "VWFA-1", "VWFA-2", "mfs-words", "mTL-words"]
  695. hemis = ["lh", "rh"]
  696. layers = [
  697. "features.2", "features.5", "features.7", "features.9",
  698. "features.12", "classifier.2", "classifier.5", "classifier.6"
  699. ]
  700. layer_dim_dict = {"features.2": 46656, "features.5": 32448, "features.7": 64896, "features.9": 43264,
  701. "features.12": 9216, "classifier.2": 4096, "classifier.5": 4096, "classifier.6": 1000}
  702. if tuning_method == 'CL' or tuning_method == 'reg':
  703. if use_best_intermediate_layer:
  704. _, results, _, _, _ = load_test_cv_single_subj_results_all_layers(project_dir, subj_num)
  705. results_dict = {}
  706. counter = 0
  707. for hemi in hemis:
  708. for roi_name in roi_names:
  709. current_roi = hemi + "_" + roi_name
  710. if not np.isnan(results[counter]).any():
  711. results_dict[current_roi] = results[counter]
  712. counter += 1
  713. best_layer = layers[np.argmax(results_dict[hemisphere_abbr + "h_" + roi])]
  714. print("Using layer:", best_layer)
  715. feature_extractor = tx.Extractor(
  716. model, ["alex." + best_layer]).to(device)
  717. output_dim = layer_dim_dict[best_layer]
  718. else:
  719. feature_extractor = tx.Extractor(
  720. model, ["alex.classifier.5"]).to(device)
  721. output_dim = 4096
  722. elif tuning_method == 'reg':
  723. if use_best_intermediate_layer:
  724. _, _, results, _, _ = load_test_cv_single_subj_results_all_layers(project_dir, subj_num)
  725. results_dict = {}
  726. counter = 0
  727. for hemi in hemis:
  728. for roi_name in roi_names:
  729. current_roi = hemi + "_" + roi_name
  730. if not np.isnan(results[counter]).any():
  731. results_dict[current_roi] = results[counter]
  732. counter += 1
  733. best_layer = layers[np.argmax(results_dict[hemisphere_abbr + "h_" + roi])]
  734. print("Using layer:", best_layer)
  735. feature_extractor = tx.Extractor(
  736. model, ["alex." + best_layer]).to(device)
  737. output_dim = layer_dim_dict[best_layer]
  738. else:
  739. feature_extractor = tx.Extractor(
  740. model, ["alex.classifier.5"]).to(device)
  741. output_dim = 4096
  742. elif tuning_method == 'untuned':
  743. feature_extractor = tx.Extractor(
  744. model, ["classifier.5"]).to(device)
  745. output_dim = 4096
  746. features = np.zeros((test_size, output_dim))
  747. ids = np.zeros(test_size)
  748. for batch_index, data in tqdm(enumerate(test_dataloader), total=len(test_dataloader)):
  749. batch_size = data[0].shape[0]
  750. if batch_index == 0:
  751. low_idx = 0
  752. high_idx = batch_size
  753. else:
  754. low_idx = high_idx
  755. high_idx += batch_size
  756. # Extract features
  757. with torch.no_grad():
  758. if tuning_method == 'CL':
  759. fmri_dummy = torch.zeros(
  760. (batch_size, num_voxels)).to(device)
  761. _, alex_out_dict = feature_extractor(fmri_dummy, data[1])
  762. # _, alex_out_dict = feature_extractor(fmri_dummy, data[0].to(device))
  763. elif tuning_method == 'reg':
  764. _, alex_out_dict = feature_extractor(data[1].to(device))
  765. elif tuning_method == 'untuned':
  766. _, alex_out_dict = feature_extractor(data[1].to(device))
  767. if tuning_method == 'CL' or tuning_method == 'reg':
  768. if use_best_intermediate_layer:
  769. ft = alex_out_dict["alex." + best_layer].detach().cpu().numpy().reshape(high_idx-low_idx, -1)
  770. else:
  771. ft = alex_out_dict['alex.classifier.5'].detach().cpu().numpy()
  772. elif tuning_method == 'untuned':
  773. ft = alex_out_dict['classifier.5'].detach().cpu().numpy()
  774. features[low_idx:high_idx] = ft
  775. ids[low_idx:high_idx] = data[3]
  776. del ft
  777. # Save features and ids
  778. np.save(features_save_path, features)
  779. np.save(ids_save_path, ids)
  780. # Save fmri responses for test images. Save corresponding NSD IDs of images.
  781. # Use to match format of save_embeddings function.
  782. def save_test_fmri_responses(project_dir, subj_num, hemisphere, roi, device):
  783. # Strip whitespace from roi to handle cases where it comes from files with trailing spaces
  784. roi = roi.strip()
  785. hemisphere_abbr = 'l' if hemisphere == 'left' else 'r'
  786. fmri_responses_save_path = os.path.join(project_dir, "results", "Subj" + str(subj_num), "subj" + str(subj_num) + "_" +
  787. hemisphere_abbr + "h_" + roi + "_test_fmri_responses.npy")
  788. _, test_dataloader, _, test_size, num_voxels = get_dataloaders(project_dir,
  789. device, subj_num, hemisphere, roi, batch_size=1024, use_all_data=False, shuffle=False, return_nsd_id=True)
  790. fmri_responses = np.zeros((test_size, num_voxels))
  791. for batch_index, data in tqdm(enumerate(test_dataloader), total=len(test_dataloader)):
  792. batch_size = data[0].shape[0]
  793. if batch_index == 0:
  794. low_idx = 0
  795. high_idx = batch_size
  796. else:
  797. low_idx = high_idx
  798. high_idx += batch_size
  799. fmri_responses[low_idx:high_idx] = data[0]
  800. del data
  801. np.save(fmri_responses_save_path, fmri_responses)

results_utils.py at commit af1f772, no license · at the source

Overview

Authors: Alex Mulrooney1, Zhi Li1, Austin J Brockmeier1,2
ORCID iDs: Alex Mulrooney
  1. Department of Electrical and Computer Engineering, University of Delaware, Newark, Delaware, United States of America
  2. Department of Computer and Information Sciences, University of Delaware, Newark, Delaware, United States of America
Institutions: University of Delaware (United States)
Journal: PLoS computational biology, volume 22, issue 8, article e1014656
Dates: received 28 August 2025; accepted 3 August 2026; published online 17 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014656 · PMID 42607104 · PMCID PMC13492995 · OpenAlex W7203628548
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: fMRI (modality), human (organism), systems (subfield)
Methods: Smoothing, state filtering, decompositions, Machine learning, Statistics, Connectivity, fMRI & imaging
MeSH: Models, Neurological*, Visual Cortex*, Algorithms, Brain Mapping, Computational Biology, Convolutional Neural Networks, Humans, Image Processing, Computer-Assisted, Magnetic Resonance Imaging (* major topic)
Topic: Face Recognition and Perception (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: University of Delaware Department of Electrical and Computer Engineering First Year Fellowship; University of Delaware College of Engineering (None); University of Delaware Research Foundation; University of Delaware Electrical and Computer Engineering Department (None); Office of Naval Research (N00014-24-1-2259); University of Delaware General University Research fund
Citations: cited by 1 paper (Europe PMC); 59 references in the paper

Abstract

Predicting the neural response to natural images in the visual cortex requires extracting relevant features from the images and relating those feature to the observed responses. In this work, we optimize the feature extraction in order to maximize the information shared between the image features and the neural response across voxels in a given region of interest (ROI) extracted from the BOLD signal measured by functional magnetic resonance imaging (fMRI). We adapt contrastive learning (CL) to fine-tune a convolutional neural network, which was pretrained for image classification, such that a mapping of a given image’s features are more similar to the corresponding fMRI response than to the responses to other images. We exploit the Natural Scenes Dataset as organized for the Algonauts Project, which contains the high-resolution fMRI responses of eight subjects to tens of thousands of naturalistic images. We show that CL fine-tuning creates feature extraction models that enable higher encoding accuracy in both early and higher visual ROIs as compared to the features from the pretrained network. Quantitatively, the performance is similar to a baseline approach that directly uses a regression loss at the output of the network to tune it for fMRI response encoding. We investigate inter-subject transfer of the CL fine-tuned models, including subjects from the Natural Object Dataset, another lower-resolution dataset with 9 subjects. We also pool subjects for fine-tuning, which further improves encoding performance in early ROIs. Finally, we examine the performance of the fine-tuned models on common image classification tasks, explore the landscape of ROI-specific models by applying dimensionality reduction on the Bhattacharya dissimilarity matrix created using the predictions on those tasks, and show that these landscapes match those based on representational similarity analysis. Finally, we generate images via Stable Diffusion based on vector-space prompts created by aligning the CL-tuned models embeddings for different ROIs, showing that generated images have similar embeddings to the original but that estimates of the intrinsic dimension are lower for generated versus original representations.

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 5 matches between paragraphs and lines of code.

alexmul1114/fmri_encoding_contrastive_learning

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: af1f772661f2c5789154d2f1078b077569fceba4, 18 August 2026
Languages: Python (13), Jupyter (1), Shell (1)
Size: 564 files, 15 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, environment (requirements.txt), 1 notebook
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (13 files), NumPy (8 files), scikit-learn (4 files), Pillow (3 files), SciPy (3 files), DataLad (1 file), Matplotlib (1 file), NiBabel (1 file), pandas (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
16 files

The paper's code and data availability statement is in the Data section.

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;
  • 15 scripts, each with its path and the digest of its content;
  • 5 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

Data Availability

All data we use in our paper is publicly available. The Natural Scenes Dataset is available at https://naturalscenesdataset.org/. The Natural Object Dataset is available at https://openneuro.org/datasets/ds004496/versions/2.1.2. The ImageNet dataset is available at https://www.image-net.org/. The Caltech256 dataset is available at https://data.caltech.edu/records/nyy15-4j048. The Places365 dataset is available at https://data.caltech.edu/records/nyy15-4j048. The MS-COCO dataset is available at https://cocodataset.org/#home. We provide code for reproducing our results as well as instructions for obtaining the data on Github at https://github.com/alexmul1114/fmri_encoding_contrastive_learning.

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, 3 authors, 9 MeSH terms, 6 funders, 36 references.

Cite

This paper

Mulrooney, A., Li, Z., & Brockmeier, A. J. (2026). Contrastive learning to fine-tune feature extraction models for the visual cortex. PLoS computational biology, 22(8), e1014656. https://doi.org/10.1371/journal.pcbi.1014656

BibTeX

@article{mulrooney2026contrastive,
author = {Mulrooney, Alex and Li, Zhi and Brockmeier, Austin J},
title = {{Contrastive learning to fine-tune feature extraction models for the visual cortex}},
journal = {PLoS computational biology},
year = {2026},
month = aug,
volume = {22},
number = {8},
pages = {e1014656},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014656},
url = {https://doi.org/10.1371/journal.pcbi.1014656},
pmid = {42607104},
pmcid = {PMC13492995}
}

RIS

TY - JOUR
AU - Mulrooney, Alex
AU - Li, Zhi
AU - Brockmeier, Austin J
TI - Contrastive learning to fine-tune feature extraction models for the visual cortex
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/08/17
VL - 22
IS - 8
SP - e1014656
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014656
UR - https://doi.org/10.1371/journal.pcbi.1014656
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014656",
"type": "article-journal",
"title": "Contrastive learning to fine-tune feature extraction models for the visual cortex",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Mulrooney",
"given": "Alex"
},
{
"family": "Li",
"given": "Zhi"
},
{
"family": "Brockmeier",
"given": "Austin J"
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "8",
"page": "e1014656",
"DOI": "10.1371/journal.pcbi.1014656",
"PMID": "42607104",
"PMCID": "PMC13492995",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014656",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
17
]
]
}
}

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.1038/s41597-026-07248-6 [code]
A large-scale fMRI dataset for vision-language semantic association.
Journal: Scientific data
In common: Pillow, NiBabel, PyTorch, 5 other tools, fMRI, 8 references
[2] doi:10.1038/s42003-026-10169-0 [code]
Shared representations in brains and models reveal a two-route cortical organization during scene perception.
Journal: Communications biology
In common: Pillow, NiBabel, PyTorch, 5 other tools, 7 references
[3] doi:10.7554/elife.107933 [code]
Modality-agnostic decoding of vision and language from fMRI.
Journal: eLife
In common: Pillow, NiBabel, PyTorch, 5 other tools, fMRI, 5 references
[4] doi:10.1371/journal.pcbi.1014263 [code]
MIRAGE: Robust multi-modal architectures translate fMRI-to-image models from vision to mental imagery.
Journal: PLoS computational biology
In common: Pillow, NiBabel, PyTorch, 5 other tools, fMRI, 4 references
[5] 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, NiBabel, PyTorch, 5 other tools, fMRI, 4 references
[6] doi:10.1162/imag.a.1309 [code]
Probing the content of semantic representations in body-selective regions.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Pillow, NiBabel, PyTorch, 5 other tools, 4 references
[7] doi:10.1162/imag.a.1256 [code]
Gamer in the scanner: Event-related analysis of fMRI activity during retro videogame play guided by automated annotations of game content.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: DataLad, Pillow, NiBabel, 6 other tools, fMRI, 1 reference
[8] doi:10.1038/s41467-026-76098-y [code]
A single computational objective can produce specialization of streams in visual cortex.
Journal: Nature communications
In common: Pillow, NiBabel, PyTorch, 5 other tools, 3 references
[9] doi:10.1523/jneurosci.0038-26.2026 [code]
Multidimensional Feature Tuning in Category Selective Areas of Human Visual Cortex.
Journal: The Journal of neuroscience : the official journal of the Society for Neuroscience
In common: Pillow, NiBabel, PyTorch, 5 other tools, fMRI, systems, 2 references
[10] doi:10.1167/jov.26.5.7 [code]
Representations in vision and language converge in a shared, multidimensional space of perceived similarities.
Journal: Journal of vision
In common: Pillow, NiBabel, PyTorch, 5 other tools, 3 references

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.