Skip to content

Restore Python 3.9 compatibility - #49

Merged
msperryucsd merged 1 commit into
mainfrom
fix/py39-compat
Sep 29, 2026
Merged

msperryucsd merged 1 commit into
mainfrom
fix/py39-compat

Conversation

@HgXe

@HgXe HgXe commented Sep 29, 2026

Copy link
Copy Markdown
Collaborator

The Python 3.9 CI job has been failing at pip install. This PR fixes that failure and a JAX incompatibility that would have come up right after it. It also fixes two related issues found along the way.

Changes

sphinx-collections requires Python 3.10+. requirements.txt installed the unpinned fork git+https://github.com/anugrahjo/sphinx-collections.git. That fork recently merged upstream v0.3.1, which declares requires-python >=3.10. On 3.9, pip fails with requires a different Python: 3.9.25 not in '>=3.10'. The package is only used for building the docs, so it now has a python_version >= "3.10" marker. ReadTheDocs builds on 3.11 and still installs it.

pure_callback(..., vmap_method="sequential") fails on older JAX. Newer JAX needs this argument for custom ops (and other callbacks) inside a vmap. Older JAX, including 0.4.30 (the newest release for Python 3.9), doesn't accept it and fails with TypeError: Value 'sequential' ... is not a valid JAX type. Older JAX already runs callbacks sequentially under vmap by default. The new helper sequential_pure_callback in csdl_alpha/backends/jax/utils.py passes vmap_method only when jax.pure_callback accepts it. All four call sites use it:

  • CustomExplicitOperation.compute_jax
  • CustomJacOperation.compute_jax
  • SubOperation.compute_jax
  • fallback_to_inline_jax

Docs extension renamed. The same upstream merge removed the sphinxcontrib.collections namespace package, so docs/conf.py now loads sphinx_collections. The collections config and the writer_function driver are unchanged.

pytest 9 compatibility. conftest.py declared the path argument in pytest_collect_file, which pytest 9 removed. file_path was already being used, so path is dropped.

Testing

A custom explicit op inside a BLoop (jax.vmap), including its derivative, through the JAX backend:

Environment Before After
Python 3.10, JAX 0.4.29 TypeError (vmap_method) passes
Python 3.9, JAX 0.4.30 — passes
Python 3.10, JAX 0.6.2 — passes; fails with NotImplementedError if vmap_method is left out

Full suite, installed the same way CI installs it:

  • Python 3.9 / JAX 0.4.30: 302/302 on the default backend; 301/302 with --backend jax.
  • Python 3.10 / JAX 0.6.2: 301/302 and 300/302.

The failures are TestProduct::test_functionality (--backend jax) and test_docstrings (3.10 only). Both also fail on main locally.

  • Test collection gives 302 tests on both pytest 7.3.1 and 9.1.1.
  • A docs build from a fresh requirements.txt install on Python 3.10 succeeds, and the collections steps run. sphinxcontrib.collections no longer imports with the current fork.

- Skip the sphinx-collections fork on Python < 3.10; it now requires >=3.10
  and is only needed for docs builds.
- Add sequential_pure_callback, which passes vmap_method="sequential" only
  when the installed JAX supports it. Older JAX (the newest available on
  Python 3.9) rejects that argument but is sequential under vmap by default.
- Load the docs extension as sphinx_collections; the fork dissolved the
  sphinxcontrib.collections namespace package.
- Drop the py.path argument from pytest_collect_file; pytest 9 removed it.
@HgXe
HgXe requested a review from msperryucsd September 29, 2026 22:41

@msperryucsd msperryucsd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good job

@msperryucsd
msperryucsd merged commit a69686b into main Sep 29, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants