OSCR

Decoding behavior with minimal and interpretable agent models.

Code ↔ Paper

6 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 6 matches
  1. [1] § Methods › Metric-adaptive particle swarm optimization ↔ src/DiscreteObs/inference.py, lines 287–346 · score 0.79 · social coefficient, cognitive coefficient, inertia, swarming, global, neighbors
  2. [2] § Methods › Metric-adaptive particle swarm optimization ↔ src/MAPSO.py, lines 341–418 · score 0.65 · nearest, coefficient, inertia, social, swarming, cognitive
  3. [3] § Methods › Metric-adaptive particle swarm optimization ↔ src/DiscreteObs/inference.py, lines 287–346 · score 0.63 · global best, parameter space, mutation, adaptive, swarm, neighbors
  4. [4] § Methods › MAPSO training schedule ↔ src/FSC.py, lines 852–915 · score 0.55 · best inferred, negative log likelihood, MAPSO, epochs, optimization, particle
  5. [5] § Methods › Metric-adaptive particle swarm optimization ↔ src/MAPSO.py, lines 119–187 · score 0.55 · mutated, mutation, adaptively, swarming, MAPSO, velocity
  6. [6] § Methods › MAPSO training schedule ↔ src/FSC.py, lines 713–771 · score 0.52 · Gaussian distribution, particles initial, multivariate, covariance, MAPSO

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 · 929 lines · 41 KB · no license · 2 matches

  1. import torch
  2. from torch import nn
  3. import numpy as np
  4. import numba as nb
  5. import random
  6. import utils
  7. import MAPSO
  8. class InferenceDiscreteObs:
  9. def __init__(self, FSC):
  10. """
  11. Initialize the inference backend for a discrete-observation FSC.
  12. This constructor moves all FSC trainable quantities to the best
  13. available torch device (CUDA, MPS, or CPU), converts numpy arrays to
  14. tensors when needed, and wraps parameters in ``nn.Parameter`` if the
  15. FSC is not already trained.
  16. Parameters:
  17. --- FSC: FSC
  18. Parent FSC object configured for discrete observations.
  19. """
  20. self.FSC = FSC
  21. if torch.cuda.is_available():
  22. self.device = torch.device("cuda")
  23. elif torch.backends.mps.is_available():
  24. self.device = torch.device("mps")
  25. else:
  26. self.device = torch.device("cpu")
  27. if not isinstance(self.FSC.psi, torch.Tensor):
  28. self.FSC.psi = torch.tensor(self.FSC.psi.astype(np.float32), device=self.device)
  29. if not isinstance(self.FSC.psi, nn.Parameter) and FSC.trained == False:
  30. self.FSC.psi = nn.Parameter(self.FSC.psi)
  31. for idx, param in enumerate(self.FSC.GPModel.params):
  32. param_name = self.FSC.GPModel.param_names[idx]
  33. if not isinstance(param, torch.Tensor):
  34. param = torch.tensor(param.astype(np.float32), device=self.device)
  35. if not isinstance(param, nn.Parameter) and FSC.trained == False:
  36. param = nn.Parameter(param)
  37. self.FSC.GPModel.__setattr__(param_name, param)
  38. self.InternalMemSpace = torch.arange(self.FSC.M)
  39. self.InternalActSpace = torch.arange(self.FSC.A)
  40. self.InternalObsSpace = torch.arange(self.FSC.Y)
  41. self.trajectories_loaded = False
  42. self.optimizer_initialized = False
  43. self.trained = False
  44. def get_policy_params(self):
  45. """
  46. Return the current policy parameters in model-defined order.
  47. Returns:
  48. --- tuple of torch.Tensor
  49. Tuple ``(param_0, ..., param_k)`` matching
  50. ``self.FSC.GPModel.param_names``.
  51. """
  52. return tuple([self.FSC.GPModel.__getattribute__(param) for param in self.FSC.GPModel.param_names])
  53. def get_TMat(self):
  54. """
  55. Return the joint transition matrix used during inference.
  56. Returns:
  57. --- torch.Tensor
  58. Transition tensor with shape ``(Y, A, M, M)``.
  59. """
  60. return self.FSC.GPModel.get_TMat_torch()
  61. def get_memory_transition(self):
  62. """
  63. Return the memory transition component of the policy.
  64. Returns:
  65. --- torch.Tensor
  66. Memory transition tensor ``g(m' | a, m, y)`` with shape
  67. ``(Y, A, M, M)``.
  68. """
  69. return self.FSC.GPModel.get_memory_transition_torch()
  70. def get_action_policy(self):
  71. """
  72. Return the marginal action policy.
  73. Returns:
  74. --- torch.Tensor
  75. Action policy tensor ``pi(a | m)`` with shape ``(M, A)``.
  76. """
  77. return self.FSC.GPModel.get_action_policy_torch()
  78. def load_trajectories(self, trajectories):
  79. """
  80. Loads a set of trajectories to be used for training the FSC.
  81. Parameters:
  82. --- trajectories: list of dicts
  83. List of dictionaries containing the actions and observations for each trajectory.
  84. """
  85. self.ObsAct_trajectories = []
  86. self.observations_trajectories_np = []
  87. self.actions_trajectories_np = []
  88. self.n_trajectories = len(trajectories)
  89. self.pStart_ya_emp = np.zeros((self.FSC.Y, self.FSC.A))
  90. for trajectory in trajectories:
  91. observations = self._map_obs_to_internal_space(trajectory["observations"])
  92. actions = self._map_act_to_internal_space(trajectory["actions"])
  93. self.observations_trajectories_np.append(observations)
  94. self.actions_trajectories_np.append(actions)
  95. self.ObsAct_trajectories.append([torch.tensor(observations), torch.tensor(actions)])
  96. y0 = observations[0]
  97. a0 = actions[0]
  98. self.pStart_ya_emp[y0, a0] += 1
  99. self.pStart_ya_emp /= np.sum(self.pStart_ya_emp)
  100. self.trajectories_loaded = True
  101. def _map_obs_to_internal_space(self, obs):
  102. """
  103. Helper method to map an observation sequence to the internal observation space.
  104. Parameters:
  105. --- obs: np.array
  106. Observation sequence to map.
  107. Returns:
  108. --- obs_internal: np.array
  109. Observation sequence in the internal observation space.
  110. """
  111. return np.array([self.get_obs_idx(o) for o in obs])
  112. def _map_act_to_internal_space(self, act):
  113. """
  114. Helper method to map an action sequence to the internal action space.
  115. Parameters:
  116. --- act: np.array
  117. Action sequence to map.
  118. Returns:
  119. --- act_internal: np.array
  120. Action sequence in the internal action space.
  121. """
  122. return np.array([self.get_act_idx(a) for a in act])
  123. def get_obs_idx(self, obs):
  124. """
  125. Helper method to get the index of an observation in the observation space.
  126. Parameters:
  127. --- obs: int
  128. Observation to get the index of.
  129. Returns:
  130. --- idx: int
  131. Index of the observation in the observation space.
  132. """
  133. return np.where(obs == self.FSC.ObsSpace)[0][0]
  134. def get_act_idx(self, act):
  135. """
  136. Helper method to get the index of an action in the action space.
  137. Parameters:
  138. --- act: int
  139. Action to get the index of.
  140. Returns:
  141. --- idx: int
  142. Index of the action in the action space.
  143. """
  144. return np.where(self.FSC.ActSpace == act)[0][0]
  145. def get_mem_idx(self, mem):
  146. """
  147. Helper method to get the index of a memory state in the memory space.
  148. Parameters:
  149. --- mem: int
  150. Memory state to get the index of.
  151. Returns:
  152. --- idx: int
  153. Index of the memory state in the memory space
  154. """
  155. return np.where(self.FSC.MemSpace == mem)[0][0]
  156. def evaluate_nloglikelihood(self, idx_traj, grad_required=False):
  157. """
  158. Wrapper method to evaluate the negative log-likelihood of a given trajectory.
  159. It distinguishes between the case of custom observation and action spaces and the case
  160. of default observation and action spaces.
  161. Parameters:
  162. --- idx_traj: int
  163. Index of the trajectory to evaluate.
  164. --- grad_required: bool (default = False)
  165. Flag indicating whether the gradient is required or not.
  166. Returns:
  167. --- nLL: float
  168. Negative log-likelihood of the trajectory.
  169. """
  170. observations, actions = self.ObsAct_trajectories[idx_traj]
  171. return self.loss(observations, actions, grad_required = grad_required)
  172. def loss(self, observations, actions, grad_required=True):
  173. """
  174. Method to compute the negative log-likelihood of a given trajectory with default observation and action spaces.
  175. The gradients of the loss are computed if the grad_required flag is set to True.
  176. Parameters:
  177. --- observations: torch.tensor
  178. Array of observations.
  179. --- actions: torch.tensor
  180. Array of actions.
  181. --- grad_required: bool (default = True)
  182. Flag indicating whether the gradient is required or not.
  183. Returns:
  184. --- nLL: float
  185. Negative log-likelihood of the trajectory.
  186. """
  187. nLL = torch.tensor(0.0, requires_grad = grad_required)
  188. TMat = self.FSC.GPModel.get_TMat_torch()
  189. if self.FSC._init_memory_obs_dependent:
  190. rho = self.FSC.rho[observations[0]]
  191. else:
  192. rho = self.FSC.rho
  193. for t in range(observations.size(0)):
  194. idx_a = actions[t]
  195. idx_obs = observations[t]
  196. transition_probs = TMat[idx_obs, idx_a].T
  197. if torch.sum(transition_probs) == 0:
  198. break
  199. if t == 0:
  200. if transition_probs.device.type == 'mps':
  201. # MPS-specific workaround
  202. transition_probs_safe = transition_probs.clone().detach().requires_grad_(transition_probs.requires_grad)
  203. rho_safe = rho.clone().detach().requires_grad_(rho.requires_grad)
  204. m = torch.matmul(transition_probs_safe, rho_safe)
  205. else:
  206. m = torch.matmul(transition_probs, rho)
  207. else:
  208. m = torch.matmul(transition_probs, m)
  209. if torch.sum(m) == 0:
  210. break
  211. mv = torch.sum(m)
  212. nLL = nLL - torch.log(mv)
  213. m /= mv
  214. if torch.sum(m) == 0:
  215. return nLL
  216. else:
  217. return nLL - torch.log(torch.sum(m))
  218. def optimize_w_MAPSO(self, trainable_params, trainable_params_mask,
  219. n_particles, NEpochs,
  220. init_particles, init_velocities,
  221. c1_init, c2_init, w_init,
  222. sigma_min, sigma_max,
  223. dynamic_topology, n_neighbors_init, n_neighbors_final, num_neighbors_mid,
  224. verbose, verbose_epochs):
  225. """
  226. Optimize FSC parameters with MAPSO (Adaptive Particle Swarm).
  227. This method builds a flattened parameter space containing all policy
  228. parameters plus ``psi``, applies optional element-wise trainability
  229. masks, initializes particle positions/velocities from the configured
  230. distributions, and runs either global-best MAPSO or kNN-local MAPSO.
  231. Parameters:
  232. --- trainable_params: dict
  233. Dictionary mapping parameter names (including ``psi``) to booleans
  234. indicating whether each parameter block is trainable.
  235. --- trainable_params_mask: dict
  236. Optional element-wise masks. Values are either ``None`` or boolean
  237. arrays with the same shape as the corresponding parameter.
  238. --- n_particles: int
  239. Number of particles in the swarm.
  240. --- NEpochs: int
  241. Number of MAPSO iterations.
  242. --- init_particles: dict
  243. Particle initialization settings (distribution and hyperparameters).
  244. --- init_velocities: dict
  245. Velocity initialization settings (distribution and hyperparameters).
  246. --- c1_init: float
  247. Initial cognitive coefficient.
  248. --- c2_init: float
  249. Initial social coefficient.
  250. --- w_init: float
  251. Initial inertia weight.
  252. --- sigma_min: float
  253. Minimum mutation scale used by MAPSO convergence strategy.
  254. --- sigma_max: float
  255. Maximum mutation scale used by MAPSO convergence strategy.
  256. --- dynamic_topology: bool
  257. If True, use kNN local-best topology; otherwise use global-best.
  258. --- n_neighbors_init: int or None
  259. Initial number of neighbors for dynamic topology.
  260. --- n_neighbors_final: int or None
  261. Final number of neighbors for dynamic topology.
  262. --- num_neighbors_mid: int or None
  263. Midpoint number of neighbors for dynamic topology.
  264. --- verbose: bool
  265. Print per-iteration diagnostics.
  266. --- verbose_epochs: bool
  267. Print per-epoch summary.
  268. Returns:
  269. --- np.ndarray
  270. Best objective value at each MAPSO iteration.
  271. """
  272. assert self.trajectories_loaded, "No trajectories have been loaded. Load trajectories with the load_trajectories method."
  273. assert not self.trained, "The model has already been trained. If you want to train it again, reinitialize it or set the flag self.trained to False."
  274. spacedim = sum([param.numel() for param in self.get_policy_params()]) + self.FSC.psi.numel()
  275. trainable_mask = np.zeros(spacedim, dtype=bool)
  276. init_pos = np.zeros((n_particles, spacedim))
  277. init_vel = np.zeros((n_particles, spacedim))
  278. if init_particles["distribution"] == "multivariate_normal":
  279. random_pos_mv = np.random.multivariate_normal(init_particles["mean"], init_particles["cov"],
  280. n_particles)
  281. start_idx = 0
  282. for idx, param in enumerate(self.get_policy_params()):
  283. param_name = self.FSC.GPModel.param_names[idx]
  284. end_idx = start_idx + param.numel()
  285. if trainable_params_mask[param_name] is not None:
  286. trainable_mask[start_idx:end_idx] = trainable_params_mask[param_name].flatten()
  287. else:
  288. trainable_mask[start_idx:end_idx] = trainable_params[param_name]
  289. flatten_param = self.FSC.GPModel.__getattribute__(self.FSC.GPModel.param_names[idx]).detach().cpu().numpy().flatten()
  290. init_pos[:, start_idx:end_idx] = np.tile(flatten_param, (n_particles, 1))
  291. init_vel[:, start_idx:end_idx] = np.zeros((n_particles, param.numel()))
  292. if trainable_params[param_name]:
  293. num_trainable = np.sum(trainable_mask[start_idx:end_idx])
  294. if init_particles["distribution"] == "uniform":
  295. random_pos = np.random.uniform(init_particles["xmin"], init_particles["xmax"],
  296. (n_particles, num_trainable))
  297. elif init_particles["distribution"] == "normal":
  298. random_pos = np.random.normal(init_particles["mean"], init_particles["std"],
  299. (n_particles, num_trainable))
  300. elif init_particles["distribution"] == "multivariate_normal":
  301. random_pos = random_pos_mv[:, start_idx:end_idx][:, trainable_mask[start_idx:end_idx]]
  302. elif init_particles["distribution"] == "uniform_with_biases":
  303. random_pos = np.random.uniform(init_particles["xmin"], init_particles["xmax"],
  304. (n_particles, num_trainable))
  305. random_pos += init_particles["biases"][start_idx:end_idx][trainable_mask[start_idx:end_idx]]
  306. else:
  307. raise ValueError("Invalid position distribution.")
  308. if init_velocities["distribution"] == "uniform":
  309. random_vel = np.random.uniform(init_velocities["vmin"], init_velocities["vmax"],
  310. (n_particles, num_trainable))
  311. elif init_velocities["distribution"] == "normal":
  312. random_vel = np.random.normal(init_velocities["mean"], init_velocities["std"],
  313. (n_particles, num_trainable))
  314. else:
  315. raise ValueError("Invalid velocity distribution.")
  316. init_pos[:, start_idx:end_idx][:, trainable_mask[start_idx:end_idx]] = random_pos
  317. init_vel[:, start_idx:end_idx][:, trainable_mask[start_idx:end_idx]] = random_vel
  318. start_idx = end_idx
  319. if trainable_params_mask["psi"] is not None:
  320. trainable_mask[-self.FSC.psi.numel():] = trainable_params_mask["psi"].flatten()
  321. else:
  322. trainable_mask[-self.FSC.psi.numel():] = trainable_params["psi"]
  323. #print(self.FSC.psi.shape, self.FSC.psi.numel())
  324. init_pos[:, -self.FSC.psi.numel():] = np.tile(self.FSC.psi.detach().cpu().numpy().flatten(), (n_particles, 1))
  325. init_vel[:, -self.FSC.psi.numel():] = np.zeros((n_particles, self.FSC.psi.numel()))
  326. if trainable_params["psi"]:
  327. num_trainable = np.sum(trainable_mask[-self.FSC.psi.numel():])
  328. if init_particles["distribution"] == "uniform":
  329. random_pos = np.random.uniform(init_particles["xmin"], init_particles["xmax"],
  330. (n_particles, num_trainable))
  331. elif init_particles["distribution"] == "normal":
  332. random_pos = np.random.normal(init_particles["mean"], init_particles["std"],
  333. (n_particles, num_trainable))
  334. elif init_particles["distribution"] == "multivariate_normal":
  335. random_pos = random_pos_mv[:, -self.FSC.psi.numel():][:, trainable_mask[-self.FSC.psi.numel():]]
  336. elif init_particles["distribution"] == "uniform_with_biases":
  337. random_pos = np.random.uniform(init_particles["xmin"], init_particles["xmax"],
  338. (n_particles, num_trainable))
  339. random_pos += init_particles["biases"][-self.FSC.psi.numel():][trainable_mask[-self.FSC.psi.numel():]]
  340. else:
  341. raise ValueError("Invalid position distribution.")
  342. if init_velocities["distribution"] == "uniform":
  343. random_vel = np.random.uniform(init_velocities["vmin"], init_velocities["vmax"],
  344. (n_particles, num_trainable))
  345. elif init_velocities["distribution"] == "normal":
  346. random_vel = np.random.normal(init_velocities["mean"], init_velocities["std"],
  347. (n_particles, num_trainable))
  348. else:
  349. raise ValueError("Invalid velocity distribution.")
  350. init_pos[:, -self.FSC.psi.numel():][:, trainable_mask[-self.FSC.psi.numel():]] = random_pos
  351. init_vel[:, -self.FSC.psi.numel():][:, trainable_mask[-self.FSC.psi.numel():]] = random_vel
  352. if dynamic_topology:
  353. gbests, gbest_values = MAPSO.particle_swarm_optimization_discrete_kNN(self.FSC.GPModel._nb_get_TMat_flat,
  354. trainable_mask,
  355. spacedim, n_particles, NEpochs,
  356. self.FSC.M, self.FSC.A, self.FSC.Y,
  357. self.observations_trajectories_np,
  358. self.actions_trajectories_np,
  359. init_pos, init_vel,
  360. num_neighbors_init = n_neighbors_init,
  361. num_neighbors_final = n_neighbors_final,
  362. num_neighbors_mid = num_neighbors_mid,
  363. c1 = c1_init, c2 = c2_init, w = w_init,
  364. sigma_min = sigma_min, sigma_max = sigma_max,
  365. verbose = verbose, verbose_epochs = verbose_epochs,
  366. init_memory_obs_dependent = self.FSC._init_memory_obs_dependent)
  367. else:
  368. gbests, gbest_values = MAPSO.particle_swarm_optimization_discrete(self.FSC.GPModel._nb_get_TMat_flat,
  369. trainable_mask,
  370. spacedim, n_particles, NEpochs,
  371. self.FSC.M, self.FSC.A, self.FSC.Y,
  372. self.observations_trajectories_np,
  373. self.actions_trajectories_np,
  374. init_pos, init_vel,
  375. c1 = c1_init, c2 = c2_init, w = w_init,
  376. sigma_min = sigma_min, sigma_max = sigma_max,
  377. verbose = verbose, verbose_epochs = verbose_epochs,
  378. init_memory_obs_dependent = self.FSC._init_memory_obs_dependent)
  379. for idx, param in enumerate(self.get_policy_params()):
  380. param_name = self.FSC.GPModel.param_names[idx]
  381. start_idx = 0
  382. for idx, param in enumerate(self.get_policy_params()):
  383. param_name = self.FSC.GPModel.param_names[idx]
  384. end_idx = start_idx + param.numel()
  385. new_param = torch.tensor(gbests[-1, start_idx:end_idx].reshape(param.shape).astype(np.float32), device=self.device)
  386. self.FSC.GPModel.__setattr__(param_name,
  387. nn.Parameter(new_param))
  388. start_idx = end_idx
  389. new_psi = gbests[-1, -self.FSC.psi.numel():].astype(np.float32)
  390. new_psi = new_psi.reshape(self.FSC.psi.shape)
  391. self.FSC.psi = nn.Parameter(torch.tensor(new_psi, device=self.device))
  392. return gbest_values
  393. def optimize_w_gradient(self, use_ccopt,
  394. trainable_params, trainable_params_mask, any_masked,
  395. NEpochs, NBatch, lr,
  396. train_split, optimizer, scheduler_dict,
  397. maxiter, rho0, th, c_gauge,
  398. verbose, verbose_epochs):
  399. """
  400. Optimize FSC parameters with gradient-based training.
  401. Supports per-parameter learning rates, optional parameter masking, an
  402. optional convex-concave optimization step for ``rho`` (when enabled),
  403. mini-batch training, and optional validation split with best-epoch
  404. model selection.
  405. Parameters:
  406. --- use_ccopt: bool
  407. If True, update ``psi`` through convex-concave optimization of rho
  408. at each batch (only for observation-independent initialization).
  409. --- trainable_params: dict
  410. Parameter-level trainability flags.
  411. --- trainable_params_mask: dict
  412. Optional element-wise trainability masks.
  413. --- any_masked: bool
  414. True if at least one element-wise mask is active.
  415. --- NEpochs: int
  416. Number of gradient epochs.
  417. --- NBatch: int
  418. Batch size in number of trajectories.
  419. --- lr: float or dict
  420. Global learning rate or per-parameter learning-rate dictionary.
  421. --- train_split: float
  422. Fraction of trajectories used for training.
  423. --- optimizer: str
  424. Optimizer name (expected: ``ADAM`` or ``SGD``).
  425. --- scheduler_dict: dict
  426. Scheduler configuration dictionary.
  427. --- maxiter: int or None
  428. Maximum iterations for ccopt.
  429. --- rho0: np.ndarray or None
  430. Initial rho for ccopt.
  431. --- th: float or None
  432. Convergence threshold for ccopt.
  433. --- c_gauge: float or None
  434. Additive gauge applied to ``log(rho)`` when reconstructing psi.
  435. --- verbose: bool
  436. Verbosity flag (kept for interface consistency).
  437. --- verbose_epochs: bool
  438. If True, print per-epoch losses and learning rate.
  439. Returns:
  440. --- tuple
  441. ``(losses_train, losses_val)`` where validation losses are ``None``
  442. when no validation split is used.
  443. """
  444. assert self.trajectories_loaded, "No trajectories have been loaded. Load trajectories with the load_trajectories method."
  445. assert not self.trained, "The model has already been trained. If you want to train it again, reinitialize it or set the flag self.trained to False."
  446. lr_dict = {}
  447. if isinstance(lr, float):
  448. single_lr = True
  449. for param in self.FSC.GPModel.param_names:
  450. lr_dict[param] = lr
  451. lr_dict["psi"] = lr
  452. elif isinstance(lr, dict):
  453. for param in self.FSC.GPModel.param_names:
  454. if param != "psi" and trainable_params[param]:
  455. if param in lr:
  456. lr_dict[param] = lr[param]
  457. else:
  458. raise ValueError(f"Missing learning rate for parameter {param}.")
  459. if "psi" not in lr and trainable_params["psi"] and not use_ccopt:
  460. raise ValueError("Missing learning rate for psi.")
  461. else:
  462. lr_dict["psi"] = lr["psi"]
  463. else:
  464. raise ValueError("Invalid learning rate. The learning rate must be a float or a dictionary with the parameters as keys.")
  465. parkey_optimizer = []
  466. for idx, param in enumerate(self.FSC.GPModel.param_names):
  467. if trainable_params[param]:
  468. parkey_optimizer.append({"params": self.FSC.GPModel.__getattribute__(param), 'lr': lr_dict[param]})
  469. if trainable_params["psi"] and not use_ccopt:
  470. parkey_optimizer.append({'params': self.FSC.psi, 'lr': lr_dict["psi"]})
  471. if optimizer == "ADAM":
  472. self.optimizer = torch.optim.Adam(parkey_optimizer)
  473. elif optimizer == "SDG":
  474. self.optimizer = torch.optim.SGD(parkey_optimizer)
  475. if scheduler_dict["type"] == "exponential":
  476. scheduler = torch.optim.lr_scheduler.ExponentialLR(self.optimizer, gamma=scheduler_dict["decay_rate"])
  477. elif scheduler_dict["type"] == "fixed":
  478. scheduler = None
  479. else:
  480. raise ValueError("Invalid scheduler.")
  481. NTrain = int(train_split * len(self.ObsAct_trajectories))
  482. NVal = len(self.ObsAct_trajectories) - NTrain
  483. trjs_train = self.ObsAct_trajectories[:NTrain]
  484. best_params = [self.FSC.GPModel.__getattribute__(param) for param in self.FSC.GPModel.param_names]
  485. best_psi = self.FSC.psi
  486. best_epoch = 0
  487. losses_train = []
  488. init_loss = 0
  489. for idx_traj in range(NTrain):
  490. init_loss += self.loss(trjs_train[idx_traj][0], trjs_train[idx_traj][1], grad_required=False).item()
  491. losses_train.append(init_loss / NTrain)
  492. init_msg = f"Training with {NTrain} trajectories"
  493. if NVal != 0:
  494. trjs_val = self.FeatAct_trajectories[NTrain:]
  495. losses_val = []
  496. init_loss_val = 0
  497. for idx_traj in range(NVal):
  498. init_loss_val += self.loss(trjs_val[idx_traj][0], trjs_val[idx_traj][1], grad_required=False).item()
  499. losses_val.append(init_loss_val / NVal)
  500. init_msg += f" and validating with {NVal} trajectories."
  501. init_msg += " Initial training loss: " + str(losses_train[0]) + ". Initial validation loss: " + str(losses_val[0]) + "."
  502. else:
  503. init_msg += ". Initial loss: " + str(losses_train[0]) + "."
  504. if verbose_epochs:
  505. if single_lr:
  506. print(init_msg + f" Using a single learning rate of {lr}.")
  507. else:
  508. for idx, param in enumerate(self.FSC.GPModel.param_names):
  509. print(init_msg + f" Using learning rate {lr[param]} for {param}.")
  510. print(init_msg + f" Using learning rate {lr['psi']} for psi.")
  511. for epoch in range(NEpochs):
  512. running_loss = 0.0
  513. running_count = 0
  514. random.shuffle(trjs_train)
  515. for idx in range(0, NTrain, NBatch):
  516. if any_masked:
  517. pre_loss_params = {}
  518. for key in self.FSC.GPModel.param_names:
  519. mask = trainable_params_mask[key]
  520. if mask is not None:
  521. pre_loss_params[key] = self.FSC.GPModel.__getattribute__(key).detach().clone()
  522. if trainable_params_mask["psi"] is not None:
  523. pre_loss_psi = self.FSC.psi.detach().clone()
  524. self.optimizer.zero_grad()
  525. loss = torch.tensor(0.0, requires_grad=True)
  526. TMat = self.FSC.GPModel.get_TMat_torch()
  527. if use_ccopt and not self.FSC._init_memory_obs_dependent and trainable_params["psi"]:
  528. if rho0 is None:
  529. rho0 = np.ones(self.FSC.M)/self.FSC.M
  530. rho, _ = InferenceDiscreteObs.optimize_rho(self.FSC.Y, self.FSC.M, self.FSC.A,
  531. TMat.detach().cpu().numpy(), self.pStart_ya_emp,
  532. rho0, maxiter, th = th)
  533. rho = torch.tensor(rho.astype(np.float32), device = self.device)
  534. self.FSC.psi = nn.Parameter(torch.log(rho) + c_gauge)
  535. count = 0
  536. for idx_traj in range(idx, idx + NBatch):
  537. if idx_traj < NTrain:
  538. loss_traj = self.loss(trjs_train[idx_traj][0], trjs_train[idx_traj][1])
  539. if torch.isnan(loss_traj):
  540. continue
  541. loss = loss + loss_traj
  542. count += 1
  543. if count == 0:
  544. err_msg = "Gradient optimization failed because no valid trajectories were found in a batch. This means that the loss could not be evaluated due to forbidden transition, and that the current parameters are not compatible with some trajectories."
  545. err_msg += " Either improve initialization or choose a smaller learning rate. Overwriting with the best parameters found so far."
  546. self.FSC.psi = best_psi
  547. for idx, param in enumerate(self.FSC.GPModel.param_names):
  548. self.FSC.GPModel.__setattr__(param, nn.Parameter(best_params[idx]))
  549. raise RuntimeError(err_msg)
  550. loss.backward()
  551. self.optimizer.step()
  552. running_loss += loss.item()
  553. running_count += count
  554. if any_masked:
  555. for key in self.FSC.GPModel.param_names:
  556. mask = ~trainable_params_mask[key]
  557. if mask is not None:
  558. self.FSC.GPModel.__getattribute__(key).data[mask] = pre_loss_params[key].data[mask].clone()
  559. mask_psi = ~trainable_params_mask["psi"]
  560. if mask_psi is not None:
  561. self.FSC.psi.data[mask_psi] = pre_loss_psi.data[mask_psi].clone()
  562. running_loss = running_loss / running_count
  563. losses_train.append(running_loss)
  564. if NVal != 0:
  565. running_loss_val = 0.0
  566. for idx_traj in range(NVal):
  567. loss_val = torch.tensor(0.0, requires_grad=False)
  568. loss_traj_val = self.loss(trjs_val[idx_traj][0], trjs_val[idx_traj][1], grad_required=False)
  569. loss_val = loss_val + loss_traj_val
  570. running_loss_val += loss_val.item()
  571. running_loss_val = running_loss_val / NVal
  572. losses_val.append(running_loss_val)
  573. if running_loss_val < min(losses_val[:-1]):
  574. best_params = [self.FSC.GPModel.__getattribute__(param).detach().clone() for param in self.FSC.GPModel.param_names]
  575. best_psi = self.FSC.psi.detach().clone()
  576. best_epoch = epoch + 1
  577. if verbose_epochs:
  578. print(f"Epoch {epoch + 1} - Training loss: {round(running_loss, 5)}, Validation loss: {round(running_loss_val, 5)} - Learning rate: {round(self.optimizer.param_groups[0]['lr'], 5)}")
  579. else:
  580. if running_loss < min(losses_train[:-1]):
  581. best_params = [self.FSC.GPModel.__getattribute__(param).detach().clone() for param in self.FSC.GPModel.param_names]
  582. best_psi = self.FSC.psi.detach().clone()
  583. best_epoch = epoch + 1
  584. if verbose_epochs:
  585. print(f"Epoch {epoch + 1} - Training loss: {round(running_loss, 5)} - Learning rate: {round(self.optimizer.param_groups[0]['lr'], 5)}")
  586. if scheduler is not None:
  587. scheduler.step()
  588. if verbose_epochs:
  589. print("Training complete. Best parameters found at epoch", best_epoch)
  590. self.FSC.psi = best_psi
  591. for idx, param in enumerate(self.FSC.GPModel.param_names):
  592. self.FSC.GPModel.__setattr__(param, nn.Parameter(best_params[idx]))
  593. if NVal != 0:
  594. return losses_train, losses_val
  595. else:
  596. return losses_train, None
  597. def optimize(self, inference_params, verbose, verbose_epochs):
  598. """
  599. Run the full inference pipeline according to ``inference_params``.
  600. Depending on the configuration, this method executes MAPSO,
  601. gradient-based optimization, or both in sequence. It aggregates loss
  602. histories, computes ``best_loss``, and sets ``self.trained=True``.
  603. Parameters:
  604. --- inference_params: dict
  605. Validated inference configuration produced by
  606. ``FSC.set_inference_params``.
  607. --- verbose: bool
  608. Verbosity flag forwarded to optimization backends.
  609. --- verbose_epochs: bool
  610. If True, print per-epoch optimization updates.
  611. Returns:
  612. --- dict
  613. Loss history dictionary containing at least ``train`` and
  614. optionally ``MAPSO`` and ``val``.
  615. """
  616. assert self.trajectories_loaded, "No trajectories have been loaded. Load trajectories with the load_trajectories method."
  617. assert not self.trained, "The model has already been trained. If you want to train it again, reinitialize it or set the flag self.trained to False."
  618. loss_epochs = {}
  619. trainable_params = inference_params["trainable_parameters"]
  620. trainable_mask = inference_params["trainable_mask"]
  621. # check if any value of the trainable_mask dictionary is not None
  622. any_masked = False
  623. for key, val in trainable_mask.items():
  624. if val is not None:
  625. any_masked = True
  626. break
  627. if inference_params["use_MAPSO"]:
  628. n_particles = inference_params['n_particles_MAPSO']
  629. NEpochs = inference_params['NEpochs_MAPSO']
  630. c1_init = inference_params['c1_init_MAPSO']
  631. c2_init = inference_params['c2_init_MAPSO']
  632. w_init = inference_params['w_init_MAPSO']
  633. sigma_min = inference_params['sigma_min_MAPSO']
  634. sigma_max = inference_params['sigma_max_MAPSO']
  635. dynamic_topology = inference_params['dynamic_topology_MAPSO']
  636. num_neighbors_init = inference_params['num_neighbors_init_MAPSO']
  637. num_neighbors_final = inference_params['num_neighbors_final_MAPSO']
  638. num_neighbors_mid = inference_params['num_neighbors_mid_MAPSO']
  639. init_particles = inference_params['init_particles_MAPSO']
  640. init_velocities = inference_params['init_velocities_MAPSO']
  641. loss_MAPSO = self.optimize_w_MAPSO(trainable_params, trainable_mask,
  642. n_particles, NEpochs,
  643. init_particles, init_velocities,
  644. c1_init, c2_init, w_init,
  645. sigma_min, sigma_max,
  646. dynamic_topology,
  647. num_neighbors_init, num_neighbors_final, num_neighbors_mid,
  648. verbose, verbose_epochs)
  649. loss_epochs["MAPSO"] = loss_MAPSO
  650. if inference_params["use_gradient"]:
  651. NEpochs = inference_params['NEpochs_gradient']
  652. NBatch = inference_params['NBatch_gradient']
  653. lr = inference_params['lr_gradient']
  654. train_split = inference_params['train_split_gradient']
  655. scheduler = inference_params['scheduler_gradient']
  656. optimizer = inference_params['optimizer_gradient']
  657. if inference_params["use_ccopt"]:
  658. if self.FSC._init_memory_obs_dependent:
  659. raise ValueError("CCOpt cannot be used with memory observation dependent initialization.")
  660. maxiter = inference_params['maxiter_ccopt']
  661. rho0 = inference_params['rho0_ccopt']
  662. th = inference_params['th_ccopt']
  663. c_gauge = inference_params['c_gauge_ccopt']
  664. else:
  665. maxiter = None
  666. rho0 = None
  667. th = None
  668. c_gauge = None
  669. losses_gradient = self.optimize_w_gradient(inference_params["use_ccopt"],
  670. trainable_params, trainable_mask, any_masked,
  671. NEpochs, NBatch, lr,
  672. train_split, optimizer, scheduler,
  673. maxiter, rho0, th, c_gauge,
  674. verbose, verbose_epochs)
  675. if inference_params["use_MAPSO"]:
  676. loss_epochs["train"] = np.concatenate([loss_epochs["MAPSO"], losses_gradient[0][1:]])
  677. if train_split < 1.0:
  678. loss_epochs["val"] = np.concatenate([loss_epochs["MAPSO"], losses_gradient[1][1:]])
  679. self.best_loss = np.min(loss_epochs["val"])
  680. else:
  681. self.best_loss = np.min(loss_epochs["train"])
  682. else:
  683. loss_epochs["train"] = losses_gradient[0]
  684. if train_split < 1.0:
  685. loss_epochs["val"] = losses_gradient[1]
  686. self.best_loss = np.min(loss_epochs["val"])
  687. else:
  688. self.best_loss = np.min(loss_epochs["train"])
  689. else:
  690. loss_epochs["train"] = loss_MAPSO
  691. self.best_loss = np.min(loss_epochs["train"])
  692. self.trained = True
  693. return loss_epochs
  694. def get_inferred_policy_params(self):
  695. """
  696. Return the current inferred policy parameters.
  697. Returns:
  698. --- list of torch.Tensor
  699. Policy parameters in the order of
  700. ``self.FSC.GPModel.param_names``.
  701. """
  702. return [self.FSC.GPModel.__getattribute__(param) for param in self.FSC.GPModel.param_names]
  703. @staticmethod
  704. @nb.njit
  705. def optimize_rho(Y, M, A, TMat, pya, rhok, maxiter, th):
  706. """
  707. Numba-accelerated convex-concave fixed-point update for rho.
  708. Given the current transition matrix and empirical start distribution
  709. over observation-action pairs, iteratively updates ``rho`` until
  710. convergence or the maximum number of iterations is reached.
  711. Parameters:
  712. --- Y: int
  713. Number of observations.
  714. --- M: int
  715. Number of memory states.
  716. --- A: int
  717. Number of actions.
  718. --- TMat: np.ndarray
  719. Transition tensor with shape ``(Y, A, M, M)``.
  720. --- pya: np.ndarray
  721. Empirical start distribution over ``(y, a)`` with shape ``(Y, A)``.
  722. --- rhok: np.ndarray
  723. Current rho iterate with shape ``(M,)``.
  724. --- maxiter: int
  725. Maximum number of fixed-point iterations.
  726. --- th: float
  727. Convergence threshold on ``||rho_{k+1} - rho_k||``.
  728. Returns:
  729. --- tuple
  730. ``(rho, err)`` where ``rho`` is the final iterate and ``err`` is
  731. the final norm difference between consecutive iterates.
  732. """
  733. TMat = np.transpose(TMat, (0, 2, 3, 1))
  734. wVec = np.zeros((Y, A, M))
  735. for y in range(Y):
  736. for a in range(A):
  737. for m in range(M):
  738. wVec[y, a, m] = np.sum(TMat[y, m, :, a])
  739. for _ in range(maxiter):
  740. wsumexp_test_k = np.zeros((Y, A))
  741. for y in range(Y):
  742. for a in range(A):
  743. wsumexp_test_k[y, a] = np.sum(wVec[y, a] * rhok)
  744. grad = wVec * rhok / wsumexp_test_k[..., None]
  745. rhok_new = np.zeros(M)
  746. for y in range(Y):
  747. for a in range(A):
  748. rhok_new += pya[y, a] * grad[y, a]
  749. if np.linalg.norm(rhok_new - rhok) < th:
  750. break
  751. rhok = rhok_new
  752. return rhok, np.linalg.norm(rhok_new - rhok)

inference.py at commit 219eb08, no license · at the source

Overview

  1. Quantitative Life Sciences section, The Abdus Salam International Center for Theoretical Physics (ICTP), Trieste, Italy
  2. Department of Oncology, Università degli Studi di Torino, Italy
Journal: PLoS computational biology, volume 22, issue 8, article e1014585
Dates: received 17 February 2026; accepted 16 July 2026; published online 7 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014585 · PMID 42566485 · PMCID PMC13475988 · OpenAlex W7201868829
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), mouse (organism), rat (organism), cognitive (subfield)
Methods: Machine learning, fMRI & imaging
MeSH: Behavior, Animal*, Decision Making*, Models, Neurological*, Animals, Computational Biology, Computer Simulation, Mice, Rats (* major topic)
Topic: Zebrafish Biomedical Research Applications (Cell Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: not cited yet (Europe PMC); 62 references in the paper

Abstract

Understanding how living organisms process sensory information from their surroundings and translate it into decisions is a fundamental problem across biological scales – from biochemical signalling in single-cells to neural computations in animal brains. In this work, we address this challenge by introducing a method to reconstruct general decision processes directly from behavioral observations alone. Our approach is applicable to any biological agent and does not require prior knowledge of its internal mechanisms or its environment. Our agent model is defined by a recurrent dynamics over a discrete set of internal states which encode and process sensory information, and dictate which actions to execute. We validate our method on synthetic agents and demonstrate that we can exactly recover the agent’s behavior for non-trivial tasks. Then, we infer agent models from experimental data of rats performing evidence accumulation and of mice making decisions under uncertainty and in changing environments. In both cases, very few internal states suffice to reproduce the observed behavior with high accuracy. Crucially, the immediate interpretability of the inferred dynamics allows to understand the computational process underlying decision-making.

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

Repository

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

giorgionicoletti/FSC-inference-MAPSO

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 219eb08e8f79777470dbcd0896eaa035cb693c2f, 15 May 2026
Languages: Python (13), Jupyter (5)
Size: 260 files, 18 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, 5 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (16 files), Numba (9 files), Matplotlib (8 files), PyTorch (5 files), pandas (3 files), NetworkX (1 file), SciPy (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
19 files

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

Tracing map

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

What the map holds:

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

The code to infer finite state controllers from behavioral trajectories is available at https://github.com/giorgionicoletti/FSC-inference-MAPSO.

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, 2 authors, 8 MeSH terms, 51 references.

Cite

This paper

Nicoletti, G., & Celani, A. (2026). Decoding behavior with minimal and interpretable agent models. PLoS computational biology, 22(8), e1014585. https://doi.org/10.1371/journal.pcbi.1014585

BibTeX

@article{nicoletti2026decoding,
author = {Nicoletti, Giorgio and Celani, Antonio},
title = {{Decoding behavior with minimal and interpretable agent models}},
journal = {PLoS computational biology},
year = {2026},
month = aug,
volume = {22},
number = {8},
pages = {e1014585},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014585},
url = {https://doi.org/10.1371/journal.pcbi.1014585},
pmid = {42566485},
pmcid = {PMC13475988}
}

RIS

TY - JOUR
AU - Nicoletti, Giorgio
AU - Celani, Antonio
TI - Decoding behavior with minimal and interpretable agent models
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/08/07
VL - 22
IS - 8
SP - e1014585
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014585
UR - https://doi.org/10.1371/journal.pcbi.1014585
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014585",
"type": "article-journal",
"title": "Decoding behavior with minimal and interpretable agent models",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Nicoletti",
"given": "Giorgio"
},
{
"family": "Celani",
"given": "Antonio"
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "8",
"page": "e1014585",
"DOI": "10.1371/journal.pcbi.1014585",
"PMID": "42566485",
"PMCID": "PMC13475988",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014585",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
7
]
]
}
}

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-75924-7 [code]
Data-driven reduced modeling of neural dynamics.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, computational modeling (no new data), 6 references
[2] doi:10.3390/biomimetics11080569 [code]
Pretraining of Embodied Recurrent Networks Bridges the Gap Between Artificial and Cortical Neural Activities.
Journal: Biomimetics (Basel, Switzerland)
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 5 references
[3] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: Numba, PyTorch, pandas, 3 other tools, 3 references
[4] doi:10.1371/journal.pbio.3003831 [code]
Disinhibitory signaling enables flexible coding of top-down information in cortical networks.
Journal: PLoS biology
In common: NetworkX, PyTorch, pandas, 3 other tools, mouse, 3 references
[5] doi:10.7554/elife.109313 [code]
Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making.
Journal: eLife
In common: Numba, PyTorch, pandas, 3 other tools, computational modeling (no new data), cognitive, mouse, 1 reference
[6] doi:10.1016/j.celrep.2026.117262 [code]
Cerebellar acceleration of learning in an evidence-accumulation task.
Journal: Cell reports
In common: pandas, SciPy, Matplotlib, 1 other tool, cognitive, mouse, 3 references
[7] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: Numba, NetworkX, PyTorch, 4 other tools, mouse, 1 reference
[8] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: Numba, NetworkX, PyTorch, 4 other tools, mouse, 1 reference
[9] doi:10.3389/fnsys.2026.1822122 [code]
Convergence-divergence circuits for multimodal integration of innate and learned opponent valences.
Journal: Frontiers in systems neuroscience
In common: Numba, NetworkX, PyTorch, 4 other tools, 1 reference
[10] doi:10.1371/journal.pcbi.1014162 [code]
Exploring neural manifolds across a wide range of intrinsic dimensions.
Journal: PLoS computational biology
In common: pandas, SciPy, Matplotlib, 1 other tool, 4 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.