Neural sampling from cognitive maps enables goal-directed imagination and planning.
The 2 matches
- [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] § 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
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- # Environment settings
- env_min = 0
- env_max = 4
- # Random walk settings
- n_steps = 50000
- dt = 0.05
- speed = 0.1 # constant speed
- turn_noise = 0.1 # radians per step
- n_traj = 1 # number of trajectories to plot
- # Helper: is a point valid?
- def is_valid(pos):
- x, y = pos
- if not (env_min <= x <= env_max and env_min <= y <= env_max):
- return False
- return True
- # Generate trajectories
- trajectories = []
- for _ in range(n_traj):
- pos = np.array([np.random.uniform(0, 4), np.random.uniform(0, 4)])
- angle = np.random.uniform(0, 2*np.pi)
- path = [pos.copy()]
- for _ in range(n_steps):
- # Add small random turn
- angle += np.random.randn() * turn_noise
- velocity = speed * np.array([np.cos(angle), np.sin(angle)])
- new_pos = pos + velocity * dt
- if is_valid(new_pos):
- pos = new_pos
- else:
- # Resample a new random direction if invalid move
- angle = np.random.uniform(0, 2*np.pi)
- path.append(pos.copy())
- trajectories.append(np.array(path))
- # Plot
- fig, ax = plt.subplots(figsize=(6, 6),dpi=300)
- for traj in trajectories:
- # Plot trajectory
- ax.plot(traj[:,0], traj[:,1], color='gray', alpha=0.5, linewidth=1)
- # for i in range(1, len(traj)):
- # ax.plot(traj[i-1:i+1,0], traj[i-1:i+1,1], color=‘gray’, alpha=0.5, linewidth=1)
- # Mark start
- start = traj[0]
- ax.plot(start[0], start[1], 'o', markersize=20, c='black')
- # Mark end
- end = traj[-1]
- ax.plot(end[0], end[1], 's', markersize=20, c='black')
- # Forbidden region
- # ax.add_patch(plt.Rectangle((0, 0), forbidden_x, forbidden_y,
- # color='red', alpha=0.2, zorder=0))
- # Formatting
- ax.set_xlim(env_min, env_max)
- ax.set_ylim(env_min, env_max)
- ax.set_aspect('equal')
- ax.set_xticks([])
- ax.set_yticks([])
- # ax.set_title(‘Random Walk with Forbidden Region’)
- plt.show()
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- # Environment settings
- env_min = 0
- env_max = 4
- forbidden_x = 2
- forbidden_y = 2
- # Random walk settings
- n_steps = 50000
- dt = 0.05
- speed = 0.1 # constant speed
- turn_noise = 0.1 # radians per step
- n_traj = 1 # number of trajectories to plot
- # Helper: is a point valid?
- def is_valid(pos):
- x, y = pos
- if not (env_min <= x <= env_max and env_min <= y <= env_max):
- return False
- if x < forbidden_x and y < forbidden_y:
- return False
- return True
- # Generate trajectories
- trajectories = []
- for _ in range(n_traj):
- pos = np.array([np.random.uniform(2, 4), np.random.uniform(2, 4)]) # start outside forbidden
- angle = np.random.uniform(0, 2*np.pi)
- path = [pos.copy()]
- for _ in range(n_steps):
- # Add small random turn
- angle += np.random.randn() * turn_noise
- velocity = speed * np.array([np.cos(angle), np.sin(angle)])
- new_pos = pos + velocity * dt
- if is_valid(new_pos):
- pos = new_pos
- else:
- # Resample a new random direction if invalid move
- angle = np.random.uniform(0, 2*np.pi)
- path.append(pos.copy())
- trajectories.append(np.array(path))
- # Plot
- fig, ax = plt.subplots(figsize=(6, 6),dpi=300)
- for traj in trajectories:
- # Plot trajectory
- ax.plot(traj[:,0], traj[:,1], color='gray', alpha=0.99, linewidth=1)
- # for i in range(1, len(traj)):
- # ax.plot(traj[i-1:i+1,0], traj[i-1:i+1,1], color='gray', alpha=0.5, linewidth=1)
- # Mark start
- start = traj[0]
- # ax.plot(start[0], start[1], 'o', markersize=20, c='black')
- # Mark end
- end = traj[-1]
- # ax.plot(end[0], end[1], 's', markersize=20, c='black')
- # Forbidden region
- # ax.add_patch(plt.Rectangle((0, 0), forbidden_x, forbidden_y,
- # color='red', alpha=0.2, zorder=0))
- # Formatting
- ax.set_xlim(env_min, env_max)
- ax.set_ylim(env_min, env_max)
- ax.set_aspect('equal')
- ax.set_xticks([])
- ax.set_yticks([])
- # ax.set_title('Random Walk with Forbidden Region')
- plt.show()
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- # -------------------------
- # Function Definitions
- # -------------------------
- def is_valid(pos, env_min=0, env_max=4):
- x, y = pos
- return env_min <= x <= env_max and env_min <= y <= env_max
- def generate_grid_cells(n_cells=1000, arena_size=4, seed=None):
- if seed is not None:
- np.random.seed(seed)
- scales = np.random.uniform(0.02, 8, size=n_cells)
- phases = np.random.uniform(0, arena_size, size=(n_cells, 2))
- orientations = np.random.uniform(0, np.pi/3, size=n_cells)
- return scales, phases, orientations
- def grid_activation(x, y, scales, phases, orientations):
- positions = np.array([x, y])
- activations = []
- angles = np.array([0, np.pi / 3, 2 * np.pi / 3])
- for scale, phase, orientation in zip(scales, phases, orientations):
- # Rotate and compute grating
- R = np.array([[np.cos(orientation), -np.sin(orientation)],
- [np.sin(orientation), np.cos(orientation)]])
- pos_rot = R @ (positions - phase)
- proj = pos_rot[0] * np.cos(angles) + pos_rot[1] * np.sin(angles)
- grating = np.sum(np.cos((4 * np.pi / (scale * np.sqrt(3))) * proj))
- activations.append((2/3) * grating)
- return np.array(activations)
- def grid_activation_vectorized(X, Y, scales, phases, orientations):
- H, W = X.shape
- n_cells = len(scales)
- pos = np.stack([X.ravel(), Y.ravel()], axis=1) # (H*W, 2)
- activations = np.zeros((H * W, n_cells))
- angles = np.array([0, np.pi / 3, 2 * np.pi / 3])
- for i in range(n_cells):
- scale = scales[i]
- phase = phases[i]
- orientation = orientations[i]
- c, s = np.cos(orientation), np.sin(orientation)
- R = np.array([[c, -s], [s, c]])
- rel_pos = pos - phase
- pos_rot = rel_pos @ R.T
- proj = np.outer(pos_rot[:, 0], np.cos(angles)) + np.outer(pos_rot[:, 1], np.sin(angles))
- grating = np.sum(np.cos((4 * np.pi / (scale * np.sqrt(3))) * proj), axis=1)
- activations[:, i] = (2/3) * grating
- return activations.reshape(H, W, n_cells)
- # -------------------------
- # Random Walk Generation
- # -------------------------
- env_min, env_max = 0, 4
- n_steps = 50000
- dt = 0.05
- speed = 0.2
- turn_noise = 0.1
- pos = np.random.uniform(env_min, env_max, size=2)
- angle = np.random.uniform(0, 2*np.pi)
- trajectory = [pos.copy()]
- for _ in range(n_steps):
- angle += np.random.randn() * turn_noise
- vel = speed * np.array([np.cos(angle), np.sin(angle)])
- new_pos = pos + vel * dt
- if is_valid(new_pos, env_min, env_max):
- pos = new_pos
- else:
- angle = np.random.uniform(0, 2*np.pi)
- trajectory.append(pos.copy())
- trajectory = np.array(trajectory)
- # -------------------------
- # Grid Cell Setup
- # -------------------------
- n_cells = 1000
- scales, phases, orientations = generate_grid_cells(n_cells, env_max)
- # -------------------------
- # Plateau Events & Weights
- # -------------------------
- n_plateau = 1000
- plateau_times = np.sort(np.random.choice(np.arange(len(trajectory)), size=n_plateau, replace=False))
- weights = np.zeros((n_plateau, n_cells))
- for i, t in enumerate(plateau_times):
- x, y = trajectory[t]
- weights[i] = grid_activation(x, y, scales, phases, orientations)
- # -------------------------
- # Plot 1: Random Walk + Plateau Fires
- # -------------------------
- fig, ax = plt.subplots(figsize=(6, 6), dpi=300)
- ax.plot(trajectory[:, 0], trajectory[:, 1], color='gray', alpha=0.5, linewidth=1)
- ax.scatter(trajectory[plateau_times, 0], trajectory[plateau_times, 1],
- c='black', s=10, label='Plateau Fires')
- ax.plot(trajectory[0, 0], trajectory[0, 1], 'o', markersize=12, c='black')
- ax.plot(trajectory[-1, 0], trajectory[-1, 1], 's', markersize=12, c='black')
- ax.set_xlim(env_min, env_max)
- ax.set_ylim(env_min, env_max)
- ax.set_aspect('equal')
- ax.axis('off')
- plt.show()
- # -------------------------
- # Compute Reactive Fields
- # -------------------------
- grid_res = 50
- x_vals = np.linspace(env_min, env_max, grid_res)
- y_vals = np.linspace(env_min, env_max, grid_res)
- X, Y = np.meshgrid(x_vals, y_vals)
- grid_embeddings = grid_activation_vectorized(X, Y, scales, phases, orientations)
- # resulting shape: (grid_res, grid_res, n_cells)
- # reactive_fields(h, w, neuron) = sum_i grid_embeddings(h,w,i) * weights(neuron,i)
- reactive_fields = np.tensordot(grid_embeddings, weights, axes=([2], [1])) # shape: (H, W, n_plateau)
- # -------------------------
- # Plot 2: First 100 Reactive Fields
- # -------------------------
- fig2, axes = plt.subplots(10, 10, figsize=(20, 20))
- for idx, ax in enumerate(axes.flatten()):
- im = ax.imshow(reactive_fields[:, :, idx],
- extent=[env_min, env_max, env_min, env_max],
- origin='lower')
- ax.axis('off')
- plt.tight_layout()
- plt.show()
- # %%
- from sklearn.decomposition import PCA
- import numpy as np
- import matplotlib.pyplot as plt
- # Identify top 1% largest-scale grid cells
- threshold = np.percentile(scales,85)
- large_indices = np.where(scales >= threshold)[0]
- # Extract embeddings for only those cells and flatten
- H, W, _ = grid_embeddings.shape
- emb_large = grid_embeddings[:, :, large_indices].reshape(H * W, len(large_indices))
- # PCA on filtered embeddings
- pca_large = PCA(n_components=2)
- pc_large = pca_large.fit_transform(emb_large)
- # Plot
- plt.figure(figsize=(6, 6), dpi=300)
- plt.scatter(pc_large[:, 0], pc_large[:, 1], s=5)
- # plt.title('PCA of Top 1%‑Scale Grid Cell Embeddings')
- # plt.xlabel('PC1')
- plt.axis('off')
- # plt.ylabel('PC2')
- plt.tight_layout()
- plt.show()
- # %%
- from sklearn.decomposition import PCA
- import numpy as np
- import matplotlib.pyplot as plt
- from matplotlib.colors import Normalize, ListedColormap
- # Identify top 15% largest-scale grid cells (you had 85th percentile)
- threshold = np.percentile(scales, 85)
- large_indices = np.where(scales >= threshold)[0]
- # Extract embeddings for only those cells and flatten
- H, W, _ = grid_embeddings.shape
- emb_large = grid_embeddings[:, :, large_indices].reshape(H * W, len(large_indices))
- # PCA on filtered embeddings
- pca_large = PCA(n_components=2)
- pc_large = pca_large.fit_transform(emb_large)
- # --- Colormap: Blues truncated to [0.1, 0.3] ---
- blues = plt.cm.Blues
- trunc_blues = ListedColormap(blues(np.linspace(0.7, 0.8, 256)))
- # Normalize by x-values (PC1)
- norm = Normalize(vmin=pc_large[:, 0].min(), vmax=pc_large[:, 0].max())
- # Plot
- plt.figure(figsize=(6, 6), dpi=300)
- plt.scatter(pc_large[:, 0], pc_large[:, 1],
- c=pc_large[:, 0], cmap=trunc_blues, norm=norm,
- s=40, linewidths=0)
- plt.axis('off')
- plt.tight_layout()
- plt.show()
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- from sklearn.decomposition import PCA
- grid_res = 50
- x_vals = np.linspace(env_min, env_max, grid_res)
- y_vals = np.linspace(env_min, env_max, grid_res)
- X, Y = np.meshgrid(x_vals, y_vals)
- # Grid embedding and coords from workspace
- H, W, n_cells = grid_embeddings.shape
- dx = X[0, 1] - X[0, 0]
- dy = Y[1, 0] - Y[0, 0]
- # Top 10% largest-scale cells
- threshold_10 = np.percentile(scales, 80)
- top10_idx = np.where(scales >= threshold_10)[0]
- n_top = len(top10_idx)
- # Extract embeddings for top cells
- emb_top10 = grid_embeddings.take(top10_idx, axis=2)
- # Compute finite differences
- partial_x = (emb_top10[:, 1:, :] - emb_top10[:, :-1, :]) / dx
- partial_y = (emb_top10[1:, :, :] - emb_top10[:-1, :, :]) / dy
- # Average over space
- avg_px = partial_x.reshape(-1, n_top).mean(axis=0)
- avg_py = partial_y.reshape(-1, n_top).mean(axis=0)
- # W matrix (#top,4): [Up, Down, Left, Right]
- W = np.stack([ avg_py, -avg_py, -avg_px, avg_px ], axis=1)
- # Build flattened embedding matrix manually (H*W, n_top)
- emb_flat_top = np.stack([emb_top10[:, :, i].ravel() for i in range(n_top)], axis=1)
- # PCA on embeddings of top cells
- pca_top10 = PCA(n_components=2)
- pc_top10 = pca_top10.fit_transform(emb_flat_top)
- # Project action vectors
- vectors = pca_top10.components_.dot(W) # (2,4)
- # Plot PCA scatter + arrows
- plt.figure(figsize=(6, 6),dpi=300)
- plt.scatter(pc_top10[:, 0], pc_top10[:, 1], s=0.2,c='black')
- actions = ['Up', 'Down', 'Left', 'Right']
- for i, action in enumerate(actions):
- vx, vy = vectors[0, i], vectors[1, i]
- plt.arrow(20, 0, vx, vy, head_width=0.4,fc='black')
- plt.text(vx * 1.3, vy * 1.3, action)
- # plt.title('PCA of Top-10%-Scale Grid Embeddings with Action Vectors')
- plt.xlabel('PC1')
- plt.ylabel('PC2')
- plt.axis('equal')
- plt.tight_layout()
- plt.axis('off')
- plt.show()
- print("W matrix shape:", W.shape)
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- import matplotlib.patches as patches
- # -------------------------
- # Helpers from earlier
- # -------------------------
- def grid_activation(x, y, scales, phases, orientations):
- pos = np.array([x, y])
- angles = np.array([0, np.pi/3, 2*np.pi/3])
- acts = []
- for scale, phase, orient in zip(scales, phases, orientations):
- R = np.array([[np.cos(orient), -np.sin(orient)],
- [np.sin(orient), np.cos(orient)]])
- pr = R @ (pos - phase)
- proj = pr[0]*np.cos(angles) + pr[1]*np.sin(angles)
- grating = np.sum(np.cos((4*np.pi/(scale*np.sqrt(3))) * proj))
- acts.append((2/3)*grating)
- return np.array(acts)
- def is_inside_rect(pos, rect):
- x, y, l, orient = rect
- if orient == 'h':
- x0, y0 = x - l/2, y - 0.05
- x1, y1 = x + l/2, y + 0.05
- else:
- x0, y0 = x - 0.05, y - l/2
- x1, y1 = x + 0.05, y + l/2
- return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
- def seg_intersect(A, B, C, D):
- def ccw(P, Q, R):
- return (R[1]-P[1])*(Q[0]-P[0]) > (Q[1]-P[1])*(R[0]-P[0])
- return ccw(A, C, D) != ccw(B, C, D) and ccw(A, B, C) != ccw(A, B, D)
- def segment_intersects_rect(p1, p2, rect):
- x0, x1, y0, y1 = (rect[0] - rect[2]/2, rect[0] + rect[2]/2,
- rect[1] - (0.01 if rect[3]=='h' else rect[2]/2),
- rect[1] + (0.01 if rect[3]=='h' else rect[2]/2))
- if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
- return True
- edges = [
- (np.array([x0, y0]), np.array([x1, y0])),
- (np.array([x1, y0]), np.array([x1, y1])),
- (np.array([x1, y1]), np.array([x0, y1])),
- (np.array([x0, y1]), np.array([x0, y0])),
- ]
- for C, D in edges:
- if seg_intersect(p1, p2, C, D):
- return True
- return False
- def is_valid_position(pos, walls):
- if not (0 <= pos[0] <= 4 and 0 <= pos[1] <= 4):
- return False
- for w in walls:
- if is_inside_rect(pos, w):
- return False
- return True
- def compute_valid_step(pos, vel, dt, walls):
- full = vel * dt
- cand = pos + full
- if is_valid_position(cand, walls) and not any(segment_intersects_rect(pos, cand, w) for w in walls):
- return cand
- # slide
- for dx, dy in [(full[0], 0), (0, full[1])]:
- cand2 = pos + np.array([dx, dy])
- if is_valid_position(cand2, walls) and not any(segment_intersects_rect(pos, cand2, w) for w in walls):
- return cand2
- # fallback smaller
- for f in [0.5, 0.25, 0.1]:
- cand3 = pos + full * f
- if is_valid_position(cand3, walls) and not any(segment_intersects_rect(pos, cand3, w) for w in walls):
- return cand3
- return pos
- # -------------------------
- # Setup grid code & W
- # -------------------------
- # Use pre-generated scales, phases, orientations, X, Y from workspace
- # Select top 10%
- threshold = np.percentile(scales, 90)
- top_idx = np.where(scales >= threshold)[0]
- sc_top = scales[top_idx]
- ph_top = phases[top_idx]
- or_top = orientations[top_idx]
- n_top = len(top_idx)
- # Compute W by finite differences averaged
- # reuse grid_embeddings, X, Y
- H, Wg, _ = grid_embeddings.shape
- dx = X[0,1] - X[0,0]
- dy = Y[1,0] - Y[0,0]
- emb_top = grid_embeddings[:,:,top_idx]
- pd_x = (emb_top[:,1:,:] - emb_top[:,:-1,:]) / dx
- pd_y = (emb_top[1:,:,:] - emb_top[:-1,:,:]) / dy
- avg_px = pd_x.reshape(-1, n_top).mean(axis=0)
- avg_py = pd_y.reshape(-1, n_top).mean(axis=0)
- Wmat = np.stack([avg_py, -avg_py, -avg_px, avg_px], axis=1)
- # -------------------------
- # Navigation via grid code
- # -------------------------
- walls = [] # no walls
- start = np.array([0.5, 2.5])
- goal = np.array([2.5, 0.5])
- noise_level = 80
- dt = 0.05
- threshold = 0.1
- max_steps = 2000
- # Precompute goal embedding
- goal_emb = grid_activation(goal[0], goal[1], sc_top, ph_top, or_top)
- paths = []
- for _ in range(5):
- pos = start.copy()
- path = [pos.copy()]
- for _ in range(max_steps):
- # current embedding
- cur_emb = grid_activation(pos[0], pos[1], sc_top, ph_top, or_top)
- delta = goal_emb - cur_emb # high-D target diff
- scores = delta @ Wmat + noise_level * np.random.randn(4) # (4,) for [Up,Down,Left,Right]
- # compute movement
- step = np.array([scores[3] - scores[2], scores[0] - scores[1]])
- # normalize to unit-length
- if np.linalg.norm(step) > 1e-6:
- step = step / np.linalg.norm(step)
- pos = compute_valid_step(pos, step, dt, walls)
- path.append(pos.copy())
- if np.linalg.norm(pos - goal) < threshold:
- break
- paths.append(np.array(path))
- # -------------------------
- # Plotting
- # -------------------------
- fig, ax = plt.subplots(figsize=(6,6))
- ax.plot(start[0], start[1], 'o', ms=20, c='black')
- ax.plot(goal[0], goal[1], '*', ms=30, c='black')
- colors = plt.cm.rainbow(np.linspace(0,1,len(paths)))
- for p, c in zip(paths, colors):
- ax.plot(p[:,0], p[:,1], color=c, lw=2, alpha=0.8)
- ax.set_xlim(0,4); ax.set_ylim(0,4); ax.set_aspect('equal')
- ax.set_xticks([])
- ax.set_yticks([])
- # ax.axis('off')
- plt.show()
- # %%
- probs.max()
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- # — your existing grid_activation, scales, phases, orientations, and BTSP weights —
- # grid_activation(x,y,scales,phases,orientations) → (n_grid_cells,)
- # weights: shape (n_place_cells, n_grid_cells)
- # paths: list of trajectories, each shape [T,2]
- # 1) Grab first trajectory
- traj = paths[0]
- T = traj.shape[0]
- n_pc = weights.shape[0]
- # 2) Sample spikes from learned place-cell activations
- spikes = np.zeros((T, n_pc), dtype=int)
- for t, (x, y) in enumerate(traj):
- gact = grid_activation(x, y, scales, phases, orientations) # pre-syn grid code
- p_act = weights @ gact # learned place-cell drive
- exps = np.exp(p_act - p_act.max())
- probs = exps / exps.sum() * 1 + 0.0005*np.random.randn(exps.shape[0]) # softmax → firing probs
- spikes[t] = (np.random.rand(n_pc) < probs).astype(int)
- # 3) Pick top-50 most active, rank by mean spike time
- counts = spikes.sum(axis=0)
- avg_time = (np.arange(T)[:,None] * spikes).sum(axis=0) / (counts + 1e-9)
- top50 = np.argsort(counts)[-50:] # highest-firing 50
- order = np.argsort(avg_time[top50]) # earliest mean time first
- selected = top50[order]
- # 4) Raster plot (spikes only)
- fig, ax = plt.subplots(figsize=(3,5), dpi=300)
- for r, neuron in enumerate(selected):
- times = np.where(spikes[:, neuron])[0]
- ax.vlines(times, r + 0.2, r + 0.8, color='black', linewidth=0.7)
- ax.set_xlim(0, T)
- ax.set_ylim(0, 50)
- # ax.set_xlabel('Time step')
- ticks = ax.get_xticks() # e.g. [0, 100, 200, …]
- ax.set_xticks(ticks)
- ax.set_xticklabels([(t/100) for t in ticks])
- ax.set_yticks([0,10,20,30,40,50]) # hide neuron labels
- plt.tight_layout()
- plt.show()
- # %%
- ticks
- # %%
- fig, ax = plt.subplots(figsize=(1.5,2.5), dpi=300)
- for r, neuron in enumerate(selected):
- times = np.where(spikes[:, neuron])[0]
- ax.vlines(times, r + 0.2, r + 0.8, color='black', linewidth=0.4)
- ax.set_xlim(0, T)
- ax.set_ylim(0, 50)
- # ax.set_xlabel('Time step')
- ticks = ax.get_xticks() # e.g. [0, 100, 200, …]
- # ticks = [0,100,200,300,400,500]
- ax.set_xticks(ticks)
- ax.set_xticklabels([(t/100) for t in ticks])
- ax.set_yticks([0,10,20,30,40,50]) # hide neuron labels
- plt.tight_layout()
- plt.show()
- # %%
- import numpy as np
- import matplotlib.pyplot as plt
- import matplotlib.patches as patches
- # -------------------------
- # Helpers from earlier
- # -------------------------
- def grid_activation(x, y, scales, phases, orientations):
- pos = np.array([x, y])
- angles = np.array([0, np.pi/3, 2*np.pi/3])
- acts = []
- for scale, phase, orient in zip(scales, phases, orientations):
- R = np.array([[np.cos(orient), -np.sin(orient)],
- [np.sin(orient), np.cos(orient)]])
- pr = R @ (pos - phase)
- proj = pr[0]*np.cos(angles) + pr[1]*np.sin(angles)
- grating = np.sum(np.cos((4*np.pi/(scale*np.sqrt(3))) * proj))
- acts.append((2/3)*grating)
- return np.array(acts)
- def rect_bounds(rect, eps=0.01):
- x, y, l, orient = rect
- if orient == 'h':
- # horizontal: thin in y
- return (x - l/2, x + l/2,
- y - eps, y + eps)
- else:
- # vertical: thin in x
- return (x - eps, x + eps,
- y - l/2, y + l/2)
- def is_inside_rect(pos, rect):
- x0, x1, y0, y1 = rect_bounds(rect)
- return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
- def segment_intersects_rect(p1, p2, rect):
- x0, x1, y0, y1 = rect_bounds(rect)
- # If endpoint lies within rect, it intersects
- if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
- return True
- # Otherwise check each edge
- edges = [
- (np.array([x0, y0]), np.array([x1, y0])),
- (np.array([x1, y0]), np.array([x1, y1])),
- (np.array([x1, y1]), np.array([x0, y1])),
- (np.array([x0, y1]), np.array([x0, y0]))
- ]
- for C, D in edges:
- if seg_intersect(p1, p2, C, D):
- return True
- return False
- # def is_inside_rect(pos, rect):
- # x, y, l, orient = rect
- # if orient == 'h':
- # x0, y0 = x - l/2, y - 0.05
- # x1, y1 = x + l/2, y + 0.05
- # else:
- # x0, y0 = x - 0.05, y - l/2
- # x1, y1 = x + 0.05, y + l/2
- # return (x0 <= pos[0] <= x1) and (y0 <= pos[1] <= y1)
- def ccw(A, B, C):
- return (C[1]-A[1])*(B[0]-A[0]) > (B[1]-A[1])*(C[0]-A[0])
- def seg_intersect(A, B, C, D):
- return ccw(A, C, D) != ccw(B, C, D) and ccw(A, B, C) != ccw(A, B, D)
- # def segment_intersects_rect(p1, p2, rect):
- # x0, x1, y0, y1 = (
- # rect[0] - rect[2]/2, rect[0] + rect[2]/2,
- # rect[1] - (0.01 if rect[3]=='h' else rect[2]/2),
- # rect[1] + (0.01 if rect[3]=='h' else rect[2]/2)
- # )
- # if x0 <= p2[0] <= x1 and y0 <= p2[1] <= y1:
- # return True
- # edges = [
- # (np.array([x0, y0]), np.array([x1, y0])),
- # (np.array([x1, y0]), np.array([x1, y1])),
- # (np.array([x1, y1]), np.array([x0, y1])),
- # (np.array([x0, y1]), np.array([x0, y0])),
- # ]
- # for C, D in edges:
- # if seg_intersect(p1, p2, C, D):
- # return True
- return False
- def is_valid_position(pos, walls):
- if not (0 <= pos[0] <= 4 and 0 <= pos[1] <= 4):
- return False
- for w in walls:
- if is_inside_rect(pos, w):
- return False
- return True
- def compute_valid_step(pos, vel, dt, walls):
- full = vel * dt
- cand = pos + full
- if is_valid_position(cand, walls) and not any(segment_intersects_rect(pos, cand, w) for w in walls):
- return cand
- # slide
- for dx, dy in [(full[0], 0), (0, full[1])]:
- cand2 = pos + np.array([dx, dy])
- if is_valid_position(cand2, walls) and not any(segment_intersects_rect(pos, cand2, w) for w in walls):
- return cand2
- # fallback smaller
- for f in [0.5, 0.25, 0.1]:
- cand3 = pos + full * f
- if is_valid_position(cand3, walls) and not any(segment_intersects_rect(pos, cand3, w) for w in walls):
- return cand3
- return pos
- # -------------------------
- # Setup grid code & W
- # -------------------------
- # Top 10% scales
- threshold = np.percentile(scales, 90)
- top_idx = np.where(scales >= threshold)[0]
- sc_top = scales[top_idx]
- ph_top = phases[top_idx]
- or_top = orientations[top_idx]
- n_top = len(top_idx)
- # Finite diff to build Wmat
- H, Wg, _ = grid_embeddings.shape
- dx = X[0,1] - X[0,0]
- dy = Y[1,0] - Y[0,0]
- emb_top = grid_embeddings[:,:,top_idx]
- pd_x = (emb_top[:,1:,:] - emb_top[:,:-1,:]) / dx
- pd_y = (emb_top[1:,:,:] - emb_top[:-1,:,:]) / dy
- avg_px = pd_x.reshape(-1, n_top).mean(axis=0)
- avg_py = pd_y.reshape(-1, n_top).mean(axis=0)
- Wmat = np.stack([avg_py, -avg_py, -avg_px, avg_px], axis=1)
- # -------------------------
- # Navigation with repulsion
- # -------------------------
- walls = [
- (1,1.5,1,'h'),
- (1.5,1,1,'v'),
- (1.5,2,1,'v'),
- (2,2.5,1,'h'),
- (3,2.5,1,'h'),
- (2.5,3,1,'v')
- ]
- # Precompute repulsion points
- rep_points = []
- for x, y, l, orient in walls:
- rep_points.append(np.array([x, y]))
- if orient == 'h':
- rep_points += [np.array([x-l/2,y]), np.array([x+l/2,y])]
- else:
- rep_points += [np.array([x,y-l/2]), np.array([x,y+l/2])]
- start = np.array([0.2, 0.2])
- goal = np.array([2.0, 2.0])
- k_wall = 0.08
- noise_level = 30.0
- dt = 0.05
- threshold_goal = 0.1
- max_steps = 1000
- n_paths = 5
- max_dist = np.linalg.norm(start - goal)
- goal_emb = grid_activation(goal[0], goal[1], sc_top, ph_top, or_top)
- def compute_velocity(pos):
- # Repulsive 2D
- f_rep = np.zeros(2)
- for wc in rep_points:
- delta = pos - wc
- dsq = np.sum(delta**2)
- if dsq > 1e-4:
- f_rep += k_wall * delta / dsq
- rep_scale = np.clip(np.linalg.norm(pos-goal) / max_dist, 0, 1)
- f_rep *= rep_scale
- # Grid-code guidance
- cur_emb = grid_activation(pos[0], pos[1], sc_top, ph_top, or_top)
- delta_emb = goal_emb - cur_emb
- scores = delta_emb @ Wmat + noise_level * np.random.randn(4) # [Up,Down,Left,Right]
- step = np.array([scores[3]-scores[2], scores[0]-scores[1]])
- if np.linalg.norm(step) > 1e-6:
- step = step / np.linalg.norm(step)
- return step + f_rep # combine
- # Generate paths
- paths = []
- for _ in range(n_paths):
- pos = start.copy()
- path = [pos.copy()]
- for _ in range(max_steps):
- vel = compute_velocity(pos)
- pos = compute_valid_step(pos, vel, dt, walls)
- path.append(pos.copy())
- if np.linalg.norm(pos - goal) < threshold_goal:
- break
- paths.append(np.array(path))
- # -------------------------
- # Plotting
- # -------------------------
- fig, ax = plt.subplots(figsize=(6,6))
- # draw walls
- for x, y, l, orient in walls:
- if orient == 'h':
- ax.add_patch(patches.Rectangle((x-l/2, y-0.02), l, 0.04, color='black'))
- else:
- ax.add_patch(patches.Rectangle((x-0.02, y-l/2), 0.04, l, color='black'))
- # start & goal
- ax.plot(start[0], start[1], 'o', ms=12, c='black')
- ax.plot(goal[0], goal[1], '*', ms=15, c='black')
- # paths
- colors = plt.cm.rainbow(np.linspace(0,1,n_paths))
- for p, c in zip(paths, colors):
- ax.plot(p[:,0], p[:,1], color=c, lw=2, alpha=0.8)
- ax.set_xlim(0,4); ax.set_ylim(0,4); ax.set_aspect('equal')
- ax.set_xticks([]); ax.set_yticks([])
- plt.show()
- # %%
- vel
- # %%
gcml_grid_cell.ipynb at commit ff76859, under MIT · at the source
Overview
- Department of Precision Instruments, Center for Brain-Inspired Computing Research (CBICR), Tsinghua University,Beijing, China
- Institute of Machine Learning and Neural Computation, Graz University of Technology,Graz, Austria
- Institute of Cognitive Sciences and Technologies, National Research Council,Rome, Italy
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
ff76859b71a2bc2056b50f5e052475351c007f76, 1 April 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
6 files
- gcml_abstract_graph.ipyn
b , Jupyter, 644 lines - gcml_grid_cell.ipynb, Jupyter, 874 lines, 2 matches
- gcml_tiling.ipynb, Jupyter, 643 lines
- utils/
dataset_tiling.py , Python, 303 lines - LICENSE, License, 21 lines
- README.md, Text, 11 lines
Zenodo 19370442
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
- 27 September 2026: the link answers (HTTP 200)
6 files
- gcml_abstract_graph.ipyn
b , Jupyter, 644 lines - gcml_grid_cell.ipynb, Jupyter, 874 lines
- gcml_tiling.ipynb, Jupyter, 643 lines
- utils/
dataset_tiling.py , Python, 303 lines - LICENSE, License, 21 lines
- README.md, Text, 11 lines
Code availability
The code used for training and evaluating the GCML is publicly available via GitHub at https://
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://
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://
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/
url = {https://
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/
VL - 8
IS - 7
SP - 1045
EP - 1065
SN - 2522-5839
PB - Nature Portfolio
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "8",
"issue": "7",
"page": "1045-1065",
"DOI": "10.1038/
"PMID": "42499994",
"PMCID": "PMC13395624",
"ISSN": "2522-5839",
"publisher": "Nature Portfolio",
"URL": "https://
"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 biologyIn 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 communicationsIn 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 : CBIn 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: HippocampusIn 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 communicationsIn 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 biologyIn 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 communicationsIn 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 advancesIn 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 neuroscienceIn 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 neuroscienceIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 2 repositories of the authors' code, each at its verified commit and with its license, 8 scripts, and 2 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:6b8876420c725172…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
