Skip to content

Commit e55d40e

Browse files
committed
add make_dataset function for constructing RecallDataset structures from sparse inputs
1 parent c77d647 commit e55d40e

16 files changed

Lines changed: 193 additions & 337 deletions

‎jaxcmr/helpers.py‎

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@
3333
"find_max_list_length",
3434
"apply_by_subject",
3535
"has_repeats_per_row",
36+
"make_dataset",
3637
]
3738

3839

@@ -260,3 +261,119 @@ def check_row(row):
260261
return jnp.any(valid_repeats)
261262

262263
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

‎tests/test_backreplagrank.py‎

Lines changed: 2 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -13,23 +13,7 @@
1313
plot_back_rep_lagrank,
1414
)
1515
from jaxcmr.analyses.replagrank import replagrank
16-
from jaxcmr.typing import RecallDataset
17-
18-
19-
def _make_dataset(recalls, presentations, subjects=None) -> RecallDataset:
20-
recalls = jnp.array(recalls)
21-
presentations = jnp.array(presentations)
22-
n = recalls.shape[0]
23-
if subjects is None:
24-
subjects = jnp.zeros(n, dtype=int)
25-
else:
26-
subjects = jnp.array(subjects)
27-
return {
28-
"recalls": recalls,
29-
"pres_itemnos": presentations,
30-
"subject": subjects,
31-
"listLength": jnp.full(n, presentations.shape[1]),
32-
}
16+
from jaxcmr.helpers import make_dataset
3317

3418

3519
def _rep_dataset():
@@ -45,7 +29,7 @@ def _rep_dataset():
4529
[3, 1, 2, 4, 5, 0, 0, 0],
4630
[1, 3, 5, 2, 0, 0, 0, 0],
4731
])
48-
return _make_dataset(recalls, pres, subjects=[0, 0, 1, 1])
32+
return make_dataset(recalls, pres, subject=jnp.array([0, 0, 1, 1]))
4933

5034

5135
class TestReversalChangesResults:

‎tests/test_compound_cueing_crp.py‎

Lines changed: 4 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -22,29 +22,7 @@
2222
tabulate_trial,
2323
compound_cueing_crp,
2424
)
25-
from jaxcmr.typing import RecallDataset
26-
27-
28-
# -----------------------------------------------------------------------------
29-
# Helpers
30-
# -----------------------------------------------------------------------------
31-
32-
33-
def _make_dataset(
34-
recalls: jnp.ndarray,
35-
presentations: jnp.ndarray,
36-
) -> RecallDataset:
37-
"""Return minimal dataset keyed for the compound cueing analysis."""
38-
recalls = jnp.asarray(recalls, dtype=jnp.int32)
39-
presentations = jnp.asarray(presentations, dtype=jnp.int32)
40-
n_trials, _ = recalls.shape
41-
list_length = presentations.shape[1]
42-
return {
43-
"subject": jnp.ones((n_trials, 1), dtype=jnp.int32),
44-
"listLength": jnp.full((n_trials, 1), list_length, dtype=jnp.int32),
45-
"pres_itemnos": presentations,
46-
"recalls": recalls,
47-
}
25+
from jaxcmr.helpers import make_dataset
4826

4927

5028
# -----------------------------------------------------------------------------
@@ -421,7 +399,7 @@ def test_crp_computation_single_trial():
421399
# Arrange
422400
presentation = jnp.array([[1, 2, 3, 4, 5, 6, 7, 8, 9, 3, 10, 11]], dtype=jnp.int32)
423401
trial = jnp.array([[1, 2, 3, 0, 0, 0, 0, 0, 0, 0, 0, 0]], dtype=jnp.int32)
424-
dataset = _make_dataset(trial, presentation)
402+
dataset = make_dataset(trial, presentation)
425403

426404
# Act
427405
result = compound_cueing_crp(dataset, min_spacing=6, size=2)
@@ -447,7 +425,7 @@ def test_crp_nan_when_no_opportunities():
447425
# Arrange
448426
presentation = jnp.array([[1, 2, 3, 4, 5, 6, 7, 8]], dtype=jnp.int32)
449427
trial = jnp.array([[1, 2, 3, 4, 5, 6, 7, 8]], dtype=jnp.int32)
450-
dataset = _make_dataset(trial, presentation)
428+
dataset = make_dataset(trial, presentation)
451429

452430
# Act
453431
result = compound_cueing_crp(dataset, min_spacing=6, size=2)
@@ -483,7 +461,7 @@ def test_crp_aggregates_across_trials():
483461
],
484462
dtype=jnp.int32,
485463
)
486-
dataset = _make_dataset(trials, presentation)
464+
dataset = make_dataset(trials, presentation)
487465

488466
# Act
489467
result = compound_cueing_crp(dataset, min_spacing=6, size=2)

‎tests/test_crp.py‎

Lines changed: 6 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from matplotlib.figure import Figure
88
import pytest
99
from jaxcmr.analyses import crp
10+
from jaxcmr.helpers import make_dataset
1011
from jaxcmr.typing import RecallDataset
1112

1213

@@ -28,23 +29,6 @@ def _lag_index(lags_len: int, lag: int) -> int:
2829
center = lags_len // 2
2930
return center + lag
3031

31-
def _make_dataset(
32-
recalls: jnp.ndarray,
33-
presentations: jnp.ndarray,
34-
) -> RecallDataset:
35-
"""Return minimal dataset keyed for the CRP analysis."""
36-
37-
recalls = jnp.asarray(recalls, dtype=jnp.int32)
38-
presentations = jnp.asarray(presentations, dtype=jnp.int32)
39-
n_trials, _ = recalls.shape
40-
list_length = presentations.shape[1]
41-
return {
42-
"subject": jnp.ones((n_trials, 1), dtype=jnp.int32),
43-
"listLength": jnp.full((n_trials, 1), list_length, dtype=jnp.int32),
44-
"pres_itemnos": presentations,
45-
"recalls": recalls,
46-
} # type: ignore
47-
4832
# -----------------------------------------------------------------------------
4933
# set_false_at_index tests
5034
# -----------------------------------------------------------------------------
@@ -444,8 +428,8 @@ def test_produces_same_crp_when_repeats_removed():
444428
presentations = jnp.array([[1, 2, 3, 4]], dtype=jnp.int32)
445429

446430
# Act / When
447-
dataset_with_repeat = _make_dataset(trials_with_repeat, presentations)
448-
dataset_without_repeat = _make_dataset(trials_without_repeat, presentations)
431+
dataset_with_repeat = make_dataset(trials_with_repeat, presentations)
432+
dataset_without_repeat = make_dataset(trials_without_repeat, presentations)
449433

450434
with_repeat = crp.crp(dataset_with_repeat, size=1)
451435
without_repeat = crp.crp(dataset_without_repeat, size=1)
@@ -471,7 +455,7 @@ def test_returns_array_when_single_trial_item():
471455
presentations = jnp.array([[1]], dtype=jnp.int32)
472456

473457
# Act / When
474-
dataset = _make_dataset(trials, presentations)
458+
dataset = make_dataset(trials, presentations)
475459
out = crp.crp(dataset, size=1)
476460

477461
# Assert / Then
@@ -495,7 +479,7 @@ def test_matches_uncompiled_result_when_jitted_with_static_argnames():
495479
trials = jnp.array([[1, 2, 3], [1, 3, 2]], dtype=jnp.int32)
496480
presentations = jnp.array([[1, 2, 3], [1, 2, 3]], dtype=jnp.int32)
497481
size = 1
498-
dataset = _make_dataset(trials, presentations)
482+
dataset = make_dataset(trials, presentations)
499483
expected = crp.crp(dataset, size)
500484

501485
# Act / When
@@ -523,7 +507,7 @@ def test_runs_with_larger_size_when_jitted():
523507
trials = jnp.array([[1, 2, 3]], dtype=jnp.int32)
524508
presentations = jnp.array([[1, 2, 3]], dtype=jnp.int32)
525509
size = 2
526-
dataset = _make_dataset(trials, presentations)
510+
dataset = make_dataset(trials, presentations)
527511
expected = crp.crp(dataset, size)
528512

529513
# Act / When

‎tests/test_cue_centered_lagrank.py‎

Lines changed: 5 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -16,28 +16,7 @@
1616
test_cue_centered_lagrank as run_test,
1717
)
1818
from jaxcmr.analyses.lagrank import LagRankTestResult
19-
from jaxcmr.typing import RecallDataset
20-
21-
22-
def _make_dataset(recalls, presentations, cues, should_tabulate, subjects=None) -> RecallDataset:
23-
"""Wrap arrays into a RecallDataset dict with cue fields."""
24-
recalls = jnp.array(recalls)
25-
presentations = jnp.array(presentations)
26-
cues = jnp.array(cues)
27-
should_tab = jnp.array(should_tabulate, dtype=bool)
28-
n = recalls.shape[0]
29-
if subjects is None:
30-
subjects = jnp.zeros(n, dtype=int)
31-
else:
32-
subjects = jnp.array(subjects)
33-
return {
34-
"recalls": recalls,
35-
"pres_itemnos": presentations,
36-
"cue_clips": cues,
37-
"_should_tabulate": should_tab,
38-
"subject": subjects,
39-
"listLength": jnp.full(n, presentations.shape[1]),
40-
}
19+
from jaxcmr.helpers import make_dataset
4120

4221

4322
# ---- Tabulation: no cue → skipped ----
@@ -98,7 +77,7 @@ def dataset(self):
9877
[3, 2, 1, 0],
9978
])
10079
should_tab = recalls > 0
101-
return _make_dataset(recalls, pres, cues, should_tab)
80+
return {**make_dataset(recalls, pres), "cue_clips": cues, "_should_tabulate": should_tab}
10281

10382
def test_returns_1d(self, dataset):
10483
"""cue_centered_lagrank returns (n_trials,)."""
@@ -140,7 +119,7 @@ def test_subject_shape(self):
140119
])
141120
should_tab = recalls > 0
142121
subjects = jnp.array([0, 0, 1, 1])
143-
dataset = _make_dataset(recalls, pres, cues, should_tab, subjects)
122+
dataset = {**make_dataset(recalls, pres, subject=subjects), "cue_clips": cues, "_should_tabulate": should_tab}
144123
mask = jnp.ones(4, dtype=bool)
145124
result = subject_cue_centered_lagrank(dataset, mask, size=1)
146125
assert result.shape == (2,)
@@ -178,7 +157,7 @@ def test_plot_returns_axes(self):
178157
[1, 2, 0, 0],
179158
])
180159
should_tab = recalls > 0
181-
dataset = _make_dataset(recalls, pres, cues, should_tab)
160+
dataset = {**make_dataset(recalls, pres), "cue_clips": cues, "_should_tabulate": should_tab}
182161
mask = jnp.ones(2, dtype=bool)
183162
ax = plot_cue_centered_lagrank(
184163
dataset, mask, should_tabulate=should_tab, size=1, labels=["Test"]
@@ -206,7 +185,7 @@ def test_jit_compatible(self):
206185
[1, 2, 0, 0],
207186
])
208187
should_tab = recalls > 0
209-
dataset = _make_dataset(recalls, pres, cues, should_tab)
188+
dataset = {**make_dataset(recalls, pres), "cue_clips": cues, "_should_tabulate": should_tab}
210189
result_nojit = cue_centered_lagrank(dataset, size=1)
211190
result_jit = jit(cue_centered_lagrank, static_argnames=("size",))(
212191
dataset, size=1

0 commit comments

Comments
 (0)