Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 24 additions & 1 deletion .github/scripts/smoke_install.sh
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,14 @@
set -e

pip install ./PyAutoNerves ./PyAutoFit ./PyAutoArray ./PyAutoGalaxy
pip install "jax<0.7" "jaxlib<0.7"
# NOTE: no jax pin here, deliberately. One lived on this line until it was
# found to be vestigial: autolens_workspace_test#82 added
# `jax<0.7 jaxlib<0.7` solely to keep `tensorflow-probability==0.25.0`
# importable, and autolens_workspace_test#184 removed that
# dependency when the stack moved to tfp-nightly. jax's supported range is
# owned by autonerves' base dependencies (jax>=0.7.0,<0.12.0, PyAutoLens#702)
# and is installed by the line above -- do not restate it here, it drifts.
# The assertion at the end of this script checks the resolved version.
pip install "./PyAutoArray[optional]" "./PyAutoGalaxy[optional]"
# NOTE: do NOT `pip install tensorflow-probability==0.25.0` here. The stable
# release crashes at import under the resolved modern JAX
Expand All @@ -20,3 +27,19 @@ pip install "./PyAutoArray[optional]" "./PyAutoGalaxy[optional]"
# last time so site-packages has skip_latents() and other recent
# autonerves APIs available at import time.
pip install --force-reinstall --no-deps ./PyAutoNerves

# Assert the resolved jax rather than inferring it from a green smoke run. This
# install previously landed on a supported jax only because the [optional]
# re-resolution above happened to undo the pin removed here -- correct by line
# ordering, not by constraint. Reordering these lines now fails loudly at install
# time instead of silently dropping the smoke suite onto an unsupported jax.
python - <<'JAXCHECK'
import jax

major, minor = (int(part) for part in jax.__version__.split(".")[:2])
assert (0, 7) <= (major, minor) < (0, 12), (
f"resolved jax {jax.__version__} is outside autonerves' supported range "
"(>=0.7.0,<0.12.0)"
)
print(f"resolved jax {jax.__version__}")
JAXCHECK
Loading