Motivating instance (found 2026-08-30, biosensors project)
scripts/analysis/cimist_pilot.py's run_pooled_condition implements its n_inits>1 multi-seed-init path as a plain Python loop:
for extra_seed in range(base_seed + 1, base_seed + n_inits):
... # calls run_demistify_pipeline (a JAX/JIT pipeline) once per seed
The function's own docstring admits this is "linear-in-n_inits wall time" by construction — each iteration reruns a JIT pipeline sequentially instead of vmapping/scanning over the seed axis. A sibling incident the same week had a researcher farm the identical seed axis out as an 80-task SLURM array instead of an in-process vmap (the seed axis is a textbook case for xtrax's AxisSpec/BatchPlanner composition, already used in denxity's plan_loo_axis/plan_pairs_axis).
jaxlint has no rule for this direction of the loop-vs-vectorize problem today. JL003 flags a Python for/while loop inside a @jax.jit/@jax.grad function (a tracing/compile hazard) — the opposite case. None of the current rules (JL001/002/003/004/005/010/011/012, JD*, JM*) flag a Python loop outside jit that sequentially re-invokes a jittable/vmappable function across an axis that could instead be vmapped or lax.scanned.
Scope: two tiers
1. Hotfix (cheap, narrow, high-precision) — a new sensor (alongside src/jaxlint/core/sensors/{purity,scatter_gather,shapes,best_practices}.py) that pattern-matches the specific "RNG-seed sweep loop" shape:
- a
for/while loop, NOT inside a jit context,
- whose body (a) calls a function decorated with
@jax.jit or importable from a known-jittable module, (b) constructs/passes a jax.random.PRNGKey/fold_in/split keyed on the loop variable as the only varying argument, and (c) accumulates each iteration's return value via .append()/list comprehension with no other cross-iteration state read (i.e. provably embarrassingly parallel, not a real scan recurrence).
This narrow shape is exactly the cimist_pilot.py pattern and should be mechanically detectable via AST without full dataflow analysis. Suggest a new rule ID in the JL0xx family (e.g. JL020) at warning severity given the false-positive risk on genuinely path-dependent loops.
2. Systematic feature (larger scope, needs its own design pass) — a general "missed-vectorization" detector that identifies any Python loop over an axis whose iterations are shape-invariant and vary only a leaf argument (not just PRNG keys — any scalar/array axis: replicate index, condition, hyperparameter sweep value), and suggests the loop can be replaced by a vmap/lax.scan/xtrax AxisSpec-driven tiled_vmap. Open design questions before implementing:
- Should jaxlint special-case xtrax-aware suggestions (project-specific), or stay framework-agnostic ("this loop looks vmappable" only)?
- Needs real dataflow analysis to prove no cross-iteration dependency beyond the accumulator — a loop that reads a variable mutated by a previous iteration (e.g. a running best-fit warm-start) must NOT be flagged.
- False-positive calibration against jaxlint's existing test corpus plus a real-world corpus (demistify, denxity, biosensors'
scripts/analysis/) before enabling by default.
Recommend scoping the dataflow-analysis question as a design spike before writing the tier-2 detector — it's the part most likely to make tier-1's narrow hotfix already provably safe as a byproduct.
Cross-refs
- Tracked as praxia debt #1563 in this repo's own tracker.
- biosensors debt #1558 (the in-project seed-vmap tech debt this gap was discovered alongside).
- biosensors memory
feedback_xtrax_vmap_campaign_variables.md.
Motivating instance (found 2026-08-30, biosensors project)
scripts/analysis/cimist_pilot.py'srun_pooled_conditionimplements itsn_inits>1multi-seed-init path as a plain Python loop:The function's own docstring admits this is "linear-in-n_inits wall time" by construction — each iteration reruns a JIT pipeline sequentially instead of vmapping/scanning over the seed axis. A sibling incident the same week had a researcher farm the identical seed axis out as an 80-task SLURM array instead of an in-process vmap (the seed axis is a textbook case for
xtrax'sAxisSpec/BatchPlannercomposition, already used indenxity'splan_loo_axis/plan_pairs_axis).jaxlint has no rule for this direction of the loop-vs-vectorize problem today.
JL003flags a Pythonfor/whileloop inside a@jax.jit/@jax.gradfunction (a tracing/compile hazard) — the opposite case. None of the current rules (JL001/002/003/004/005/010/011/012,JD*,JM*) flag a Python loop outside jit that sequentially re-invokes a jittable/vmappable function across an axis that could instead be vmapped orlax.scanned.Scope: two tiers
1. Hotfix (cheap, narrow, high-precision) — a new sensor (alongside
src/jaxlint/core/sensors/{purity,scatter_gather,shapes,best_practices}.py) that pattern-matches the specific "RNG-seed sweep loop" shape:for/whileloop, NOT inside a jit context,@jax.jitor importable from a known-jittable module, (b) constructs/passes ajax.random.PRNGKey/fold_in/splitkeyed on the loop variable as the only varying argument, and (c) accumulates each iteration's return value via.append()/list comprehension with no other cross-iteration state read (i.e. provably embarrassingly parallel, not a real scan recurrence).This narrow shape is exactly the
cimist_pilot.pypattern and should be mechanically detectable via AST without full dataflow analysis. Suggest a new rule ID in theJL0xxfamily (e.g.JL020) atwarningseverity given the false-positive risk on genuinely path-dependent loops.2. Systematic feature (larger scope, needs its own design pass) — a general "missed-vectorization" detector that identifies any Python loop over an axis whose iterations are shape-invariant and vary only a leaf argument (not just PRNG keys — any scalar/array axis: replicate index, condition, hyperparameter sweep value), and suggests the loop can be replaced by a vmap/
lax.scan/xtraxAxisSpec-driventiled_vmap. Open design questions before implementing:scripts/analysis/) before enabling by default.Recommend scoping the dataflow-analysis question as a design spike before writing the tier-2 detector — it's the part most likely to make tier-1's narrow hotfix already provably safe as a byproduct.
Cross-refs
feedback_xtrax_vmap_campaign_variables.md.