diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index f0ac892..798d414 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -15,7 +15,7 @@ repos: args: [--pytest-test-first] - id: end-of-file-fixer - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.16.2 + rev: v0.16.9 hooks: - id: ruff-check args: [--fix] diff --git a/docs/explanation/architecture.md b/docs/explanation/architecture.md index afc4ab8..e0f9ace 100644 --- a/docs/explanation/architecture.md +++ b/docs/explanation/architecture.md @@ -99,11 +99,15 @@ sibling libraries' constructors, then pass them in: ```python # Scene-building lives in skyscapes from skyscapes import from_exovista + scene = from_exovista("path/to/exovista_system.fits") # Optical-path building lives in optixstuff from optixstuff import ( - OpticalPath, SimplePrimary, ConstantThroughput, IdealDetector, + OpticalPath, + SimplePrimary, + ConstantThroughput, + IdealDetector, ) from yippy import EqxCoronagraph @@ -116,8 +120,10 @@ optical_path = OpticalPath( # coronagraphoto is what turns these into a 2D image from coronagraphoto import system_rate, system_readout + rate_map = system_rate(scene, optical_path, ...) import jax + image = system_readout(scene, optical_path, jax.random.PRNGKey(0), ...) ``` diff --git a/docs/explanation/performance.md b/docs/explanation/performance.md index 03ff3ac..95fc892 100644 --- a/docs/explanation/performance.md +++ b/docs/explanation/performance.md @@ -31,9 +31,7 @@ array rather than being memcopied into the compiled binary. @eqx.filter_jit def simulate_frame(mjd, wavelength_nm, key): rate = planet_rate(planet, optical_path, ...) - return optical_path.detector.readout_source_electrons( - rate, EXPOSURE_S, key - ) + return optical_path.detector.readout_source_electrons(rate, EXPOSURE_S, key) ``` ```python @@ -41,9 +39,7 @@ def simulate_frame(mjd, wavelength_nm, key): @eqx.filter_jit def simulate_frame(optical_path, planet, mjd, wavelength_nm, key): rate = planet_rate(planet, optical_path, ...) - return optical_path.detector.readout_source_electrons( - rate, EXPOSURE_S, key - ) + return optical_path.detector.readout_source_electrons(rate, EXPOSURE_S, key) ``` ## Measured impact @@ -83,6 +79,7 @@ set the report environment variable before the first compile: ```python import os + os.environ["JAX_CAPTURED_CONSTANTS_REPORT_FRAMES"] = "-1" ``` diff --git a/docs/explanation/rate_vs_readout.md b/docs/explanation/rate_vs_readout.md index 9e254d4..ae772de 100644 --- a/docs/explanation/rate_vs_readout.md +++ b/docs/explanation/rate_vs_readout.md @@ -58,19 +58,23 @@ import jax import equinox as eqx from coronagraphoto import system_rate + # Differentiable rate pipeline -- gradient flows end-to-end @eqx.filter_jit def total_electrons(wavelength_nm, scene, optical_path): rate_map = system_rate( - scene, optical_path, + scene, + optical_path, start_time_jd=2_460_000.0, wavelength_nm=wavelength_nm, bin_width_nm=50.0, telescope_pa_deg=0.0, - ecliptic_lat_deg=0.0, solar_lon_deg=135.0, + ecliptic_lat_deg=0.0, + solar_lon_deg=135.0, ) return rate_map.sum() * EXPOSURE_S + grad_fn = eqx.filter_grad(total_electrons) ``` @@ -101,10 +105,16 @@ import jax key = jax.random.PRNGKey(0) image = system_readout( - scene, optical_path, key, - start_time_jd=..., exposure_time_s=..., wavelength_nm=..., - bin_width_nm=..., telescope_pa_deg=..., - ecliptic_lat_deg=..., solar_lon_deg=..., + scene, + optical_path, + key, + start_time_jd=..., + exposure_time_s=..., + wavelength_nm=..., + bin_width_nm=..., + telescope_pa_deg=..., + ecliptic_lat_deg=..., + solar_lon_deg=..., ) ``` diff --git a/docs/index.md b/docs/index.md index ec40a5b..8aa92f2 100644 --- a/docs/index.md +++ b/docs/index.md @@ -56,7 +56,10 @@ post-processing, and analysis live in sibling libraries. import jax from coronagraphoto import system_readout from optixstuff import ( - OpticalPath, SimplePrimary, IdealDetector, ConstantThroughput, + OpticalPath, + SimplePrimary, + IdealDetector, + ConstantThroughput, ) from skyscapes import from_exovista from yippy import EqxCoronagraph @@ -69,12 +72,16 @@ optical_path = OpticalPath( detector=IdealDetector(pixel_scale_arcsec=0.01, shape=(512, 512)), ) image = system_readout( - scene, optical_path, jax.random.PRNGKey(0), + scene, + optical_path, + jax.random.PRNGKey(0), start_time_jd=2_460_000.0, exposure_time_s=3600.0, - wavelength_nm=550.0, bin_width_nm=50.0, + wavelength_nm=550.0, + bin_width_nm=50.0, telescope_pa_deg=0.0, - ecliptic_lat_deg=0.0, solar_lon_deg=135.0, + ecliptic_lat_deg=0.0, + solar_lon_deg=135.0, ) ``` diff --git a/docs/installation.md b/docs/installation.md index 363ef47..71566bf 100644 --- a/docs/installation.md +++ b/docs/installation.md @@ -39,8 +39,9 @@ loaded: ```python import jax -print(jax.default_backend()) # 'gpu' -print(jax.config.read("jax_enable_x64")) # True if you enabled it + +print(jax.default_backend()) # 'gpu' +print(jax.config.read("jax_enable_x64")) # True if you enabled it ``` If `default_backend()` reports `gpu` but the CUDA plugin warned about @@ -81,6 +82,7 @@ get the workspace-editable versions. ```python import coronagraphoto + print(coronagraphoto.__version__) ```