OSCR

Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training.

Code ↔ Paper

11 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 11 matches
  1. [1] § Methods › Simulation of monkey training using artificial neural networks › Base models ↔ analysis/paperFigures.py, lines 1514–1582 · score 0.80 · SimCLR, barlowTwins, MoCov2, ResNet18, VGG16, AlexNet
  2. [2] § Results › IT-like models best predict differences along untrained dimensions in IT ↔ analysis/paperFigures.py, lines 1692–1785 · score 0.72 · Horizontal position, Vertical position, mapped layers, rx, ry, rz
  3. [3] § Methods › Simulation of monkey training using artificial neural networks › Training strategies ↔ training/train_wandb_binaryTask.py, lines 382–522 · score 0.71 · cross entropy loss, learning rate scheduler, decay, epochs, validation, optimizing
  4. [4] § Methods › Simulation of monkey training using artificial neural networks › Training strategies ↔ training/train_wandb_v2.py, lines 18–145 · score 0.71 · cross entropy loss, learning rate scheduler, decay, epochs, validation, optimizing
  5. [5] § Methods › Neural and behavioral data analyses › Population-level object decoding ↔ analysis/categoryDecodes.py, lines 24–54 · score 0.61 · logistic regression, L2, stratified, fold, classifiers, splits
  6. [6] § Methods › Simulation of monkey training using artificial neural networks › Aligning base models to task-naïve IT responses ↔ analysis/paperFigures.py, lines 1216–1312 · score 0.59 · shifted log, 1–100, fit, accuracies, decodes, models
  7. [7] § Methods › Neural and behavioral data analyses › Neural recording quality metrics and pooling ↔ analysis/categoryDecodes.py, lines 638–740 · score 0.58 · spearman brown corrected, reliability criterion, split, repetition, stem, untrained
  8. [8] § Methods › Neural and behavioral data analyses › Behavioral metrics ↔ analysis/consistency.py, lines 15–81 · score 0.56 · Spearman Brown corrected, random splits, repetitions, reliability, behavioral, metric
  9. [9] § Results › Task-optimized ANNs reproduce IT’s task-training-induced differences ↔ analysis/paperFigures.py, lines 1788–1916 · score 0.55 · signal rotation, metrics matched, covariance, selectivity, models, training
  10. [10] § Methods › Simulation of monkey training using artificial neural networks › Aligning base models to task-naïve IT responses ↔ analysis/compute_modelFeatures.py, lines 79–195 · score 0.52 · brain score, mapped layer, cross, validated, predictivity, models
  11. [11] § Results › Task-optimized ANNs reproduce IT’s task-training-induced differences ↔ analysis/paperFigures.py, lines 1788–1916 · score 0.51 · Signal rotation, metrics matched, covariance, ratio, models, training

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 · 2,242 lines · 113 KB · MIT · 5 matches

  1. from pathlib import Path
  2. import joblib
  3. import matplotlib.pyplot as plt
  4. import numpy as np
  5. import pandas as pd
  6. import seaborn as sns
  7. from matplotlib_venn import venn3
  8. from pingouin import kruskal
  9. from scipy.optimize import curve_fit
  10. from scipy.stats import gaussian_kde
  11. from scipy.integrate import quad
  12. from analysis.categoryDecodes import load_ModelIT_categoryDecodes_allModels, load_MonkeyIT_Decodes_deltas
  13. from analysis.categoryDecodes import load_MonkeyIT_Decodes, get_ModelIT_categoryDecoding_UnitSiteScaling, \
  14. load_MonkeyIT_Decodes_permutation, load_MonkeyIT_subregion_Decodes, load_MonkeyIT_Decodes_novel500, \
  15. load_MonkeyIT_Decodes_novel500_HVMcontrol, load_MonkeyIT_Decodes_singleAnimal
  16. from analysis.consistency import get_consistency_modelIT_modelBehavior, get_consistency_behavior_ITneurons
  17. from analysis.orthogonalDecodes import load_ModelITorthogonalDecodes_selection, load_neuralITorthogonalDecodes,load_ModelLayer_orthogonalDecodes_selection
  18. from analysis.rsa import load_ModelRDMs_allModels, load_neuralRDMs_deltas
  19. from analysis.rsa import load_neuralRDMs, load_neuralRDMs_permutation, load_neuralRDMs_singleAnimal
  20. from analysis.selectivity import load_ModelSelectivity_allModels, load_neuralITselectivity_deltas
  21. from analysis.selectivity import load_neuralITselectivity, load_neuralITselectivity_permutation, load_neuralITselectivity_singleAnimal
  22. from analysis.trainingResults import loadTrainingResults
  23. from analysis.LFI import get_model_change_metrics
  24. from analysis.utils_behavioralData import load_MonkeyBehavior_individual, load_MonkeyBehavior_pooled, \
  25. getPerformanceThreshold, get_MonkeyBehavioralMetrics_draws, load_MonkeyTraining_dPrime, load_MonkeyTraining_individual
  26. from analysis.utils_neuralData import get_SiteReliability, load_MonkeyNeurons_pooled_rates, load_orthogonalFeatures
  27. from analysis.metrics import dPrime_monkey, computeDelta
  28. cm = 1 / 2.54
  29. pretraining_mapping = {'resnet18_v1': 'supervised',
  30. 'resnet34_v1': 'supervised',
  31. 'resnet50_v1': 'supervised',
  32. 'resnet50_MoCov2_200epochs': 'self-supervised',
  33. 'resnet50_barlowTwins_300epochs': 'self-supervised',
  34. 'resnet50_simclr_100epochs': 'self-supervised',
  35. 'resnet101_v1': 'supervised',
  36. 'resnet152_v1': 'supervised',
  37. 'alexnet': 'supervised',
  38. 'resnext101_32x8d_wsl': 'WSL-supervised',
  39. 'resnext101_32x16d_wsl': 'WSL-supervised',
  40. 'vgg16': 'supervised',
  41. 'vgg19': 'supervised'}
  42. def FigureS1_TrainTestSets(colors=None, figure_dir=None, source_file_dir=None):
  43. if figure_dir == None:
  44. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  45. if source_file_dir == None:
  46. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  47. train_df = load_orthogonalFeatures(split='train')
  48. train_df['category'] = train_df['obj'].map({'bear': 'bear','ELEPHANT_M': 'elephant','face0001': 'person',
  49. 'alfa155': 'car','breed_pug': 'dog','Apple_Fruit_obj': 'apple',
  50. '001': 'chair','f16': 'plane'})
  51. train_df['xpos'] = train_df['ty'] * -1
  52. train_df['ypos'] = train_df['tz']
  53. train_df['ecc'] = np.sqrt(train_df['xpos']**2 + train_df['ypos']**2)
  54. train_df['filename'] = 'im'+ train_df['legacy_mkturk_image_num'].astype(str)
  55. train_df = train_df.rename(columns={'s':'objSize'})
  56. train_df['split'] = 'train'
  57. test_df = load_orthogonalFeatures(split='test')
  58. test_df['split'] = 'test'
  59. df = pd.concat([train_df, test_df], ignore_index=True)
  60. title_mapping={'xpos': 'Vertical position',
  61. 'ypos': 'Horizontal position',
  62. 'objSize': 'Object size',
  63. 'ecc': 'Eccentricity',
  64. 'ryz': 'Rotation (rx)',
  65. 'rxz': 'Rotation (ry)',
  66. 'rxy': 'Rotation (rz)',}
  67. g, ax = plt.subplots(2,4, sharex=False, sharey=True, figsize=(15*cm,10*cm), squeeze=False)
  68. ax = ax.flatten()
  69. for i, var in enumerate(['xpos', 'ypos', 'objSize','ecc', 'ryz', 'rxz', 'rxy']):
  70. sns.histplot(data=df, hue='split', kde=True, x=var, bins=100,ax=ax[i], legend=False)
  71. ax[i].set_xlabel(f"{title_mapping[var]}")
  72. ax[i].set_ylabel(' ')
  73. ax[0].set_ylabel('# images')
  74. plt.tight_layout()
  75. g.savefig(figure_dir + 'FigureS1_datasets_v0.png', dpi=300)
  76. g.savefig(figure_dir + 'FigureS1_datasets_v0.pdf', dpi=300, transparent=True)
  77. df.to_csv(source_file_dir + 'FigureS1_datasets.csv', index=False)
  78. def FigureS2_trainingCurves(figure_dir=None, source_file_dir = None):
  79. if figure_dir == None:
  80. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  81. if source_file_dir == None:
  82. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  83. df_trials = load_MonkeyTraining_individual()
  84. num_trials_day = df_trials.groupby(['day', 'subject'], as_index=False)['trial'].nunique()
  85. g = sns.catplot(data=num_trials_day, x='day', hue='subject', y='trial', kind='bar',
  86. height=6 * cm, legend=False, aspect=1.6)
  87. g.set_xlabels('Training day')
  88. g.set_ylabels("# trials")
  89. g.savefig(figure_dir + 'FigureS2_LearningCurves_v2.png', dpi=300)
  90. g.savefig(figure_dir + 'FigureS2_LearningCurves_v2.pdf', dpi=300, transparent=True)
  91. num_trials_day.to_csv(source_file_dir + 'FigureS2_LearningCurves_v2.csv', index=False)
  92. # compute duration in number of days
  93. print(df_trials.groupby('subject')['date'].max() - df_trials.groupby('subject')['date'].min())
  94. df_dPrime = load_MonkeyTraining_dPrime()
  95. g = sns.relplot(data=df_dPrime, x= 'start_trial', y='dPrime', hue='subject', kind='line',
  96. height=6*cm, legend=False, aspect=1.2)
  97. g.set_xlabels('Trial')
  98. g.set_ylabels("Performance (d')")
  99. g.refline(x=0)
  100. g.refline(x=17400)
  101. g.savefig(figure_dir + 'FigureS2_LearningCurves_v0.png', dpi=300)
  102. g.savefig(figure_dir + 'FigureS2_LearningCurves_v0.pdf', dpi=300, transparent=True)
  103. df_dPrime.to_csv(source_file_dir + 'FigureS2_LearningCurves_v0.csv', index=False)
  104. g = sns.catplot(data=df_dPrime.loc[df_dPrime['start_trial'].isin([0, 17400])],
  105. x='start_trial', y='dPrime', hue='subject', kind='point',
  106. height=6 * cm, legend=False, aspect=0.8, dodge=0.2, join=False)
  107. g.set_ylabels("Performance (d')")
  108. g.set_xlabels("Training phase")
  109. g.set(ylim=[0, 4.5])
  110. g.set_xticklabels(['Early', 'Late'])
  111. g.savefig(figure_dir + 'FigureS2_LearningCurves_v1.png', dpi=300)
  112. g.savefig(figure_dir + 'FigureS2_LearningCurves_v1.pdf', dpi=300, transparent=True)
  113. df_dPrime.loc[df_dPrime['start_trial'].isin([0, 17400])].to_csv(source_file_dir + 'FigureS2_LearningCurves_v1.csv', index=False)
  114. def Figure1_monkeyPerformance(colors=None, figure_dir=None, data_dir=None, source_file_dir = None):
  115. if figure_dir == None:
  116. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  117. if source_file_dir == None:
  118. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  119. monkey_performance = load_MonkeyBehavior_individual(data_dir=data_dir)
  120. performance_subject = monkey_performance.groupby('subject').mean()
  121. performance_subject['trained?'] = 'yes'
  122. order = ['no', 'yes']
  123. g = sns.catplot(data=performance_subject, x='trained?', y='i1', hue='trained?', kind='point', errorbar=("ci", 95),
  124. order=order,
  125. palette={'no': colors[0], 'yes': sns.xkcd_palette(['raspberry'])[0]}, height=4.5 * cm, aspect=0.85,
  126. legend=False)
  127. g.refline(y=0)
  128. g.set_ylabels("Performance (d')")
  129. g.set(ylim=[-1, 3.5])
  130. g.savefig(figure_dir + 'Figure1_performance_i1.png', dpi=300)
  131. g.savefig(figure_dir + 'Figure1_performance_i1.pdf', dpi=300, transparent=True)
  132. performance_subject.to_csv(source_file_dir + 'Figure1_performance_i1.csv')
  133. def Figure1_trainingCurves(figure_dir = None, source_file_dir = None):
  134. if figure_dir == None:
  135. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  136. if source_file_dir == None:
  137. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  138. df_dPrime = load_MonkeyTraining_dPrime()
  139. df_dPrime_mean = df_dPrime.groupby(['class', 'start_trial'], as_index=False)['dPrime'].mean()
  140. g = sns.relplot(data=df_dPrime_mean.loc[df_dPrime_mean['start_trial'] <= 17400], x='start_trial', y='dPrime', kind='line',
  141. legend=False, height=4.5*cm, aspect=1.2, color=sns.xkcd_palette(['raspberry'])[0])
  142. g.set_xlabels('Trial')
  143. g.set_ylabels("Performance (d')")
  144. g.refline(y=0)
  145. g.savefig(figure_dir + 'Figure1_LearningCurves.png', dpi=300)
  146. g.savefig(figure_dir + 'Figure1_LearningCurves.pdf', dpi=300, transparent=True)
  147. df_dPrime_mean.loc[df_dPrime_mean['start_trial'] <= 17400].to_csv(source_file_dir + 'Figure1_LearningCurves.csv', index=False)
  148. def FigureS4_SiteReliability(figure_dir=None, source_file_dir=None, colors=None, recompute=False, criterion=0.3, n_sites_per_sub=53,
  149. draws=1000, version='v0',
  150. mode='empirical_conditioned_v2'):
  151. if figure_dir == None:
  152. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  153. if source_file_dir == None:
  154. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  155. if version in ['v1', 'v2']:
  156. rates = load_MonkeyNeurons_pooled_rates(subsampled_neurons=n_sites_per_sub, draws=draws,
  157. mode=mode, recompute=recompute, reliability_criterion=criterion)
  158. if version == 'v0':
  159. reliabilities = get_SiteReliability(recompute=recompute)
  160. df = []
  161. for state in ['trained', 'untrained']:
  162. subs = reliabilities[state].keys()
  163. for sub in subs:
  164. df.append(pd.DataFrame({'reliabilities': np.mean(reliabilities[state][sub], axis=1),
  165. 'state': state,
  166. 'subject': sub}))
  167. df = pd.concat(df, ignore_index=True)
  168. g = sns.displot(data=df, x='reliabilities', hue='state', palette={'trained': colors[1], 'untrained': colors[0]},
  169. height=5 * cm, aspect=1.2,
  170. bins=30)
  171. g.set_ylabels("# sites")
  172. g.refline(x=criterion)
  173. g.savefig(figure_dir + 'FigureS4_reliability.png', dpi=300)
  174. g.savefig(figure_dir + 'FigureS4_reliability.pdf', dpi=300, transparent=True)
  175. df.to_csv(source_file_dir + 'FigureS4_reliability.csv', index=False)
  176. elif version == 'v1':
  177. d = np.random.randint(draws)
  178. df = []
  179. for state in ['trained', 'untrained']:
  180. df.append(pd.DataFrame({'reliabilities': rates[f'reliability_{state}'][d],
  181. 'state': state,
  182. }))
  183. df = pd.concat(df, ignore_index=True)
  184. g = sns.displot(data=df, x='reliabilities', hue='state', palette={'trained': colors[1], 'untrained': colors[0]},
  185. height=5 * cm, aspect=1.2,
  186. bins=30)
  187. g.refline(x=criterion)
  188. g.savefig(figure_dir + 'FigureS4_reliability_v1.png', dpi=300)
  189. g.savefig(figure_dir + 'FigureS4_reliability_v1.pdf', dpi=300, transparent=True)
  190. elif version == 'v2':
  191. diff_dist = np.median(rates[f'reliability_trained'], axis=1) - np.median(rates[f'reliability_untrained'],
  192. axis=1)
  193. g = sns.displot(x=diff_dist,
  194. height=4.5 * cm, aspect=1.2,
  195. bins=30)
  196. g.refline(x=0, c='white')
  197. g.set_ylabels("# draws")
  198. g.set_xlabels("reliability differences\n(median in pools)")
  199. g.savefig(figure_dir + 'FigureS4_reliability_v2.png', dpi=300)
  200. g.savefig(figure_dir + 'FigureS4_reliability_v2.pdf', dpi=300, transparent=True)
  201. elif version == 'v3':
  202. df = []
  203. for hemi in ['LH', 'RH']:
  204. reliabilities = get_SiteReliability(recompute=recompute, area='IT', hemisphere=hemi)
  205. for state in ['trained', 'untrained']:
  206. subs = reliabilities[state].keys()
  207. for sub in subs:
  208. df.append(pd.DataFrame({'reliabilities': np.mean(reliabilities[state][sub], axis=1),
  209. 'state': state,
  210. 'hemisphere': hemi,
  211. 'subject': sub}))
  212. df = pd.concat(df, ignore_index=True)
  213. g = sns.displot(data=df.loc[df['state'] == 'untrained'], x='reliabilities', col='hemisphere', row='subject',
  214. height=3.5 * cm, aspect=1.2,
  215. bins=30, col_order=['LH', 'RH'], row_order=['monkeyT', 'monkeyC', 'monkeyS'])
  216. g.set_titles(' ')
  217. g.set(ylim=[0, 40])
  218. for condition in g.axes_dict:
  219. subject, hemi = condition[0], condition[1]
  220. site_yield = (df.loc[(df['state'] == 'untrained') & (df['subject'] == subject) & (df['hemisphere'] == hemi), 'reliabilities'] > criterion).sum()
  221. total = len(df.loc[(df['state'] == 'untrained') & (df['subject'] == subject) & (df['hemisphere'] == hemi), 'reliabilities'])
  222. g.axes_dict[condition].text(-0.2, 35, f"{site_yield}/{total}", size=8)
  223. g.set_ylabels("# sites")
  224. g.refline(x=criterion)
  225. g.savefig(figure_dir + 'FigureS3_reliability_naive_v3.png', dpi=300)
  226. g.savefig(figure_dir + 'FigureS3_reliability_naive_v3.pdf', dpi=300, transparent=True)
  227. df.loc[df['state'] == 'untrained'].to_csv(source_file_dir + 'FigureS3_reliability_naive_v3.csv', index=False)
  228. g = sns.displot(data=df.loc[df['state'] == 'trained'], x='reliabilities', col='hemisphere', row='subject',
  229. height=3.5 * cm, aspect=1.2,
  230. col_order=['LH', 'RH'], row_order=['monkeyN', 'monkeyM', 'monkeyB'],
  231. bins=30)
  232. g.set_titles(' ')
  233. g.set(ylim=[0, 40])
  234. for condition in g.axes_dict:
  235. subject, hemi = condition[0], condition[1]
  236. site_yield = (df.loc[(df['state'] == 'trained') & (df['subject'] == subject) & (
  237. df['hemisphere'] == hemi), 'reliabilities'] > criterion).sum()
  238. total = len(df.loc[(df['state'] == 'trained') & (df['subject'] == subject) & (
  239. df['hemisphere'] == hemi), 'reliabilities'])
  240. g.axes_dict[condition].text(-0.2, 35, f"{site_yield}/{total}", size=8)
  241. g.set_ylabels("# sites")
  242. g.refline(x=criterion)
  243. g.savefig(figure_dir + 'FigureS3_reliability_trained_v3.png', dpi=300)
  244. g.savefig(figure_dir + 'FigureS3_reliability_trained_v3.pdf', dpi=300, transparent=True)
  245. df.loc[df['state'] == 'trained'].to_csv(source_file_dir + 'FigureS3_reliability_trained_v3.csv', index=False)
  246. def Figure2_selectivity(figure_dir=None, source_file_dir=None, colors=None,
  247. recompute=False, num_sites=159,
  248. num_subs=3, draws=1000,
  249. mode='empirical_conditioned_v2',
  250. version='v0', n_jobs=-1):
  251. if figure_dir == None:
  252. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  253. if source_file_dir == None:
  254. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  255. if version == 'v0':
  256. neural_df = load_neuralITselectivity(recompute=recompute, draws=1, n_jobs=n_jobs, include='combinedCategories',
  257. mode='empirical_all')
  258. df_units = []
  259. for state in ['untrained', 'trained']:
  260. df_units.append(
  261. pd.DataFrame({'state': state, 'category_selectivity': np.array(neural_df[state]).flatten()}))
  262. df_units = pd.concat(df_units, ignore_index=True)
  263. print(df_units.groupby('state').size())
  264. g = sns.displot(data=df_units, x='category_selectivity', hue='state', stat='percent', common_norm=False,
  265. kde=True, palette=colors, height=5 * cm, aspect=1.25, legend=False)
  266. g.refline(x=0, ls='--', color='white')
  267. g.set(xticks=[0, 0.1, 0.2], yticks=[0, 10, 20])
  268. g.set_xlabels('Category selectivity')
  269. g.set_ylabels('% neural sites')
  270. g.savefig(figure_dir + 'Figure2_selectivity_v0.png', dpi=300)
  271. g.savefig(figure_dir + 'Figure2_selectivity_v0.pdf', dpi=300, transparent=True)
  272. df_units.index.name = 'unit'
  273. df_units.to_csv(source_file_dir + 'Figure2_selectivity_v0.csv')
  274. elif version == 'v1':
  275. neural_df_draws = load_neuralITselectivity(recompute=recompute, draws=draws, n_jobs=n_jobs,
  276. include='combinedCategories',
  277. mode=mode, neural_sites_perSubject=num_sites // num_subs)
  278. draws = len(neural_df_draws['trained'])
  279. df = []
  280. for state in ['untrained', 'trained']:
  281. for d in range(draws):
  282. df.append(
  283. {'state': state, 'draw': d, 'selectivity': np.median(neural_df_draws[f'{state}'][d])})
  284. df = pd.DataFrame(df)
  285. g = sns.catplot(data=df, x='state', y='selectivity', hue='state', kind='point', errorbar=("pi", 95),
  286. palette=colors, height=5 * cm, aspect=1,
  287. legend=False)
  288. g.set(ylim=[0.02, 0.06], yticks=[0.02, 0.04, 0.06])
  289. g.set_xticklabels(["naive", "trained"])
  290. g.set_xlabels(f'IT neuron pool\n({num_sites} sites)')
  291. g.set_ylabels("Selectivity")
  292. g.savefig(figure_dir + 'Figure2_selectivity_v1.png', dpi=300)
  293. g.savefig(figure_dir + 'Figure2_selectivity_v1.pdf', dpi=300, transparent=True)
  294. print('Medians:')
  295. print(df.groupby('state').median())
  296. print(
  297. f"{((df.groupby('state').median().loc['trained', 'selectivity'] / df.groupby('state').median().loc['untrained', 'selectivity']) - 1) * 100}% increase")
  298. print(
  299. f"CI (95%): {((df.loc[df['state'] == 'trained', 'selectivity'].reset_index() - df.loc[df['state'] == 'untrained', 'selectivity'].reset_index())['selectivity']).quantile(0.025)} - {((df.loc[df['state'] == 'trained', 'selectivity'].reset_index() - df.loc[df['state'] == 'untrained', 'selectivity'].reset_index())['selectivity']).quantile(0.975)}")
  300. df.to_csv(source_file_dir + 'Figure2_selectivity_v1.csv', index=False)
  301. def Figure2_rsa(figure_dir=None, source_file_dir=None,
  302. neural_sites_perSubject=53, draws=1000,
  303. colors=None, recompute=False,
  304. version='v0', mode='empirical_conditioned_v2'):
  305. if figure_dir == None:
  306. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  307. if source_file_dir == None:
  308. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  309. n_sites = np.unique(np.geomspace(1, neural_sites_perSubject, num=12, dtype='int'))
  310. print(n_sites)
  311. neural_df = load_neuralRDMs(recompute=recompute, draws=draws, mode=mode,
  312. neural_sites_perSubject=n_sites)
  313. if version == 'v0':
  314. df = []
  315. for state in ['untrained', 'trained']:
  316. for d in range(draws):
  317. df.append(
  318. {'state': state, 'draw': d,
  319. 'tau': np.mean(neural_df[f'{neural_sites_perSubject * 3}_sites'][f'{state}_category'][d])})
  320. df = pd.DataFrame(df)
  321. g = sns.catplot(data=df, x='state', y='tau', hue='state', kind='point', errorbar=("pi", 95),
  322. palette=colors, height=5 * cm, aspect=1,
  323. legend=False)
  324. g.set(ylim=[0.065, 0.13], yticks=[0.08, 0.1, 0.12])
  325. g.set_xlabels('IT neuron pool')
  326. g.set_ylabels("Object-level\ncorrelation (τ)")
  327. g.savefig(figure_dir + 'Figure2_rsa_v0.png', dpi=300)
  328. g.savefig(figure_dir + 'Figure2_rsa_v0.pdf', dpi=300, transparent=True)
  329. print('Means:')
  330. print(df.groupby('state').mean())
  331. print('SD:')
  332. print(df.groupby('state').std())
  333. print(
  334. f"{((df.groupby('state').median().loc['trained', 'tau'] / df.groupby('state').median().loc['untrained', 'tau']) - 1) * 100}% increase")
  335. print(
  336. f"CI (95%): {((df.loc[df['state'] == 'trained', 'tau'].reset_index() - df.loc[df['state'] == 'untrained', 'tau'].reset_index())['tau']).quantile(0.025)} - {((df.loc[df['state'] == 'trained', 'tau'].reset_index() - df.loc[df['state'] == 'untrained', 'tau'].reset_index())['tau']).quantile(0.975)}")
  337. print(
  338. f"Mean difference: {((df.loc[df['state'] == 'trained', 'tau'].reset_index() - df.loc[df['state'] == 'untrained', 'tau'].reset_index())['tau']).mean()}")
  339. print(
  340. f"SD (diff): {((df.loc[df['state'] == 'trained', 'tau'].reset_index() - df.loc[df['state'] == 'untrained', 'tau'].reset_index())['tau']).std()}")
  341. df.to_csv(source_file_dir + 'Figure2_rsa_v0.csv', index=False)
  342. elif version == 'v1':
  343. df = []
  344. for n in n_sites:
  345. for state in ['untrained', 'trained']:
  346. for d in range(draws):
  347. df.append(
  348. {'sites': 3 * n, 'state': state, 'draw': d,
  349. 'tau': np.mean(neural_df[f'{n * 3}_sites'][f'{state}_category'][d])})
  350. df = pd.DataFrame(df)
  351. g = sns.relplot(data=df, x='sites', y='tau', hue='state', kind='line', errorbar=("pi", 95),
  352. palette=colors, height=5 * cm, aspect=1.3,
  353. legend=False)
  354. g.set(xlim=[-5, 200], xticks=[0, 100, 200], ylim=[0, 0.15])
  355. g.set_ylabels("Object-level\ncorrelation (τ)")
  356. g.set_xlabels("IT neuron pool size")
  357. g.savefig(figure_dir + 'Figure2_rsa_v1.png', dpi=300)
  358. g.savefig(figure_dir + 'Figure2_rsa_v1.pdf', dpi=300, transparent=True)
  359. df.to_csv(source_file_dir + 'Figure2_rsa_v1.csv', index=False)
  360. def Figure2_decoding(figure_dir=None, source_file_dir=None, colors=None, version='v1', recompute=False,
  361. performance_threshold=None, num_sites=159, mode='empirical_conditioned_v2', draws=1000):
  362. if figure_dir == None:
  363. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  364. if source_file_dir == None:
  365. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  366. if version == 'v0':
  367. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False,
  368. recompute=recompute,
  369. mode=mode, draws=draws)
  370. df = []
  371. for state in ['untrained', 'trained']:
  372. for d in range(draws):
  373. df.append(
  374. {'state': state, 'draw': d, 'i1': np.mean(IT_neurons[f'{num_sites}_neurons'][state]['dprimes'][d])})
  375. df = pd.DataFrame(df)
  376. g = sns.catplot(data=df, x='state', y='i1', hue='state', kind='point', errorbar=("pi", 95),
  377. palette=colors, height=5 * cm, aspect=1,
  378. legend=False)
  379. g.set(ylim=[2.5, 3.5], yticks=[2.5, 3, 3.5])
  380. g.set_titles('Image-by-image\ndecoding')
  381. g.set_xticklabels(["naive", "trained"])
  382. g.set_ylabels("Performance (d')")
  383. g.set_xlabels("IT neuron pool\n(159 sites)")
  384. g.savefig(figure_dir + 'Figure2_decoding_v0.png', dpi=300)
  385. g.savefig(figure_dir + 'Figure2_decoding_v0.pdf', dpi=300, transparent=True)
  386. print('Means:')
  387. print(df.groupby('state').mean())
  388. print('SD:')
  389. print(df.groupby('state').std())
  390. print(
  391. f"{((df.groupby('state').median().loc['trained', 'i1'] / df.groupby('state').median().loc['untrained', 'i1']) - 1) * 100}% increase")
  392. print(
  393. f"CI (95%): {((df.loc[df['state'] == 'trained', 'i1'].reset_index() - df.loc[df['state'] == 'untrained', 'i1'].reset_index())['i1']).quantile(0.025)} - {((df.loc[df['state'] == 'trained', 'i1'].reset_index() - df.loc[df['state'] == 'untrained', 'i1'].reset_index())['i1']).quantile(0.975)}")
  394. print(
  395. f"Mean difference: {((df.loc[df['state'] == 'trained', 'i1'].reset_index() - df.loc[df['state'] == 'untrained', 'i1'].reset_index())['i1']).mean()}")
  396. print(
  397. f"SD (diff): {((df.loc[df['state'] == 'trained', 'i1'].reset_index() - df.loc[df['state'] == 'untrained', 'i1'].reset_index())['i1']).std()}")
  398. df.to_csv(source_file_dir + 'Figure2_decoding_v0.csv', index=False)
  399. elif version == 'v1':
  400. mean_performance = getPerformanceThreshold('test_d-prime', kind=('mean', None))
  401. n_sites = np.unique((np.geomspace(3, num_sites, num=12, dtype='int') // 3) * 3)
  402. dPrimes = load_MonkeyIT_Decodes(resultFile_stem='CategoryDecodes_ITneurons_manySizes_', num_neurons=n_sites,
  403. includeReliability=False,
  404. recompute=recompute, mode=mode, draws=draws)
  405. df = []
  406. for n in n_sites:
  407. for state in ['untrained', 'trained']:
  408. for d in range(draws):
  409. df.append(
  410. {'state': state, 'draw': d, 'sites': n,
  411. 'i1': np.mean(dPrimes[f'{n}_neurons'][state]['dprimes'][d])})
  412. df = pd.DataFrame(df)
  413. def func(x, a, b, c): # x-shifted log
  414. return a * np.log(x + b) + c
  415. fitting_df_trained = df[df['state'] == 'trained'].groupby('sites', as_index=False)['i1'].mean()
  416. fitting_df_untrained = df[df['state'] == 'untrained'].groupby('sites', as_index=False)['i1'].mean()
  417. popt_t, pcov_t = curve_fit(func, fitting_df_trained.sites.values, fitting_df_trained.i1.values)
  418. popt_u, pcov_u = curve_fit(func, fitting_df_untrained.sites.values, fitting_df_untrained.i1.values)
  419. fitting_df_trained['predicted_i1'] = func(fitting_df_trained.sites.values, *popt_t)
  420. fitting_df_untrained['predicted_i1'] = func(fitting_df_untrained.sites.values, *popt_u)
  421. extrapolation_df_trained = pd.DataFrame(
  422. {'sites': np.arange(1, 200), 'i1_predicted': func(np.arange(1, 200), *popt_t)})
  423. extrapolation_df_untrained = pd.DataFrame(
  424. {'sites': np.arange(1, 200), 'i1_predicted': func(np.arange(1, 200), *popt_u)})
  425. g = sns.relplot(data=df, x='sites', y='i1', hue='state', kind='line', errorbar=("pi", 95),
  426. palette=colors, height=5 * cm, aspect=1.2,
  427. legend=False)
  428. g.ax.plot(extrapolation_df_trained.sites.values, extrapolation_df_trained['i1_predicted'].values, ls=':',
  429. color=sns.xkcd_palette(['jade'])[0])
  430. g.ax.plot(extrapolation_df_untrained.sites.values, extrapolation_df_untrained['i1_predicted'].values, ls=':',
  431. color=sns.xkcd_palette(['slate grey'])[0])
  432. g.refline(y=mean_performance['min_value'], color='grey', zorder=1)
  433. g.map(plt.axhspan, ymin=performance_threshold['min_value'], ymax=performance_threshold['max_value'], zorder=0,
  434. color='grey', alpha=0.4)
  435. print(
  436. f"Untrained: {extrapolation_df_untrained.loc[extrapolation_df_untrained['i1_predicted'] >= performance_threshold['min_value'], 'sites'].min()}")
  437. print(
  438. f"Trained: {extrapolation_df_trained.loc[extrapolation_df_trained['i1_predicted'] >= performance_threshold['min_value'], 'sites'].min()}")
  439. g.set(xlim=[-5, 200])
  440. g.set_titles('Image-by-image\ndecoding')
  441. g.set_ylabels("Performance (d')")
  442. g.set_xlabels("IT neuron pool")
  443. g.savefig(figure_dir + 'Figure2_decoding_v1.png', dpi=300)
  444. g.savefig(figure_dir + 'Figure2_decoding_v1.pdf', dpi=300, transparent=True)
  445. df.to_csv(source_file_dir + 'Figure2_decoding_v1.csv', index=False)
  446. elif version == 'v2':
  447. monkey_behavior = get_MonkeyBehavioralMetrics_draws(recompute=recompute, draws=draws)
  448. monkey_behavior = monkey_behavior.rename(columns={'test_d-prime': 'delta_i1'})
  449. monkey_behavior.drop(columns=['test_acc_binary', 'test_d-prime_o1'], inplace=True)
  450. monkey_behavior['data'] = 'Behavior'
  451. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False,
  452. recompute=recompute,
  453. mode=mode, draws=draws)
  454. df = []
  455. for d in range(draws):
  456. df.append(
  457. {'data': 'IT', 'draw': d,
  458. 'delta_i1': np.mean(IT_neurons[f'{num_sites}_neurons']['trained']['dprimes'][d]) - np.mean(
  459. IT_neurons[f'{num_sites}_neurons']['untrained']['dprimes'][d])})
  460. df = pd.DataFrame(df)
  461. df = pd.concat([df, monkey_behavior])
  462. g = sns.catplot(data=df, x='data', y='delta_i1', kind='point', errorbar=("pi", 95),
  463. palette={'Behavior': sns.xkcd_palette(['raspberry'])[0],
  464. 'IT': sns.xkcd_palette(['royal blue'])[0]}, height=5 * cm, aspect=0.6, hue='data',
  465. legend=False, dodge=False)
  466. g.set(ylim=[0, 4], yticks=[0, 2, 4])
  467. g.set_xticklabels(rotation=90)
  468. g.set_ylabels("Δ performance (d')")
  469. g.savefig(figure_dir + 'Figure2_decoding_v2.png', dpi=300)
  470. g.savefig(figure_dir + 'Figure2_decoding_v2.pdf', dpi=300, transparent=True)
  471. df.to_csv(source_file_dir + 'Figure2_decoding_v2.csv', index=False)
  472. def FigureS5_decoding_reliability(reliabilities=[0.3, 0.5, 0.7], num_sites_all=[159, 147, 102],
  473. figure_dir=None, source_file_dir=None, colors=None, recompute=False,
  474. mode='empirical_conditioned_v2', draws=1000):
  475. if figure_dir == None:
  476. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  477. if source_file_dir == None:
  478. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  479. df = []
  480. for reliability, num_sites in zip(reliabilities, num_sites_all):
  481. if reliability == 0.3:
  482. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False,
  483. recompute=recompute,
  484. mode=mode, draws=draws)
  485. else:
  486. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False,
  487. recompute=recompute,
  488. mode=mode, draws=draws, reliability_criterion=reliability)
  489. for state in ['untrained', 'trained']:
  490. for d in range(draws):
  491. df.append(
  492. {'state': state, 'draw': d, 'i1': np.mean(IT_neurons[f'{num_sites}_neurons'][state]['dprimes'][d]),
  493. 'reliability criterion': reliability})
  494. df = pd.DataFrame(df)
  495. g = sns.catplot(data=df, x='state', y='i1', hue='state', kind='point', errorbar=("pi", 95),
  496. palette=colors, height=5 * cm, aspect=1, col= 'reliability criterion',
  497. legend=False)
  498. g.set_titles('{col_name}')
  499. g.set_xticklabels(["naive", "trained"])
  500. g.set_ylabels("Performance (d')")
  501. g.set_xlabels("IT neuron pool")
  502. g.savefig(figure_dir + 'FigureS5_decoding_reliabilityCriterion.png', dpi=300)
  503. g.savefig(figure_dir + 'FigureS5_decoding_reliabilityCriterion.pdf', dpi=300, transparent=True)
  504. df.to_csv(source_file_dir + 'FigureS5_decoding_reliabilityCriterion.csv', index=False)
  505. def FigureS6_singleSubject_analyses(figure_dir=None, source_file_dir=None, colors=None, recompute=False, reliability_criterion=0.3):
  506. if figure_dir == None:
  507. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  508. if source_file_dir == None:
  509. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  510. IT_decoding = load_MonkeyIT_Decodes_singleAnimal(recompute=recompute,reliability_criterion=reliability_criterion)
  511. IT_selectivity = load_neuralITselectivity_singleAnimal(recompute=recompute,reliability_criterion=reliability_criterion)
  512. IT_rsa = load_neuralRDMs_singleAnimal(recompute=recompute,reliability_criterion=reliability_criterion)
  513. df = []
  514. for state in ['untrained', 'trained']:
  515. subs = list(IT_decoding[state].keys())
  516. for sub in subs:
  517. df.append(
  518. {'state': state, 'subject': sub[-1:],
  519. 'i1': np.mean(IT_decoding[state][sub]['i1']),
  520. 'selectivity': np.median(IT_selectivity[state][sub]),
  521. 'tau': float(np.squeeze(IT_rsa[state][sub])),
  522. 'num_sites': np.mean(IT_decoding[state][sub]['n_sites'])})
  523. df = pd.DataFrame(df)
  524. g = sns.catplot(data=df, x='subject', y='i1', hue='state', kind='point',
  525. errorbar=("pi", 95),
  526. palette=colors, height=5 * cm, aspect=1,
  527. legend=False, join=False)
  528. g.set(ylim=[1, 4])
  529. g.set_ylabels("Performance (d')\n ")
  530. g.refline(y=df.loc[df['state'] == 'trained']['i1'].mean(), color=colors[1])
  531. g.refline(y=df.loc[df['state'] == 'untrained']['i1'].mean(), color=colors[0])
  532. g.set_xlabels("Subject")
  533. g.savefig(figure_dir + 'FigureS6_singleSubject_decoding_i1.png', dpi=300)
  534. g.savefig(figure_dir + 'FigureS6_singleSubject_decoding_i1.pdf', dpi=300, transparent=True)
  535. g = sns.catplot(data=df, x='subject', y='selectivity', hue='state', kind='point',
  536. errorbar=("pi", 95),
  537. palette=colors, height=5 * cm, aspect=1,
  538. legend=False, join=False)
  539. g.set(ylim=[0, 0.08])
  540. g.refline(y=df.loc[df['state'] == 'trained']['selectivity'].mean(), color=colors[1])
  541. g.refline(y=df.loc[df['state'] == 'untrained']['selectivity'].mean(), color=colors[0])
  542. g.set_ylabels("Selectivity")
  543. g.set_xlabels("Subject")
  544. g.savefig(figure_dir + 'FigureS6_singleSubject_selectivity.png', dpi=300)
  545. g.savefig(figure_dir + 'FigureS6_singleSubject_selectivity.pdf', dpi=300, transparent=True)
  546. g = sns.catplot(data=df, x='subject', y='tau', hue='state', kind='point',
  547. errorbar=("pi", 95),
  548. palette=colors, height=5 * cm, aspect=1,
  549. legend=False, join=False)
  550. g.set(ylim=[0.04, 0.12])
  551. g.set_ylabels("Correlation (τ)")
  552. g.set_xlabels("Subject")
  553. g.refline(y=df.loc[df['state'] == 'trained']['tau'].mean(), color=colors[1])
  554. g.refline(y=df.loc[df['state'] == 'untrained']['tau'].mean(), color=colors[0])
  555. g.savefig(figure_dir + 'FigureS6_singleSubject_rsa.png', dpi=300)
  556. g.savefig(figure_dir + 'FigureS6_singleSubject_rsa.pdf', dpi=300, transparent=True)
  557. df.to_csv(source_file_dir + 'FigureS6_singleSubject.csv', index=False)
  558. def FigureS7_selectivity(figure_dir=None, source_file_dir=None, recompute=False, n_jobs=-1,
  559. num_sites=159, mode='empirical_conditioned_v2', draws=200):
  560. if figure_dir == None:
  561. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  562. if source_file_dir == None:
  563. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  564. IT_neurons_null = load_neuralITselectivity_permutation(neural_sites_perSubject=num_sites // 3,
  565. include='combinedCategories',
  566. recompute=recompute, mode=mode, draws=draws, n_jobs=n_jobs)
  567. IT_neurons_observed = load_neuralITselectivity(recompute=False, n_jobs=n_jobs,
  568. include='combinedCategories',
  569. mode=mode, neural_sites_perSubject=num_sites // 3)
  570. deltas = pd.DataFrame()
  571. for permutation in IT_neurons_null.keys():
  572. if permutation == 'sites':
  573. pass
  574. else:
  575. deltas = pd.concat([deltas,
  576. pd.DataFrame({'permutation': str(permutation),
  577. 'delta': np.median(IT_neurons_null[permutation]['group1'],
  578. axis=1) - np.median(
  579. IT_neurons_null[permutation]['group2'], axis=1), 'kind': 'null'})],
  580. ignore_index=True)
  581. deltas = pd.concat([deltas, pd.DataFrame({'permutation': 'None',
  582. 'delta': np.median(IT_neurons_observed['trained'], axis=1)[
  583. :draws] - np.median(IT_neurons_observed['untrained'], axis=1)[
  584. :draws], 'kind': 'empirical'})], ignore_index=True)
  585. g = sns.displot(data=deltas[deltas['kind'] == 'null'], x='delta', hue='kind', palette='colorblind', height=5 * cm,
  586. aspect=1.2)
  587. g.refline(x=deltas.groupby('kind')['delta'].mean()['empirical'], c='red')
  588. g.set_xlabels("differences (Δ)")
  589. print(
  590. f"p-value (two-sided): {(deltas.groupby('kind')['delta'].mean().loc['empirical'] <= deltas.loc[deltas['kind'] == 'null', 'delta']).mean()}")
  591. g.savefig(figure_dir + 'FigureS7_permutation_selectivity.png', dpi=300)
  592. g.savefig(figure_dir + 'FigureS7_permutation_selectivity.pdf', dpi=300, transparent=True)
  593. deltas.to_csv(source_file_dir + 'FigureS7_permutation_selectivity.csv', index=False)
  594. def FigureS7_rsa(figure_dir=None, source_file_dir=None, recompute=False, subs=3,
  595. neural_sites_perSubject=53, mode='empirical_conditioned_v2', draws=200):
  596. if figure_dir == None:
  597. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  598. if source_file_dir == None:
  599. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  600. IT_neurons_null = load_neuralRDMs_permutation(neural_sites_perSubject=np.array([neural_sites_perSubject]),
  601. recompute=recompute,
  602. mode=mode, draws=draws)
  603. IT_neurons_observed = load_neuralRDMs(recompute=False, mode=mode,
  604. neural_sites_perSubject=np.array([neural_sites_perSubject]))
  605. num_sites = neural_sites_perSubject * subs
  606. deltas = pd.DataFrame()
  607. for permutation in IT_neurons_null[f'{num_sites}_sites'].keys():
  608. if permutation == 'sites':
  609. pass
  610. else:
  611. deltas = pd.concat([deltas,
  612. pd.DataFrame(
  613. {'delta': IT_neurons_null[f'{num_sites}_sites'][permutation]['group1_category']
  614. - IT_neurons_null[f'{num_sites}_sites'][permutation]['group2_category'],
  615. 'kind': 'null'})], ignore_index=True)
  616. deltas = pd.concat(
  617. [deltas, pd.DataFrame({'delta': IT_neurons_observed[f'{num_sites}_sites']['trained_category'][:draws] -
  618. IT_neurons_observed[f'{num_sites}_sites']['untrained_category'][:draws],
  619. 'kind': 'empirical'})], ignore_index=True)
  620. g = sns.displot(data=deltas[deltas['kind'] == 'null'], x='delta', hue='kind', palette='colorblind', height=5 * cm,
  621. aspect=1.2)
  622. g.refline(x=deltas.groupby('kind')['delta'].mean()['empirical'], c='red')
  623. g.set_xlabels("differences (Δ)")
  624. print(
  625. f"p-value (two-sided): {(deltas.groupby('kind')['delta'].mean().loc['empirical'] <= deltas.loc[deltas['kind'] == 'null', 'delta']).mean()}")
  626. g.savefig(figure_dir + 'FigureS7_permutation_rsa.png', dpi=300)
  627. g.savefig(figure_dir + 'FigureS7_permutation_rsa.pdf', dpi=300, transparent=True)
  628. deltas.to_csv(source_file_dir + 'FigureS7_permutation_rsa.csv', index=False)
  629. def FigureS7_decoding(figure_dir=None, source_file_dir=None,
  630. recompute=False,
  631. num_sites=159, mode='empirical_conditioned_v2', draws=200):
  632. if figure_dir == None:
  633. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  634. if source_file_dir == None:
  635. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  636. IT_neurons_null = load_MonkeyIT_Decodes_permutation(num_neurons=np.array([num_sites]),
  637. recompute=recompute, mode=mode, draws=draws)
  638. IT_neurons_observed = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False,
  639. recompute=False, mode=mode)
  640. deltas = pd.DataFrame()
  641. for permutation in IT_neurons_null[f'{num_sites}_neurons'].keys():
  642. if permutation == 'sites':
  643. pass
  644. else:
  645. deltas = pd.concat([deltas,
  646. pd.DataFrame({'delta': np.mean(
  647. IT_neurons_null[f'{num_sites}_neurons'][permutation]['group1']['dprimes'],
  648. axis=1) - np.mean(
  649. IT_neurons_null[f'{num_sites}_neurons'][permutation]['group2']['dprimes'], axis=1)
  650. , 'kind': 'null'})], ignore_index=True)
  651. deltas = pd.concat(
  652. [deltas, pd.DataFrame({'delta': np.mean(IT_neurons_observed[f'{num_sites}_neurons']['trained']['dprimes'],
  653. axis=1)[:draws] - np.mean(
  654. IT_neurons_observed[f'{num_sites}_neurons']['untrained']['dprimes'], axis=1)[:draws]
  655. , 'kind': 'empirical'})], ignore_index=True)
  656. g = sns.displot(data=deltas[deltas['kind'] == 'null'], x='delta', hue='kind', palette='colorblind', height=5 * cm,
  657. aspect=1.2)
  658. g.refline(x=deltas.groupby('kind')['delta'].mean()['empirical'], c='red')
  659. g.set_xlabels("differences (Δ)")
  660. g.savefig(figure_dir + 'FigureS7_permutation_decoding.png', dpi=300)
  661. g.savefig(figure_dir + 'FigureS7_permutation_decoding.pdf', dpi=300, transparent=True)
  662. deltas.to_csv(source_file_dir + 'FigureS7_permutation_decoding.csv', index=False)
  663. print(
  664. f"p-value (two-sided): {(deltas.groupby('kind')['delta'].mean().loc['empirical'] <= deltas.loc[deltas['kind'] == 'null', 'delta']).mean()}")
  665. def Figure3_decoding_anatomy(mode='empirical_conditioned_v2', draws=300, num_sites=159, recompute=False,
  666. figure_dir=None, source_file_dir=None, colors=None):
  667. if figure_dir == None:
  668. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  669. if source_file_dir == None:
  670. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  671. IT_decodes_AIT = load_MonkeyIT_subregion_Decodes('AIT', mode=mode, draws=draws, num_neurons=np.array([32]),
  672. includeReliability=False, recompute=recompute)
  673. IT_decodes_AIT['num_subs'] = 1
  674. IT_decodes_CIT = load_MonkeyIT_subregion_Decodes('CIT', mode=mode, draws=draws, num_neurons=np.array([19]),
  675. includeReliability=False, recompute=recompute)
  676. IT_decodes_CIT['num_subs'] = 3
  677. IT_decodes_PIT = load_MonkeyIT_subregion_Decodes('PIT', mode=mode, draws=draws, num_neurons=np.array([17]), includeReliability=False, recompute=recompute)
  678. IT_decodes_PIT['num_subs'] = 2
  679. df = []
  680. for area in ['AIT', 'CIT', 'PIT']:
  681. for state in ['untrained', 'trained']:
  682. for d in range(draws):
  683. num_sites = np.squeeze(eval(f'IT_decodes_{area}')['num_subs'] * eval(f'IT_decodes_{area}')['num_neurons'])
  684. df.append(
  685. {'area': area, 'sites': num_sites, 'subs':eval(f'IT_decodes_{area}')['num_subs'], 'state': state, 'draw': d, 'i1': np.mean(eval(f'IT_decodes_{area}')[f'{num_sites}_neurons'][state]['dprimes'][d])})
  686. df = pd.DataFrame(df)
  687. g = sns.catplot(data=df, col='area', x='state', y='i1', hue='state', kind='point', errorbar=("pi", 95),
  688. palette=colors, height=5 * cm, aspect=0.75,
  689. legend=False, sharey=False)
  690. g.set_titles('{col_name}')
  691. g.set_xticklabels(["naïve", "trained"])
  692. g.set_ylabels("Performance (d')")
  693. g.set_xlabels('IT neuron pool')
  694. g.savefig(figure_dir + 'Figure3_subregions_decoding_i1.png', dpi=300)
  695. g.savefig(figure_dir + 'Figure3_subregions_decoding_i1.pdf', dpi=300, transparent=True)
  696. df.to_csv(source_file_dir + 'Figure3_subregions_decoding_i1.csv', index=False)
  697. print('Means:')
  698. print(df.groupby(['area', 'state']).mean())
  699. print('SD:')
  700. print(df.groupby(['area', 'state']).std())
  701. for area in ['AIT', 'CIT', 'PIT']:
  702. print(area)
  703. print(
  704. f"{((df.loc[df['area'] == area].groupby(['state'])['i1'].mean()['trained'] / df.loc[df['area'] == area].groupby(['state'])['i1'].mean()['untrained']) - 1) * 100}% increase")
  705. df_dist = ((df.loc[(df['area'] == area) & (df['state'] == 'trained'), 'i1'].reset_index() - df.loc[
  706. (df['area'] == area) & (df['state'] == 'untrained'), 'i1'].reset_index())['i1'])
  707. print(
  708. f"CI (95%): {df_dist.quantile(0.025)} - {df_dist.quantile(0.975)}")
  709. print(
  710. f"p-val: {np.min([(df_dist < 0).mean(), (df_dist > 0).mean()]) * 2}")
  711. print(
  712. f"Mean diff: {((df.loc[(df['state'] == 'trained') & (df['area'] == area), 'i1'].reset_index() - df.loc[(df['state'] == 'untrained') & (df['area'] == area), 'i1'].reset_index())['i1']).mean()}")
  713. print(
  714. f"Mean SD: {((df.loc[(df['state'] == 'trained') & (df['area'] == area), 'i1'].reset_index() - df.loc[(df['state'] == 'untrained') & (df['area'] == area), 'i1'].reset_index())['i1']).std()}")
  715. IT_decodes_LH = load_MonkeyIT_subregion_Decodes('IT', hemisphere='LH', mode=mode, draws=draws,
  716. num_neurons=np.array([89]), # n_jobs=1,
  717. includeReliability=False, recompute=False)
  718. IT_decodes_LH['num_subs'] = 2
  719. area = 'LH'
  720. df_LH = []
  721. for state in ['untrained', 'trained']:
  722. for d in range(draws):
  723. num_sites = np.squeeze(
  724. eval(f'IT_decodes_{area}')['num_subs'] * eval(f'IT_decodes_{area}')['num_neurons'])
  725. df_LH.append(
  726. {'area': area, 'sites': num_sites, 'subs': eval(f'IT_decodes_{area}')['num_subs'], 'state': state,
  727. 'draw': d, 'i1': np.mean(eval(f'IT_decodes_{area}')[f'{num_sites}_neurons'][state]['dprimes'][d])})
  728. df_LH = pd.DataFrame(df_LH)
  729. g = sns.catplot(data=df_LH, col='area', x='state', y='i1', hue='state', kind='point',
  730. errorbar=("pi", 95),
  731. palette=colors, height=5 * cm, aspect=0.9,
  732. legend=False, sharey=False)
  733. g.set_titles('Left hemisphere')
  734. g.set_xticklabels(["naïve", "trained"])
  735. g.set_ylabels("Performance (d')")
  736. g.set_xlabels('IT neuron pool')
  737. g.savefig(figure_dir + 'Figure3_hemisphere-LH_decoding_i1.png', dpi=300)
  738. g.savefig(figure_dir + 'Figure3_hemisphere-LH_decoding_i1.pdf', dpi=300, transparent=True)
  739. df_LH.to_csv(source_file_dir + 'Figure3_hemisphere-LH_decoding_i1.csv', index=False)
  740. print('Means:')
  741. print(df_LH.groupby(['state'])['i1'].mean())
  742. print('SD:')
  743. print(df_LH.groupby(['state'])['i1'].std())
  744. print(f"{((df_LH.groupby(['state'])['i1'].mean()['trained'] / df_LH.groupby(['state'])['i1'].mean()['untrained']) - 1) * 100}% increase")
  745. df_dist = ((df_LH.loc[(df_LH['state'] == 'trained'), 'i1'].reset_index() - df_LH.loc[
  746. (df['state'] == 'untrained'), 'i1'].reset_index())['i1'])
  747. print(
  748. f"CI (95%): {df_dist.quantile(0.025)} - {df_dist.quantile(0.975)}")
  749. print(
  750. f"p-val: {np.min([(df_dist < 0).mean(), (df_dist > 0).mean()]) * 2}")
  751. print(
  752. f"Mean diff: {df_dist.mean()}")
  753. print(
  754. f"SD diff: {df_dist.std()}")
  755. n_sites = np.unique((np.geomspace(3, num_sites, num=12, dtype='int') // 3) * 3)
  756. dPrimes = load_MonkeyIT_Decodes(resultFile_stem='CategoryDecodes_ITneurons_manySizes_', num_neurons=n_sites,
  757. includeReliability=False,
  758. recompute=recompute, mode=mode, draws=draws)
  759. n_sites = dPrimes['num_neurons']
  760. df_multi = []
  761. for n in n_sites:
  762. for state in ['untrained', 'trained']:
  763. for d in range(draws):
  764. df_multi.append(
  765. {'state': state, 'draw': d, 'sites': n,
  766. 'i1': np.mean(dPrimes[f'{n}_neurons'][state]['dprimes'][d])})
  767. df_multi = pd.DataFrame(df_multi)
  768. def func(x, a, b, c): # x-shifted log
  769. return a * np.log(x + b) + c
  770. fitting_df_trained = df_multi[df_multi['state'] == 'trained'].groupby('sites')['i1'].mean()
  771. fitting_df_untrained = df_multi[df_multi['state'] == 'untrained'].groupby('sites')['i1'].mean()
  772. popt_t, pcov_t = curve_fit(func, fitting_df_trained.index.values, fitting_df_trained.values)
  773. popt_u, pcov_u = curve_fit(func, fitting_df_untrained.index.values, fitting_df_untrained.values)
  774. fitting_df_trained['predicted_i1'] = func(fitting_df_trained.index.values, *popt_t)
  775. fitting_df_untrained['predicted_i1'] = func(fitting_df_untrained.index.values, *popt_u)
  776. extrapolation_df_trained = pd.DataFrame(
  777. {'sites': np.arange(1, 200), 'i1_predicted': func(np.arange(1, 200), *popt_t)})
  778. extrapolation_df_untrained = pd.DataFrame(
  779. {'sites': np.arange(1, 200), 'i1_predicted': func(np.arange(1, 200), *popt_u)})
  780. g = sns.relplot(data=df_multi, x='sites', y='i1', hue='state', kind='line', errorbar=("pi", 95),
  781. palette=colors, height=5 * cm, aspect=1.2,
  782. legend=False)
  783. g.ax.plot(extrapolation_df_trained.sites.values, extrapolation_df_trained['i1_predicted'].values, ls=':',
  784. color=sns.xkcd_palette(['jade'])[0])
  785. g.ax.plot(extrapolation_df_untrained.sites.values, extrapolation_df_untrained['i1_predicted'].values, ls=':',
  786. color=sns.xkcd_palette(['slate grey'])[0])
  787. for area in ['AIT', 'CIT', 'PIT']:
  788. num_sites = df.loc[(df['area'] == area), 'sites'].values[0]
  789. for s, state in enumerate(['untrained', 'trained']):
  790. g.ax.scatter(num_sites,df.loc[(df['area'] == area) & (df['state'] == state), 'i1'].mean(),
  791. color=colors[s])
  792. for s, state in enumerate(['untrained', 'trained']):
  793. g.ax.scatter(df_LH['sites'].values[0], df_LH.loc[(df['state'] == state), 'i1'].mean(),
  794. color=colors[s])
  795. g.set(xlim=[-5, 200])
  796. g.set_titles('Image-by-image\ndecoding')
  797. g.set_ylabels("Performance (d')")
  798. g.set_xlabels("IT neuron pool")
  799. g.savefig(figure_dir + 'Figure3_anatomyScaling_decoding_i1.png', dpi=300)
  800. g.savefig(figure_dir + 'Figure3_anatomyScaling_decoding_i1.pdf', dpi=300, transparent=True)
  801. def Figure4_novel500_decoding(draws=1000, mode='empirical_conditioned_v2', figure_dir=None, source_file_dir=None,
  802. colors=None, recompute=False, num_sites=55):
  803. if figure_dir == None:
  804. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  805. if source_file_dir == None:
  806. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  807. n_sites = np.unique(np.geomspace(3, num_sites, num=12, dtype='int'))
  808. IT_neurons_hvm = load_MonkeyIT_Decodes_novel500_HVMcontrol(num_neurons=n_sites, includeReliability=False,
  809. recompute=recompute,
  810. mode=mode, draws=draws)
  811. IT_neurons_novel500 = load_MonkeyIT_Decodes_novel500(num_neurons=n_sites, includeReliability=False,
  812. recompute=recompute,
  813. mode=mode, draws=draws)
  814. df = []
  815. for state in ['untrained', 'trained']:
  816. for site in n_sites:
  817. for d in range(draws):
  818. df.append(
  819. {'dataset': 'Novel500', 'state': state, 'pool size': site, 'draw': d,
  820. 'i1': np.mean(IT_neurons_novel500[f'{site}_neurons'][state]['dprimes'][d])})
  821. df.append(
  822. {'dataset': 'HVM', 'state': state, 'pool size': site, 'draw': d,
  823. 'i1': np.mean(IT_neurons_hvm[f'{site}_neurons'][state]['dprimes'][d])})
  824. df = pd.DataFrame(df)
  825. # DIFFICULTY-CONTROLLED DATASETS
  826. # sample images such that HVM has the same untrained difficulty
  827. np.random.seed(2)
  828. novel500_kde = gaussian_kde(np.mean(IT_neurons_novel500[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0))
  829. hvm_kde = gaussian_kde(np.mean(IT_neurons_hvm[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0))
  830. # Function to find the minimum of two KDE evaluations
  831. shared_kde = lambda x: np.min([novel500_kde(x), hvm_kde(x)], axis=0)
  832. overlap_area, _ = quad(shared_kde, -1, 4)
  833. maximum_sample = int(500 * overlap_area)
  834. # map this to the images
  835. naive_img_hvm = pd.DataFrame({'dataset': 'HVM', 'img': np.arange(640),
  836. 'i1': np.mean(IT_neurons_hvm[f'{num_sites}_neurons']['untrained']['dprimes'],
  837. axis=0)})
  838. naive_img_hvm['pdf'] = novel500_kde(naive_img_hvm['i1']) / hvm_kde(naive_img_hvm['i1'])
  839. naive_img_novel500 = pd.DataFrame({'dataset': 'novel500', 'img': np.arange(500), 'i1': np.mean(
  840. IT_neurons_novel500[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0)})
  841. naive_img_novel500['pdf'] = hvm_kde(naive_img_novel500['i1']) / novel500_kde(naive_img_novel500['i1'])
  842. hvm_mean = 0
  843. novel500_mean = 1
  844. while np.round(abs(hvm_mean - novel500_mean), 1) != 0:
  845. sampled_hvm = naive_img_hvm.sample(maximum_sample, weights='pdf')
  846. sampled_novel500 = naive_img_novel500.sample(maximum_sample, weights='pdf')
  847. hvm_mean = sampled_hvm['i1'].mean()
  848. novel500_mean = sampled_novel500['i1'].mean()
  849. print(hvm_mean)
  850. print(novel500_mean)
  851. g, ax = plt.subplots(1, 2, figsize=(4.5 * cm * 0.9 * 2, 4.5 * cm), sharey=True, sharex=True)
  852. sns.histplot(np.mean(IT_neurons_hvm[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0), ax=ax[0],
  853. color=sns.xkcd_palette(['dull yellow'])[0], kde=True)
  854. sns.histplot(np.mean(IT_neurons_novel500[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0), ax=ax[0],
  855. color=sns.xkcd_palette(['pale lavender'])[0], kde=True)
  856. sns.lineplot(x=np.arange(-1, 4, 0.1), y=shared_kde(np.arange(-1, 4, 0.1)) * 100,
  857. color=sns.xkcd_palette(['blush'])[0], ax=ax[0])
  858. ax[0].set_xlabel("Task-naive accuracy (d')")
  859. ax[0].set_ylabel("# images")
  860. sns.histplot(sampled_hvm['i1'], ax=ax[1], color=sns.xkcd_palette(['dull yellow'])[0])
  861. sns.histplot(sampled_novel500['i1'], ax=ax[1], color=sns.xkcd_palette(['pale lavender'])[0])
  862. sns.lineplot(x=np.arange(-1, 4, 0.1), y=shared_kde(np.arange(-1, 4, 0.1)) * 100,
  863. color=sns.xkcd_palette(['blush'])[0], ax=ax[1])
  864. ax[1].set_xlabel(" ")
  865. plt.tight_layout()
  866. g.savefig(figure_dir + 'Figure4_novel500_decoding_naiveDists.png', dpi=300)
  867. g.savefig(figure_dir + 'Figure4_novel500_decoding_naiveDists.pdf', dpi=300, transparent=True)
  868. # apply sampling to the whole dataset
  869. df_matched = []
  870. for state in ['untrained', 'trained']:
  871. for d in range(draws):
  872. df_matched.append(
  873. {'dataset': 'Novel500', 'state': state, 'pool size': num_sites, 'draw': d,
  874. 'i1': np.mean(
  875. IT_neurons_novel500[f'{num_sites}_neurons'][state]['dprimes'][d, sampled_novel500['img'].values])})
  876. df_matched.append(
  877. {'dataset': 'HVM', 'state': state, 'pool size': num_sites, 'draw': d,
  878. 'i1': np.mean(IT_neurons_hvm[f'{num_sites}_neurons'][state]['dprimes'][d, sampled_hvm['img'].values])})
  879. df_matched = pd.DataFrame(df_matched)
  880. g = sns.catplot(data=df_matched, x='state', y='i1', hue='state', kind='point',
  881. errorbar=("pi", 95), col='dataset',
  882. palette=colors, height=5 * cm, aspect=0.7,
  883. legend=False)
  884. g.set(ylim=[0.5, 2.5])
  885. g.set_titles(' ')
  886. g.set_xticklabels(["naïve", "trained"])
  887. g.set_ylabels("Performance (d')")
  888. g.set_xlabels("IT neuron pool")
  889. g.refline(y=(sampled_hvm['i1'].mean() + sampled_novel500['i1'].mean()) / 2, color=sns.xkcd_palette(['blush'])[0], zorder=0)
  890. g.savefig(figure_dir + 'Figure4_novel500_decoding_i1_naiveMatched.png', dpi=300)
  891. g.savefig(figure_dir + 'Figure4_novel500_decoding_i1_naiveMatched.pdf', dpi=300, transparent=True)
  892. df_matched.to_csv(source_file_dir + 'Figure4_novel500_decoding_i1_naiveMatched.csv', index=False)
  893. print('Means:')
  894. print(df_matched.loc[df_matched['pool size'] == num_sites].groupby(['dataset', 'state']).mean())
  895. print('SD:')
  896. print(df_matched.loc[df_matched['pool size'] == num_sites].groupby(['dataset', 'state']).std())
  897. for dataset in ['HVM', 'Novel500']:
  898. print(dataset)
  899. print(
  900. f"{((df_matched.loc[(df_matched['dataset'] == dataset) & (df_matched['pool size'] == num_sites)].groupby(['state'])['i1'].mean()['trained'] / df_matched.loc[(df_matched['dataset'] == dataset) & (df_matched['pool size'] == num_sites)].groupby(['state'])['i1'].mean()['untrained']) - 1) * 100}% increase")
  901. df_dist = ((df_matched.loc[(df_matched['dataset'] == dataset) & (df_matched['pool size'] == num_sites) & (
  902. df_matched['state'] == 'trained'), 'i1'].reset_index() - df_matched.loc[
  903. (df_matched['dataset'] == dataset) & (df_matched['pool size'] == num_sites) & (
  904. df_matched['state'] == 'untrained'), 'i1'].reset_index())['i1'])
  905. print(
  906. f"CI (95%): {df_dist.quantile(0.025)} - {df_dist.quantile(0.975)}")
  907. print(
  908. f"p-val: {np.min([(df_dist < 0).mean(), (df_dist > 0).mean()]) * 2}")
  909. print(
  910. f"Mean diff: {df_dist.mean()}")
  911. print(
  912. f"SD diff: {df_dist.std()}")
  913. def FigureS9_BrainScoreLayerMapping(base_models, figure_dir=None, source_file_dir=None,
  914. result_file=None
  915. ):
  916. if figure_dir == None:
  917. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  918. if result_file == None:
  919. result_file = f'{Path(__file__).parent.parent}/Results/meta/BrainScore_ITmappings.pkl'
  920. if source_file_dir == None:
  921. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  922. df_all = pd.read_pickle(result_file)
  923. col_order = [' ']
  924. col_order.extend(base_models)
  925. g = sns.FacetGrid(data=df_all, col='baseModel', col_wrap=3, col_order=col_order, height=5 * cm, aspect=1.3,
  926. sharex=False, sharey=False)
  927. g.map(plt.errorbar, 'layer', 'scores_x', 'scores_y', color='grey', zorder=0)
  928. for i, ax in enumerate(g.axes[1:]):
  929. ax.set_xticklabels(ax.get_xticklabels(), rotation=90)
  930. xs, ys = ax.lines[-1].get_data()
  931. ax.scatter(xs[np.argmax(ys)], np.max(ys), color=sns.xkcd_palette(['royal blue'])[0], s=20)
  932. ax.axvline(xs[np.argmax(ys)], color=sns.xkcd_palette(['royal blue'])[0], ls='--')
  933. ax.set_title(base_models[i])
  934. g.set_xlabels('Model layer')
  935. g.set_ylabels('IT encoding\nscore')
  936. plt.subplots_adjust(hspace=1.75, wspace=0.2)
  937. g.savefig(f'{figure_dir}FigureS9_ITMappings.png', dpi=300)
  938. g.savefig(f'{figure_dir}FigureS9_ITMappings.pdf', transparent=True, dpi=300)
  939. df_all.to_csv(source_file_dir + 'FigureS9_ITMappings.csv', index=False)
  940. def FigureS10_UnitSiteScaling_IT(base_models, num_sites=159,
  941. recompute=False,
  942. reestimate_predictions=False,
  943. draws=1000,
  944. maximum_scaling=101,
  945. resultFile=None,
  946. figure_dir=None,
  947. source_file_dir=None,
  948. mode='empirical_conditioned_v2'):
  949. if figure_dir == None:
  950. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  951. if resultFile == None:
  952. resultFile = f'{Path(__file__).parent.parent}/Results/meta/UnitSiteScaling.pkl'
  953. if source_file_dir == None:
  954. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  955. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False, mode=mode,
  956. recompute=recompute, draws=draws)
  957. dfs = get_ModelIT_categoryDecoding_UnitSiteScaling(base_models, reestimate_predictions=reestimate_predictions,
  958. recompute=recompute, num_sites=num_sites)
  959. target_value = IT_neurons[f'{num_sites}_neurons']['untrained']['dprimes'].mean()
  960. dfs['siteScaling'] = dfs['units'] / num_sites
  961. def func(x, a, b, c): # x-shifted log
  962. return a * np.log(x + b) + c
  963. fitting_df = dfs[dfs['siteScaling'] < maximum_scaling].groupby(['base_model', 'siteScaling']).mean()
  964. sigma_df = dfs[dfs['siteScaling'] < maximum_scaling].groupby(['base_model', 'siteScaling']).std()
  965. sigma_df.loc[sigma_df['i1'] == 0] = 1
  966. g = sns.relplot(data=dfs, x='siteScaling', y='i1', kind='line', errorbar=("pi", 95),
  967. color='seagreen', height=5 * cm, aspect=1.2, style='base_model',
  968. dashes=False, col='base_model', col_wrap=3,
  969. markers=['o' for _ in base_models], legend=False)
  970. g.set_titles("{col_name}")
  971. g.refline(y=target_value, color='grey', ls='-')
  972. unit_scaling = []
  973. for model, ax in zip(base_models, g.axes):
  974. # print(model)
  975. popt, pcov = curve_fit(func, fitting_df.loc[model].index.values, fitting_df.loc[model, 'i1'].values,
  976. sigma=sigma_df.loc[model, 'i1'].values,
  977. )
  978. # print(popt)
  979. fitting_df.loc[model, 'predicted_i1'] = func(fitting_df.loc[model].index.values, *popt)
  980. interpolation_df = pd.DataFrame(
  981. {'base_model': model, 'sitesScaling': np.arange(1, 100), 'i1_predicted': func(np.arange(1, 100), *popt)})
  982. ax.plot(interpolation_df.sitesScaling.values, interpolation_df['i1_predicted'].values, ls=':', color='m')
  983. required_sites = interpolation_df.loc[interpolation_df['i1_predicted'] >= target_value, 'sitesScaling'].min()
  984. # Identify the adjacent actual datapoints
  985. required_sites_empirical_upper = fitting_df.loc[model, 'i1'][
  986. fitting_df.loc[model, 'i1'] >= target_value].index.min()
  987. required_sites_empirical_lower = fitting_df.loc[model, 'i1'][
  988. fitting_df.loc[model, 'i1'] <= target_value].index.max()
  989. if pd.isna(required_sites_empirical_upper) | pd.isna(required_sites_empirical_lower):
  990. required_sites_empirical_estimated = np.nan
  991. else:
  992. lower_val = fitting_df.loc[model].loc[required_sites_empirical_lower]['i1']
  993. upper_val = fitting_df.loc[model].loc[required_sites_empirical_upper]['i1']
  994. linearFit = np.polyfit([required_sites_empirical_lower, required_sites_empirical_upper],
  995. [lower_val, upper_val], 1)
  996. required_sites_empirical_estimated = (target_value - linearFit[1]) / linearFit[0]
  997. # ax.axvline(x=required_sites_empirical_estimated, color='orange')
  998. if required_sites != np.nan:
  999. ax.axvline(x=required_sites, color='m')
  1000. print(f'{model}: {required_sites} / {required_sites_empirical_estimated}')
  1001. unit_scaling.append(
  1002. {'base_model': model, 'siteUnitScaling': required_sites, 'units': required_sites * num_sites})
  1003. g.set(xlim=[0, maximum_scaling])
  1004. g.set_ylabels("Accuracy (d')")
  1005. g.set_xlabels("Unit/site ratio")
  1006. g.savefig(figure_dir + 'FigureS10_siteUnitScaling.png', dpi=300)
  1007. g.savefig(figure_dir + 'FigureS10_siteUnitScaling.pdf', dpi=300, transparent=True)
  1008. dfs.to_csv(source_file_dir + 'FigureS10_siteUnitScaling.csv', index=False)
  1009. unit_scaling = pd.DataFrame(unit_scaling)
  1010. if recompute:
  1011. unit_scaling.to_pickle(resultFile)
  1012. return unit_scaling.loc[unit_scaling['siteUnitScaling'].isna() == False, 'base_model'].tolist()
  1013. def Figure6_performance(training_groups, base_models, model_suffix='best', metric='test_d-prime', recompute=False,
  1014. performance_threshold=None, figure_dir=None, source_file_dir=None,
  1015. version='v0', dataset='hvm_modelTest'):
  1016. if figure_dir == None:
  1017. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1018. if source_file_dir == None:
  1019. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1020. df = pd.DataFrame()
  1021. for f in training_groups:
  1022. trained_feature_dir = training_groups[f]['feature_dir'] + f'{dataset}/{model_suffix}/'
  1023. if (recompute == True) or (
  1024. not Path(f'{training_groups[f]["feature_dir"]}/TrainingMetrics.pkl').exists()):
  1025. ### This requires retraining. Wandb meta-files not included.
  1026. tmp = loadTrainingResults(base_models, training_groups[f]['feature_dir'],
  1027. project=training_groups[f]['project'], entity=training_groups[f]['entity'],
  1028. sweep_suffix=training_groups[f]['sweep_suffix'],
  1029. model_suffix=model_suffix, recompute=recompute)
  1030. ###
  1031. else:
  1032. tmp = joblib.load(
  1033. f'{training_groups[f]["feature_dir"]}/TrainingMetrics.pkl')
  1034. tmp['Training'] = f
  1035. df = pd.concat([df, tmp], ignore_index=True)
  1036. df = df.loc[df['base_model'].isin(base_models)]
  1037. if version == 'v0':
  1038. g = sns.FacetGrid(df, col="Training", sharex=True, sharey=False, col_wrap=3, aspect=1.25, height=1.7)
  1039. g.map(sns.histplot, metric, binwidth=0.2,
  1040. clip_on=False,
  1041. fill=True, alpha=0.5, linewidth=0.5, color='white')
  1042. g.figure.subplots_adjust(hspace=-0.5)
  1043. df['selected'] = False
  1044. if 'max_value' not in performance_threshold:
  1045. g.refline(x=performance_threshold['min_value'], color=sns.xkcd_palette(['raspberry'])[0], linestyle='--')
  1046. df.loc[df[metric] >= performance_threshold['min_value'], 'selected'] = True
  1047. else:
  1048. g.map(plt.axvspan, xmin=performance_threshold['min_value'], xmax=performance_threshold['max_value'],
  1049. zorder=0, color=sns.xkcd_palette(['raspberry'])[0],
  1050. alpha=0.2)
  1051. df.loc[(df[metric] >= performance_threshold['min_value']) & (
  1052. df[metric] <= performance_threshold['max_value']), 'selected'] = True
  1053. g.set_titles("")
  1054. g.set(xlabel="Performance (d')", ylabel="# models")
  1055. g.savefig(figure_dir + 'FigureS12_performance.png', dpi=300)
  1056. g.savefig(figure_dir + 'FigureS12_performance.pdf', dpi=300, transparent=True)
  1057. print(df.groupby('Training')['selected'].sum())
  1058. df.to_csv(source_file_dir + 'FigureS12_performance.csv', index=False)
  1059. elif version == 'v1':
  1060. g = sns.displot(data=df, x=metric,
  1061. color='white', height=3.75 * cm, aspect=1.25, clip_on=False, fill=True, alpha=0.5,
  1062. linewidth=0.5,
  1063. facet_kws={'sharex': False, 'sharey': False}, common_bins=False, bins=30)
  1064. g.figure.subplots_adjust(hspace=-0.5)
  1065. df['selected'] = False
  1066. if 'max_value' not in performance_threshold:
  1067. g.refline(x=performance_threshold['min_value'], color=sns.xkcd_palette(['raspberry'])[0], linestyle='--')
  1068. df.loc[df[metric] >= performance_threshold['min_value'], 'selected'] = True
  1069. else:
  1070. g.map(plt.axvspan, xmin=performance_threshold['min_value'], xmax=performance_threshold['max_value'],
  1071. zorder=0, color=sns.xkcd_palette(['raspberry'])[0], alpha=0.2,
  1072. )
  1073. df.loc[(df[metric] >= performance_threshold['min_value']) & (
  1074. df[metric] <= performance_threshold['max_value']), 'selected'] = True
  1075. g.set_titles("")
  1076. g.set(xlabel="Performance (d')", ylabel="# models")
  1077. g.savefig(figure_dir + 'Figure6_performance.png', dpi=300)
  1078. g.savefig(figure_dir + 'Figure6_performance.pdf', dpi=300, transparent=True)
  1079. print(df['selected'].sum())
  1080. df.to_csv(source_file_dir + 'Figure6_performance.csv', index=False)
  1081. return df
  1082. def Figure6_modelITComparison_means(base_models, training_groups, performance_threshold, model_draws=100,
  1083. num_sites=159, draws=1000, mode='empirical_conditioned_v2', recompute=False,
  1084. delta='proportion',
  1085. figure_dir=None,
  1086. meta_dir=None,
  1087. source_file_dir=None,
  1088. version='v0'
  1089. ):
  1090. if figure_dir == None:
  1091. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1092. if meta_dir == None:
  1093. meta_dir = f'{Path(__file__).parent.parent}/Results/meta/'
  1094. if source_file_dir == None:
  1095. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1096. df_selectivity = load_ModelSelectivity_allModels(base_models, training_groups, performance_threshold, delta=delta,
  1097. model_draws=model_draws, recompute=recompute)
  1098. df_rsa = load_ModelRDMs_allModels(base_models, training_groups, performance_threshold, delta=delta,
  1099. model_draws=model_draws, recompute=recompute)
  1100. df_decoding = load_ModelIT_categoryDecodes_allModels(base_models, training_groups, performance_threshold,
  1101. delta=delta,
  1102. model_draws=model_draws, recompute=recompute)
  1103. # Already computed in the neural analyses
  1104. neural_selectivity = load_neuralITselectivity_deltas(recompute=False, draws=draws, mode=mode, delta=delta)
  1105. neural_rsa = load_neuralRDMs_deltas(recompute=False, draws=draws, mode=mode, num_sites=num_sites, delta=delta)
  1106. neural_decoding = load_MonkeyIT_Decodes_deltas(num_sites=num_sites, includeReliability=False, recompute=False,
  1107. mode=mode, draws=draws, delta=delta)
  1108. df_combined = pd.DataFrame()
  1109. df_combined['delta'] = df_selectivity.groupby('model_ID')['delta (selectivity)'].median()
  1110. df_combined['metric'] = 'selectivity'
  1111. df_combined = df_combined.reset_index()
  1112. df_tmp = pd.DataFrame()
  1113. df_tmp['delta'] = df_rsa.groupby('model_ID')['delta (tau)'].mean()
  1114. df_tmp['metric'] = 'rsa'
  1115. df_tmp = df_tmp.reset_index()
  1116. df_combined = pd.concat([df_combined, df_tmp])
  1117. df_tmp = pd.DataFrame()
  1118. df_tmp['delta'] = df_decoding.groupby('model_ID')['delta (i1)'].mean()
  1119. df_tmp['metric'] = 'decoding'
  1120. df_tmp = df_tmp.reset_index()
  1121. df_combined = pd.concat([df_combined, df_tmp])
  1122. # Define subsets
  1123. set_decoding = set(df_combined.loc[(df_combined.loc[
  1124. df_combined['metric'] == 'decoding', 'delta'] >= np.quantile(
  1125. eval(f'neural_decoding'), 0.025)) & (
  1126. df_combined.loc[
  1127. df_combined['metric'] == 'decoding', 'delta'] <= np.quantile(
  1128. eval(f'neural_decoding'), 0.975)), 'model_ID'])
  1129. set_rsa = set(df_combined.loc[(df_combined.loc[df_combined['metric'] == 'rsa', 'delta'] >= np.quantile(
  1130. eval(f'neural_rsa'), 0.025)) & (
  1131. df_combined.loc[df_combined['metric'] == 'rsa', 'delta'] <= np.quantile(
  1132. eval(f'neural_rsa'), 0.975)), 'model_ID'])
  1133. set_selectivity = set(df_combined.loc[(df_combined.loc[
  1134. df_combined['metric'] == 'selectivity', 'delta'] >= np.quantile(
  1135. eval(f'neural_selectivity'), 0.025)) & (
  1136. df_combined.loc[df_combined[
  1137. 'metric'] == 'selectivity', 'delta'] <= np.quantile(
  1138. eval(f'neural_selectivity'), 0.975)), 'model_ID'])
  1139. models_pt3 = pd.DataFrame({'model_ID': df_combined.model_ID.unique()})
  1140. models_pt3['match_selectivity'] = models_pt3['model_ID'].isin(set_selectivity)
  1141. models_pt3['match_rsa'] = models_pt3['model_ID'].isin(set_rsa)
  1142. models_pt3['match_decoding'] = models_pt3['model_ID'].isin(set_decoding)
  1143. models_pt3['match_all'] = models_pt3['match_selectivity'] & models_pt3['match_rsa'] & models_pt3['match_decoding']
  1144. models_pt3['match_none'] = (models_pt3['match_selectivity'] == False) & (models_pt3['match_rsa'] == False) & (
  1145. models_pt3['match_decoding'] == False)
  1146. models_pt3['match_one'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(int) +
  1147. models_pt3['match_rsa'].astype(int)) == 1
  1148. models_pt3['match_two'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(
  1149. int) + models_pt3['match_rsa'].astype(int)) == 2
  1150. for feature in ['base_model', 'Training', 'base_model_type']:
  1151. models_pt3[feature] = models_pt3['model_ID'].map(dict(zip(df_decoding['model_ID'], df_decoding[feature])))
  1152. if recompute:
  1153. models_pt3.to_pickle(meta_dir + f'SelectedModels_pt3_means.pkl')
  1154. plotting_titles = {'rsa': 'Representational\nstrength',
  1155. 'selectivity': 'Selectivity',
  1156. 'decoding': 'Linear decodability'}
  1157. if version == 'v0':
  1158. g = sns.displot(data=df_combined, x='delta', row='metric',
  1159. color=sns.xkcd_palette(['raspberry'])[0], height=3.75 * cm, aspect=1.25, clip_on=False,
  1160. fill=True, alpha=0.2, linewidth=0.5,
  1161. facet_kws={'sharex': False, 'sharey': False}, common_bins=False, bins=30)
  1162. for title, ax in g.axes_dict.items():
  1163. print(title)
  1164. ax.axvspan(np.quantile(eval(f'neural_{title}'), 0.025),
  1165. np.quantile(eval(f'neural_{title}'), 0.975),
  1166. color=sns.xkcd_palette(['royal blue'])[0], zorder=0, alpha=0.2)
  1167. ax.set_title(plotting_titles[title], fontweight='bold', fontsize=10)
  1168. print(
  1169. f"# models falling below neural CI: {(df_combined.loc[df_combined['metric'] == title, 'delta'] < np.quantile(eval(f'neural_{title}'), 0.025)).sum()}")
  1170. print(
  1171. f"# models falling above neural CI: {(df_combined.loc[df_combined['metric'] == title, 'delta'] > np.quantile(eval(f'neural_{title}'), 0.975)).sum()}")
  1172. print(
  1173. f"# models within neural CI: {((df_combined.loc[df_combined['metric'] == title, 'delta'] >= np.quantile(eval(f'neural_{title}'), 0.025)) & (df_combined.loc[df_combined['metric'] == title, 'delta'] <= np.quantile(eval(f'neural_{title}'), 0.975))).sum()}")
  1174. g.set_ylabels('# models')
  1175. g.set_xlabels('Δ%')
  1176. g.savefig(figure_dir + f'Figure6_modelITComparison_means_v0.png', dpi=300)
  1177. g.savefig(figure_dir + f'Figure6_modelITComparison_means_v0.pdf', dpi=300, transparent=True)
  1178. df_combined.to_csv(source_file_dir + 'Figure6_modelITComparison_means_v0.csv', index=False)
  1179. elif version == 'v1':
  1180. # make a venn diagram
  1181. fig, ax = plt.subplots()
  1182. venn3([set_selectivity, set_decoding, set_rsa], ('Selectivity', 'Decoding', 'RSA'),
  1183. ax=ax)
  1184. fig.savefig(figure_dir + f'Figure6_modelITComparison_means_v1.png', dpi=300)
  1185. fig.savefig(figure_dir + f'Figure6_modelITComparison_means_v1.pdf', dpi=300, transparent=True)
  1186. elif version == 'v2':
  1187. df_plot = models_pt3.groupby('Training', as_index=False)['match_all'].mean()
  1188. df_plot['match_all_%'] = df_plot['match_all'] * 100
  1189. df_plot['dummy'] = 100
  1190. g = sns.catplot(data=df_plot, x='Training', y='match_all_%',
  1191. kind='bar', color=sns.xkcd_palette(['royal blue'])[0],
  1192. height=3.75 * cm, aspect=0.9, order=['FT', '2stepFT', 'binary'],
  1193. edgecolor=".5", linewidth=0.5,
  1194. )
  1195. g.map(sns.barplot, 'Training', 'dummy', order=['FT', '2stepFT', 'binary'], zorder=0, color='white',
  1196. linewidth=0.5, edgecolor=".5")
  1197. g.set_xticklabels(['Standard', 'Step-wise', 'Binary choice'], rotation=90)
  1198. g.set(ylim=[0, 100])
  1199. g.set_ylabels('% models')
  1200. g.savefig(figure_dir + f'FigureS12_modelITComparison_means_Training.png')
  1201. g.savefig(figure_dir + f'FigureS12_modelITComparison_means_Training.pdf', dpi=300, transparent=True)
  1202. df_plot.to_csv(source_file_dir + 'FigureS12_modelITComparison_means_Training.csv', index=False)
  1203. df_plot = models_pt3.groupby('base_model', as_index=False)['match_all'].mean()
  1204. df_plot['match_all_%'] = df_plot['match_all'] * 100
  1205. df_plot['dummy'] = 100
  1206. model_order = ['resnet18_v1', 'resnet34_v1', 'resnet50_v1', 'resnet101_v1', 'resnet152_v1',
  1207. 'alexnet', 'vgg16', 'vgg19', ' ',
  1208. 'resnet50_MoCov2_200epochs', 'resnet50_simclr_100epochs', 'resnet50_barlowTwins_300epochs',
  1209. ]
  1210. model_titles = ['18', '34', '50', '101', '152',
  1211. 'alexnet', 'vgg16', 'vgg19', ' ',
  1212. 'MoCov2', 'SimCLR', 'barlowTwins']
  1213. g = sns.catplot(data=df_plot, x='base_model', y='match_all_%',
  1214. kind='bar', color=sns.xkcd_palette(['royal blue'])[0],
  1215. height=3.75 * cm, aspect=2.4,
  1216. order=model_order,
  1217. edgecolor=".5", linewidth=0.5,
  1218. )
  1219. g.map(sns.barplot, 'base_model', 'dummy', order=model_order,
  1220. zorder=0, color='white',
  1221. linewidth=0.5, edgecolor=".5")
  1222. g.set_xticklabels(model_titles, rotation=90)
  1223. g.set(ylim=[0, 100])
  1224. g.set_ylabels('% models')
  1225. g.savefig(figure_dir + f'FigureS12_modelITComparison_means_baseModel.png')
  1226. g.savefig(figure_dir + f'FigureS12_modelITComparison_means_baseModel.pdf', dpi=300, transparent=True)
  1227. df_plot.to_csv(source_file_dir + 'FigureS12_modelITComparison_means_baseModel.csv', index=False)
  1228. def FigureS11_modelITComparison_controlLayers(layers, base_models, training_groups, performance_threshold, model_draws=100,
  1229. num_sites=159, draws=1000, mode='empirical_conditioned_v2', recompute=False,
  1230. delta='proportion', figure_dir=None, source_file_dir=None,):
  1231. if figure_dir == None:
  1232. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1233. if source_file_dir == None:
  1234. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1235. df_combined_all = []
  1236. for layer in layers:
  1237. print(layer)
  1238. area = f'Layer-{layer}'
  1239. df_selectivity = load_ModelSelectivity_allModels(base_models, training_groups, performance_threshold,
  1240. delta=delta, area=area, model_draws=model_draws,
  1241. recompute=recompute)
  1242. df_rsa = load_ModelRDMs_allModels(base_models, training_groups, performance_threshold, delta=delta, area=area,
  1243. model_draws=model_draws, recompute=recompute)
  1244. df_decoding = load_ModelIT_categoryDecodes_allModels(base_models, training_groups, performance_threshold,
  1245. delta=delta, area=area, model_draws=model_draws,
  1246. recompute=recompute)
  1247. df_combined = pd.DataFrame()
  1248. df_combined['delta'] = df_selectivity.groupby('model_ID')['delta (selectivity)'].median()
  1249. df_combined['metric'] = 'selectivity'
  1250. df_combined = df_combined.reset_index()
  1251. df_tmp = pd.DataFrame()
  1252. df_tmp['delta'] = df_rsa.groupby('model_ID')['delta (tau)'].mean()
  1253. df_tmp['metric'] = 'rsa'
  1254. df_tmp = df_tmp.reset_index()
  1255. df_combined = pd.concat([df_combined, df_tmp])
  1256. df_tmp = pd.DataFrame()
  1257. df_tmp['delta'] = df_decoding.groupby('model_ID')['delta (i1)'].mean()
  1258. df_tmp['metric'] = 'decoding'
  1259. df_tmp = df_tmp.reset_index()
  1260. df_combined = pd.concat([df_combined, df_tmp])
  1261. df_combined['area'] = area
  1262. df_combined['layer'] = layer
  1263. df_combined_all.append(df_combined)
  1264. df_combined_all = pd.concat(df_combined_all, ignore_index=True)
  1265. neural_selectivity = load_neuralITselectivity_deltas(recompute=False, draws=draws, mode=mode, delta=delta)
  1266. neural_decoding = load_MonkeyIT_Decodes_deltas(num_sites=num_sites, includeReliability=False, recompute=False,
  1267. mode=mode, draws=draws, delta=delta)
  1268. neural_rsa = load_neuralRDMs_deltas(recompute=False, draws=draws, mode=mode, num_sites=num_sites, delta=delta)
  1269. plotting_titles = {'rsa': 'Representational\nstrength',
  1270. 'selectivity': 'Selectivity',
  1271. 'decoding': 'Linear decodability'}
  1272. #model_IDs = df_combined_all.loc[df_combined_all['area'] == 'EarlyLayerControl', 'model_ID'].unique()
  1273. df_combined_plotting = df_combined_all.loc[
  1274. df_combined_all['layer'].isin(['layer2[0].relu', 'layer4[0].relu', 'avgpool'])]
  1275. df_combined_plotting.loc[df_combined_plotting['layer'] == 'layer4[0].relu', 'area'] = 'IT-mapped layer'
  1276. df_combined_plotting.loc[df_combined_plotting['layer'] == 'layer2[0].relu', 'area'] = 'Early control layer'
  1277. df_combined_plotting.loc[df_combined_plotting['layer'] == 'avgpool', 'area'] = 'Late control layer'
  1278. g = sns.displot(data=df_combined_plotting, x='delta', col='metric',
  1279. hue='area', height=5 * cm, aspect=1, clip_on=False, fill=True, alpha=0.2, linewidth=0.5,
  1280. facet_kws={'sharex': False, 'sharey': False}, common_bins=False, kde=True,
  1281. palette=sns.xkcd_palette(['tangerine', 'raspberry', 'pastel purple']))
  1282. for title, ax in g.axes_dict.items():
  1283. print(title)
  1284. ax.axvspan(np.quantile(eval(f'neural_{title}'), 0.025),
  1285. np.quantile(eval(f'neural_{title}'), 0.975),
  1286. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=0.3)
  1287. ax.set_title(plotting_titles[title])
  1288. g.set_ylabels('# models')
  1289. g.set_xlabels('Δ%')
  1290. #plt.show()
  1291. g.savefig(figure_dir + f'FigureS11_modelITComparison_controlLayers.png', dpi=300)
  1292. g.savefig(figure_dir + f'FigureS11_modelITComparison_controlLayers.pdf', dpi=300, transparent=True)
  1293. df_combined_plotting.to_csv(source_file_dir + 'FigureS11_modelITComparison_controlLayers.csv', index=False)
  1294. g = sns.catplot(data=df_combined_all, x='layer', y='delta', col='metric', kind='point', order=layers,
  1295. color=sns.xkcd_palette(['eggplant'])[0],
  1296. col_order=['selectivity', 'rsa', 'decoding'], errorbar=('pi', 95), height=5.5 * cm, aspect=1.3,
  1297. sharey=False)
  1298. for title, ax in g.axes_dict.items():
  1299. ax.axhspan(np.quantile(eval(f'neural_{title}'), 0.025),
  1300. np.quantile(eval(f'neural_{title}'), 0.975),
  1301. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=0.3)
  1302. ax.set_title(plotting_titles[title])
  1303. g.refline(x=layers.index('layer4[0].relu'), c=sns.xkcd_palette(['raspberry'])[0], ls='-', zorder=0)
  1304. g.refline(x=layers.index('layer2[0].relu'), c=sns.xkcd_palette(['tangerine'])[0], ls='-', zorder=0)
  1305. g.refline(x=layers.index('avgpool'), c=sns.xkcd_palette(['pastel purple'])[0], ls='-', zorder=0)
  1306. g.refline(y=0)
  1307. g.set_xticklabels([])
  1308. g.set_ylabels('Δ%')
  1309. g.set_xlabels('Model layer')
  1310. g.savefig(figure_dir + f'FigureS11_modelITComparison_controlLayers_v1.png', dpi=300)
  1311. g.savefig(figure_dir + f'FigureS11_modelITComparison_controlLayers_v1.pdf', dpi=300, transparent=True)
  1312. df_combined_all.to_csv(
  1313. source_file_dir + 'FigureS11_modelITComparison_controlLayers_v1.csv', index=False)
  1314. def FigureS11_modelITcomparison_controlLayers_orthogonalDecoding(training_groups, layers=None, seed=6,
  1315. recompute=False, metric='corrs', model_draws=100, draws=1000,
  1316. delta='proportion', mode='empirical_conditioned',reestimate_predictions=False,
  1317. figure_dir=None, meta_dir=None, source_file_dir=None):
  1318. if figure_dir == None:
  1319. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1320. if meta_dir == None:
  1321. meta_dir = f'{Path(__file__).parent.parent}/Results/meta/'
  1322. if source_file_dir == None:
  1323. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1324. if layers == None:
  1325. layers = ['layer2[0].relu', 'layer4[0].relu', 'avgpool']
  1326. np.random.seed(seed)
  1327. models_pt3 = pd.read_pickle(meta_dir + f'SelectedModels_pt3_means.pkl')
  1328. models_pt3['match_one'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(int) + models_pt3['match_rsa'].astype(int)) == 1
  1329. models_pt3['match_two'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(
  1330. int) + models_pt3['match_rsa'].astype(int)) == 2
  1331. # Only include IT-like models and vary the layer...
  1332. models_pt3 = models_pt3.loc[(models_pt3['base_model'] == 'resnet152_v1') & (models_pt3['match_all'])]
  1333. df = []
  1334. for layer in layers:
  1335. print(layer)
  1336. area = f'Layer-{layer}'
  1337. df.append(load_ModelLayer_orthogonalDecodes_selection(models_pt3, training_groups,area=area, draws=model_draws, recompute=recompute,
  1338. reestimate_predictions=reestimate_predictions))
  1339. df[-1]['layer'] = layer
  1340. df = pd.concat(df, ignore_index=True)
  1341. df.loc[df['layer'] == 'layer4[0].relu', 'area'] = 'IT-mapped layer'
  1342. df.loc[df['layer'] == 'layer2[0].relu', 'area'] = 'Early control layer'
  1343. df.loc[df['layer'] == 'avgpool', 'area'] = 'Late control layer'
  1344. neural_df = load_neuralITorthogonalDecodes(recompute=recompute, draws=draws, mode=mode, metric=metric)
  1345. neural_df = neural_df.reset_index()
  1346. neural_df['delta'] = computeDelta(neural_df['corrs', 'untrained'], neural_df['corrs', 'trained'], kind=delta)
  1347. neural_df_plot = neural_df.copy()
  1348. neural_df_plot.columns = ["_".join(a) for a in neural_df_plot.columns.to_flat_index()]
  1349. neural_df_plot = neural_df_plot.melt(value_vars=['corrs_trained', 'corrs_untrained'], id_vars=['draws_', 'targetFeature_'])
  1350. df['delta'] = computeDelta(df['corrs_pre'], df['corrs_post'])
  1351. df['# matched IT-category metrics'] = None
  1352. df.loc[df['model_ID'].isin(models_pt3.loc[models_pt3['match_all'], 'model_ID']), '# matched IT-category metrics'] = '3/3'
  1353. df.loc[df['model_ID'].isin(models_pt3.loc[models_pt3['match_one'], 'model_ID']), '# matched IT-category metrics'] = '1/3'
  1354. df.loc[df['model_ID'].isin(models_pt3.loc[models_pt3['match_two'], 'model_ID']), '# matched IT-category metrics'] = '2/3'
  1355. df.loc[df['model_ID'].isin(models_pt3.loc[models_pt3['match_none'], 'model_ID']), '# matched IT-category metrics'] = '0/3'
  1356. df_means = df.groupby(['model_ID', 'targetFeature', 'base_model', 'base_model_type', '# matched IT-category metrics', 'area'], as_index=False)['delta'].mean()
  1357. df_means['targetFeature'] = df_means['targetFeature'].str.split('_', expand=True)[2]
  1358. col_order = ['xpos', 'ypos', 'objSize', 'ecc', 'ryz', 'rxz', 'rxy']
  1359. title_mapping = {'xpos': 'Vertical position',
  1360. 'ypos': 'Horizontal position',
  1361. 'objSize': 'Object size',
  1362. 'ecc': 'Eccentricity',
  1363. 'rxz': 'Rotation (ry)',
  1364. 'ryz': 'Rotation (rx)',
  1365. 'rxy': 'Rotation (rz)'}
  1366. colors = sns.xkcd_palette(['tangerine', 'raspberry', 'pastel purple'])
  1367. g = sns.catplot(df_means, col='targetFeature', height=4.5 * cm, aspect=0.8,
  1368. y='delta', x='area', kind='strip', order=['Early control layer', 'IT-mapped layer', 'Late control layer'],
  1369. palette=colors,
  1370. sharey=False, sharex=False, col_order=col_order, alpha=0.5, zorder=1)
  1371. g.refline(y=0)
  1372. for title, ax in g.axes_dict.items():
  1373. ax.axhspan(np.quantile(neural_df.loc[neural_df['targetFeature'] == title, 'delta'], 0.025),
  1374. np.quantile(neural_df.loc[neural_df['targetFeature'] == title, 'delta'], 0.975),
  1375. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=.4)
  1376. ax.set_title(title_mapping[title], fontsize=10)
  1377. g.set_ylabels('%Δ performance')
  1378. g.set_xlabels('Layer')
  1379. g.set_xticklabels('')
  1380. plt.tight_layout()
  1381. #plt.show()
  1382. g.savefig(figure_dir + f'FigureS11_orthogonalDecodes_change_controlLayers.pdf', dpi=300, transparent=True)
  1383. g.savefig(figure_dir + f'FigureS11_orthogonalDecodes_change_controlLayers.png', dpi=300)
  1384. df_means.to_csv(
  1385. source_file_dir + 'FigureS11_orthogonalDecodes_change_controlLayers.csv', index=False)
  1386. def Figure7_LFI_models(training_groups, recompute=False, meta_dir = None, figure_dir= None, source_file_dir=None, max_dim = 250):
  1387. if figure_dir == None:
  1388. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1389. if meta_dir == None:
  1390. meta_dir = f'{Path(__file__).parent.parent}/Results/meta/'
  1391. if source_file_dir == None:
  1392. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1393. models_pt3 = pd.read_pickle(meta_dir + f'SelectedModels_pt3_means.pkl')
  1394. models_pt3['# metrics matched'] = models_pt3['match_selectivity'].astype(int) + models_pt3['match_rsa'].astype(
  1395. int) + models_pt3['match_decoding'].astype(int)
  1396. df_model = get_model_change_metrics(models_pt3, training_groups, recompute=recompute)
  1397. categories = np.arange(8)
  1398. plot_df = []
  1399. pca_df = []
  1400. for model_ID in df_model:
  1401. for c in categories:
  1402. plot_df.append({'model_ID': model_ID,
  1403. '# metrics matched': models_pt3.loc[models_pt3['model_ID'] == model_ID, '# metrics matched'].item(),
  1404. 'base_model': df_model[model_ID]['base_model'],
  1405. 'training': df_model[model_ID]['Training'],
  1406. 'category': c,
  1407. 'state': 'naive',
  1408. 'N_units': int(df_model[model_ID][c]['N_units']),
  1409. 'Signal strength': df_model[model_ID][c]['df_abs_pre'],
  1410. 'Response variance': df_model[model_ID][c]['variance_median_pre'],
  1411. 'Noise correlations': df_model[model_ID][c]['NoiseCorrelation_median_pre'],
  1412. 'Fano Factor': df_model[model_ID][c]['FanoFactor_median_pre'],
  1413. 'aLFI': float(df_model[model_ID][c]['LFI_pre']),
  1414. })
  1415. pca_df.append(pd.DataFrame({'model_ID': model_ID,
  1416. '# metrics matched': models_pt3.loc[
  1417. models_pt3['model_ID'] == model_ID, '# metrics matched'].item(),
  1418. 'base_model': df_model[model_ID]['base_model'],
  1419. 'training': df_model[model_ID]['Training'],
  1420. 'category': c,
  1421. 'N_units': int(df_model[model_ID][c]['N_units']),
  1422. 'N_components': np.arange(len(df_model[model_ID][c]['PCA_rotation']))[:max_dim],
  1423. 'PCA rotation': np.array(df_model[model_ID][c]['PCA_rotation'])[:max_dim],
  1424. 'PC absolute (naive)': df_model[model_ID][c]['PC_abs_pre'][:max_dim],
  1425. 'PC relative (naive)': df_model[model_ID][c]['PC_abs_pre'][:max_dim]/ np.array(df_model[model_ID][c]['PC_abs_pre']).sum(),
  1426. 'PC absolute (trained)': df_model[model_ID][c]['PC_abs_post'][:max_dim],
  1427. 'PC relative (trained)': df_model[model_ID][c]['PC_abs_post'][:max_dim] / np.array(
  1428. df_model[model_ID][c]['PC_abs_post'][:max_dim]).sum(),
  1429. 'PC ratio': (df_model[model_ID][c]['PC_abs_post'][:max_dim] / np.array(
  1430. df_model[model_ID][c]['PC_abs_post'][:max_dim]).sum())/(df_model[model_ID][c]['PC_abs_pre'][:max_dim]/ np.array(df_model[model_ID][c]['PC_abs_pre'][:max_dim]).sum())
  1431. }))
  1432. plot_df.append({'model_ID': model_ID,
  1433. '# metrics matched': models_pt3.loc[
  1434. models_pt3['model_ID'] == model_ID, '# metrics matched'].item(),
  1435. 'base_model': df_model[model_ID]['base_model'],
  1436. 'training': df_model[model_ID]['Training'],
  1437. 'category': c,
  1438. 'state': 'trained',
  1439. 'N_units': int(df_model[model_ID][c]['N_units']),
  1440. 'Signal strength': df_model[model_ID][c]['df_abs_post'],
  1441. 'Response variance': df_model[model_ID][c]['variance_median_post'],
  1442. 'Noise correlations': df_model[model_ID][c]['NoiseCorrelation_median_post'],
  1443. 'Fano Factor': df_model[model_ID][c]['FanoFactor_median_post'],
  1444. 'aLFI': float(df_model[model_ID][c]['LFI_post']),
  1445. 'rotation_pre_post': float(df_model[model_ID][c]['rotation_pre_post']),
  1446. 'gain_LFI': np.log10(df_model[model_ID][c]['LFI_post']) - np.log10(df_model[model_ID][c]['LFI_pre']),
  1447. 'ratio_Signal strength': float(df_model[model_ID][c]['df_abs_post']) / float(
  1448. df_model[model_ID][c]['df_abs_pre']),
  1449. 'ratio_Response variance': float(df_model[model_ID][c]['variance_median_post']) / float(
  1450. df_model[model_ID][c]['variance_median_pre']),
  1451. 'ratio_FanoFactor': df_model[model_ID][c]['FanoFactor_median_post'] / df_model[model_ID][c]['FanoFactor_median_pre'],
  1452. 'ratio_NC': df_model[model_ID][c]['NoiseCorrelation_median_post'] / df_model[model_ID][c][
  1453. 'NoiseCorrelation_median_pre']
  1454. })
  1455. plot_df = pd.DataFrame(plot_df)
  1456. pca_df = pd.concat(pca_df, ignore_index=True)
  1457. plot_df_mean = plot_df.groupby(['model_ID', 'state', 'base_model', 'training'], as_index=False).mean()
  1458. palette = sns.color_palette(['#e9d2de', '#c0bcde', '#5369B0', '#4a3f99'])
  1459. # Plot signal separation
  1460. g = sns.catplot(data=plot_df_mean, x='# metrics matched', y='ratio_Signal strength',
  1461. aspect=0.9, height=6 * cm, kind='point', legend=False,
  1462. hue='# metrics matched', palette=palette)
  1463. g.set_xticklabels([0, 1, 2, 3])
  1464. g.set_ylabels('Signal strength ratio\n(trained/naive)')
  1465. g.refline(y=1)
  1466. g.savefig(figure_dir + f'Figure7_InfoMetrics_Signal.png', dpi=300)
  1467. g.savefig(figure_dir + f'Figure7_InfoMetrics_Signal.pdf', dpi=300, transparent=True)
  1468. plot_df_mean.to_csv(
  1469. source_file_dir + 'Figure7_InfoMetrics.csv', index=False)
  1470. g = sns.catplot(data=plot_df_mean, x='# metrics matched', y='ratio_Response variance',
  1471. aspect=0.9, height=6 * cm, kind='point', legend=False,
  1472. hue='# metrics matched', palette=palette)
  1473. g.set_xticklabels([0, 1, 2, 3])
  1474. g.set_ylabels('Response variance\nratio (trained/naive)')
  1475. g.refline(y=1)
  1476. g.savefig(figure_dir + f'Figure7_InfoMetrics_Var.png', dpi=300)
  1477. g.savefig(figure_dir + f'Figure7_InfoMetrics_Var.pdf', dpi=300, transparent=True)
  1478. g = sns.catplot(data=plot_df_mean, x='# metrics matched', y='rotation_pre_post',
  1479. aspect=0.9, height=6 * cm, kind='point', legend=False,
  1480. hue='# metrics matched', palette=palette)
  1481. g.set_xticklabels([0, 1, 2, 3])
  1482. g.refline(y=0)
  1483. g.set_ylabels('Signal rotation angle\n(trained/naive)')
  1484. g.savefig(figure_dir + f'Figure7_InfoMetrics_Rot.png', dpi=300)
  1485. g.savefig(figure_dir + f'Figure7_InfoMetrics_Rot.pdf', dpi=300, transparent=True)
  1486. g = sns.relplot(data=pca_df, x='N_components', y='PCA rotation', kind='line',
  1487. hue='# metrics matched', palette=palette, height=6 * cm, aspect=0.9, legend=False)
  1488. g.set(xlim=[0, 25])
  1489. g.set_xlabels('Principal components')
  1490. g.set_ylabels('Covariance rotation\nangle (trained/naive)')
  1491. g.savefig(figure_dir + f'Figure7_InfoMetrics_PCsRot.png', dpi=300)
  1492. g.savefig(figure_dir + f'Figure7_InfoMetrics_PCsRot.pdf', dpi=300, transparent=True)
  1493. plot_df_mean.to_csv(
  1494. source_file_dir + 'Figure7_InfoMetrics_PCsRot.csv', index=False)
  1495. def Figure8_orthogonalDecoding_means(training_groups,
  1496. recompute=False, metric='corrs', model_draws=100, draws=1000,
  1497. delta='proportion', mode='empirical_conditioned_v2',
  1498. reestimate_predictions=False,
  1499. figure_dir=None,
  1500. meta_dir=None,
  1501. source_file_dir=None,
  1502. version='v0'):
  1503. if figure_dir == None:
  1504. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1505. if meta_dir == None:
  1506. meta_dir = f'{Path(__file__).parent.parent}/Results/meta/'
  1507. if source_file_dir == None:
  1508. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1509. models_pt3 = pd.read_pickle(meta_dir + f'SelectedModels_pt3_means.pkl')
  1510. models_pt3['match_one'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(int) +
  1511. models_pt3['match_rsa'].astype(int)) == 1
  1512. models_pt3['match_two'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(
  1513. int) + models_pt3['match_rsa'].astype(int)) == 2
  1514. df = load_ModelITorthogonalDecodes_selection(models_pt3, training_groups, draws=model_draws, recompute=recompute,
  1515. reestimate_predictions=reestimate_predictions)
  1516. neural_df = load_neuralITorthogonalDecodes(recompute=recompute, draws=draws, mode=mode, metric=metric)
  1517. neural_df = neural_df.reset_index()
  1518. neural_df['delta'] = computeDelta(neural_df['corrs', 'untrained'], neural_df['corrs', 'trained'], kind=delta)
  1519. if version == 'v0':
  1520. print('Medians:')
  1521. print(neural_df.groupby(['targetFeature']).median())
  1522. for target in neural_df['targetFeature'].unique():
  1523. print(target)
  1524. print(
  1525. f"{((neural_df.loc[neural_df['targetFeature'] == target, 'corrs']['trained'].mean() / neural_df.loc[neural_df['targetFeature'] == target, 'corrs']['untrained'].mean()) - 1) * 100}% increase")
  1526. difference_dist = neural_df.loc[neural_df['targetFeature'] == target, 'corrs']['trained'] - \
  1527. neural_df.loc[neural_df['targetFeature'] == target, 'corrs']['untrained']
  1528. lower_CI = difference_dist.quantile(0.025)
  1529. upper_CI = difference_dist.quantile(0.975)
  1530. p_val = np.min([(difference_dist < 0).mean(), (difference_dist > 0).mean()]) * 2
  1531. print(f'mean diff {difference_dist.mean()}')
  1532. print(f'mean SD {difference_dist.std()}')
  1533. print(
  1534. f"CI (95%): {lower_CI} - {upper_CI}")
  1535. print(f"p-val: {p_val}")
  1536. neural_df_plot = neural_df.copy()
  1537. neural_df_plot.columns = ["_".join(a) for a in neural_df_plot.columns.to_flat_index()]
  1538. neural_df_plot = neural_df_plot.melt(value_vars=['corrs_trained', 'corrs_untrained'],
  1539. id_vars=['draws_', 'targetFeature_'])
  1540. g = sns.catplot(data=neural_df_plot, x='targetFeature_', y='value', hue='variable', kind='point',
  1541. errorbar=("pi", 95),
  1542. palette={'corrs_untrained': sns.xkcd_palette(['slate grey'])[0],
  1543. 'corrs_trained': sns.xkcd_palette(['jade'])[0]}, height=4.5 * cm, aspect=2,
  1544. legend=False, order=['xpos', 'ypos', 'ecc', 'objSize', 'ryz', 'rxz', 'rxy'], join=False,
  1545. dodge=0.25)
  1546. g.refline(y=0)
  1547. g.set_xticklabels(["x-pos.", "y-pos.", 'ecc.', 'size', 'rx', 'ry', 'rz'])
  1548. g.set_ylabels("Correlation")
  1549. g.set_xlabels("Category-orthogonal property")
  1550. g.savefig(figure_dir + 'FigureS13_orthogonal_neuralDecodes.png', dpi=300)
  1551. g.savefig(figure_dir + 'FigureS13_orthogonal_neuralDecodes.pdf', dpi=300, transparent=True)
  1552. neural_df_plot.to_csv(
  1553. source_file_dir + 'FigureS13_orthogonal_neuralDecodes.csv', index=False)
  1554. elif version == 'v1':
  1555. df['delta'] = computeDelta(df['corrs_pre'], df['corrs_post'])
  1556. df['# matched IT-category metrics'] = None
  1557. df.loc[df['model_ID'].isin(
  1558. models_pt3.loc[models_pt3['match_all'], 'model_ID']), '# matched IT-category metrics'] = '3/3'
  1559. df.loc[df['model_ID'].isin(
  1560. models_pt3.loc[models_pt3['match_one'], 'model_ID']), '# matched IT-category metrics'] = '1/3'
  1561. df.loc[df['model_ID'].isin(
  1562. models_pt3.loc[models_pt3['match_two'], 'model_ID']), '# matched IT-category metrics'] = '2/3'
  1563. df.loc[df['model_ID'].isin(
  1564. models_pt3.loc[models_pt3['match_none'], 'model_ID']), '# matched IT-category metrics'] = '0/3'
  1565. df_means = df.groupby(
  1566. ['model_ID', 'targetFeature', 'base_model', 'base_model_type', '# matched IT-category metrics'],
  1567. as_index=False).mean()
  1568. df_means['targetFeature'] = df_means['targetFeature'].str[14:]
  1569. col_order = ['xpos', 'ypos', 'objSize', 'ecc', 'ryz', 'rxz', 'rxy']
  1570. title_mapping = {'xpos': 'Vertical position',
  1571. 'ypos': 'Horizontal position',
  1572. 'objSize': 'Object size',
  1573. 'ecc': 'Eccentricity',
  1574. 'rxz': 'Rotation (ry)',
  1575. 'ryz': 'Rotation (rx)',
  1576. 'rxy': 'Rotation (rz)'}
  1577. g = sns.catplot(df_means, col='targetFeature', col_wrap=4, height=4 * cm, aspect=1,
  1578. y='delta', x='# matched IT-category metrics', kind='strip', order=['0/3', '1/3', '2/3', '3/3'],
  1579. color=sns.xkcd_palette(['cerulean'])[0],
  1580. sharey=False, sharex=False, col_order=col_order, alpha=0.1, zorder=1)
  1581. g.refline(y=0)
  1582. for title, ax in g.axes_dict.items():
  1583. ax.axhspan(np.quantile(neural_df.loc[neural_df['targetFeature'] == title, 'delta'], 0.025),
  1584. np.quantile(neural_df.loc[neural_df['targetFeature'] == title, 'delta'], 0.975),
  1585. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=.5)
  1586. ax.set_title(title_mapping[title], fontsize=10)
  1587. sns.pointplot(
  1588. data=df_means.loc[df_means['targetFeature'] == title], x='# matched IT-category metrics', y='delta',
  1589. join=False,
  1590. order=['0/3', '1/3', '2/3', '3/3'],
  1591. errorbar=None, estimator=np.median, color=sns.xkcd_palette(['cerulean'])[0],
  1592. markers="_", scale=2, ax=ax,
  1593. )
  1594. g.set_ylabels('%Δ accuracy')
  1595. g.set_xlabels('# IT-aligned\ncategory metrics')
  1596. g.savefig(figure_dir + f'Figure8_orthogonalDecodes.pdf', dpi=300, transparent=True)
  1597. g.savefig(figure_dir + f'Figure8_orthogonalDecodes.png', dpi=300)
  1598. df_means.to_csv(
  1599. source_file_dir + 'Figure8_orthogonalDecodes.csv', index=False)
  1600. k = df_means['# matched IT-category metrics'].nunique() # four model groups
  1601. n = df_means['model_ID'].nunique()
  1602. for target in col_order:
  1603. print(target)
  1604. stats = kruskal(data=df_means.loc[df_means['targetFeature'] == target], dv='delta',
  1605. between='# matched IT-category metrics')
  1606. print(stats)
  1607. epsilon_sq = (stats['H'].item() - k + 1) / (n - k)
  1608. print('Effect size (epsilon_sq):', epsilon_sq)
  1609. def Figure8_performanceStripes(figure_dir=None,
  1610. recompute=False, num_sites=159,
  1611. ):
  1612. if figure_dir == None:
  1613. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1614. IT_neurons = load_MonkeyIT_Decodes(num_neurons=np.array([num_sites]), includeReliability=False, recompute=recompute)
  1615. IT_image_trained = np.mean(IT_neurons[f'{num_sites}_neurons']['trained']['dprimes'], axis=0)
  1616. g = sns.FacetGrid(data=pd.DataFrame(np.zeros((2, 1))), height=10 * cm, aspect=0.3)
  1617. sns.heatmap(IT_image_trained[:, None], cmap='Spectral', ax=g.ax, cbar=False,
  1618. vmin=np.percentile(IT_image_trained, 5), vmax=np.percentile(IT_image_trained, 95))
  1619. g.ax.set(xlabel="", ylabel="", yticks=np.arange(1, 8) * 80, xticks=[])
  1620. plt.savefig(figure_dir + 'Figure6_performanceStripes_ITtrained.png', dpi=300)
  1621. plt.savefig(figure_dir + 'Figure6_performanceStripes_ITtrained.pdf', dpi=300, transparent=True)
  1622. IT_image_naive = np.mean(IT_neurons[f'{num_sites}_neurons']['untrained']['dprimes'], axis=0)
  1623. g = sns.FacetGrid(data=pd.DataFrame(np.zeros((2, 1))), height=10 * cm, aspect=0.3)
  1624. sns.heatmap(IT_image_naive[:, None], cmap='Spectral', ax=g.ax, cbar=False,
  1625. vmin=np.percentile(IT_image_naive, 5), vmax=np.percentile(IT_image_naive, 95))
  1626. g.ax.set(xlabel="", ylabel="", yticks=np.arange(1, 8) * 80, xticks=[])
  1627. plt.savefig(figure_dir + 'Figure6_performanceStripes_ITnaive.png', dpi=300)
  1628. plt.savefig(figure_dir + 'Figure6_performanceStripes_ITnaive.pdf', dpi=300, transparent=True)
  1629. pooled_behavior = load_MonkeyBehavior_pooled()
  1630. dPrimes, _ = dPrime_monkey(pooled_behavior)
  1631. g = sns.FacetGrid(data=pd.DataFrame(np.zeros((2, 1))), height=10 * cm, aspect=0.3)
  1632. sns.heatmap(dPrimes['i1'][:, None], cmap='Spectral', ax=g.ax, cbar=False,
  1633. vmin=np.percentile(dPrimes['i1'], 5), vmax=np.percentile(dPrimes['i1'], 95))
  1634. g.ax.set(xlabel="", ylabel="", yticks=np.arange(1, 8) * 80, xticks=[])
  1635. plt.savefig(figure_dir + 'Figure6_performanceStripes_behavior.png', dpi=300)
  1636. plt.savefig(figure_dir + 'Figure6_performanceStripes_behavior.pdf', dpi=300, transparent=True)
  1637. def Figure8_ITbehavioralConsistency_means(training_groups, seed=6,
  1638. recompute=False, metric='i1', model_draws=100, draws=1000, num_sites=159,
  1639. delta='proportion', mode='empirical_conditioned_v2',
  1640. source_file_dir=None,
  1641. figure_dir=None,
  1642. meta_dir=None, version='v1'):
  1643. if figure_dir == None:
  1644. figure_dir = f'{Path(__file__).parent.parent}/Figures/'
  1645. if meta_dir == None:
  1646. meta_dir = f'{Path(__file__).parent.parent}/Results/meta/'
  1647. if source_file_dir == None:
  1648. source_file_dir = f'{Path(__file__).parent.parent}/SourceDataFiles/'
  1649. np.random.seed(seed)
  1650. models_pt3 = pd.read_pickle(meta_dir + f'SelectedModels_pt3_means.pkl')
  1651. models_pt3['match_one'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(int) +
  1652. models_pt3['match_rsa'].astype(int)) == 1
  1653. models_pt3['match_two'] = (models_pt3['match_decoding'].astype(int) + models_pt3['match_selectivity'].astype(
  1654. int) + models_pt3['match_rsa'].astype(int)) == 2
  1655. neural_df = get_consistency_behavior_ITneurons(consistency_metric=metric, recompute=recompute, num_sites=num_sites,
  1656. mode=mode, draws=draws)
  1657. neural_df = neural_df.reset_index()
  1658. neural_df['delta'] = computeDelta(neural_df['consistency_i1_untrained'], neural_df['consistency_i1_trained'],
  1659. kind=delta)
  1660. if version == 'v0':
  1661. print('Medians:')
  1662. print(neural_df.median())
  1663. print('Means:')
  1664. print(neural_df.mean())
  1665. print('SD:')
  1666. print(neural_df.std())
  1667. print(
  1668. f"{((neural_df['consistency_i1_trained'].mean() / neural_df['consistency_i1_untrained'].mean()) - 1) * 100}% increase")
  1669. difference_dist = neural_df['consistency_i1_trained'] - \
  1670. neural_df['consistency_i1_untrained']
  1671. lower_CI = difference_dist.quantile(0.025)
  1672. upper_CI = difference_dist.quantile(0.975)
  1673. p_val = np.min([(difference_dist < 0).mean(), (difference_dist > 0).mean()]) * 2
  1674. print(
  1675. f"Behavioral consistency CI (95%): {lower_CI} - {upper_CI}")
  1676. print(f"p-val: {p_val}")
  1677. neural_df_plot = neural_df.copy()
  1678. neural_df_plot = neural_df_plot.melt(value_vars=['consistency_i1_untrained', 'consistency_i1_trained'],
  1679. id_vars='draw')
  1680. g = sns.catplot(data=neural_df_plot, x='variable', y='value', hue='variable', kind='point',
  1681. errorbar=("pi", 95),
  1682. palette={'consistency_i1_untrained': sns.xkcd_palette(['slate grey'])[0],
  1683. 'consistency_i1_trained': sns.xkcd_palette(['jade'])[0]}, height=4.5 * cm, aspect=1.05,
  1684. legend=False, join=False, dodge=0.25)
  1685. g.refline(y=0)
  1686. g.set_ylabels("IT-behavior\nconsistency")
  1687. g.set_xlabels("IT neuron pools")
  1688. g.set_xticklabels(['naïve', 'trained'])
  1689. g.savefig(figure_dir + 'FigureS13_consistency_neuralDecodes.png', dpi=300)
  1690. g.savefig(figure_dir + 'FigureS13_consistency_neuralDecodes.pdf', dpi=300, transparent=True)
  1691. neural_df_plot.to_csv(
  1692. source_file_dir + 'FigureS13_consistency_neuralDecodes.csv', index=False)
  1693. elif version == 'v1':
  1694. df = get_consistency_modelIT_modelBehavior(models_pt3, training_groups, recompute=recompute)
  1695. df['delta'] = computeDelta(df['consistency_i1_pre'], df['consistency_i1_post'])
  1696. df['Training'] = df['model_ID'].map(dict(zip(models_pt3['model_ID'], models_pt3['Training'])))
  1697. df['# matched IT-category metrics'] = None
  1698. df.loc[df['model_ID'].isin(
  1699. models_pt3.loc[models_pt3['match_all'], 'model_ID']), '# matched IT-category metrics'] = '3/3'
  1700. df.loc[df['model_ID'].isin(
  1701. models_pt3.loc[models_pt3['match_one'], 'model_ID']), '# matched IT-category metrics'] = '1/3'
  1702. df.loc[df['model_ID'].isin(
  1703. models_pt3.loc[models_pt3['match_two'], 'model_ID']), '# matched IT-category metrics'] = '2/3'
  1704. df.loc[df['model_ID'].isin(
  1705. models_pt3.loc[models_pt3['match_none'], 'model_ID']), '# matched IT-category metrics'] = '0/3'
  1706. df_means = df.groupby(['model_ID', '# matched IT-category metrics', 'Training'], as_index=False).mean()
  1707. print(kruskal(data=df_means, dv='delta', between='# matched IT-category metrics'))
  1708. g = sns.catplot(df_means, height=4 * cm, aspect=1.15,
  1709. y='delta', x='# matched IT-category metrics', kind='strip', order=['0/3', '1/3', '2/3', '3/3'],
  1710. color=sns.xkcd_palette(['cerulean'])[0],
  1711. sharey=False, sharex=False, zorder=1, alpha=0.1, dodge=True)
  1712. g.refline(y=0)
  1713. g.ax.axhspan(np.quantile(neural_df['delta'], 0.025),
  1714. np.quantile(neural_df['delta'], 0.975),
  1715. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=.5)
  1716. sns.pointplot(
  1717. data=df_means, x='# matched IT-category metrics', y='delta', join=False,
  1718. order=['0/3', '1/3', '2/3', '3/3'],
  1719. errorbar=None, estimator=np.median, color=sns.xkcd_palette(['cerulean'])[0],
  1720. markers="_", scale=2, ax=g.ax,
  1721. )
  1722. g.set_ylabels('%Δ consistency')
  1723. g.set_xlabels('# IT-aligned\ncategory metrics')
  1724. g.savefig(figure_dir + f'Figure8_behavioralConsistency_v1.pdf', dpi=300, transparent=True)
  1725. g.savefig(figure_dir + f'Figure8_behavioralConsistency_v1.png', dpi=300)
  1726. df_means.to_csv(
  1727. source_file_dir + 'Figure8_behavioralConsistency_v1.csv', index=False)
  1728. g = sns.catplot(df_means, height=4 * cm, aspect=1.15,
  1729. y='delta', x='# matched IT-category metrics', kind='strip', order=['0/3', '1/3', '2/3', '3/3'],
  1730. color=sns.xkcd_palette(['cerulean'])[0],
  1731. sharey=False, sharex=False, zorder=1, alpha=0.1,
  1732. dodge=True)
  1733. g.refline(y=0)
  1734. g.ax.axhspan(np.quantile(neural_df['delta'], 0.025),
  1735. np.quantile(neural_df['delta'], 0.975),
  1736. color=sns.xkcd_palette(['grey'])[0], zorder=0, alpha=.5)
  1737. sns.pointplot(
  1738. data=df_means, x='# matched IT-category metrics', y='delta', join=False,
  1739. order=['0/3', '1/3', '2/3', '3/3'],
  1740. errorbar=None, estimator=np.median, color=sns.xkcd_palette(['cerulean'])[0],
  1741. markers="_", scale=2, ax=g.ax,
  1742. )
  1743. g.set_ylabels('%Δ consistency')
  1744. g.set_xlabels('# IT-aligned\ncategory metrics')
  1745. g.set(ylim=[-20, 105])
  1746. g.savefig(figure_dir + f'Figure8_behavioralConsistency_v1_zoom.pdf', dpi=300, transparent=True)
  1747. g.savefig(figure_dir + f'Figure8_behavioralConsistency_v1_zoom.png', dpi=300)

paperFigures.py at commit 85fe071, under MIT · at the source

Overview

  1. McGovern Institute for Brain Research, Dept. of Brain and Cognitive Sciences, Massachusetts Institute of Technology, Cambridge, USA
  2. Center for Brains, Minds and Machines, Massachusetts Institute of Technology, Cambridge, USA
  3. MIT Quest for Intelligence, Cambridge, USA
  4. Centre for Integrative and Applied Neuroscience, York University, Toronto, Canada
  5. Department of Biology, Centre for Vision Research, York University, Toronto, Canada
Journal: Nature communications, volume 17, issue 1, article 8434
Dates: received 14 July 2025; accepted 12 June 2026; published online 8 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41467-026-74816-0 · PMID 42420264 · PMCID PMC13478300 · OpenAlex W7167740178
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: non-human primate (organism)
Methods: Spectral & time-frequency, Connectivity, Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Single-unit activity, calcium imaging, Physiology & signal measures
Keywords: Object vision, Cortex
MeSH: Learning*, Neuronal Plasticity*, Temporal Lobe*, Animals, Macaca mulatta, Male, Models, Neurological, Neural Networks, Computer, Neurons, Pattern Recognition, Visual, Photic Stimulation, Visual Pathways (* major topic)
Topic: Face Recognition and Perception (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: National Science Foundation (NSF) (2124136, CCF-1231216); Simons Foundation (AN-NC-GB-Culmination-00002986-04, SFARI, 967073); Canada Research Chairs (CRC-2021-00326); Canada First Research Excellence Fund (VISTA Program); Deutsche Forschungsgemeinschaft (German Research Foundation) (547591872); United States Department of Defense | United States Navy | Office of Naval Research (ONR) (N00014-21-1-2801, N00014-20-1-2589); Canadian Network for Research and Innovation in Machining Technology, Natural Sciences and Engineering Research Council of Canada (RGPIN 2024-06223)
Citations: cited by 1 paper (Europe PMC); 131 references in the paper

Abstract

How does the primate brain coordinate plasticity when learning to discriminate new objects? We measured consequences of object learning on inferior temporal (IT) cortex, a key waypoint supporting object recognition in the ventral visual stream, in male macaques. Neural activity in task-trained monkeys’ IT showed increased object selectivity, enhanced linear separability across objects, and more object-invariant representations compared to task-naïve monkeys. To model these differences, we developed a computational framework using anatomically-mapped artificial neural network (ANN) models of the ventral stream with various learning algorithms. Simulations revealed that gradient-based, performance-optimizing updates of ANN internal representations accurately approximated observed IT cortex changes. These models predict novel training-induced phenomena in the IT cortex, including changes independent of object identity and IT’s alignment with behavior. This convergence between empirical measurements and model predictions suggests ventral stream plasticity follows task optimization principles well-approximated by gradient descent, enabling accurate predictions about visual plasticity and generalization to test images.

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

Repositories

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

dicarlolab/mkturk

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4bdfaca7d1caa02115fb16c97b158dca1e620595, 1 September 2018
Languages: JavaScript (14)
Size: 69 files, 14 scripts
Software Heritage: archived
Found in: the text, “Object training and active binary object discrim”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
15 files

OSF dcgze

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Languages: Python (19)
Size: 170 files, 19 scripts
Software Heritage: not checked
Found in: “Code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (14 files), pandas (13 files), scikit-learn (7 files), PyTorch (6 files), SciPy (5 files), h5py (1 file), Matplotlib (1 file), Pingouin (1 file), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
21 files

lynnsoerensen/ObjectTraining_IT_ANNs

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 85fe071926ac2f44f8ac27428a33ca4a4fd1d4c4, 1 May 2026
Languages: Python (23)
Size: 2,132 files, 23 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 (17 files), pandas (13 files), PyTorch (10 files), scikit-learn (8 files), SciPy (5 files), h5py (1 file), Matplotlib (1 file), Pingouin (1 file), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
25 files

Zenodo 19955396

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
At the source:

Code availability

All code can be accessed on the Open Science Framework repository (https://osf.io/dcgze/ (https://osf.io/dcgze/?view_only=6dfe548c7ba24238932d247e65523053)), on GitHub (https://github.com/lynnsoerensen/ObjectTraining_IT_ANNs) and on Zenodo (10.5281/zenodo.19955396).

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:

  • 4 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 56 scripts, each with its path and the digest of its content;
  • 11 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

No dataset and no data link were found in the paper.

Data availability

All data has been deposited in the following Open Science Framework repository (https://osf.io/dcgze/ (https://osf.io/dcgze/?view_only=6dfe548c7ba24238932d247e65523053)). Source data are provided in this paper.

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, 2 keywords, 12 MeSH terms, 7 funders, 116 references.

Cite

This paper

Sörensen, L. K. A., DiCarlo, J. J., & Kar, K. (2026). Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training. Nature communications, 17(1), 8434. https://doi.org/10.1038/s41467-026-74816-0

BibTeX

@article{sorensen2026hierarchical,
author = {Sörensen, Lynn K A and DiCarlo, James J and Kar, Kohitij},
title = {{Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training}},
journal = {Nature communications},
year = {2026},
month = jul,
volume = {17},
number = {1},
pages = {8434},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-74816-0},
url = {https://doi.org/10.1038/s41467-026-74816-0},
pmid = {42420264},
pmcid = {PMC13478300}
}

RIS

TY - JOUR
AU - Sörensen, Lynn K A
AU - DiCarlo, James J
AU - Kar, Kohitij
TI - Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/07/08
VL - 17
IS - 1
SP - 8434
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-74816-0
UR - https://doi.org/10.1038/s41467-026-74816-0
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-74816-0",
"type": "article-journal",
"title": "Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training",
"container-title": "Nature communications",
"author": [
{
"family": "Sörensen",
"given": "Lynn K A"
},
{
"family": "DiCarlo",
"given": "James J"
},
{
"family": "Kar",
"given": "Kohitij"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "8434",
"DOI": "10.1038/s41467-026-74816-0",
"PMID": "42420264",
"PMCID": "PMC13478300",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-74816-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
8
]
]
}
}

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/s41467-026-76106-1 [code]
Facial expression discrimination emerges from partially overlapping neural subspaces of detection and identity.
Journal: Nature communications
In common: h5py, seaborn, scikit-learn, 4 other tools, non-human primate, 12 references, author Kohitij Kar
[2] 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: h5py, PyTorch, seaborn, 5 other tools, 6 references
[3] doi:10.1038/s41593-026-02207-1 [code]
Neuronal tuning aligns dynamically with object and texture manifolds across the visual hierarchy.
Journal: Nature neuroscience
In common: PyTorch, seaborn, scikit-learn, 4 other tools, non-human primate, 6 references
[4] doi:10.1038/s41467-026-72146-9 [code]
Modeling attention and binding in the brain through bidirectional recurrent gating.
Journal: Nature communications
In common: PyTorch, seaborn, SciPy, 2 other tools, 8 references
[5] 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: h5py, PyTorch, seaborn, 5 other tools, 6 references
[6] doi:10.1038/s41598-026-43946-2
Human lateral occipital complex is invariant for position and size transformations at the single-neuron level.
Journal: Scientific reports
In common: 8 references
[7] doi:10.1371/journal.pone.0347992 [code]
Rotation-tolerant representations elucidate the time-course of high-level object processing.
Journal: PloS one
In common: pandas, NumPy, 7 references
[8] doi:10.1038/s41562-026-02414-7 [code]
Optimized feature gains explain and predict successes and failures of human selective listening.
Journal: Nature human behaviour
In common: Pingouin, h5py, PyTorch, 6 other tools, 2 references
[9] doi:10.1038/s41467-026-74460-8 [code]
Spike-based alignment learning solves the weight transport problem.
Journal: Nature communications
In common: PyTorch, seaborn, scikit-learn, 4 other tools, 4 references
[10] doi:10.7554/elife.105968 [code]
Modeling the hallucinatory effects of classical psychedelics in terms of replay-dependent plasticity mechanisms.
Journal: eLife
In common: PyTorch, scikit-learn, Matplotlib, 1 other tool, 5 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.