Skip to content

Detect sequential Python loops over an embarrassingly-parallel JAX call (missed vmap/scan opportunity) #1

Description

@maraxen

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions