Skip to content
Open
Show file tree
Hide file tree
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
2 changes: 1 addition & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
8 changes: 7 additions & 1 deletion docs/explanation/architecture.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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), ...)
```

Expand Down
9 changes: 3 additions & 6 deletions docs/explanation/performance.md
Original file line number Diff line number Diff line change
Expand Up @@ -31,19 +31,15 @@ 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
# Prefer: every JAX-array-bearing object is an argument.
@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
Expand Down Expand Up @@ -83,6 +79,7 @@ set the report environment variable before the first compile:

```python
import os

os.environ["JAX_CAPTURED_CONSTANTS_REPORT_FRAMES"] = "-1"
```

Expand Down
22 changes: 16 additions & 6 deletions docs/explanation/rate_vs_readout.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
```

Expand Down Expand Up @@ -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=...,
)
```

Expand Down
15 changes: 11 additions & 4 deletions docs/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
)
```

Expand Down
6 changes: 4 additions & 2 deletions docs/installation.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -81,6 +82,7 @@ get the workspace-editable versions.

```python
import coronagraphoto

print(coronagraphoto.__version__)
```

Expand Down
Loading