|
33 | 33 | "find_max_list_length", |
34 | 34 | "apply_by_subject", |
35 | 35 | "has_repeats_per_row", |
| 36 | + "make_dataset", |
36 | 37 | ] |
37 | 38 |
|
38 | 39 |
|
@@ -260,3 +261,119 @@ def check_row(row): |
260 | 261 | return jnp.any(valid_repeats) |
261 | 262 |
|
262 | 263 | return vmap(check_row)(arr) |
| 264 | + |
| 265 | + |
| 266 | +def make_dataset( |
| 267 | + recalls: Integer[Array, "n_trials num_recalled"], |
| 268 | + pres_itemnos: Integer[Array, "n_trials num_presented"] | None = None, |
| 269 | + listLength: Integer[Array, "n_trials 1"] | int | None = None, |
| 270 | + subject: Integer[Array, "n_trials 1"] | int = 1, |
| 271 | + *, |
| 272 | + reps: int = 1, |
| 273 | +) -> RecallDataset: |
| 274 | + """Construct a ``RecallDataset`` from flexible inputs. |
| 275 | +
|
| 276 | + Parameters |
| 277 | + ---------- |
| 278 | + recalls |
| 279 | + Within-list recall events (1-indexed, 0 = padding). A 1-D vector is |
| 280 | + treated as a single trial. |
| 281 | + pres_itemnos |
| 282 | + Within-list presentation items. When omitted, generated as |
| 283 | + ``arange(1, list_length + 1)`` per trial. |
| 284 | + listLength |
| 285 | + List length per trial. Inferred from *pres_itemnos* when omitted. |
| 286 | + When both are provided, an assertion checks compatibility. |
| 287 | + When neither is provided, inferred as |
| 288 | + ``max(recalls.shape[1], recalls.max())``. |
| 289 | + subject |
| 290 | + Subject identifier per trial. Defaults to ``1`` (single subject). |
| 291 | + reps |
| 292 | + Number of times to tile the resulting trials along axis 0. |
| 293 | +
|
| 294 | + Returns |
| 295 | + ------- |
| 296 | + RecallDataset |
| 297 | + Dictionary with keys ``recalls``, ``pres_itemnos``, ``listLength``, |
| 298 | + and ``subject``, all shaped ``(n_trials * reps, ...)``. |
| 299 | +
|
| 300 | + """ |
| 301 | + # -- normalise recalls (required) ----------------------------------------- |
| 302 | + recalls_arr = jnp.atleast_2d(jnp.asarray(recalls, dtype=jnp.int32)) |
| 303 | + |
| 304 | + # -- normalise pres_itemnos (optional) ------------------------------------ |
| 305 | + pres_arr: jnp.ndarray | None = None |
| 306 | + if pres_itemnos is not None: |
| 307 | + pres_arr = jnp.atleast_2d(jnp.asarray(pres_itemnos, dtype=jnp.int32)) |
| 308 | + |
| 309 | + # -- normalise listLength (optional) -------------------------------------- |
| 310 | + ll_arr: jnp.ndarray | None = None |
| 311 | + if listLength is not None: |
| 312 | + if isinstance(listLength, int): |
| 313 | + ll_arr = jnp.array([[listLength]], dtype=jnp.int32) |
| 314 | + else: |
| 315 | + ll_arr = jnp.asarray(listLength, dtype=jnp.int32).reshape(-1, 1) |
| 316 | + |
| 317 | + # -- normalise subject ---------------------------------------------------- |
| 318 | + if isinstance(subject, int): |
| 319 | + subj_arr = jnp.array([[subject]], dtype=jnp.int32) |
| 320 | + else: |
| 321 | + subj_arr = jnp.asarray(subject, dtype=jnp.int32).reshape(-1, 1) |
| 322 | + |
| 323 | + # -- resolve n_trials ----------------------------------------------------- |
| 324 | + sizes: list[int] = [] |
| 325 | + for arr in (recalls_arr, pres_arr, ll_arr, subj_arr): |
| 326 | + if arr is not None and arr.shape[0] > 1: |
| 327 | + sizes.append(arr.shape[0]) |
| 328 | + if sizes: |
| 329 | + n_trials = sizes[0] |
| 330 | + assert all(s == n_trials for s in sizes), ( |
| 331 | + f"Multi-trial args disagree on n_trials: {sizes}" |
| 332 | + ) |
| 333 | + else: |
| 334 | + n_trials = 1 |
| 335 | + |
| 336 | + # -- tile single-trial args to n_trials ----------------------------------- |
| 337 | + if recalls_arr.shape[0] == 1 and n_trials > 1: |
| 338 | + recalls_arr = jnp.tile(recalls_arr, (n_trials, 1)) |
| 339 | + if pres_arr is not None and pres_arr.shape[0] == 1 and n_trials > 1: |
| 340 | + pres_arr = jnp.tile(pres_arr, (n_trials, 1)) |
| 341 | + if ll_arr is not None and ll_arr.shape[0] == 1 and n_trials > 1: |
| 342 | + ll_arr = jnp.tile(ll_arr, (n_trials, 1)) |
| 343 | + if subj_arr.shape[0] == 1 and n_trials > 1: |
| 344 | + subj_arr = jnp.tile(subj_arr, (n_trials, 1)) |
| 345 | + |
| 346 | + # -- infer / validate list_length ----------------------------------------- |
| 347 | + if pres_arr is not None and ll_arr is not None: |
| 348 | + assert jnp.all(ll_arr == pres_arr.shape[1]), ( |
| 349 | + f"listLength ({ll_arr.ravel()}) incompatible with " |
| 350 | + f"pres_itemnos width ({pres_arr.shape[1]})" |
| 351 | + ) |
| 352 | + list_length = int(pres_arr.shape[1]) |
| 353 | + elif pres_arr is not None: |
| 354 | + list_length = int(pres_arr.shape[1]) |
| 355 | + ll_arr = jnp.full((n_trials, 1), list_length, dtype=jnp.int32) |
| 356 | + elif ll_arr is not None: |
| 357 | + list_length = int(ll_arr[0, 0]) |
| 358 | + else: |
| 359 | + list_length = int(max(recalls_arr.shape[1], jnp.max(recalls_arr))) |
| 360 | + ll_arr = jnp.full((n_trials, 1), list_length, dtype=jnp.int32) |
| 361 | + |
| 362 | + # -- generate default pres_itemnos ---------------------------------------- |
| 363 | + if pres_arr is None: |
| 364 | + row = jnp.arange(1, list_length + 1, dtype=jnp.int32) |
| 365 | + pres_arr = jnp.tile(row[None, :], (n_trials, 1)) |
| 366 | + |
| 367 | + # -- apply reps ----------------------------------------------------------- |
| 368 | + if reps > 1: |
| 369 | + recalls_arr = jnp.tile(recalls_arr, (reps, 1)) |
| 370 | + pres_arr = jnp.tile(pres_arr, (reps, 1)) |
| 371 | + ll_arr = jnp.tile(ll_arr, (reps, 1)) |
| 372 | + subj_arr = jnp.tile(subj_arr, (reps, 1)) |
| 373 | + |
| 374 | + return { |
| 375 | + "recalls": recalls_arr, |
| 376 | + "pres_itemnos": pres_arr, |
| 377 | + "listLength": ll_arr, |
| 378 | + "subject": subj_arr, |
| 379 | + } # type: ignore |
0 commit comments