OSCR

Neural sampling from cognitive maps enables goal-directed imagination and planning.

Code ↔ Paper

2 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 2 matches
  1. [1] § Methods › Details of our model for spatial goal-directed imagination › Using the cognitive map for generating imagined paths ↔ gcml_grid_cell.ipynb, lines 1–61 · score 0.75 · constant speed, random walk, turn noise, forbidden, resampling, Invalid
  2. [2] § Methods › Details of our model for spatial goal-directed imagination › Generating a cognitive map based on grid cells ↔ gcml_grid_cell.ipynb, lines 142–285 · score 0.53 · event, plateau, arena, field, Pi, firing

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

Jupyter notebook · 874 lines · 26 KB · MIT · 2 matches

  1. # %%
  2. import numpy as np
  3. import matplotlib.pyplot as plt
  4. # Environment settings
  5. env_min = 0
  6. env_max = 4
  7. # Random walk settings
  8. n_steps = 50000
  9. dt = 0.05
  10. speed = 0.1 # constant speed
  11. turn_noise = 0.1 # radians per step
  12. n_traj = 1 # number of trajectories to plot
  13. # Helper: is a point valid?
  14. def is_valid(pos):
  15. x, y = pos
  16. if not (env_min <= x <= env_max and env_min <= y <= env_max):
  17. return False
  18. return True
  19. # Generate trajectories
  20. trajectories = []
  21. for _ in range(n_traj):
  22. pos = np.array([np.random.uniform(0, 4), np.random.uniform(0, 4)])
  23. angle = np.random.uniform(0, 2*np.pi)
  24. path = [pos.copy()]
  25. for _ in range(n_steps):
  26. # Add small random turn
  27. angle += np.random.randn() * turn_noise
  28. velocity = speed * np.array([np.cos(angle), np.sin(angle)])
  29. new_pos = pos + velocity * dt
  30. if is_valid(new_pos):
  31. pos = new_pos
  32. else:
  33. # Resample a new random direction if invalid move
  34. angle = np.random.uniform(0, 2*np.pi)
  35. path.append(pos.copy())
  36. trajectories.append(np.array(path))
  37. # Plot
  38. fig, ax = plt.subplots(figsize=(6, 6),dpi=300)
  39. for traj in trajectories:
  40. # Plot trajectory
  41. ax.plot(traj[:,0], traj[:,1], color='gray', alpha=0.5, linewidth=1)
  42. # for i in range(1, len(traj)):
  43. # ax.plot(traj[i-1:i+1,0], traj[i-1:i+1,1], color=‘gray’, alpha=0.5, linewidth=1)
  44. # Mark start
  45. start = traj[0]
  46. ax.plot(start[0], start[1], 'o', markersize=20, c='black')
  47. # Mark end
  48. end = traj[-1]
  49. ax.plot(end[0], end[1], 's', markersize=20, c='black')
  50. # Forbidden region
  51. # ax.add_patch(plt.Rectangle((0, 0), forbidden_x, forbidden_y,
  52. # color='red', alpha=0.2, zorder=0))
  53. # Formatting
  54. ax.set_xlim(env_min, env_max)
  55. ax.set_ylim(env_min, env_max)
  56. ax.set_aspect('equal')
  57. ax.set_xticks([])
  58. ax.set_yticks([])
  59. # ax.set_title(‘Random Walk with Forbidden Region’)
  60. plt.show()
  61. # %%
  62. import numpy as np
  63. import matplotlib.pyplot as plt
  64. # Environment settings
  65. env_min = 0
  66. env_max = 4
  67. forbidden_x = 2
  68. forbidden_y = 2
  69. # Random walk settings
  70. n_steps = 50000
  71. dt = 0.05
  72. speed = 0.1 # constant speed
  73. turn_noise = 0.1 # radians per step
  74. n_traj = 1 # number of trajectories to plot
  75. # Helper: is a point valid?
  76. def is_valid(pos):
  77. x, y = pos
  78. if not (env_min <= x <= env_max and env_min <= y <= env_max):
  79. return False
  80. if x < forbidden_x and y < forbidden_y:
  81. return False
  82. return True
  83. # Generate trajectories
  84. trajectories = []
  85. for _ in range(n_traj):
  86. pos = np.array([np.random.uniform(2, 4), np.random.uniform(2, 4)]) # start outside forbidden
  87. angle = np.random.uniform(0, 2*np.pi)
  88. path = [pos.copy()]
  89. for _ in range(n_steps):
  90. # Add small random turn
  91. angle += np.random.randn() * turn_noise
  92. velocity = speed * np.array([np.cos(angle), np.sin(angle)])
  93. new_pos = pos + velocity * dt
  94. if is_valid(new_pos):
  95. pos = new_pos
  96. else:
  97. # Resample a new random direction if invalid move
  98. angle = np.random.uniform(0, 2*np.pi)
  99. path.append(pos.copy())
  100. trajectories.append(np.array(path))
  101. # Plot
  102. fig, ax = plt.subplots(figsize=(6, 6),dpi=300)
  103. for traj in trajectories:
  104. # Plot trajectory
  105. ax.plot(traj[:,0], traj[:,1], color='gray', alpha=0.99, linewidth=1)
  106. # for i in range(1, len(traj)):
  107. # ax.plot(traj[i-1:i+1,0], traj[i-1:i+1,1], color='gray', alpha=0.5, linewidth=1)
  108. # Mark start
  109. start = traj[0]
  110. # ax.plot(start[0], start[1], 'o', markersize=20, c='black')
  111. # Mark end
  112. end = traj[-1]
  113. # ax.plot(end[0], end[1], 's', markersize=20, c='black')
  114. # Forbidden region
  115. # ax.add_patch(plt.Rectangle((0, 0), forbidden_x, forbidden_y,
  116. # color='red', alpha=0.2, zorder=0))
  117. # Formatting
  118. ax.set_xlim(env_min, env_max)
  119. ax.set_ylim(env_min, env_max)
  120. ax.set_aspect('equal')
  121. ax.set_xticks([])
  122. ax.set_yticks([])
  123. # ax.set_title('Random Walk with Forbidden Region')
  124. plt.show()
  125. # %%
  126. import numpy as np
  127. import matplotlib.pyplot as plt
  128. # -------------------------
  129. # Function Definitions
  130. # -------------------------
  131. def is_valid(pos, env_min=0, env_max=4):
  132. x, y = pos
  133. return env_min <= x <= env_max and env_min <= y <= env_max
  134. def generate_grid_cells(n_cells=1000, arena_size=4, seed=None):
  135. if seed is not None:
  136. np.random.seed(seed)
  137. scales = np.random.uniform(0.02, 8, size=n_cells)
  138. phases = np.random.uniform(0, arena_size, size=(n_cells, 2))
  139. orientations = np.random.uniform(0, np.pi/3, size=n_cells)
  140. return scales, phases, orientations
  141. def grid_activation(x, y, scales, phases, orientations):
  142. positions = np.array([x, y])
  143. activations = []
  144. angles = np.array([0, np.pi / 3, 2 * np.pi / 3])
  145. for scale, phase, orientation in zip(scales, phases, orientations):
  146. # Rotate and compute grating
  147. R = np.array([[np.cos(orientation), -np.sin(orientation)],
  148. [np.sin(orientation), np.cos(orientation)]])
  149. pos_rot = R @ (positions - phase)
  150. proj = pos_rot[0] * np.cos(angles) + pos_rot[1] * np.sin(angles)
  151. grating = np.sum(np.cos((4 * np.pi / (scale * np.sqrt(3))) * proj))
  152. activations.append((2/3) * grating)
  153. return np.array(activations)
  154. def grid_activation_vectorized(X, Y, scales, phases, orientations):
  155. H, W = X.shape
  156. n_cells = len(scales)
  157. pos = np.stack([X.ravel(), Y.ravel()], axis=1) # (H*W, 2)
  158. activations = np.zeros((H * W, n_cells))
  159. angles = np.array([0, np.pi / 3, 2 * np.pi / 3])
  160. for i in range(n_cells):
  161. scale = scales[i]
  162. phase = phases[i]
  163. orientation = orientations[i]
  164. c, s = np.cos(orientation), np.sin(orientation)
  165. R = np.array([[c, -s], [s, c]])
  166. rel_pos = pos - phase
  167. pos_rot = rel_pos @ R.T
  168. proj = np.outer(pos_rot[:, 0], np.cos(angles)) + np.outer(pos_rot[:, 1], np.sin(angles))
  169. grating = np.sum(np.cos((4 * np.pi / (scale * np.sqrt(3))) * proj), axis=1)
  170. activations[:, i] = (2/3) * grating
  171. return activations.reshape(H, W, n_cells)
  172. # -------------------------
  173. # Random Walk Generation
  174. # -------------------------
  175. env_min, env_max = 0, 4
  176. n_steps = 50000
  177. dt = 0.05
  178. speed = 0.2
  179. turn_noise = 0.1
  180. pos = np.random.uniform(env_min, env_max, size=2)
  181. angle = np.random.uniform(0, 2*np.pi)
  182. trajectory = [pos.copy()]
  183. for _ in range(n_steps):
  184. angle += np.random.randn() * turn_noise
  185. vel = speed * np.array([np.cos(angle), np.sin(angle)])
  186. new_pos = pos + vel * dt
  187. if is_valid(new_pos, env_min, env_max):
  188. pos = new_pos
  189. else:
  190. angle = np.random.uniform(0, 2*np.pi)
  191. trajectory.append(pos.copy())
  192. trajectory = np.array(trajectory)
  193. # -------------------------
  194. # Grid Cell Setup
  195. # -------------------------
  196. n_cells = 1000
  197. scales, phases, orientations = generate_grid_cells(n_cells, env_max)
  198. # -------------------------
  199. # Plateau Events & Weights
  200. # -------------------------
  201. n_plateau = 1000
  202. plateau_times = np.sort(np.random.choice(np.arange(len(trajectory)), size=n_plateau, replace=False))
  203. weights = np.zeros((n_plateau, n_cells))
  204. for i, t in enumerate(plateau_times):
  205. x, y = trajectory[t]
  206. weights[i] = grid_activation(x, y, scales, phases, orientations)
  207. # -------------------------
  208. # Plot 1: Random Walk + Plateau Fires
  209. # -------------------------
  210. fig, ax = plt.subplots(figsize=(6, 6), dpi=300)
  211. ax.plot(trajectory[:, 0], trajectory[:, 1], color='gray', alpha=0.5, linewidth=1)
  212. ax.scatter(trajectory[plateau_times, 0], trajectory[plateau_times, 1],
  213. c='black', s=10, label='Plateau Fires')
  214. ax.plot(trajectory[0, 0], trajectory[0, 1], 'o', markersize=12, c='black')
  215. ax.plot(trajectory[-1, 0], trajectory[-1, 1], 's', markersize=12, c='black')
  216. ax.set_xlim(env_min, env_max)
  217. ax.set_ylim(env_min, env_max)
  218. ax.set_aspect('equal')
  219. ax.axis('off')
  220. plt.show()
  221. # -------------------------
  222. # Compute Reactive Fields
  223. # -------------------------
  224. grid_res = 50
  225. x_vals = np.linspace(env_min, env_max, grid_res)
  226. y_vals = np.linspace(env_min, env_max, grid_res)
  227. X, Y = np.meshgrid(x_vals, y_vals)
  228. grid_embeddings = grid_activation_vectorized(X, Y, scales, phases, orientations)
  229. # resulting shape: (grid_res, grid_res, n_cells)
  230. # reactive_fields(h, w, neuron) = sum_i grid_embeddings(h,w,i) * weights(neuron,i)
  231. reactive_fields = np.tensordot(grid_embeddings, weights, axes=([2], [1])) # shape: (H, W, n_plateau)
  232. # -------------------------
  233. # Plot 2: First 100 Reactive Fields
  234. # -------------------------
  235. fig2, axes = plt.subplots(10, 10, figsize=(20, 20))
  236. for idx, ax in enumerate(axes.flatten()):
  237. im = ax.imshow(reactive_fields[:, :, idx],
  238. extent=[env_min, env_max, env_min, env_max],
  239. origin='lower')
  240. ax.axis('off')
  241. plt.tight_layout()
  242. plt.show()
  243. # %%
  244. from sklearn.decomposition import PCA
  245. import numpy as np
  246. import matplotlib.pyplot as plt
  247. # Identify top 1% largest-scale grid cells
  248. threshold = np.percentile(scales,85)
  249. large_indices = np.where(scales >= threshold)[0]
  250. # Extract embeddings for only those cells and flatten
  251. H, W, _ = grid_embeddings.shape
  252. emb_large = grid_embeddings[:, :, large_indices].reshape(H * W, len(large_indices))
  253. # PCA on filtered embeddings
  254. pca_large = PCA(n_components=2)
  255. pc_large = pca_large.fit_transform(emb_large)
  256. # Plot
  257. plt.figure(figsize=(6, 6), dpi=300)
  258. plt.scatter(pc_large[:, 0], pc_large[:, 1], s=5)
  259. # plt.title('PCA of Top 1%‑Scale Grid Cell Embeddings')
  260. # plt.xlabel('PC1')
  261. plt.axis('off')
  262. # plt.ylabel('PC2')
  263. plt.tight_layout()
  264. plt.show()
  265. # %%
  266. from sklearn.decomposition import PCA
  267. import numpy as np
  268. import matplotlib.pyplot as plt
  269. from matplotlib.colors import Normalize, ListedColormap
  270. # Identify top 15% largest-scale grid cells (you had 85th percentile)
  271. threshold = np.percentile(scales, 85)
  272. large_indices = np.where(scales >= threshold)[0]
  273. # Extract embeddings for only those cells and flatten
  274. H, W, _ = grid_embeddings.shape
  275. emb_large = grid_embeddings[:, :, large_indices].reshape(H * W, len(large_indices))
  276. # PCA on filtered embeddings
  277. pca_large = PCA(n_components=2)
  278. pc_large = pca_large.fit_transform(emb_large)
  279. # --- Colormap: Blues truncated to [0.1, 0.3] ---
  280. blues = plt.cm.Blues
  281. trunc_blues = ListedColormap(blues(np.linspace(0.7, 0.8, 256)))
  282. # Normalize by x-values (PC1)
  283. norm = Normalize(vmin=pc_large[:, 0].min(), vmax=pc_large[:, 0].max())
  284. # Plot
  285. plt.figure(figsize=(6, 6), dpi=300)
  286. plt.scatter(pc_large[:, 0], pc_large[:, 1],
  287. c=pc_large[:, 0], cmap=trunc_blues, norm=norm,
  288. s=40, linewidths=0)
  289. plt.axis('off')
  290. plt.tight_layout()
  291. plt.show()
  292. # %%
  293. import numpy as np
  294. import matplotlib.pyplot as plt
  295. from sklearn.decomposition import PCA
  296. grid_res = 50
  297. x_vals = np.linspace(env_min, env_max, grid_res)
  298. y_vals = np.linspace(env_min, env_max, grid_res)
  299. X, Y = np.meshgrid(x_vals, y_vals)
  300. # Grid embedding and coords from workspace
  301. H, W, n_cells = grid_embeddings.shape
  302. dx = X[0, 1] - X[0, 0]
  303. dy = Y[1, 0] - Y[0, 0]
  304. # Top 10% largest-scale cells
  305. threshold_10 = np.percentile(scales, 80)
  306. top10_idx = np.where(scales >= threshold_10)[0]
  307. n_top = len(top10_idx)
  308. # Extract embeddings for top cells
  309. emb_top10 = grid_embeddings.take(top10_idx, axis=2)
  310. # Compute finite differences
  311. partial_x = (emb_top10[:, 1:, :] - emb_top10[:, :-1, :]) / dx
  312. partial_y = (emb_top10[1:, :, :] - emb_top10[:-1, :, :]) / dy
  313. # Average over space
  314. avg_px = partial_x.reshape(-1, n_top).mean(axis=0)
  315. avg_py = partial_y.reshape(-1, n_top).mean(axis=0)
  316. # W matrix (#top,4): [Up, Down, Left, Right]
  317. W = np.stack([ avg_py, -avg_py, -avg_px, avg_px ], axis=1)
  318. # Build flattened embedding matrix manually (H*W, n_top)
  319. emb_flat_top = np.stack([emb_top10[:, :, i].ravel() for i in range(n_top)], axis=1)
  320. # PCA on embeddings of top cells
  321. pca_top10 = PCA(n_components=2)
  322. pc_top10 = pca_top10.fit_transform(emb_flat_top)
  323. # Project action vectors
  324. vectors = pca_top10.components_.dot(W) # (2,4)
  325. # Plot PCA scatter + arrows
  326. plt.figure(figsize=(6, 6),dpi=300)
  327. plt.scatter(pc_top10[:, 0], pc_top10[:, 1], s=0.2,c='black')
  328. actions = ['Up', 'Down', 'Left', 'Right']
  329. for i, action in enumerate(actions):
  330. vx, vy = vectors[0, i], vectors[1, i]
  331. plt.arrow(20, 0, vx, vy, head_width=0.4,fc='black')
  332. plt.text(vx * 1.3, vy * 1.3, action)
  333. # plt.title('PCA of Top-10%-Scale Grid Embeddings with Action Vectors')
  334. plt.xlabel('PC1')
  335. plt.ylabel('PC2')
  336. plt.axis('equal')
  337. plt.tight_layout()
  338. plt.axis('off')
  339. plt.show()
  340. print("W matrix shape:", W.shape)
  341. # %%
  342. import numpy as np
  343. import matplotlib.pyplot as plt
  344. import matplotlib.patches as patches
  345. # -------------------------
  346. # Helpers from earlier
  347. # -------------------------
  348. def grid_activation(x, y, scales, phases, orientations):
  349. pos = np.array([x, y])
  350. angles = np.array([0, np.pi/3, 2*np.pi/3])
  351. acts = []
  352. for scale, phase, orient in zip(scales, phases, orientations):
  353. R = np.array([[np.cos(orient), -np.sin(orient)],
  354. [np.sin(orient), np.cos(orient)]])
  355. pr = R @ (pos - phase)
  356. proj = pr[0]*np.cos(angles) + pr[1]*np.sin(angles)
  357. grating = np.sum(np.cos((4*np.pi/(scale*np.sqrt(3))) * proj))
  358. acts.append((2/3)*grating)
  359. return np.array(acts)
  360. def is_inside_rect(pos, rect):
  361. x, y, l, orient = rect
  362. if orient == 'h':
  363. x0, y0 = x - l/2, y - 0.05
  364. x1, y1 = x + l/2, y + 0.05
  365. else:
  366. x0, y0 = x - 0.05, y - l/2
  367. x1, y1 = x + 0.05, y + l/2
  368. return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
  369. def seg_intersect(A, B, C, D):
  370. def ccw(P, Q, R):
  371. return (R[1]-P[1])*(Q[0]-P[0]) > (Q[1]-P[1])*(R[0]-P[0])
  372. return ccw(A, C, D) != ccw(B, C, D) and ccw(A, B, C) != ccw(A, B, D)
  373. def segment_intersects_rect(p1, p2, rect):
  374. x0, x1, y0, y1 = (rect[0] - rect[2]/2, rect[0] + rect[2]/2,
  375. rect[1] - (0.01 if rect[3]=='h' else rect[2]/2),
  376. rect[1] + (0.01 if rect[3]=='h' else rect[2]/2))
  377. if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
  378. return True
  379. edges = [
  380. (np.array([x0, y0]), np.array([x1, y0])),
  381. (np.array([x1, y0]), np.array([x1, y1])),
  382. (np.array([x1, y1]), np.array([x0, y1])),
  383. (np.array([x0, y1]), np.array([x0, y0])),
  384. ]
  385. for C, D in edges:
  386. if seg_intersect(p1, p2, C, D):
  387. return True
  388. return False
  389. def is_valid_position(pos, walls):
  390. if not (0 <= pos[0] <= 4 and 0 <= pos[1] <= 4):
  391. return False
  392. for w in walls:
  393. if is_inside_rect(pos, w):
  394. return False
  395. return True
  396. def compute_valid_step(pos, vel, dt, walls):
  397. full = vel * dt
  398. cand = pos + full
  399. if is_valid_position(cand, walls) and not any(segment_intersects_rect(pos, cand, w) for w in walls):
  400. return cand
  401. # slide
  402. for dx, dy in [(full[0], 0), (0, full[1])]:
  403. cand2 = pos + np.array([dx, dy])
  404. if is_valid_position(cand2, walls) and not any(segment_intersects_rect(pos, cand2, w) for w in walls):
  405. return cand2
  406. # fallback smaller
  407. for f in [0.5, 0.25, 0.1]:
  408. cand3 = pos + full * f
  409. if is_valid_position(cand3, walls) and not any(segment_intersects_rect(pos, cand3, w) for w in walls):
  410. return cand3
  411. return pos
  412. # -------------------------
  413. # Setup grid code & W
  414. # -------------------------
  415. # Use pre-generated scales, phases, orientations, X, Y from workspace
  416. # Select top 10%
  417. threshold = np.percentile(scales, 90)
  418. top_idx = np.where(scales >= threshold)[0]
  419. sc_top = scales[top_idx]
  420. ph_top = phases[top_idx]
  421. or_top = orientations[top_idx]
  422. n_top = len(top_idx)
  423. # Compute W by finite differences averaged
  424. # reuse grid_embeddings, X, Y
  425. H, Wg, _ = grid_embeddings.shape
  426. dx = X[0,1] - X[0,0]
  427. dy = Y[1,0] - Y[0,0]
  428. emb_top = grid_embeddings[:,:,top_idx]
  429. pd_x = (emb_top[:,1:,:] - emb_top[:,:-1,:]) / dx
  430. pd_y = (emb_top[1:,:,:] - emb_top[:-1,:,:]) / dy
  431. avg_px = pd_x.reshape(-1, n_top).mean(axis=0)
  432. avg_py = pd_y.reshape(-1, n_top).mean(axis=0)
  433. Wmat = np.stack([avg_py, -avg_py, -avg_px, avg_px], axis=1)
  434. # -------------------------
  435. # Navigation via grid code
  436. # -------------------------
  437. walls = [] # no walls
  438. start = np.array([0.5, 2.5])
  439. goal = np.array([2.5, 0.5])
  440. noise_level = 80
  441. dt = 0.05
  442. threshold = 0.1
  443. max_steps = 2000
  444. # Precompute goal embedding
  445. goal_emb = grid_activation(goal[0], goal[1], sc_top, ph_top, or_top)
  446. paths = []
  447. for _ in range(5):
  448. pos = start.copy()
  449. path = [pos.copy()]
  450. for _ in range(max_steps):
  451. # current embedding
  452. cur_emb = grid_activation(pos[0], pos[1], sc_top, ph_top, or_top)
  453. delta = goal_emb - cur_emb # high-D target diff
  454. scores = delta @ Wmat + noise_level * np.random.randn(4) # (4,) for [Up,Down,Left,Right]
  455. # compute movement
  456. step = np.array([scores[3] - scores[2], scores[0] - scores[1]])
  457. # normalize to unit-length
  458. if np.linalg.norm(step) > 1e-6:
  459. step = step / np.linalg.norm(step)
  460. pos = compute_valid_step(pos, step, dt, walls)
  461. path.append(pos.copy())
  462. if np.linalg.norm(pos - goal) < threshold:
  463. break
  464. paths.append(np.array(path))
  465. # -------------------------
  466. # Plotting
  467. # -------------------------
  468. fig, ax = plt.subplots(figsize=(6,6))
  469. ax.plot(start[0], start[1], 'o', ms=20, c='black')
  470. ax.plot(goal[0], goal[1], '*', ms=30, c='black')
  471. colors = plt.cm.rainbow(np.linspace(0,1,len(paths)))
  472. for p, c in zip(paths, colors):
  473. ax.plot(p[:,0], p[:,1], color=c, lw=2, alpha=0.8)
  474. ax.set_xlim(0,4); ax.set_ylim(0,4); ax.set_aspect('equal')
  475. ax.set_xticks([])
  476. ax.set_yticks([])
  477. # ax.axis('off')
  478. plt.show()
  479. # %%
  480. probs.max()
  481. # %%
  482. import numpy as np
  483. import matplotlib.pyplot as plt
  484. # — your existing grid_activation, scales, phases, orientations, and BTSP weights —
  485. # grid_activation(x,y,scales,phases,orientations) → (n_grid_cells,)
  486. # weights: shape (n_place_cells, n_grid_cells)
  487. # paths: list of trajectories, each shape [T,2]
  488. # 1) Grab first trajectory
  489. traj = paths[0]
  490. T = traj.shape[0]
  491. n_pc = weights.shape[0]
  492. # 2) Sample spikes from learned place-cell activations
  493. spikes = np.zeros((T, n_pc), dtype=int)
  494. for t, (x, y) in enumerate(traj):
  495. gact = grid_activation(x, y, scales, phases, orientations) # pre-syn grid code
  496. p_act = weights @ gact # learned place-cell drive
  497. exps = np.exp(p_act - p_act.max())
  498. probs = exps / exps.sum() * 1 + 0.0005*np.random.randn(exps.shape[0]) # softmax → firing probs
  499. spikes[t] = (np.random.rand(n_pc) < probs).astype(int)
  500. # 3) Pick top-50 most active, rank by mean spike time
  501. counts = spikes.sum(axis=0)
  502. avg_time = (np.arange(T)[:,None] * spikes).sum(axis=0) / (counts + 1e-9)
  503. top50 = np.argsort(counts)[-50:] # highest-firing 50
  504. order = np.argsort(avg_time[top50]) # earliest mean time first
  505. selected = top50[order]
  506. # 4) Raster plot (spikes only)
  507. fig, ax = plt.subplots(figsize=(3,5), dpi=300)
  508. for r, neuron in enumerate(selected):
  509. times = np.where(spikes[:, neuron])[0]
  510. ax.vlines(times, r + 0.2, r + 0.8, color='black', linewidth=0.7)
  511. ax.set_xlim(0, T)
  512. ax.set_ylim(0, 50)
  513. # ax.set_xlabel('Time step')
  514. ticks = ax.get_xticks() # e.g. [0, 100, 200, …]
  515. ax.set_xticks(ticks)
  516. ax.set_xticklabels([(t/100) for t in ticks])
  517. ax.set_yticks([0,10,20,30,40,50]) # hide neuron labels
  518. plt.tight_layout()
  519. plt.show()
  520. # %%
  521. ticks
  522. # %%
  523. fig, ax = plt.subplots(figsize=(1.5,2.5), dpi=300)
  524. for r, neuron in enumerate(selected):
  525. times = np.where(spikes[:, neuron])[0]
  526. ax.vlines(times, r + 0.2, r + 0.8, color='black', linewidth=0.4)
  527. ax.set_xlim(0, T)
  528. ax.set_ylim(0, 50)
  529. # ax.set_xlabel('Time step')
  530. ticks = ax.get_xticks() # e.g. [0, 100, 200, …]
  531. # ticks = [0,100,200,300,400,500]
  532. ax.set_xticks(ticks)
  533. ax.set_xticklabels([(t/100) for t in ticks])
  534. ax.set_yticks([0,10,20,30,40,50]) # hide neuron labels
  535. plt.tight_layout()
  536. plt.show()
  537. # %%
  538. import numpy as np
  539. import matplotlib.pyplot as plt
  540. import matplotlib.patches as patches
  541. # -------------------------
  542. # Helpers from earlier
  543. # -------------------------
  544. def grid_activation(x, y, scales, phases, orientations):
  545. pos = np.array([x, y])
  546. angles = np.array([0, np.pi/3, 2*np.pi/3])
  547. acts = []
  548. for scale, phase, orient in zip(scales, phases, orientations):
  549. R = np.array([[np.cos(orient), -np.sin(orient)],
  550. [np.sin(orient), np.cos(orient)]])
  551. pr = R @ (pos - phase)
  552. proj = pr[0]*np.cos(angles) + pr[1]*np.sin(angles)
  553. grating = np.sum(np.cos((4*np.pi/(scale*np.sqrt(3))) * proj))
  554. acts.append((2/3)*grating)
  555. return np.array(acts)
  556. def rect_bounds(rect, eps=0.01):
  557. x, y, l, orient = rect
  558. if orient == 'h':
  559. # horizontal: thin in y
  560. return (x - l/2, x + l/2,
  561. y - eps, y + eps)
  562. else:
  563. # vertical: thin in x
  564. return (x - eps, x + eps,
  565. y - l/2, y + l/2)
  566. def is_inside_rect(pos, rect):
  567. x0, x1, y0, y1 = rect_bounds(rect)
  568. return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
  569. def segment_intersects_rect(p1, p2, rect):
  570. x0, x1, y0, y1 = rect_bounds(rect)
  571. # If endpoint lies within rect, it intersects
  572. if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
  573. return True
  574. # Otherwise check each edge
  575. edges = [
  576. (np.array([x0, y0]), np.array([x1, y0])),
  577. (np.array([x1, y0]), np.array([x1, y1])),
  578. (np.array([x1, y1]), np.array([x0, y1])),
  579. (np.array([x0, y1]), np.array([x0, y0]))
  580. ]
  581. for C, D in edges:
  582. if seg_intersect(p1, p2, C, D):
  583. return True
  584. return False
  585. # def is_inside_rect(pos, rect):
  586. # x, y, l, orient = rect
  587. # if orient == 'h':
  588. # x0, y0 = x - l/2, y - 0.05
  589. # x1, y1 = x + l/2, y + 0.05
  590. # else:
  591. # x0, y0 = x - 0.05, y - l/2
  592. # x1, y1 = x + 0.05, y + l/2
  593. # return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
  594. def ccw(A, B, C):
  595. return (C[1]-A[1])*(B[0]-A[0]) > (B[1]-A[1])*(C[0]-A[0])
  596. def seg_intersect(A, B, C, D):
  597. return ccw(A, C, D) != ccw(B, C, D) and ccw(A, B, C) != ccw(A, B, D)
  598. # def segment_intersects_rect(p1, p2, rect):
  599. # x0, x1, y0, y1 = (
  600. # rect[0] - rect[2]/2, rect[0] + rect[2]/2,
  601. # rect[1] - (0.01 if rect[3]=='h' else rect[2]/2),
  602. # rect[1] + (0.01 if rect[3]=='h' else rect[2]/2)
  603. # )
  604. # if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
  605. # return True
  606. # edges = [
  607. # (np.array([x0, y0]), np.array([x1, y0])),
  608. # (np.array([x1, y0]), np.array([x1, y1])),
  609. # (np.array([x1, y1]), np.array([x0, y1])),
  610. # (np.array([x0, y1]), np.array([x0, y0])),
  611. # ]
  612. # for C, D in edges:
  613. # if seg_intersect(p1, p2, C, D):
  614. # return True
  615. return False
  616. def is_valid_position(pos, walls):
  617. if not (0 <= pos[0] <= 4 and 0 <= pos[1] <= 4):
  618. return False
  619. for w in walls:
  620. if is_inside_rect(pos, w):
  621. return False
  622. return True
  623. def compute_valid_step(pos, vel, dt, walls):
  624. full = vel * dt
  625. cand = pos + full
  626. if is_valid_position(cand, walls) and not any(segment_intersects_rect(pos, cand, w) for w in walls):
  627. return cand
  628. # slide
  629. for dx, dy in [(full[0], 0), (0, full[1])]:
  630. cand2 = pos + np.array([dx, dy])
  631. if is_valid_position(cand2, walls) and not any(segment_intersects_rect(pos, cand2, w) for w in walls):
  632. return cand2
  633. # fallback smaller
  634. for f in [0.5, 0.25, 0.1]:
  635. cand3 = pos + full * f
  636. if is_valid_position(cand3, walls) and not any(segment_intersects_rect(pos, cand3, w) for w in walls):
  637. return cand3
  638. return pos
  639. # -------------------------
  640. # Setup grid code & W
  641. # -------------------------
  642. # Top 10% scales
  643. threshold = np.percentile(scales, 90)
  644. top_idx = np.where(scales >= threshold)[0]
  645. sc_top = scales[top_idx]
  646. ph_top = phases[top_idx]
  647. or_top = orientations[top_idx]
  648. n_top = len(top_idx)
  649. # Finite diff to build Wmat
  650. H, Wg, _ = grid_embeddings.shape
  651. dx = X[0,1] - X[0,0]
  652. dy = Y[1,0] - Y[0,0]
  653. emb_top = grid_embeddings[:,:,top_idx]
  654. pd_x = (emb_top[:,1:,:] - emb_top[:,:-1,:]) / dx
  655. pd_y = (emb_top[1:,:,:] - emb_top[:-1,:,:]) / dy
  656. avg_px = pd_x.reshape(-1, n_top).mean(axis=0)
  657. avg_py = pd_y.reshape(-1, n_top).mean(axis=0)
  658. Wmat = np.stack([avg_py, -avg_py, -avg_px, avg_px], axis=1)
  659. # -------------------------
  660. # Navigation with repulsion
  661. # -------------------------
  662. walls = [
  663. (1,1.5,1,'h'),
  664. (1.5,1,1,'v'),
  665. (1.5,2,1,'v'),
  666. (2,2.5,1,'h'),
  667. (3,2.5,1,'h'),
  668. (2.5,3,1,'v')
  669. ]
  670. # Precompute repulsion points
  671. rep_points = []
  672. for x, y, l, orient in walls:
  673. rep_points.append(np.array([x, y]))
  674. if orient == 'h':
  675. rep_points += [np.array([x-l/2,y]), np.array([x+l/2,y])]
  676. else:
  677. rep_points += [np.array([x,y-l/2]), np.array([x,y+l/2])]
  678. start = np.array([0.2, 0.2])
  679. goal = np.array([2.0, 2.0])
  680. k_wall = 0.08
  681. noise_level = 30.0
  682. dt = 0.05
  683. threshold_goal = 0.1
  684. max_steps = 1000
  685. n_paths = 5
  686. max_dist = np.linalg.norm(start - goal)
  687. goal_emb = grid_activation(goal[0], goal[1], sc_top, ph_top, or_top)
  688. def compute_velocity(pos):
  689. # Repulsive 2D
  690. f_rep = np.zeros(2)
  691. for wc in rep_points:
  692. delta = pos - wc
  693. dsq = np.sum(delta**2)
  694. if dsq > 1e-4:
  695. f_rep += k_wall * delta / dsq
  696. rep_scale = np.clip(np.linalg.norm(pos-goal) / max_dist, 0, 1)
  697. f_rep *= rep_scale
  698. # Grid-code guidance
  699. cur_emb = grid_activation(pos[0], pos[1], sc_top, ph_top, or_top)
  700. delta_emb = goal_emb - cur_emb
  701. scores = delta_emb @ Wmat + noise_level * np.random.randn(4) # [Up,Down,Left,Right]
  702. step = np.array([scores[3]-scores[2], scores[0]-scores[1]])
  703. if np.linalg.norm(step) > 1e-6:
  704. step = step / np.linalg.norm(step)
  705. return step + f_rep # combine
  706. # Generate paths
  707. paths = []
  708. for _ in range(n_paths):
  709. pos = start.copy()
  710. path = [pos.copy()]
  711. for _ in range(max_steps):
  712. vel = compute_velocity(pos)
  713. pos = compute_valid_step(pos, vel, dt, walls)
  714. path.append(pos.copy())
  715. if np.linalg.norm(pos - goal) < threshold_goal:
  716. break
  717. paths.append(np.array(path))
  718. # -------------------------
  719. # Plotting
  720. # -------------------------
  721. fig, ax = plt.subplots(figsize=(6,6))
  722. # draw walls
  723. for x, y, l, orient in walls:
  724. if orient == 'h':
  725. ax.add_patch(patches.Rectangle((x-l/2, y-0.02), l, 0.04, color='black'))
  726. else:
  727. ax.add_patch(patches.Rectangle((x-0.02, y-l/2), 0.04, l, color='black'))
  728. # start & goal
  729. ax.plot(start[0], start[1], 'o', ms=12, c='black')
  730. ax.plot(goal[0], goal[1], '*', ms=15, c='black')
  731. # paths
  732. colors = plt.cm.rainbow(np.linspace(0,1,n_paths))
  733. for p, c in zip(paths, colors):
  734. ax.plot(p[:,0], p[:,1], color=c, lw=2, alpha=0.8)
  735. ax.set_xlim(0,4); ax.set_ylim(0,4); ax.set_aspect('equal')
  736. ax.set_xticks([]); ax.set_yticks([])
  737. plt.show()
  738. # %%
  739. vel
  740. # %%

gcml_grid_cell.ipynb at commit ff76859, under MIT · at the source

Overview

  1. Department of Precision Instruments, Center for Brain-Inspired Computing Research (CBICR), Tsinghua University,Beijing, China
  2. Institute of Machine Learning and Neural Computation, Graz University of Technology,Graz, Austria
  3. Institute of Cognitive Sciences and Technologies, National Research Council,Rome, Italy
Journal: Nature machine intelligence, volume 8, issue 7, pages 1045-1065
Dates: received 6 August 2025; accepted 11 May 2026; published online 21 July 2026; in print 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s42256-026-01254-4 · PMID 42499994 · PMCID PMC13395624 · OpenAlex W7169839672
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), none (in silico) (organism), cognitive (subfield)
Methods: Connectivity, Smoothing, state filtering, decompositions, Graphs, Statistics, Machine learning, Single-unit activity, calcium imaging
Keywords: Computational science, Learning algorithms, Network models
Topic: Ferroelectric and Negative Capacitance Devices (Electrical and Electronic Engineering, Engineering), according to OpenAlex
Citations: not cited yet (Europe PMC); 104 references in the paper

Abstract

Artificial intelligence systems are becoming more intelligent, but at a very high cost in terms of energy consumption and training requirements. By contrast, our brains only require 20 W of energy, they learn online and they can instantly adjust to changing contingencies. This begs the question what data structures, algorithms and learning methods enable brains to achieve that, and whether these can be ported into artificial devices. We are addressing this question for a core feature of intelligence: the capacity to plan and solve problems, including new problems that involve states that were never encountered before. Here we examine three tools that brains are likely to use for achieving that: cognitive maps, stochastic computing and compositional coding. We integrate these tools into a transparent neural network model, and demonstrate its power for flexible planning and problem-solving. Importantly, this approach is suitable for implementation by in-memory computing and other energy-efficient neuromorphic hardware. In particular, it only requires self-supervised local synaptic plasticity that is suited for on-chip learning. Hence, a core feature of brain intelligence—the capacity to generate solutions to problems that were never encountered before—does not require deep neural networks or large language models, and can be implemented in energy-efficient edge devices.

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

LH-cbicr/GCML

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: ff76859b71a2bc2056b50f5e052475351c007f76, 1 April 2026
Languages: Jupyter (3), Python (1)
Size: 10 files, 4 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (environment.yml), 3 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: Matplotlib (4 files), NumPy (4 files), scikit-learn (4 files), NetworkX (3 files), PyTorch (3 files), h5py (2 files), seaborn (2 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
6 files

Zenodo 19370442

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
Tools: Matplotlib (4 files), NumPy (4 files), scikit-learn (4 files), NetworkX (3 files), PyTorch (3 files), h5py (2 files), seaborn (2 files)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
6 files
At the source:

Code availability

The code used for training and evaluating the GCML is publicly available via GitHub at https://github.com/LH-cbicr/GCML and via Zenodo at 10.5281/zenodo.19370442 (ref. 104).

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:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 8 scripts, each with its path and the digest of its content;
  • 2 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 used in the experiments were generated synthetically and are publicly available via GitHub at https://github.com/LH-cbicr/GCML and via Zenodo at 10.5281/zenodo.19370442 (ref. 104).

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 2, 28 September 2026

  • Publisher: n/a → Nature Portfolio

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 3 keywords, 3 funders, 87 references.

Cite

This paper

Lin, H., Yang, Y., Zhao, R., Pezzulo, G., & Maass, W. (2026). Neural sampling from cognitive maps enables goal-directed imagination and planning. Nature machine intelligence, 8(7), 1045-1065. https://doi.org/10.1038/s42256-026-01254-4

BibTeX

@article{lin2026neural,
author = {Lin, Hui and Yang, Yukun and Zhao, Rong and Pezzulo, Giovanni and Maass, Wolfgang},
title = {{Neural sampling from cognitive maps enables goal-directed imagination and planning}},
journal = {Nature machine intelligence},
year = {2026},
month = jul,
volume = {8},
number = {7},
pages = {1045--1065},
publisher = {Nature Portfolio},
issn = {2522-5839},
doi = {10.1038/s42256-026-01254-4},
url = {https://doi.org/10.1038/s42256-026-01254-4},
pmid = {42499994},
pmcid = {PMC13395624}
}

RIS

TY - JOUR
AU - Lin, Hui
AU - Yang, Yukun
AU - Zhao, Rong
AU - Pezzulo, Giovanni
AU - Maass, Wolfgang
TI - Neural sampling from cognitive maps enables goal-directed imagination and planning
T2 - Nature machine intelligence
J2 - Nat Mach Intell
PY - 2026
DA - 2026/07/21
VL - 8
IS - 7
SP - 1045
EP - 1065
SN - 2522-5839
PB - Nature Portfolio
DO - 10.1038/s42256-026-01254-4
UR - https://doi.org/10.1038/s42256-026-01254-4
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s42256-026-01254-4",
"type": "article-journal",
"title": "Neural sampling from cognitive maps enables goal-directed imagination and planning",
"container-title": "Nature machine intelligence",
"author": [
{
"family": "Lin",
"given": "Hui"
},
{
"family": "Yang",
"given": "Yukun"
},
{
"family": "Zhao",
"given": "Rong"
},
{
"family": "Pezzulo",
"given": "Giovanni"
},
{
"family": "Maass",
"given": "Wolfgang"
}
],
"container-title-short": "Nat Mach Intell",
"volume": "8",
"issue": "7",
"page": "1045-1065",
"DOI": "10.1038/s42256-026-01254-4",
"PMID": "42499994",
"PMCID": "PMC13395624",
"ISSN": "2522-5839",
"publisher": "Nature Portfolio",
"URL": "https://doi.org/10.1038/s42256-026-01254-4",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
21
]
]
}
}

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.1371/journal.pbio.3003755 [code]
Action information is integrated into entorhinal representations of conceptual space and is reflected in eye movements.
Journal: PLoS biology
In common: seaborn, Matplotlib, NumPy, cognitive, 16 references
[2] doi:10.1038/s41467-026-76102-5 [code]
The inherent capacity of neurons to learn order relations and support abstract reasoning.
Journal: Nature communications
In common: PyTorch, scikit-learn, Matplotlib, 1 other tool, computational modeling (no new data), cognitive, 6 references, 2 authors
[3] doi:10.1016/j.cub.2026.05.068 [code]
An abstract relational map emerges in the human medial prefrontal cortex with consolidation.
Journal: Current biology : CB
In common: seaborn, Matplotlib, NumPy, 9 references
[4] doi:10.1002/hipo.70131 [code]
Decoding Medial Entorhinal Cortical Dynamics Produces Planning-Like Alternations in Hippocampal theta Sequences.
Journal: Hippocampus
In common: h5py, scikit-learn, Matplotlib, 1 other tool, 7 references
[5] doi:10.1038/s41467-026-74357-6 [code]
Hippocampo-neocortical interaction as compressive retrieval-augmented generation.
Journal: Nature communications
In common: NetworkX, PyTorch, seaborn, 3 other tools, computational modeling (no new data), 4 references
[6] doi:10.1371/journal.pcbi.1013487 [code]
Flexible navigation with neuromodulated cognitive maps.
Journal: PLoS computational biology
In common: Matplotlib, NumPy, none (in silico), 6 references
[7] doi:10.1038/s41467-026-74358-5 [code]
Brain-inspired spatial intelligence for embodied agents.
Journal: Nature communications
In common: NetworkX, h5py, PyTorch, 4 other tools, cognitive, 2 references
[8] doi:10.1126/sciadv.aeg6797 [code]
Dorsoventral gradient of theta sweeps in the medial entorhinal cortex.
Journal: Science advances
In common: NetworkX, h5py, seaborn, 3 other tools, none (in silico), 2 references
[9] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: NetworkX, h5py, PyTorch, 4 other tools, 2 references
[10] doi:10.1038/s41593-026-02365-2 [code]
Hippocampal theta sweeps indicate goal direction during navigation.
Journal: Nature neuroscience
In common: Matplotlib, NumPy, computational modeling (no new data), 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.