diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 6a8c407..f82368e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -62,13 +62,13 @@ jobs: run: | python -m pip install --upgrade pip pip install -r requirements.txt - pip install pytest maturin onnxruntime + pip install pytest pytest-xdist maturin onnxruntime cd ext/astro_native maturin build --release pip install --force-reinstall --no-deps target/wheels/astro_native-*.whl - name: Native kernel + originvision parity tests - run: pytest -q tests/test_native.py tests/test_originvision.py + run: pytest -q -n auto tests/test_native.py tests/test_originvision.py test: runs-on: ubuntu-latest @@ -97,7 +97,7 @@ jobs: run: | python -m pip install --upgrade pip pip install -r requirements.txt - pip install pytest pip-audit bandit + pip install pytest pytest-xdist pip-audit bandit - name: Audit dependencies run: pip-audit -r requirements.txt @@ -106,7 +106,7 @@ jobs: run: bandit -r src/ -ll -q - name: Run tests - run: pytest -q + run: pytest -q -n auto - name: Smoke run (synthetic) run: | diff --git a/.gitignore b/.gitignore index 669e751..225bbc5 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,12 @@ orion-live/ synthetic_data/ synthetic_data_mixed/ +# transient-triage training data (tools/gen_transient_triage_data.py, +# tools/mine_real_transient_data.py) -- generated, not committed +transient_triage_data.npz +transient_triage_real_data*.npz +transient_triage_real_work*/ + # Claude Code directories .claude/ **/claude/ diff --git a/CHANGELOG.md b/CHANGELOG.md index e8f4cf1..4ac48aa 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,42 @@ match the `VERSION` file and `v*` git tags. ## [Unreleased] +## [2.3.0] - 2026-09-23 + +### Added + +- **`--transient-triage`**: scores each `--transient-detect` candidate with a small CNN (new/ref/diff + stamp triplet, native `tract` inference) for a `real_probability` -- the same real/bogus triage role + ZTF's BTSbot / Rubin's DIA play downstream of classical image differencing, and the first amateur + stacking tool to do it. Advisory only, never drops a candidate. The bundled model is trained entirely + on synthetic data (`tools/gen_transient_triage_data.py` + `tools/train_transient_triage.py`); a + companion `tools/mine_real_transient_data.py` mines real cross-session negatives and real-epoch + injection positives from a user's own multi-night sessions for future retraining. Native-only for now + (no numpy/onnxruntime fallback yet). +- **Cancel button** in the desktop app. Cooperative (a shared `threading.Event`, not a thread kill): + noticed between Phase 1 frames -- usually the longest phase -- and between targets in a multi-session + run. Phases 2-4 of a single target aren't interruptible yet. + +### Changed + +- **`--originvision` is on by default now** (`--no-originvision` to disable -- a single action, same + shape as `--auto`/`--no-auto`). A bare `--originvision` on an old command line now errors instead of + being a no-op, deliberately: the desktop app's auto-generated form keys purely on argparse dest with + no dest-collision handling, so a redundant positive flag alongside the new negative one would have + silently duplicated or dropped that field. +- **`--originvision-workers` default raised 2 -> 8**, measured (not guessed): near-linear scaling on a + real full-resolution frame, 2751 ms/frame at 1 worker down to 473 ms at 8 -- the previous default left + most of the free GIL-released parallelism the code already claimed on the table. + +### Fixed + +- **`--transient-detect` no longer hard-refuses two epochs with different pixel dimensions.** Two + independently-stacked sessions of the same target routinely differ in shape (different dither pattern, + different Phase 3 crop) even though they cover the same field; `_align_reference` now reconciles them + onto a common grid first (`src.utils.embed_to_shape`, the same trick `--merge` already used for its own + differently-shaped previous stacks), carrying a real valid-data mask through the warp so the padded + border reads as uncovered, not real reference data. + ## [2.2.6] - 2026-09-22 ### Changed diff --git a/CLAUDE.md b/CLAUDE.md index f02ceab..e15b5d7 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -17,9 +17,16 @@ pip install maturin && (cd ext/astro_native && maturin develop --release) ### Run tests ```bash pytest -q +# Full suite, parallel (pytest-xdist, in requirements-dev.txt) -- measured 3.2x on a +# 16-core machine (1665 tests: 169s -> 52s, no isolation issues found) +pytest -q -n auto # Run a single test pytest tests/test_core.py::test_calculate_shift_recovery -v ``` +`-n auto` is not the default (no `addopts` in `pyproject.toml`) -- deliberately: xdist runs each +test in a forked worker process, which breaks `--pdb`/`breakpoint()` interactive debugging (a +worker's stdin isn't wired up for it) and interleaves `-s` print output across workers. CI uses +`-n auto` on both jobs (`.github/workflows/ci.yml`) where that tradeoff doesn't apply. ### Lint ```bash @@ -97,7 +104,7 @@ The pipeline is split across `src/` modules. [originstack.py](originstack.py) is | [src/lightcurve_analysis.py](src/lightcurve_analysis.py) | `--lightcurve-analysis` on `--photometry-timeseries` CSVs: astropy `LombScargle` (periods 4 cadences to half the baseline, with FAP) and `BoxLeastSquares` to place a dip, then a trapezoid `least_squares` fit with Jacobian errors and a BIC test against flat. Uses `astropy.timeseries` (verify it survives a PyInstaller build with `packaging/verify_build.ps1` before relying on it in the packaged app) | | [src/frame_processor.py](src/frame_processor.py) | Parallel workers, `execute_frame_processing`, `quality_gate` | | [src/postprocess.py](src/postprocess.py) | Full post-processing chain: `postprocess_stack`. **Step 1 (hot pixel removal) routes its 5x5 median filter through native per-channel calls** (`_median_filter_per_channel`, `median_filter_native`): the original single `scipy.ndimage.median_filter(stacked, size=(5,5,1))` call -- the first thing Phase 4 does, on every stack -- measured 5.1s on a real full-res frame; scipy's generic N-D rank-filter machinery has no fast path for a size-1 axis, so it silently bypassed this codebase's own already-existing native 2D median kernel (the same class of miss as `_fix_hot_bayer`'s pre-native-routing bug in `debayer.py`, just never caught here). 3 independent native per-channel calls are equivalent by construction and validated against the combined-axis scipy call in `tests/test_native.py` | -| [src/difference_imaging.py](src/difference_imaging.py) | Proper image subtraction and transient detection (`--transient-detect REF.fits`). Answers "did anything change?" rather than "what does my target look like?" -- novae, outbursts, supernovae, asteroids, variables. `zogy()` implements Zackay, Ofek & Gal-Yam (2016): rather than degrading one epoch to match the other (Alard-Lupton), it cross-convolves each image with the **other's** PSF, so both sides acquire the same effective PSF and stellar residuals cancel in closed form even across differing seeing -- a plain subtraction leaves a dipole at every star, scaling with brightness, i.e. worst exactly where transients hide. Returns `D` (difference), `S` (match-filtered score) and `S_corr` (score in units of its own propagated sigma, so a threshold is a real significance). **`S_corr`'s astrometric noise term is not optional**: a sub-pixel registration slip leaves a residual proportional to the local gradient, largest at bright stars, and with `astrometric_sigma=0` every bright star reports as a transient (asserted in [tests/test_difference_imaging.py](tests/test_difference_imaging.py)). `_prepare_psf` zero-pads and **rolls the PSF so its centre lands on index [0,0]** -- the FFT's origin; omitting the roll shifts every output by half the frame. `estimate_background_sigma` uses *symmetric* iterative sigma clipping even though stars are a one-sided contaminant: the obvious "keep pixels below the 80th percentile, take their MAD" truncates the Gaussian core and reads 5.7 against an injected 7.0, a 19% underestimate -- and that sigma is the denominator of every significance, so underestimating it manufactures false positives at exactly the threshold users trust. **Runs on the LINEAR stack** (`fits_stacked` in `pipeline.py`, the `RAWSTACK` product the output FITS contains), never the post-processed array: Phase 4's nonlinear stretches/denoise/local-contrast break photometric linearity, and comparing a post-processed frame against a linear reference mismatches the flux scale by a fraction of a percent -- several sigma on a bright star, reporting a transient at every star in the field. **`--transient-detect` refuses a reference without `RAWSTACK=True`** (same rule as `--merge`). The warped reference's uncovered area -- field rotation leaves empty corners -- is masked via a footprint (`_align_reference` warps an all-ones image too; `_erode` uses `border_value=1` so the frame's own edge is not treated as uncovered, which measured 66% "covered" for a 93%-covered frame): unmasked, stars in those corners came out as confident 'brightenings' (6 false candidates on a 9 deg synthetic rotation, 0 masked). The astrometric sigma is the *measured* matched-star RMS residual with `_ASTROMETRIC_SIGMA_FLOOR_PX = 0.3` as a floor -- it used to be a hard-coded 0.3 reported as a measurement. `zogy` pads to `scipy.fft.next_fast_len` (3.4x on the Origin's 1096 = 8x137 axis, whose prime factor pushes pocketfft onto Bluestein; `S_corr` identical to 2e-5 sigma in the interior). `detect_transients` is a single sorted pass, verified identical to the old per-candidate argmax loop on 300 randomized fields including ties and NaNs. `psf_difference` was dropped from `ZogyResult` (no reader). The catalogue's WCS must be built with `naxis=2` (`pipeline.py`): a bare `WCS(header)` on the `(3, H, W)` cube is 3-axis, still passes `has_celestial`, and blanks every RA/Dec | +| [src/difference_imaging.py](src/difference_imaging.py) | Proper image subtraction and transient detection (`--transient-detect REF.fits`). Answers "did anything change?" rather than "what does my target look like?" -- novae, outbursts, supernovae, asteroids, variables. `zogy()` implements Zackay, Ofek & Gal-Yam (2016): rather than degrading one epoch to match the other (Alard-Lupton), it cross-convolves each image with the **other's** PSF, so both sides acquire the same effective PSF and stellar residuals cancel in closed form even across differing seeing -- a plain subtraction leaves a dipole at every star, scaling with brightness, i.e. worst exactly where transients hide. Returns `D` (difference), `S` (match-filtered score) and `S_corr` (score in units of its own propagated sigma, so a threshold is a real significance). **`S_corr`'s astrometric noise term is not optional**: a sub-pixel registration slip leaves a residual proportional to the local gradient, largest at bright stars, and with `astrometric_sigma=0` every bright star reports as a transient (asserted in [tests/test_difference_imaging.py](tests/test_difference_imaging.py)). `_prepare_psf` zero-pads and **rolls the PSF so its centre lands on index [0,0]** -- the FFT's origin; omitting the roll shifts every output by half the frame. `estimate_background_sigma` uses *symmetric* iterative sigma clipping even though stars are a one-sided contaminant: the obvious "keep pixels below the 80th percentile, take their MAD" truncates the Gaussian core and reads 5.7 against an injected 7.0, a 19% underestimate -- and that sigma is the denominator of every significance, so underestimating it manufactures false positives at exactly the threshold users trust. **Runs on the LINEAR stack** (`fits_stacked` in `pipeline.py`, the `RAWSTACK` product the output FITS contains), never the post-processed array: Phase 4's nonlinear stretches/denoise/local-contrast break photometric linearity, and comparing a post-processed frame against a linear reference mismatches the flux scale by a fraction of a percent -- several sigma on a bright star, reporting a transient at every star in the field. **`--transient-detect` refuses a reference without `RAWSTACK=True`** (same rule as `--merge`). The warped reference's uncovered area -- field rotation leaves empty corners -- is masked via a footprint (`_align_reference` warps an all-ones image too; `_erode` uses `border_value=1` so the frame's own edge is not treated as uncovered, which measured 66% "covered" for a 93%-covered frame): unmasked, stars in those corners came out as confident 'brightenings' (6 false candidates on a 9 deg synthetic rotation, 0 masked). The astrometric sigma is the *measured* matched-star RMS residual with `_ASTROMETRIC_SIGMA_FLOOR_PX = 0.3` as a floor -- it used to be a hard-coded 0.3 reported as a measurement. `zogy` pads to `scipy.fft.next_fast_len` (3.4x on the Origin's 1096 = 8x137 axis, whose prime factor pushes pocketfft onto Bluestein; `S_corr` identical to 2e-5 sigma in the interior). `detect_transients` is a single sorted pass, verified identical to the old per-candidate argmax loop on 300 randomized fields including ties and NaNs. `psf_difference` was dropped from `ZogyResult` (no reader). The catalogue's WCS must be built with `naxis=2` (`pipeline.py`): a bare `WCS(header)` on the `(3, H, W)` cube is 3-axis, still passes `has_celestial`, and blanks every RA/Dec. **`--transient-triage`** ([src/transient_triage.py](src/transient_triage.py)) scores each `detect_transients` candidate with a small CNN for a `real_probability` -- advisory only, never drops a candidate, the same role ZTF's BTSbot / Rubin's DIA triage play downstream of classical image differencing, and (per the 2026-09 deep-research pass into recent stacking techniques) not something any amateur tool does today. Stamp input is the standard "triplet" (`new`, warped `ref`, `difference`), each channel normalised by its own frame-level `estimate_background_sigma` rather than `originvision`'s percentile stretch, which is for photographic display and would destroy the physical sigma units ZOGY's own significance relies on. Native-only for now (`astro_native.transient_triage_score`, `ext/astro_native/src/lib.rs`'s `mod transient_triage`) -- deliberately no numpy/onnxruntime fallback yet, unlike every other native kernel in this project; self-disables with a warning when unavailable. The bundled model (`src/data/transient_triage.onnx`) is trained entirely on synthetic data -- `tools/gen_transient_triage_data.py` renders synthetic star fields through the *real* `zogy()` + `detect_transients()` (real transients, cosmic rays, undersuppressed sub-pixel registration-slip dipoles, hot pixels) and `tools/train_transient_triage.py` (a script-local `torch` dependency, not in `requirements.txt` -- model training happens outside the shipped package, same stance as `originvision.onnx` itself) trains and exports it -- so treat it as a first cut, not a production classifier; no labelled real transients exist yet | | [src/sky_model.py](src/sky_model.py) | **EXPERIMENTAL** physics-based sky background model (`--bg-method physical`). Fits `sky = c0 + c_air*airglow(z) + c_moon*moonlight(rho,z) + c_zodi*zodiacal(lambda,beta) + c_lp*skyglow(az,z)`, where every component's *spatial shape* is fixed by geometry and only a scalar amplitude is free, so unlike mesh/DBE/wavelet it has nowhere to put a nebula. **Real-data verdict: it does not work on a typical ~1 deg deep-sky field, and now detects that and declines.** On a real Lagoon session the zenith angle varies by only 0.94 deg across the whole frame and the azimuth by 1.5 deg, so every component map is essentially flat, the fit has nothing to grip, and subtracting it made the corner-to-corner gradient *worse* (67 -> 111 ADU) where DBE removed 68%; sweeping `--light-pollution-azimuth` through all 360 deg moved the residual by under 0.01 ADU. Structural, not a tuning problem: physical components vary on ten-degree scales, so a narrow field's gradient is dominated by *instrumental* effects (vignetting, amp glow, filter gradients) a sky model cannot represent. `remove_physical_sky` therefore measures `_corner_gradient` before and after and returns None unless it improved things, letting `postprocess.py::_apply_physical_sky` fall back to DBE with an accurate reason (it must distinguish 'no GPS' from 'could not help' -- reporting the former on a session that had full GPS sent the reader hunting for metadata that was already present). The guard is empirical, not a field-size rule, and its corner-based proxy is weakest on radially symmetric gradients whose four corners are equal by construction. **Two earlier claims are corrected by this testing**: DBE does *not* eat nebulosity (99.5% retained on the Lagoon -- the damage that motivated this work came from the `sky_residual` residual passes, a different step `--auto` already skips), and the '98% synthetic nebula preserved' figure holds only for a gradient built from the model's own basis. Coefficients are constrained non-negative (`scipy.optimize.nnls`) and that *is* load-bearing: across a real field the maps are nearly collinear (condition ~1e5), so an unbounded fit builds an interior maximum from cancelling coefficients and inverts a nebula (-25% preservation). Clipping is upward-only (stars sit above sky). Ephemerides are closed-form, not astropy: `EarthLocation`/`AltAz` pull in the IERS tables [packaging/originstack.spec](packaging/originstack.spec) excludes -- validated against astropy at **0.009 deg (sun) / 0.05 deg (moon)**. `julian_date` honours timezone offsets: a Celestron Origin `info.json` stamps local time (`2026-08-31T20:40:32-0700`), which the first version failed to parse at all -- silently disabling the model on exactly its target data -- and which loses 7 hours of moon position if the offset is merely stripped. Positions are mean-equinox-of-date while WCS pixels are J2000, so separations carry a ~0.35 deg precession offset (uncorrected, deliberately; and worth knowing, since that offset *looks* like an ephemeris bug -- the real one found this way was 19.6 deg, from Schlyter's lunar elements being epoched at 1999-12-31.0 rather than J2000.0). `describe_fit` gates component attribution on the basis condition number: removal can be valid while the *split* between components is unidentifiable. **Later changes:** the FOV gate now runs *before* `build_geometry` (a telescope field is declined in ~0 s / 0 MB instead of ~1.3 s / 578 MB); `fit_sky_model` fits on a strided sample (`_FIT_MAX_SAMPLES`, ~100k pixels; 65x faster, dominant coefficient within 0.25%) but evaluates the model at full resolution; `remove_physical_sky_with_reason` returns `(result, reason)` with the *real* reason (narrow field / failed fit / no improvement / no geometry) and `remove_physical_sky` is a thin wrapper; `_nnls` has **no fallback** -- it used to catch every exception (incl. scipy's max-iteration `RuntimeError`) and substitute a clamped `lstsq`, i.e. the unbounded fit this module documents as inverting nebulae, and the guard cannot see radially symmetric damage. A failure now declines to DBE with a warning. Azimuth and helio-ecliptic longitude are upsampled via sin/cos (`_up_angle`): interpolating across the 360/0 seam ramped the long way round, measured as a 22.5 deg/px step where the truth is ~0.05. The unreleased multi-frame joint fit and `amp_glow_basis` were removed -- nothing in `src/` called them | | [src/uncertainty.py](src/uncertainty.py) | End-to-end uncertainty propagation (`--uncertainty-propagate`) and confidence mapping. Phase 3's `--uncertainty-map` sigma describes the *linear* stack; Phase 4 then reshapes the noise field, so this carries the error bars through to the delivered image. `propagate_uncertainty` is **Monte Carlo, not analytic**: it draws K noise realizations (`--uncertainty-realizations`, default 8) from the Phase 3 sigma map and pushes each through the *unmodified* `postprocess_stack`, taking the per-pixel spread. Deliberate — nine Phase 4 steps are nonlinear denoisers and four are iterative deconvolvers (RL, FISTA, anisotropic diffusion), several spatially adaptive (BayesShrink thresholds, the structure-tensor coherence map), so no closed-form Jacobian exists for most of the chain; the MC estimator is exact up to `~1/sqrt(2K)` *for a chain that does not adapt to its input's noise level*, and stays correct automatically when a denoiser is added. **It is slightly biased low for adaptive steps:** each realization is `stacked + N(0, sigma)` but `stacked` already carries ~sigma, so Phase 4 sees sqrt(2)x the real noise and steps that estimate parameters from the data (BayesShrink, DBE sky sigma) denoise it harder. Measured against this project's own `wavelet_denoise` over sigma {1,4,12} x threshold {2,3,5}: ratio 0.93-1.03 -- a few percent, inside the ~25% MC error at K=8. A first-principles argument predicted ~2x understatement; it did not reproduce (shrinkage and threshold move together), and a reviewer's toy-chain figure was likewise wrong -- measure against the real denoiser before believing either. `propagate_uncertainty` returns `(sigma_post, mean_post, adaptivity)`, where `adaptivity` is sigma(full amp)/sigma(half amp): 2.0 = scale-invariant, and `pipeline.py` warns below 1.5. Realizations run under `_quiet_args` (file-writing/network steps off — `_QUIET_OFF`: `remove_stars`, `nmf_separate`, `photometric_calibration`, `annotate`, `aberration_report`, `diagnostic`, `export_masks`, `keep_intermediates`, `comet_radial_renorm`, `comet_larson_sekanina`, and `denoise_strength_calibrate`, a 9-point sweep. **Any Phase 4 step that writes a file or calls out belongs in that list**: every realization runs under `redirect_stdout`, so a step left on prints its own "Saved:" into a swallowed buffer and the file left on disk is the *last noise realization*, silently) with stdout swallowed, so K passes don't emit K sets of sidecars or K Gaia queries; everything that shapes the *noise* is left exactly as the real pass ran it. `confidence_map` turns the propagated sigma into per-pixel SNR above sky and returns **`NaN` where the propagated sigma is exactly zero** — those pixels were clamped to a constant by the chain (sky pedestal lift, non-negativity clips), so they carry no measurement, and dividing by ~0 would otherwise rank the pipeline's own floor artifacts as the most confident pixels in the frame (observed at ~25% of a real synthetic frame). `error_aware_black_point` (`--error-aware-stretch SIGMA`) returns the faintest pixel clearing SIGMA confidence, used as the preview black point so sub-threshold content clips to black instead of being stretched into apparent structure. Validated against chains with known variance transformation — identity recovers the input sigma, a x3 scale scales it x3, a 3x3 box blur divides it by exactly 3 ([tests/test_uncertainty_propagation.py](tests/test_uncertainty_propagation.py)) | | [src/psf_deconvolution.py](src/psf_deconvolution.py) | PSF estimation, Richardson-Lucy (global + spatially-variant `richardson_lucy_svpsf`), TV, and `sparse_wavelet_deconvolve` (`--deconvolve sparse`) -- FISTA (Beck & Teboulle 2009), L1-regularised in this project's own wavelet basis (`src/wavelet.py`) rather than TV's spatial-gradient basis; a different sparsity prior, same forward/adjoint PSF convolution and positivity-pedestal pattern as the RL/TV paths | @@ -111,7 +118,7 @@ The pipeline is split across `src/` modules. [originstack.py](originstack.py) is | [src/live_stack.py](src/live_stack.py) | Real-time stacking (`--live`): watches the capture directory and folds each new sub into a running weighted-mean stack, pushing the growing result + running SNR to whatever UI is attached (console-only on a plain CLI run) | | [src/ui_events.py](src/ui_events.py) | In-process UI event/state sink for the desktop app (replaced `webview.py`'s HTTP/SSE dashboard, 2026-08): `UIEvents` holds live phase progress, log stream, per-frame quality ticker, and preview state (named milestone slots with a retained downsized float16 source for on-demand re-stretch, a per-frame thumbnail ring) — the same state model the old dashboard served over SSE, now polled directly by `desktop_app.py`'s `root.after()` timer via `snapshot()`/`version` instead of pushed over a socket. Every `safe_print()`/`print_phase()` call in the codebase tees into it for free (no-op unless `attach()`ed, i.e. true no-op on a plain CLI run). `restretch()` re-renders a retained milestone at new stretch params on demand | | [src/desktop_control.py](src/desktop_control.py) | Desktop-app control layer: `get_form_schema()` introspects `cli.build_parser()` live (no hand-maintained duplicate of ~120 flags to drift; toolkit-agnostic plain dicts, consumed by `desktop_app.py`'s tkinter form builder), `build_argv_from_form()` turns a submitted form back into a synthetic argv fed through the real `cli.parse_args()` (so preset/config/`--auto` precedence and `_explicit_cli_dests` come out exactly as from a real command line), and `RunManager`/`get_run_manager()` runs one pipeline job at a time on a background thread, publishing through `ui_events.py`. `RunManager.is_running()` backs `desktop_app.py`'s close-confirmation (warns before quitting mid-run) | -| [src/desktop_app.py](src/desktop_app.py) | `python desktop_app.py`: a native `tkinter` window (stdlib — no `pywebview`/WebView2 Runtime dependency). A Setup form auto-built from `desktop_control.get_form_schema()` (one `Notebook` tab per argparse group), a live progress/log panel, and a `Canvas`-based preview viewer (zoom via mouse wheel, pan via drag, before/after wipe-slider compare via `PIL.Image.paste` compositing, per-frame thumbnail ring) fed by polling `ui_events.py`'s `UIEvents.snapshot()`. Layout is two columns: left = Setup form + pipeline bar + log (the log gets the whole lower-left); right = preview + thumbnail strip + a RECENT FRAMES quality table. (The old GHS-slider STRETCH panel was removed 2026-09; `ui_events.restretch()` / `PreviewCanvas.replace_pixels()` stay as unused-from-GUI API.) File/folder pickers are stdlib `tkinter.filedialog` (no `Api`/JS bridge needed — same process, no IPC). `--verify-headless ` skips the GUI/mainloop entirely and runs a real stack through the same frozen entry point, for `packaging/verify_build.ps1`'s multiprocessing regression check. Every failure path routes through `_fatal()` (log to `%LOCALAPPDATA%\OriginStack\logs\` + a `native_dialog.py` error dialog) — load-bearing for the packaged build (see below), which runs windowed with no console for a bare `print()` to reach. `root.protocol('WM_DELETE_WINDOW', ...)` confirms via `RunManager.is_running()` before quitting mid-run | +| [src/desktop_app.py](src/desktop_app.py) | `python desktop_app.py`: a native `tkinter` window (stdlib — no `pywebview`/WebView2 Runtime dependency). A Setup form auto-built from `desktop_control.get_form_schema()` (one `Notebook` tab per argparse group), a live progress/log panel, and a `Canvas`-based preview viewer (zoom via mouse wheel, pan via drag, before/after wipe-slider compare via `PIL.Image.paste` compositing, per-frame thumbnail ring) fed by polling `ui_events.py`'s `UIEvents.snapshot()`. Layout is two columns: left = Setup form + pipeline bar + log (the log gets the whole lower-left); right = preview + thumbnail strip + a RECENT FRAMES quality table. (The old GHS-slider STRETCH panel was removed 2026-09; `ui_events.restretch()` / `PreviewCanvas.replace_pixels()` stay as unused-from-GUI API.) File/folder pickers are stdlib `tkinter.filedialog` (no `Api`/JS bridge needed — same process, no IPC). `--verify-headless ` skips the GUI/mainloop entirely and runs a real stack through the same frozen entry point, for `packaging/verify_build.ps1`'s multiprocessing regression check. Every failure path routes through `_fatal()` (log to `%LOCALAPPDATA%\OriginStack\logs\` + a `native_dialog.py` error dialog) — load-bearing for the packaged build (see below), which runs windowed with no console for a bare `print()` to reach. `root.protocol('WM_DELETE_WINDOW', ...)` confirms via `RunManager.is_running()` before quitting mid-run. **Cancel button** (next to Start, enabled only while a run is active): cooperative, not a thread kill -- `RunManager.cancel()` sets a `threading.Event` shared onto the run's `args` as `_cancel_event`, noticed at a handful of checkpoints (`frame_processor._check_cancel`, called between Phase 1 frames in all three dispatch paths -- ProcessPool, GPU/thread pool, sequential -- since Phase 1 is usually the longest phase; and between targets in `cli.process_directory`'s per-target loop, for multi-session/hierarchical runs). Since every Phase 1 future is submitted upfront, the ProcessPool/thread-pool paths explicitly `shutdown(wait=False, cancel_futures=True)` on a caught `RunCancelled` so the `with` block's own exit only waits for whatever's already mid-frame, not the full remainder of the session. Phases 2-4 of a single target aren't interruptible yet -- once one starts it runs to completion, same as before this existed. `RunCancelled` (`src/models.py`) propagates up to `RunManager._run`, which sets `status = 'cancelled'` (a status distinct from `'ok'`/`'error'`) rather than reporting it as a failure | | [src/notify.py](src/notify.py) | `notify_windows(title, message)`: best-effort native Windows balloon-tip notification via a small `ctypes`/`Shell_NotifyIcon` wrapper (no dependency — matches this project's preference for small native/stdlib implementations over dependency weight, e.g. over `plyer`). No-op on non-Windows; any failure is swallowed, since this is cosmetic and must never affect run state | | [src/stream_stack.py](src/stream_stack.py) | Two-pass streaming stack of an already-complete directory (`--stream`): O(1) full-resolution memory via an online (single-pass) sigma-clip Welford accumulator (`online_sigma_clip_seed_burnin`/`online_sigma_clip_fold_frame` in `stacking.py`) instead of materializing the whole `(N,H,W,C)` aligned stack. Sibling to `--live` (shares `build_session_masters`); v1 limitations: reference frame picked by quality score alone, hard-limit quality gating only, no `--elastic-registration`/drizzle/patch-weighted combine/`--merge`. Mutually exclusive with `--live` | | [src/channel_combine.py](src/channel_combine.py) | `combine` subcommand: LRGB + narrowband palettes (SHO/HOO), SCNR green removal (`--scnr`), magenta-star fix (`--star-recolor`). `--continuum FILE --continuum-target {ha,oiii,sii}` subtracts a broadband/OSC continuum reference before combination; the subtraction scale is fit automatically (`optimal_continuum_scale`) by sweeping candidate scales and picking the one that *maximises* the background residual's pixel skewness -- verified against synthetic ground truth in `tests/test_continuum_subtraction.py` (a first "zero-crossing" design was tried and empirically shown wrong by that same test before shipping: skewness peaks at the true scale, it doesn't cross zero there). Since the subtraction residual is linear in the scale factor, its skewness at every swept scale is a closed-form polynomial of just 7 central moments of the masked `(narrowband, continuum)` pixel pair -- computed once (native Rust `continuum_scale_moments`, numpy fallback) instead of re-scanning the full pixel array once per scale; see "Native (Rust) acceleration" below | @@ -127,7 +134,7 @@ The pipeline is split across `src/` modules. [originstack.py](originstack.py) is | [src/photometry_timeseries.py](src/photometry_timeseries.py) | Per-frame differential light curves (`--photometry-timeseries`): runs right after Phase 3 while the registered frames are still in `mem_rgb`, warps + crops each sub, aperture-photometers a fixed `match_gaia_field` star list on every frame (`aperture_photometry_batch`), then iteratively ensemble-calibrates a per-frame per-channel zero point (comparison stars = well-detected, unsaturated, mid-brightness; high-scatter members clipped out). Writes `_lightcurves.csv` (one row per frame×star: MJD, airmass, per-channel mag/magerr, flag) and `_lightcurve_stats.csv` (per star: mean/rms/MAD/ptp/reduced-χ² + a `variable` flag). `--photometry-target "RA,DEC"` / `"px:X,Y"` marks one star and prints its stats. Needs a session `info.json` WCS (`--plate-solve` runs after this point); differential only — no absolute ZP, extinction cancels in the ensemble | | [src/gain_ptc.py](src/gain_ptc.py) | Photon-transfer gain / read-noise from raw calibration frames (`estimate_gain_ptc`, auto-run in `cli._build_masters` when `--photometry`/`--photometry-timeseries` is set and ≥2 bias + ≥2 flat frames exist and no `--photometry-gain` was given). Janesick two-frame difference: `gain = (Σmean_flat - Σmean_bias) / (var(flat₁-flat₂) - var(bias₁-bias₂))`, `read_noise_adu = std(bias₁-bias₂)/√2`, on a central sigma-clipped window (dodges vignetting/amp-glow). One flat level → a single (signal, variance) point, not a full PTC curve; OSC frames reduced to luma. Result stashed as `masters['ptc_gain_e_per_adu']` / `args._ptc_gain_e_per_adu` and consumed by `photometry._read_gain` | | [src/annotation.py](src/annotation.py) | Object annotation (`--annotate`): circles + labels bright stars and named deep-sky objects (galaxies, nebulae, clusters) on a copy of the preview, via the header's WCS and live SIMBAD cone-search queries. Needs a WCS (`--plate-solve` or a session solve); fails soft otherwise | -| [src/originvision.py](src/originvision.py) | originvision integration: defect/quality/category scoring by a separately-trained vision classifier, **run in-process** (no subprocess, no external folder, no venv, no network). Inference runs through the **native `astro_native.originvision_score` kernel** ([ext/astro_native/src/lib.rs](ext/astro_native/src/lib.rs), `mod originvision`): the whole path — per-channel percentile stretch, gaussian-prefiltered bilinear resize/centre-crop, and the ONNX forward pass via the **pure-Rust `tract` runtime** — runs inside `py.allow_threads` (a thread pool of scoring calls genuinely parallelises), panic-guarded (`catch_unwind` → `RuntimeError`, so a malformed `--originvision-model` never SIGABRTs), with no Python ONNX dependency and no extra DLL. [src/originvision_infer.py](src/originvision_infer.py) is now a thin dispatcher: native when `astro_native` is built (the only path in the packaged app), else a numpy/scipy + Python-`onnxruntime` fallback for a source checkout without the crate. The exported model **ships inside the package** at `src/data/originvision.onnx` (~11 MB, "v4" — a from-scratch, no-SSL 7-task run, best-by-category checkpoint at epoch 15); the native kernel loads it by path (session cached per `(path, size)`). `--originvision` self-disables with a warning when neither backend nor the model file is available. The graph emits **8** outputs but ONNX metadata `tasks` lists **7** (`reject`/`quality`/`category`/`exposure`/`sky_brightness`/`stray_light_gradient`/`background_grid`) — `trailing` is a real graph output with *untrained* weights on this checkpoint, so `score_rgb` gates every head on `tasks` membership, never on output presence. `trailing` and `background_grid` are deliberately not surfaced (the latter trained but unused — OriginStack runs its own DBE). A model whose graph-output count doesn't match its `head_order` metadata is rejected (both backends) rather than scored around. `category` is a 4-class head — galaxy/nebula/star_cluster/comet — but `comet` is distrusted on the current checkpoint, so a top `comet` pick is demoted to the runner-up (`shape_gate=False` disables that); `exposure` is a 7-class classifier. `src/originvision_infer.py::_load_image_any` debayers a raw `.fits` light itself (single-frame Bayer → RGB) then feeds the array straight to `score_rgb` — no temp file; TIFF/PNG/JPG pass through `tifffile`/`Pillow`. **Model provenance / re-sync**: [vendor/originvision/](vendor/originvision/) holds the upstream snapshot the port + `src/data/originvision.onnx` are copied from (see `vendor/originvision/VENDORED_FROM.txt`); it is **not on the runtime path** (numpy-mirror pattern). `--originvision-model PATH` overrides the bundled model; legacy `--originvision-dir`/`--originvision-checkpoint` still resolve a path (`--originvision-python`/`--originvision-script`/`--originvision-timeout` were removed with the subprocess). `--originvision` alone, with `--auto` active (the default), samples 3 light frames spread through the session — the sampled category feeds `--auto`'s target-classification prior the same way SIMBAD/header metadata does, and a defect flag nudges settings defensively (`--trail-reject` on, stronger chroma denoising) via `_originvision_defect_flagged` in `auto_settings.py`. `--originvision-score-all` (opt-in, needs `--originvision` too — a no-op and warned about otherwise) additionally scores every accepted light frame after Phase 1 (before `quality_gate`) and once more on the final stacked master, gated in two places (`pipeline.py`'s call site and inside `score_lights_with_originvision` itself). Advisory/logging only while the model is still finishing its first training run — results are stored in `FrameInfo.metrics['originvision']` and logged (defective/stray-light flags, session-relative below-average `quality_score`, master category vs. the pipeline's own inferred target type) but never set `accepted` or feed `metrics['score']`, so nothing is auto-dropped | +| [src/originvision.py](src/originvision.py) | originvision integration: defect/quality/category scoring by a separately-trained vision classifier, **run in-process** (no subprocess, no external folder, no venv, no network). Inference runs through the **native `astro_native.originvision_score` kernel** ([ext/astro_native/src/lib.rs](ext/astro_native/src/lib.rs), `mod originvision`): the whole path — per-channel percentile stretch, gaussian-prefiltered bilinear resize/centre-crop, and the ONNX forward pass via the **pure-Rust `tract` runtime** — runs inside `py.allow_threads` (a thread pool of scoring calls genuinely parallelises), panic-guarded (`catch_unwind` → `RuntimeError`, so a malformed `--originvision-model` never SIGABRTs), with no Python ONNX dependency and no extra DLL. [src/originvision_infer.py](src/originvision_infer.py) is now a thin dispatcher: native when `astro_native` is built (the only path in the packaged app), else a numpy/scipy + Python-`onnxruntime` fallback for a source checkout without the crate. The exported model **ships inside the package** at `src/data/originvision.onnx` (~11 MB, "v4" — a from-scratch, no-SSL 7-task run, best-by-category checkpoint at epoch 15); the native kernel loads it by path (session cached per `(path, size)`). **On by default since 2026-09** (`--no-originvision` to disable — a single action, same no-positive-flag shape as `--auto`/`--no-auto`; bare `--originvision` on an old command line now errors rather than being a no-op, deliberately, since `desktop_control.py`'s form-schema/argv machinery keys purely on dest with no dest-collision handling). Self-disables with a warning when neither backend nor the model file is available, so a source checkout with neither `astro_native` nor `onnxruntime` built still runs cleanly. The graph emits **8** outputs but ONNX metadata `tasks` lists **7** (`reject`/`quality`/`category`/`exposure`/`sky_brightness`/`stray_light_gradient`/`background_grid`) — `trailing` is a real graph output with *untrained* weights on this checkpoint, so `score_rgb` gates every head on `tasks` membership, never on output presence. `trailing` and `background_grid` are deliberately not surfaced (the latter trained but unused — OriginStack runs its own DBE). A model whose graph-output count doesn't match its `head_order` metadata is rejected (both backends) rather than scored around. `category` is a 4-class head — galaxy/nebula/star_cluster/comet — but `comet` is distrusted on the current checkpoint, so a top `comet` pick is demoted to the runner-up (`shape_gate=False` disables that); `exposure` is a 7-class classifier. `src/originvision_infer.py::_load_image_any` debayers a raw `.fits` light itself (single-frame Bayer → RGB) then feeds the array straight to `score_rgb` — no temp file; TIFF/PNG/JPG pass through `tifffile`/`Pillow`. **Model provenance / re-sync**: [vendor/originvision/](vendor/originvision/) holds the upstream snapshot the port + `src/data/originvision.onnx` are copied from (see `vendor/originvision/VENDORED_FROM.txt`); it is **not on the runtime path** (numpy-mirror pattern). `--originvision-model PATH` overrides the bundled model; legacy `--originvision-dir`/`--originvision-checkpoint` still resolve a path (`--originvision-python`/`--originvision-script`/`--originvision-timeout` were removed with the subprocess). `--originvision` alone, with `--auto` active (the default), samples 3 light frames spread through the session — the sampled category feeds `--auto`'s target-classification prior the same way SIMBAD/header metadata does, and a defect flag nudges settings defensively (`--trail-reject` on, stronger chroma denoising) via `_originvision_defect_flagged` in `auto_settings.py`. `--originvision-score-all` (opt-in, needs `--originvision` too — a no-op and warned about otherwise) additionally scores every accepted light frame after Phase 1 (before `quality_gate`) and once more on the final stacked master, gated in two places (`pipeline.py`'s call site and inside `score_lights_with_originvision` itself). Advisory/logging only while the model is still finishing its first training run — results are stored in `FrameInfo.metrics['originvision']` and logged (defective/stray-light flags, session-relative below-average `quality_score`, master category vs. the pipeline's own inferred target type) but never set `accepted` or feed `metrics['score']`, so nothing is auto-dropped. **Performance (2026-09 profiling pass)**: `--originvision-workers` default raised 2 → 8 — measured near-linear scaling on a real full-res frame (2751 ms/frame at 1 worker → 1415 at 2 → 473 at 8; the previous default of 2 left most of the free GIL-released parallelism the docstring already claimed on the table). Separately, `score_rgb`'s preprocessing (percentile stretch + resize) scales with *input* pixel count even though only a 256x256 crop is ever used — 85% of a full-res (1936x1096) call was spent on pixels the model never sees (1315 ms → 196 ms once already at 256x256). A pre-downsample step (`_downsample_if_large`, `originvision_infer.py`) fixes that (up to 7x on a real frame) but was **measured and found unsafe as a default**: unlike `_resize_center_crop`'s own cv2→scipy swap (class/flag heads unaffected to ~0.002), the `reject`/`quality` heads are genuinely resolution-sensitive — `defect_probability` swung +0.04 to +0.40 on one real frame across every tested aggressiveness, enough to flip `is_defective` near the 0.5 boundary, which feeds `auto_settings.py`'s defensive nudges, not just a log line. Kept as a dormant, undocumented-to-the-CLI opt-in (`score_rgb(..., fast_preprocess=True)`) rather than wired to a flag or shipped as default — see its docstring for the full factor-vs-delta sweep | | [src/pipeline.py](src/pipeline.py) | Thin orchestrator: `stack_target` wires all four phases | | [src/health_check.py](src/health_check.py) | `run_health_check` | | [src/cli.py](src/cli.py) | `process_directory`, `parse_args`, `main`. `save_effective_config` writes strings through `_toml_str`: an unescaped Windows path (`log_file = "C:\Users\..."`) made every GUI-saved config unparseable ("Invalid hex value"), and the run then silently fell back to defaults with only a warning; regression-tested by round-tripping through `tomllib`. `tools/lint_conventions.py`'s `_git` decodes as UTF-8 for the same platform-codepage reason | @@ -163,7 +170,7 @@ The old local-normalisation step (`--local-normalize`) was **removed**: it did l The preview JPEG black point is set per target by the auto-advisor (`preview_black_sigma`, overridable with `--preview-black-sigma`); higher values (2–3) clip the sky-noise tail to black for a small target on empty sky. ### Incremental stacking (`--merge`) -The main output FITS is the linear pre-post-processing stack (`RAWSTACK=True`) with `NFRAMES`/`INTGTIME`/`TOTEXP` headers. `--merge PREV.fits [...]` processes only the new session through Phases 1-3, registers each previous stack onto the new grid (blind rigid star-pattern match in [src/blind_match.py](src/blind_match.py) first — nights differ by arbitrary field rotation on alt-az mounts, and this makes no assumption about the angle — translation-seeded star-match affine and translation-only fallbacks, hard error on failure or <25% overlap), and combines them before Phase 4 runs once on the result: each previous stack is first mapped onto the current stack's flux scale (`merge._match_flux_scale`: per-channel gain + sky offset, Tukey-IRLS on smoothed pixels bright in both -- a 25 s ISO 500 stack is not comparable to a 10 s ISO 200 one in raw ADU, and averaging them as-is mixed two brightness scales), then a per-pixel **inverse-noise-variance** weighted mean inside each warped footprint (`merge._pixel_noise`, MAD of lag-4 pixel differences; falls back to `NFRAMES` weights if any stack's noise is unmeasurable). `NFRAMES` is only a proxy for noise when every session has the same per-frame exposure. Header aggregates are summed, so the output chains into future merges. There is no cross-session outlier rejection (each session already rejected internally); not supported with `--drizzle-scale > 1`. +The main output FITS is the linear pre-post-processing stack (`RAWSTACK=True`) with `NFRAMES`/`INTGTIME`/`TOTEXP` headers. `--merge PREV.fits [...]` processes only the new session through Phases 1-3, registers each previous stack onto the new grid (blind rigid star-pattern match in [src/blind_match.py](src/blind_match.py) first — nights differ by arbitrary field rotation on alt-az mounts, and this makes no assumption about the angle — translation-seeded star-match affine and translation-only fallbacks, hard error on failure or <25% overlap), and combines them before Phase 4 runs once on the result: each previous stack is first mapped onto the current stack's flux scale (`merge._match_flux_scale`: per-channel gain + sky offset, Tukey-IRLS on smoothed pixels bright in both -- a 25 s ISO 500 stack is not comparable to a 10 s ISO 200 one in raw ADU, and averaging them as-is mixed two brightness scales), then a per-pixel **inverse-noise-variance** weighted mean inside each warped footprint (`merge._pixel_noise`, MAD of lag-4 pixel differences; falls back to `NFRAMES` weights if any stack's noise is unmeasurable). `NFRAMES` is only a proxy for noise when every session has the same per-frame exposure. Header aggregates are summed, so the output chains into future merges. There is no cross-session outlier rejection (each session already rejected internally); not supported with `--drizzle-scale > 1`. A previous stack's own shape rarely matches the current run's (different dither pattern, different Phase 3 common-crop), so it's zero-padded/cropped onto the current grid first (`src.utils.embed_to_shape`, top-left, pixel coordinates preserved) before registration -- the blind matcher needs no positional correspondence between the two canvases, only relative star geometry, so this is purely a shape-reconciliation step, not an alignment hint. `--transient-detect`'s `_align_reference` ([src/difference_imaging.py](src/difference_imaging.py)) shares this same helper for the identical reason: two independently-stacked sessions of the same target routinely differ in pixel dimensions too. It used to hard-refuse a shape mismatch there; the fix also had to carry a real *valid-data* mask through the warp (not an all-ones one) so the zero-padded border reads as uncovered, not real reference data -- otherwise stars in that padding would report as false 'brightening' transients, the same failure mode the rotation-wedge footprint mask already guards against. ### Streaming memory model Frames are processed one at a time: load → process → accumulate → free. Memory usage stays at ~1-2 frames regardless of total frame count. This is the core design constraint — never accumulate all frames in memory. @@ -215,6 +222,7 @@ with `bayerPattern` to override. ### Native (Rust) acceleration [ext/astro_native/](ext/astro_native/) is a PyO3/maturin crate of hot-path kernels, all with a numpy fallback. Coverage: - **originvision inference** (`ext/astro_native/src/lib.rs` `mod originvision`, `src/originvision_infer.py::score_rgb`, `--originvision`): the one kernel that isn't a numpy-accelerator — it runs the whole `--originvision` scoring path natively (preprocessing + ONNX forward pass via the pure-Rust **`tract`** runtime), so a source checkout with `astro_native` built and the packaged app need **no Python ONNX dependency at all**. Preprocessing mirrors the Python fallback exactly (per-channel [0.5, 99.5] stretch → gaussian-prefiltered bilinear resize/centre-crop → NCHW /255); result heads are gated on the ONNX metadata `tasks` list, and a graph-output count that doesn't match `head_order` is a hard error (both backends). The compute runs inside `py.allow_threads` and under `std::panic::catch_unwind` — `tract` parses an external `.onnx` and a malformed one can *panic* inside the parser (a `pyo3_runtime.PanicException`, which is a `BaseException` and slips past callers' `except Exception`), so it's converted to a plain `RuntimeError` here; `panic = "unwind"` (not `"abort"`) in the release profile is what lets that unwind happen. Session cached per `(model_path, size)` (`OnceLock>`; the mutex is dropped before the forward pass). `tract` pulls ~120 transitive crates and lengthens the crate's first build noticeably. Fallback is `src/originvision_infer.py`'s numpy/scipy + `onnxruntime` port (see "originvision" module row above). Parity (native vs the onnxruntime fallback: identical category/flags, scalars within tolerance) in `tests/test_native.py`, and CI's `native` job (`.github/workflows/ci.yml`) builds the crate + installs `onnxruntime` so that parity actually runs on every push. +- **transient triage inference** (`ext/astro_native/src/lib.rs` `mod transient_triage`, `src/transient_triage.py::score_candidates`, `--transient-triage`): a much smaller sibling of the originvision kernel above -- no percentile stretch, no resize (candidate stamps arrive already extracted at a fixed size and per-channel sigma-normalized in Python), no multi-head metadata decode, just a batched tract forward pass over `(N, 3, size, size)` returning one sigmoid probability per candidate in a single call rather than one call per candidate. Same panic-guard (`catch_unwind` + `py.detach`) and `(path, size)`-keyed session cache pattern as originvision. **Deliberately has no numpy/onnxruntime fallback yet** -- unlike every other native kernel in this file, a source checkout without `astro_native` built simply cannot use `--transient-triage` (self-disables with a warning) until one is added. - **Stacking combines** (`src/stacking.py`): `sigma_clip_combine` (~37×), `esd_combine` (~24×), `percentile_clip_combine` (~13×), `median_combine` (~6×). ESD's Student-t critical-value table is precomputed in Python (`_esd_lambda_table`) and passed to Rust — exact parity, no stats crate. Native path is taken when a rejection mask is not requested and the input is a C-contiguous float32 `(N,H,W,C)` array; the aligned stack memmap qualifies, so Rust views it zero-copy and the streaming memmap model is preserved. (`trimmed_mean_combine` was removed 2026-08 — Python and Rust — once recognized as functionally identical to `percentile_clip_combine` at matching params, just parameterized by trim-fraction instead of percentile bounds.) - **Linear Fit Clipping** (`src/stacking.py` `linear_fit_clip_combine`, `--stack-method linear_fit`): PixInsight's algorithm — sorts each pixel's per-frame stack ascending (order statistics), fits a line to value-vs-rank by least squares, rejects samples whose residual from that fit exceeds `sigma_low`/`sigma_high` times the residual scale, refits on survivors, iterates. More robust to non-Gaussian tails than sigma-clip's mean/std test since it doesn't assume the per-pixel distribution across frames is Gaussian, just that it's locally near-linear in sorted order. No independent open-source reference implementation exists to bit-validate against (unlike Malvar/Menon2007) since PixInsight is closed-source — validated instead via a numpy mirror (`_linear_fit_clip_tile`) checked for native/numpy parity plus a synthetic-outlier-injection test (`tests/test_native.py`) confirming a single wild sample per pixel doesn't survive into the combined result. - **Inverse-variance-weighted combine** (`src/stacking.py` `ivw_combine`, `--stack-method ivw`): the Gauss-Markov-optimal linear combiner — weights each frame by `1/noise²` using that frame's own measured Phase 1 background sigma (a real statistical optimum under a per-frame-homoscedastic noise model, unlike `--weight-noise`'s ad-hoc multiplicative heuristic). Optionally adds a per-pixel Poisson shot-noise term (`--config ivw_gain`, electrons/ADU) so brighter regions are correctly down-weighted relative to sky background in noisier frames — this is why it still routes through the native kernel's per-pixel loop rather than a single static-weight numpy broadcast, even though the no-gain case has no iteration or sorting at all. Does not reject outliers (cosmic rays/trails get a small nonzero weight, not zero) — meant to complement this pipeline's existing per-frame pre-filters (`--cosmic-ray-rejection`, `--trail-reject`), not replace them. diff --git a/VERSION b/VERSION index bda8fbe..276cbf9 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -2.2.6 +2.3.0 diff --git a/ext/astro_native/Cargo.lock b/ext/astro_native/Cargo.lock index 67de6aa..031c263 100644 --- a/ext/astro_native/Cargo.lock +++ b/ext/astro_native/Cargo.lock @@ -49,7 +49,7 @@ checksum = "fb5dfbc6d8d2675589ccbe4d0fd61df2419075625f8c1a62325e718e2b0049f9" [[package]] name = "astro_native" -version = "0.35.0" +version = "0.36.0" dependencies = [ "numpy", "pyo3", diff --git a/ext/astro_native/Cargo.toml b/ext/astro_native/Cargo.toml index 2b84ac3..a40dd11 100644 --- a/ext/astro_native/Cargo.toml +++ b/ext/astro_native/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "astro_native" -version = "0.35.0" +version = "0.36.0" edition = "2021" description = "Native (Rust) hot-path kernels for OriginStack: stacking combine, etc." diff --git a/ext/astro_native/pyproject.toml b/ext/astro_native/pyproject.toml index f380506..24850f1 100644 --- a/ext/astro_native/pyproject.toml +++ b/ext/astro_native/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "maturin" [project] name = "astro_native" -version = "0.35.0" +version = "0.36.0" description = "Native Rust hot-path kernels for OriginStack" requires-python = ">=3.10" classifiers = ["Programming Language :: Rust"] diff --git a/ext/astro_native/src/lib.rs b/ext/astro_native/src/lib.rs index 121d913..abde232 100644 --- a/ext/astro_native/src/lib.rs +++ b/ext/astro_native/src/lib.rs @@ -6469,6 +6469,151 @@ mod originvision { } } +// --------------------------------------------------------------------------- +// ZOGY transient triage: real/bogus classification of --transient-detect +// candidates (src/transient_triage.py, --transient-triage) +// --------------------------------------------------------------------------- +// +// A much smaller sibling of `mod originvision` above: no percentile stretch, +// no resize -- stamps arrive already extracted at a fixed size and +// per-channel sigma-normalized in Python (there is no perf-relevant work to +// move into Rust for a few hundred 31x31x3 stamps), and no multi-head +// metadata decode, just one scalar sigmoid probability per candidate. +// Batched: all of a frame's candidates score in a single tract forward pass +// over (N, 3, size, size), not one call per candidate. Advisory only -- +// never filters `detect_transients`' output, just attaches a +// `real_probability`. +// +// The bundled model (src/data/transient_triage.onnx) is trained entirely on +// synthetic data (tools/gen_transient_triage_data.py + +// tools/train_transient_triage.py) -- no labelled real transients exist yet +// -- so treat it as a first cut, not a production classifier. Deliberately +// has no numpy/onnxruntime fallback yet, unlike every other native kernel in +// this file: a source checkout without astro_native built simply can't use +// --transient-triage until one is added (self-disables with a warning, see +// src/transient_triage.py). +mod transient_triage { + use pyo3::prelude::*; + use std::collections::HashMap; + use std::sync::{Arc, Mutex, OnceLock}; + use tract_onnx::prelude::*; + + type Runnable = TypedRunnableModel; + + struct Session { + model: Runnable, + } + + fn cache() -> &'static Mutex>> { + static C: OnceLock>>> = OnceLock::new(); + C.get_or_init(|| Mutex::new(HashMap::new())) + } + + fn load(path: &str, _size: usize) -> TractResult { + // Unlike `originvision::load`, the batch axis here must stay + // symbolic: this kernel scores a whole frame's candidates in one + // batched call (N varies per frame), while originvision always + // calls with N=1. The channel/H/W axes are already concrete in the + // exported graph (torch.onnx.export's dynamic_axes only marks axis + // 0 as dynamic), so no `with_input_fact` override is needed -- one + // was tried and made every N != 1 call fail with a tract symbol + // resolution clash against the fixed batch=1 it forced. + let proto = tract_onnx::onnx().proto_model_for_path(path)?; + let model = tract_onnx::onnx() + .model_for_proto_model(&proto)? + .into_optimized()? + .into_runnable()?; + Ok(Session { model }) + } + + fn get_session(path: &str, size: usize) -> TractResult> { + // Key on (path, size): the input size is baked into the compiled + // graph by `load`'s `with_input_fact` + `into_optimized`, same + // reasoning as `originvision`'s cache above. + let key = format!("{path}\u{0}{size}"); + { + let c = cache().lock().unwrap(); + if let Some(s) = c.get(&key) { + return Ok(Arc::clone(s)); + } + } + let s = Arc::new(load(path, size)?); + cache().lock().unwrap().insert(key, Arc::clone(&s)); + Ok(s) + } + + fn sigmoid(x: f64) -> f64 { + 1.0 / (1.0 + (-x).exp()) + } + + /// Pure compute: batched forward pass over N pre-normalized stamps, no + /// `Python` token -- runs inside `py.detach`. + fn compute(stamps: &[f32], n: usize, size: usize, model_path: &str) -> Result, String> { + if n == 0 { + return Ok(Vec::new()); + } + let sess = get_session(model_path, size).map_err(|e| e.to_string())?; + // `stamps` is already NCHW-ordered per candidate: (n, 3, size, size). + let input = tract_ndarray::Array4::from_shape_vec((n, 3, size, size), stamps.to_vec()) + .map_err(|e| e.to_string())? + .into_tensor(); + let outputs = sess.model.run(tvec!(input.into())).map_err(|e| e.to_string())?; + let raw = outputs + .first() + .ok_or_else(|| "model produced no output".to_string())? + .to_array_view::() + .map_err(|e| e.to_string())?; + if raw.len() != n { + return Err(format!( + "model produced {} outputs for {n} candidates -- refusing to guess", + raw.len() + )); + } + Ok(raw.iter().map(|&v| sigmoid(v as f64)).collect()) + } + + #[pyfunction] + #[pyo3(signature = (stamps, model_path, size=31))] + pub fn transient_triage_score( + py: Python<'_>, + stamps: numpy::PyReadonlyArray4, + model_path: &str, + size: usize, + ) -> PyResult> { + let a = stamps.as_array(); + let sh = a.shape(); + if sh.len() != 4 || sh[1] != 3 || sh[2] != size || sh[3] != size { + return Err(pyo3::exceptions::PyValueError::new_err(format!( + "expected stamps shaped (N, 3, {size}, {size}), got {:?}", + sh + ))); + } + let n = sh[0]; + let flat: Vec = a.iter().cloned().collect(); + let model_path = model_path.to_string(); + + // Heavy work off the GIL, panic-guarded -- same reasoning as + // `originvision_score`: `tract` parses an external .onnx file and a + // malformed one can panic inside the parser, which would otherwise + // unwind past callers' `except Exception` as a bare PanicException. + let outcome = py.detach(|| { + std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + compute(&flat, n, size, &model_path) + })) + }); + + match outcome { + Ok(Ok(probs)) => Ok(probs), + Ok(Err(msg)) => Err(pyo3::exceptions::PyRuntimeError::new_err(format!( + "transient_triage native inference failed: {msg}" + ))), + Err(_) => Err(pyo3::exceptions::PyRuntimeError::new_err( + "transient_triage native inference panicked (malformed model?)".to_string(), + )), + } + } +} + // --------------------------------------------------------------------------- // CFA drizzle: splat one frame's measured Bayer samples (src/cfa_drizzle.py) // --------------------------------------------------------------------------- @@ -7913,6 +8058,7 @@ fn white_balance_grayworld_inplace<'py>( #[pymodule] fn astro_native(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_function(wrap_pyfunction!(originvision::originvision_score, m)?)?; + m.add_function(wrap_pyfunction!(transient_triage::transient_triage_score, m)?)?; m.add_function(wrap_pyfunction!(sigma_clip_combine, m)?)?; m.add_function(wrap_pyfunction!(online_sigma_clip_combine, m)?)?; m.add_function(wrap_pyfunction!(online_sigma_clip_seed_burnin, m)?)?; diff --git a/packaging/verify_build.ps1 b/packaging/verify_build.ps1 index 5ac17c1..7832572 100644 --- a/packaging/verify_build.ps1 +++ b/packaging/verify_build.ps1 @@ -88,12 +88,14 @@ if (-not (Test-Path $synthDir)) { throw "synthetic_data was not created -- canno $outPath = "$env:TEMP\originstack_verify_out.fits" $headlessLog = "$env:TEMP\originstack_verify_stdout.txt" if (Test-Path $headlessLog) { Remove-Item $headlessLog -Force } -# --originvision exercises the bundled native scorer (astro_native.originvision_score -# + src/data/originvision.onnx). It self-disables with a warning if either is -# missing from the frozen build -- asserted absent below. +# originvision runs by default now (--no-originvision to disable; no positive +# flag exists, see cli.py) and exercises the bundled native scorer +# (astro_native.originvision_score + src/data/originvision.onnx). It +# self-disables with a warning if either is missing from the frozen build -- +# asserted absent below. $headlessArgs = @('--verify-headless', '-d', (Resolve-Path $synthDir).Path, '-o', $outPath, '--parallel', '4', '--debayer-method', 'malvar', - '--white-balance', 'grayworld', '--stack-method', 'median', '--originvision') + '--white-balance', 'grayworld', '--stack-method', 'median') $headlessProc = Start-Process -FilePath $ExePath -ArgumentList $headlessArgs -PassThru ` -RedirectStandardOutput $headlessLog diff --git a/requirements-dev.txt b/requirements-dev.txt index 97e7203..7c7044e 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -2,5 +2,6 @@ # pip install -r requirements-dev.txt ruff==0.16.8 pytest +pytest-xdist pip-audit bandit diff --git a/src/cli.py b/src/cli.py index e80b0ee..14cb170 100644 --- a/src/cli.py +++ b/src/cli.py @@ -942,6 +942,10 @@ def process_directory(directory: str, output: str, args: argparse.Namespace): produced = [] _effective_args: dict = {} # produced stack path -> the args it was made with for target_idx, (d, outp) in enumerate(targets, 1): + _cancel_event = getattr(args, '_cancel_event', None) + if _cancel_event is not None and _cancel_event.is_set(): + from src.models import RunCancelled + raise RunCancelled(f"cancelled before target {target_idx}/{len(targets)}") if _baseline is not None: _restore_args(args, _baseline) @@ -1828,6 +1832,21 @@ def build_parser() -> argparse.ArgumentParser: 'corrected score image (default: 5.0). The score is calibrated ' '(source + astrometric noise are propagated), so this is a real ' 'significance, not an arbitrary cut.') + g_post.add_argument('--transient-triage', action='store_true', + help='Score each --transient-detect candidate with a small CNN for a ' + 'real_probability (real transient vs. cosmic ray / registration-slip ' + 'dipole / hot pixel) -- the same role ZTF\'s BTSbot / Rubin\'s DIA ' + 'triage play downstream of classical image differencing. Advisory ' + 'only: never drops a candidate, just adds a column to ' + '_transients.csv. Native-only (astro_native.transient_triage_score, ' + 'no onnxruntime fallback yet) -- self-disables with a warning if ' + 'unavailable. The bundled model is trained entirely on synthetic data ' + '(tools/gen_transient_triage_data.py + tools/train_transient_triage.py, ' + 'no labelled real transients exist yet), so treat it as a first cut. ' + 'Requires --transient-detect.') + g_post.add_argument('--transient-triage-model', default=None, metavar='PATH', + help='Override the bundled transient-triage model ' + '(src/data/transient_triage.onnx). Requires --transient-triage.') g_post.add_argument('--light-pollution-azimuth', type=float, default=0.0, metavar='DEG', help='physical bg-method only: compass azimuth (degrees east of north) ' @@ -1915,21 +1934,33 @@ def build_parser() -> argparse.ArgumentParser: 'accepted, rejection_reason)') g_debug.add_argument('--export-frames-dir', default=None, metavar='PATH', help='Directory to write a stretched JPEG for every accepted frame after Phase 1') - g_originvision.add_argument('--originvision', action='store_true', - help='Score the final stacked master with originvision (separately-trained ' - 'defect/quality/category classifier), run in-process against the ' + # Single action, no positive `--originvision` flag -- same reason --auto + # (src/cli.py:_UNSUPPORTED... see the --no-auto entry above) has none: + # desktop_control.py's get_form_schema()/build_argv_from_form() key + # purely on argparse dest, with no dest-collision handling, so two + # actions sharing one dest would silently double up the GUI field (a + # second tk.Variable, one of them dropped) or pick one arbitrarily in + # _dest_action_map's dest->action dict. `--originvision` text on an + # existing command line now errors (unrecognized argument) rather than + # being a redundant no-op -- deliberate, matching --auto's own + # no-positive-flag precedent, not an oversight. + g_originvision.add_argument('--no-originvision', dest='originvision', action='store_false', + default=True, + help='Disable originvision scoring (defect/quality/category classifier, on ' + 'by default). Scores the final stacked master in-process against the ' 'bundled model (src/data/originvision.onnx -- no external folder or ' 'venv). Inference is the native astro_native kernel (pure-Rust tract, ' 'nothing extra to install); a source checkout without astro_native ' - 'falls back to a Python onnxruntime path. When --auto is also active ' - '(the default -- pass --no-auto to disable), also samples 3 light ' - 'frames spread through the session: the sampled category feeds the ' - 'same target-classification prior SIMBAD/header metadata uses, and a ' - 'defect flag nudges settings defensively (trail-reject, stronger ' - 'chroma denoising) -- never auto-rejects a frame, this model is still ' - 'finishing its first training run. Pair with --originvision-score-all ' - 'to also score every accepted frame (slower on a large session). ' - 'Self-disables with a warning when no backend is available.') + 'falls back to a Python onnxruntime path -- and self-disables with a ' + 'warning if neither backend nor the model file is available, so this ' + 'runs cleanly either way. When --auto is also active (the default -- ' + 'pass --no-auto to disable), also samples 3 light frames spread through ' + 'the session: the sampled category feeds the same target-classification ' + 'prior SIMBAD/header metadata uses, and a defect flag nudges settings ' + 'defensively (trail-reject, stronger chroma denoising) -- never ' + 'auto-rejects a frame, this model is still finishing its first training ' + 'run. Pair with --originvision-score-all to also score every accepted ' + 'frame (slower on a large session).') g_originvision.add_argument('--originvision-score-all', action='store_true', help='Also score every accepted light frame with originvision (not just ' 'the fast 3-frame sample --originvision always does), logging advisory ' @@ -1939,11 +1970,13 @@ def build_parser() -> argparse.ArgumentParser: g_originvision.add_argument('--originvision-model', default=None, metavar='PATH', help='Path to an exported originvision ONNX model, overriding the bundled ' 'src/data/originvision.onnx (e.g. to test a newer checkpoint).') - g_originvision.add_argument('--originvision-workers', type=int, default=2, metavar='N', - help='Thread-pool size for per-frame originvision scoring calls (default: 2). ' + g_originvision.add_argument('--originvision-workers', type=int, default=8, metavar='N', + help='Thread-pool size for per-frame originvision scoring calls (default: 8). ' 'Both backends release the GIL during the forward pass (the native ' 'tract kernel via py.allow_threads), so a thread pool parallelises it ' - 'without a ProcessPoolExecutor.') + 'without a ProcessPoolExecutor -- measured near-linear scaling 1->8 ' + 'workers on a real frame (2751 -> 473 ms/frame effective at 8), so the ' + 'previous default of 2 (1415 ms/frame) left real throughput on the table.') # Back-compat, hidden: --originvision-dir / --originvision-checkpoint still # resolve a model path for command lines written against the pre-in-process # layout. @@ -2472,6 +2505,11 @@ def _passed_on_cli(action) -> bool: # for "ran but found nothing" rather than "didn't run at all". safe_print(" WARNING: --originvision-score-all has no effect without --originvision") + if getattr(args, 'transient_triage', False) and not getattr(args, 'transient_detect', None): + safe_print(" WARNING: --transient-triage has no effect without --transient-detect") + if getattr(args, 'transient_triage_model', None) and not getattr(args, 'transient_triage', False): + safe_print(" WARNING: --transient-triage-model has no effect without --transient-triage") + return args diff --git a/src/data/transient_triage.onnx b/src/data/transient_triage.onnx new file mode 100644 index 0000000..a73e74d Binary files /dev/null and b/src/data/transient_triage.onnx differ diff --git a/src/desktop_app.py b/src/desktop_app.py index 232a2d0..b452ae5 100644 --- a/src/desktop_app.py +++ b/src/desktop_app.py @@ -883,6 +883,9 @@ def _build_left(self, parent: ttk.Frame) -> None: self.start_btn = ttk.Button(run_row, text='Start', style='Accent.TButton', command=self._on_start) self.start_btn.pack(side='left') + self.cancel_btn = ttk.Button(run_row, text='Cancel', command=self._on_cancel) + self.cancel_btn.pack(side='left', padx=(6, 0)) + self.cancel_btn.state(['disabled']) self.status_var = tk.StringVar(value='Idle') ttk.Label(run_row, textvariable=self.status_var, style='Dim.TLabel').pack( side='left', padx=10) @@ -978,11 +981,21 @@ def _on_start(self) -> None: self.status_var.set(f"Error: {result.get('error')}") return self.start_btn.state(['disabled']) + self.cancel_btn.state(['!disabled']) self.status_var.set('Running…') self._shown_log_lines = 0 self.frames_tree.delete(*self.frames_tree.get_children()) self.summary_var.set('') + def _on_cancel(self) -> None: + # Cooperative, not instant -- takes effect at the next checkpoint + # (RunManager/frame_processor._check_cancel's docstrings). Disabling + # the button immediately is the honest signal: pressing it again + # wouldn't make the pipeline notice any sooner. + self.rm.cancel() + self.cancel_btn.state(['disabled']) + self.status_var.set('Cancelling…') + def _on_closing(self) -> None: if self.rm.is_running(): from src.native_dialog import ask_yes_no @@ -1135,13 +1148,18 @@ def _refresh_run_button(self, snap: Optional[Dict[str, Any]] = None) -> None: running = self.rm.is_running() if running: self.start_btn.state(['disabled']) - self.header_status_var.set('Running…') + self.header_status_var.set('Cancelling…' if self.rm.is_cancelling() else 'Running…') else: self.start_btn.state(['!disabled']) - if snap is not None and snap['run_status'] in ('ok', 'error'): + self.cancel_btn.state(['disabled']) + if snap is not None and snap['run_status'] in ('ok', 'error', 'cancelled'): done_ok = snap['run_status'] == 'ok' - self.status_var.set('Done' if done_ok else f"Failed: {snap['run_error']}") - self.header_status_var.set('Complete' if done_ok else 'Failed') + if snap['run_status'] == 'cancelled': + self.status_var.set('Cancelled') + self.header_status_var.set('Cancelled') + else: + self.status_var.set('Done' if done_ok else f"Failed: {snap['run_error']}") + self.header_status_var.set('Complete' if done_ok else 'Failed') # Edge-triggered (not every poll tick) so it only pops once # per completed run, not repeatedly while the status holds. if done_ok and self._last_run_status != 'ok' and self.open_folder_var.get(): diff --git a/src/desktop_control.py b/src/desktop_control.py index 209d1a1..1f8ad06 100644 --- a/src/desktop_control.py +++ b/src/desktop_control.py @@ -146,16 +146,35 @@ class RunManager: progress through the ``UIEvents`` singleton (``src/ui_events.py``) -- the desktop app attaches it before entering the tkinter mainloop, so it's already active by the time a run starts; only the pipeline work - itself needs to move off the GUI's own (main) thread.""" + itself needs to move off the GUI's own (main) thread. + + Cancellation is cooperative, not a thread kill (CPython threads can't be + force-stopped, and Phase 1 workers are real OS processes/subprocesses + mid-computation anyway): ``cancel()`` sets a ``threading.Event`` shared + onto the run's ``args`` as ``_cancel_event``, and the pipeline notices it + at a handful of checkpoints (``frame_processor._check_cancel`` between + Phase 1 frames -- usually the longest phase -- and between targets in + ``cli.process_directory``). Phases 2-4 of a single target aren't + interruptible yet: once one starts, it runs to completion.""" def __init__(self) -> None: self._lock = threading.Lock() self.status = 'idle' self.thread: Optional[threading.Thread] = None + self._cancel_event = threading.Event() def is_running(self) -> bool: return self.status == 'running' + def is_cancelling(self) -> bool: + return self.status == 'running' and self._cancel_event.is_set() + + def cancel(self) -> None: + """Request a stop. A no-op if nothing is running; safe to call more + than once. Takes effect at the next checkpoint, not instantly.""" + if self.status == 'running': + self._cancel_event.set() + def start(self, form: Dict[str, Any]) -> Dict[str, Any]: with self._lock: if self.status == 'running': @@ -165,6 +184,7 @@ def start(self, form: Dict[str, Any]) -> Dict[str, Any]: except Exception as e: return {'ok': False, 'error': f'invalid form: {e}'} self.status = 'running' + self._cancel_event = threading.Event() # fresh flag per run self.thread = threading.Thread(target=self._run, args=(argv,), name='desktop-run', daemon=True) self.thread.start() @@ -175,6 +195,7 @@ def _run(self, argv: List[str]) -> None: import tempfile from src.cli import apply_post_parse_setup, parse_args, process_directory + from src.models import RunCancelled from src.ui_events import get_ui_events from src.utils import get_logger, safe_print @@ -182,6 +203,7 @@ def _run(self, argv: List[str]) -> None: status, error = 'ok', None try: args = parse_args(argv) + args._cancel_event = self._cancel_event # Default a durable log file for GUI-triggered runs specifically # (not in apply_post_parse_setup, which is also the plain CLI's @@ -199,6 +221,9 @@ def _run(self, argv: List[str]) -> None: wv.run_started() process_directory(args.directory, args.output, args) + except RunCancelled: + status = 'cancelled' + safe_print(" Run cancelled.") except (Exception, SystemExit) as e: status = 'error' error = str(e) or e.__class__.__name__ diff --git a/src/difference_imaging.py b/src/difference_imaging.py index 8f7b059..5681702 100644 --- a/src/difference_imaging.py +++ b/src/difference_imaging.py @@ -78,6 +78,7 @@ class Transient(NamedTuple): x: float significance: float # S_corr value, in sigma kind: str # 'brightening' | 'fading' + real_probability: Optional[float] = None # --transient-triage; None if not scored def _to_luminance(img: np.ndarray) -> np.ndarray: @@ -372,15 +373,20 @@ def write_transient_catalog(path: str, transients: Sequence[Transient], has_wcs = wcs is not None n_failed = 0 first_error = None + has_triage = any(t.real_probability is not None for t in transients) with open(path, 'w', newline='', encoding='utf-8') as fh: writer = csv.writer(fh) header = ['x', 'y', 'significance_sigma', 'kind'] + if has_triage: + header += ['real_probability'] if has_wcs: header += ['ra_deg', 'dec_deg'] writer.writerow(header) for t in transients: row = [f"{t.x:.2f}", f"{t.y:.2f}", f"{t.significance:.2f}", t.kind] + if has_triage: + row += [f"{t.real_probability:.3f}" if t.real_probability is not None else ''] if has_wcs: try: ra, dec = wcs.all_pix2world(t.x, t.y, 0) @@ -410,54 +416,45 @@ def write_transient_catalog(path: str, transients: Sequence[Transient], _RESIDUAL_MATCH_TOL_PX = 3.0 -def run_transient_detection(stacked: np.ndarray, reference_path: str, - output_path: str, threshold: float = 5.0, - wcs=None) -> Optional[dict]: - """Orchestrate a two-epoch comparison: align, subtract, detect, report. - - Returns a summary dict, or None when the comparison could not be set up - (missing or non-linear reference, too few stars to estimate a PSF, no - overlap). Setup failures are reported and return None; an I/O failure - writing the outputs (disk full, read-only directory) *does* raise, so the - caller must still guard this -- ``pipeline.py`` does. +class EpochComparison(NamedTuple): + """The core two-epoch ZOGY comparison result -- everything + ``run_transient_detection`` needs to write its outputs, and everything + ``tools/gen_transient_triage_data.py``'s real-data mining mode needs to + build training stamps, without going through that function's file I/O.""" + transients: List[Transient] + difference: np.ndarray # D, NaN outside the reference footprint + score_corr: np.ndarray # S_corr, NaN outside the reference footprint + new_lum: np.ndarray # pedestal-subtracted, not yet warped (it's the reference frame) + ref_lum: np.ndarray # pedestal-subtracted AND warped onto new_lum's grid + covered: float + measured_axis: Optional[float] + astro_sigma_px: float + flux_ratio: float + + +def _compare_epochs(stacked: np.ndarray, ref: np.ndarray, + threshold: float = 5.0) -> Optional[EpochComparison]: + """Align two RGB (or 2D) epochs, run ZOGY, and detect candidates. + + ``stacked``/``ref`` are raw pixel arrays (RGB or luminance), not yet + reduced to luminance or background-subtracted -- this does both, then + registration, PSF estimation, ``zogy()`` and ``detect_transients``. + Returns ``None`` on any setup failure (logged here), same conditions + ``run_transient_detection`` always reported at this call site. """ - from src.io_fits import load_fits - - if not os.path.exists(reference_path): - safe_print(f" WARNING: transient reference not found: {reference_path}") - return None - - try: - ref, ref_header = load_fits(reference_path) - except Exception as exc: - safe_print(f" WARNING: could not read transient reference: {exc}") - return None - - # The comparison is only meaningful between two LINEAR stacks. Phase 4's - # stretches, denoisers and local contrast break photometric linearity, so - # a post-processed reference mismatches the flux scale by a fraction of a - # percent -- several sigma on a bright star -- and every star in the field - # reports as a confident transient. --merge refuses the same file for the - # same reason. - if not bool((ref_header or {}).get('RAWSTACK', False)): - safe_print(f" WARNING: {os.path.basename(reference_path)} is not a linear " - f"(pre-post-processing) stack: header RAWSTACK is missing or " - f"False. Pass the main output FITS of a previous run, not the " - f"_processed one -- skipping difference imaging.") - return None - - # The pipeline writes RGB planes as (C, H, W); load_fits hands them back - # in that order, while everything here works in (H, W, C). - ref = np.asarray(ref) - if ref.ndim == 3 and ref.shape[0] in (3, 4) and ref.shape[0] < ref.shape[-1]: - ref = np.transpose(ref, (1, 2, 0)) - new_lum = _to_luminance(stacked) ref_lum = _to_luminance(ref) if new_lum.shape != ref_lum.shape: - safe_print(f" WARNING: reference epoch is {ref_lum.shape}, this stack is " - f"{new_lum.shape} -- cannot compare different frame sizes") - return None + # Two independently-stacked sessions of the same target routinely + # differ in pixel dimensions -- different dither pattern, different + # Phase 3 common-crop -- even though they're the same field. This + # used to be a hard failure; `_align_reference` now embeds the + # reference onto this stack's own grid (same trick `merge.py` uses + # for its differently-shaped previous stacks) before the blind + # star-pattern match, which doesn't need or assume equal shapes or + # any positional correspondence between the two canvases anyway. + safe_print(f" reference epoch is {ref_lum.shape}, this stack is " + f"{new_lum.shape} -- reconciling onto a common grid") # ZOGY assumes background-subtracted inputs, so remove each epoch's own # sky level rather than trusting them to share one. This is not a @@ -505,6 +502,86 @@ def run_transient_detection(stacked: np.ndarray, reference_path: str, covered = float(valid.mean()) transients = detect_transients(score_corr, threshold=threshold) + return EpochComparison(transients=transients, difference=difference, + score_corr=score_corr, new_lum=new_lum, ref_lum=ref_lum, + covered=covered, measured_axis=measured_axis, + astro_sigma_px=astro_sigma_px, flux_ratio=flux_ratio) + + +def run_transient_detection(stacked: np.ndarray, reference_path: str, + output_path: str, threshold: float = 5.0, + wcs=None, triage: bool = False, + triage_model_path: Optional[str] = None) -> Optional[dict]: + """Orchestrate a two-epoch comparison: align, subtract, detect, report. + + ``triage`` (``--transient-triage``) additionally scores each candidate + with a small CNN (``src/transient_triage.py``) for a ``real_probability`` + -- advisory only, never drops a candidate. + + Returns a summary dict, or None when the comparison could not be set up + (missing or non-linear reference, too few stars to estimate a PSF, no + overlap). Setup failures are reported and return None; an I/O failure + writing the outputs (disk full, read-only directory) *does* raise, so the + caller must still guard this -- ``pipeline.py`` does. + """ + from src.io_fits import load_fits + + if not os.path.exists(reference_path): + safe_print(f" WARNING: transient reference not found: {reference_path}") + return None + + try: + ref, ref_header = load_fits(reference_path) + except Exception as exc: + safe_print(f" WARNING: could not read transient reference: {exc}") + return None + + # The comparison is only meaningful between two LINEAR stacks. Phase 4's + # stretches, denoisers and local contrast break photometric linearity, so + # a post-processed reference mismatches the flux scale by a fraction of a + # percent -- several sigma on a bright star -- and every star in the field + # reports as a confident transient. --merge refuses the same file for the + # same reason. + if not bool((ref_header or {}).get('RAWSTACK', False)): + safe_print(f" WARNING: {os.path.basename(reference_path)} is not a linear " + f"(pre-post-processing) stack: header RAWSTACK is missing or " + f"False. Pass the main output FITS of a previous run, not the " + f"_processed one -- skipping difference imaging.") + return None + + # The pipeline writes RGB planes as (C, H, W); load_fits hands them back + # in that order, while everything here works in (H, W, C). + ref = np.asarray(ref) + if ref.ndim == 3 and ref.shape[0] in (3, 4) and ref.shape[0] < ref.shape[-1]: + ref = np.transpose(ref, (1, 2, 0)) + + comparison = _compare_epochs(stacked, ref, threshold=threshold) + if comparison is None: + return None + transients = comparison.transients + difference = comparison.difference + score_corr = comparison.score_corr + new_lum = comparison.new_lum + ref_lum = comparison.ref_lum + covered = comparison.covered + measured_axis = comparison.measured_axis + astro_sigma_px = comparison.astro_sigma_px + flux_ratio = comparison.flux_ratio + + n_triaged = 0 + if triage and transients: + from src.transient_triage import score_candidates + # Recomputed rather than threaded out of `zogy()`'s ZogyResult -- + # cheap (same robust-sigma estimator, run on arrays already in hand) + # and avoids widening that return type for an opt-in feature. + sigma_new = estimate_background_sigma(new_lum) + sigma_ref = estimate_background_sigma(ref_lum) + sigma_diff = estimate_background_sigma(difference) + probs = score_candidates(new_lum, ref_lum, difference, transients, + sigma_new, sigma_ref, sigma_diff, + model_path=triage_model_path) + transients = [t._replace(real_probability=p) for t, p in zip(transients, probs)] + n_triaged = sum(1 for p in probs if p is not None) stem = os.path.splitext(output_path)[0] _write_fits_plane(stem + '_difference.fits', difference, @@ -520,6 +597,16 @@ def run_transient_detection(stacked: np.ndarray, reference_path: str, safe_print(f" Difference imaging: {len(transients)} candidate(s) above " f"{threshold:g} sigma ({n_bright} brightening, " f"{len(transients) - n_bright} fading)") + if triage: + if n_triaged: + n_likely = sum(1 for t in transients + if t.real_probability is not None and t.real_probability > 0.5) + safe_print(f" triage: {n_triaged}/{len(transients)} scored, " + f"{n_likely} likely real (real_probability > 0.5) -- " + f"advisory only, nothing was dropped") + else: + safe_print(" triage: requested but unavailable (see warning above) " + "-- candidates left unscored") if measured_axis is None: reg = (f"registration residual not measurable -- assumed " f"{astro_sigma_px:.2f} px/axis") @@ -563,13 +650,22 @@ def _align_reference(new_lum: np.ndarray, ref_lum: np.ndarray): Cross-night pairs differ by arbitrary field rotation on an alt-az mount, so this goes through the same blind star-pattern matcher ``--merge`` uses - rather than assuming a pure translation. + rather than assuming a pure translation. The matcher itself needs no + positional correspondence between ``new_lum``/``ref_lum`` -- it matches on + relative star geometry -- so unequal shapes are reconciled first by + embedding ``ref_lum`` onto ``new_lum``'s grid (top-left, zero-padded; + ``src.utils.embed_to_shape``, the same trick ``merge.py`` uses for a + previous stack whose own shape rarely matches the current run's): a + no-op when the shapes already match. Returns ``(warped_ref, footprint, residual_px, new_stars)`` or None: - ``footprint`` is the warped reference's coverage in [0, 1] -- the same - transform applied to an all-ones image. Outside it the reference is - fill, not data. + transform applied to a mask of where ``ref_lum`` had real data (ones + only inside its own original extent, before any embedding). Outside it + the reference is fill, not data -- and that now covers both the warp's + own uncovered wedges (field rotation) and any embed-padding border, so + neither reads as a bogus "transient" the way an all-ones mask would. - ``residual_px`` is the RMS 2D distance between matched star pairs after the transform: a real measurement of how well the epochs line up, which feeds ZOGY's astrometric noise term. None when too few pairs match to @@ -579,6 +675,13 @@ def _align_reference(new_lum: np.ndarray, ref_lum: np.ndarray): """ from src.registration import apply_transform from src.star_detect import detect_stars_matched_filter + from src.utils import embed_to_shape + + ref_valid_mask = np.ones_like(ref_lum, dtype=np.float32) + if ref_lum.shape != new_lum.shape: + H, W = new_lum.shape + ref_valid_mask = embed_to_shape(ref_valid_mask, H, W) + ref_lum = embed_to_shape(ref_lum, H, W) try: new_stars = detect_stars_matched_filter(new_lum.astype(np.float32)) @@ -602,8 +705,7 @@ def _align_reference(new_lum: np.ndarray, ref_lum: np.ndarray): try: warped = apply_transform(ref_lum.astype(np.float32), transform=transform) - footprint = apply_transform(np.ones_like(ref_lum, dtype=np.float32), - transform=transform) + footprint = apply_transform(ref_valid_mask, transform=transform) except Exception as exc: _log.debug("transient alignment: warp failed (%s)", exc) return None diff --git a/src/frame_processor.py b/src/frame_processor.py index 3abc13d..40a061a 100644 --- a/src/frame_processor.py +++ b/src/frame_processor.py @@ -33,7 +33,7 @@ from src.frame_discovery import is_nebula_filter from src.gpu_context import get_gpu from src.io_fits import load_frame -from src.models import Config, FrameInfo, ProcessingStats +from src.models import Config, FrameInfo, ProcessingStats, RunCancelled from src.quality import compute_quality_metrics, estimate_bortle, validate_image_data from src.stacking import lacosmic_reject from src.utils import format_time, mp_context, print_quality_table, safe_print @@ -915,6 +915,17 @@ def _parallel_frame_worker( return (frame_idx, metrics_clean, None, timings) +def _check_cancel(args: argparse.Namespace) -> None: + """Raise ``RunCancelled`` when the GUI's cancel button has been pressed + (``RunManager.cancel()``, ``args._cancel_event``). A plain CLI run never + sets this, so it's a no-op there. Called between frames -- not inside a + worker, which has already committed to processing the frame it picked + up -- so cancelling stops new work starting, not work in flight.""" + ev = getattr(args, '_cancel_event', None) + if ev is not None and ev.is_set(): + raise RunCancelled("cancelled during Phase 1 frame processing") + + @_with_session_cfa def execute_frame_processing( lights: List[FrameInfo], @@ -1025,40 +1036,51 @@ def _accum(timings: Optional[dict]) -> None: futures = {pool.submit(_parallel_frame_worker, t): t[1] for t in tasks} _wv = _get_ui_events() _wv_done = 0 - for future in tqdm(as_completed(futures), total=n, - desc=" Processing", unit="frame", - disable=args.verbose): - idx = futures[future] - frame_idx, metrics, error, timings = future.result() - _accum(timings) - f = lights[frame_idx] - _wv_done += 1 - _wv.progress('Processing frames', _wv_done, n) - _wv.frame_metrics(os.path.basename(f.path), metrics, - accepted=error is None) - if error: - f.accepted = False - f.metrics = {'error': error} - rejected_reasons[f.path] = error - stats.add_error(f.path, error) - if args.verbose: - print(f' REJECT {os.path.basename(f.path)}: {error}') - else: - f.metrics = metrics - _publish_frame_thumb(_wv, args, - os.path.basename(f.path), - mem_rgb[frame_idx], _wv_thumb_count) - if args.verbose: - m = f.metrics - safe_print(f' {os.path.basename(f.path)}: ' - f'score={m["score"]:.0f} SNR={m["snr"]:.1f} ' - f'stars={m["star_count"]} FWHM={m.get("fwhm",0):.1f} ' - f'sharpness={m.get("sharpness",0):.0f}') - safe_print(f' bg={m.get("background",0):.1f} ' - f'noise={m.get("noise",0):.2f} ' - f'brightness={m.get("brightness",0):.1f} ' - f'contrast={m.get("contrast",0):.1f} ' - f'dynamic_range={m.get("dynamic_range",0):.0f}') + try: + for future in tqdm(as_completed(futures), total=n, + desc=" Processing", unit="frame", + disable=args.verbose): + _check_cancel(args) + idx = futures[future] + frame_idx, metrics, error, timings = future.result() + _accum(timings) + f = lights[frame_idx] + _wv_done += 1 + _wv.progress('Processing frames', _wv_done, n) + _wv.frame_metrics(os.path.basename(f.path), metrics, + accepted=error is None) + if error: + f.accepted = False + f.metrics = {'error': error} + rejected_reasons[f.path] = error + stats.add_error(f.path, error) + if args.verbose: + safe_print(f' REJECT {os.path.basename(f.path)}: {error}') + else: + f.metrics = metrics + _publish_frame_thumb(_wv, args, + os.path.basename(f.path), + mem_rgb[frame_idx], _wv_thumb_count) + if args.verbose: + m = f.metrics + safe_print(f' {os.path.basename(f.path)}: ' + f'score={m["score"]:.0f} SNR={m["snr"]:.1f} ' + f'stars={m["star_count"]} FWHM={m.get("fwhm",0):.1f} ' + f'sharpness={m.get("sharpness",0):.0f}') + safe_print(f' bg={m.get("background",0):.1f} ' + f'noise={m.get("noise",0):.2f} ' + f'brightness={m.get("brightness",0):.1f} ' + f'contrast={m.get("contrast",0):.1f} ' + f'dynamic_range={m.get("dynamic_range",0):.0f}') + except RunCancelled: + # Cancel every not-yet-started future so the `with` block's + # own shutdown(wait=True) on the way out only waits for + # whatever's already mid-frame in the worker processes -- + # not the full remainder of the session (every future was + # submitted upfront, so a plain exit here would otherwise + # wait for all n frames regardless of the cancel). + pool.shutdown(wait=False, cancel_futures=True) + raise finally: # Release shared memory after all workers are done. for shm in shm_blocks: @@ -1150,54 +1172,66 @@ def _thread_process_frame(i, f): # quality metrics are computed asynchronously and only printed AFTER # this loop, so disabling the bar under -v would leave the whole GPU # processing loop with no output at all (looks hung). - for future in tqdm(as_completed(futures), total=n, - desc=" Processing", unit="frame", - disable=False): - i, metrics, error, lum_arr, timings = future.result() - _accum(timings) - f = lights[i] - if error: - f.accepted = False - f.metrics = {'error': error} - rejected_reasons[f.path] = error - stats.add_error(f.path, error) - if args.verbose: - safe_print(f' REJECT {os.path.basename(f.path)}: {error}') - else: - cached_lums[i] = lum_arr - if _use_qpool and lum_arr is not None: - # Submit quality to CPU pool; GPU thread is already freed. - # Its compute time runs concurrently with other frames' - # GPU work and isn't attributable to a single frame here, - # so it is not folded into the per-step totals below — - # the GPU path's timing breakdown is best-effort. - _qfuts[i] = _qpool.submit( - compute_quality_metrics, lum_arr, advanced_metrics=_adv) - else: - f.metrics = metrics + try: + for future in tqdm(as_completed(futures), total=n, + desc=" Processing", unit="frame", + disable=False): + _check_cancel(args) + i, metrics, error, lum_arr, timings = future.result() + _accum(timings) + f = lights[i] + if error: + f.accepted = False + f.metrics = {'error': error} + rejected_reasons[f.path] = error + stats.add_error(f.path, error) if args.verbose: - m = f.metrics - safe_print(f' {os.path.basename(f.path)}: ' - f'score={m["score"]:.0f} SNR={m["snr"]:.1f} ' - f'stars={m["star_count"]} FWHM={m.get("fwhm",0):.1f} ' - f'sharpness={m.get("sharpness",0):.0f}') - safe_print(f' bg={m.get("background",0):.1f} ' - f'noise={m.get("noise",0):.2f} ' - f'brightness={m.get("brightness",0):.1f} ' - f'contrast={m.get("contrast",0):.1f} ' - f'dynamic_range={m.get("dynamic_range",0):.0f}') - _completed += 1 - _wv = _get_ui_events() - _wv.progress('Processing frames', _completed, n) - if error is None and f.metrics: - _wv.frame_metrics(os.path.basename(f.path), f.metrics) - if error is None: - _publish_frame_thumb(_wv, args, os.path.basename(f.path), - mem_rgb[i], _wv_thumb_count) - # Periodically free CuPy's cached memory pool to prevent VRAM exhaustion - # from accumulating unused cached blocks across many completed frames. - if gpu.active and (_completed % _free_interval == 0): - gpu.free_pool() + safe_print(f' REJECT {os.path.basename(f.path)}: {error}') + else: + cached_lums[i] = lum_arr + if _use_qpool and lum_arr is not None: + # Submit quality to CPU pool; GPU thread is already freed. + # Its compute time runs concurrently with other frames' + # GPU work and isn't attributable to a single frame here, + # so it is not folded into the per-step totals below — + # the GPU path's timing breakdown is best-effort. + _qfuts[i] = _qpool.submit( + compute_quality_metrics, lum_arr, advanced_metrics=_adv) + else: + f.metrics = metrics + if args.verbose: + m = f.metrics + safe_print(f' {os.path.basename(f.path)}: ' + f'score={m["score"]:.0f} SNR={m["snr"]:.1f} ' + f'stars={m["star_count"]} FWHM={m.get("fwhm",0):.1f} ' + f'sharpness={m.get("sharpness",0):.0f}') + safe_print(f' bg={m.get("background",0):.1f} ' + f'noise={m.get("noise",0):.2f} ' + f'brightness={m.get("brightness",0):.1f} ' + f'contrast={m.get("contrast",0):.1f} ' + f'dynamic_range={m.get("dynamic_range",0):.0f}') + _completed += 1 + _wv = _get_ui_events() + _wv.progress('Processing frames', _completed, n) + if error is None and f.metrics: + _wv.frame_metrics(os.path.basename(f.path), f.metrics) + if error is None: + _publish_frame_thumb(_wv, args, os.path.basename(f.path), + mem_rgb[i], _wv_thumb_count) + # Periodically free CuPy's cached memory pool to prevent VRAM exhaustion + # from accumulating unused cached blocks across many completed frames. + if gpu.active and (_completed % _free_interval == 0): + gpu.free_pool() + except RunCancelled: + # Same reasoning as the ProcessPool path above: every future + # was submitted upfront, so cancel the ones not yet started + # before letting the `with` block's own shutdown wait only on + # whatever's already in flight. + if _use_qpool: + _qpool.shutdown(wait=False, cancel_futures=True) + _io_pool.shutdown(wait=False, cancel_futures=True) + executor.shutdown(wait=False, cancel_futures=True) + raise # Collect deferred quality results (CPU pool runs while GPU was active) if _qfuts: @@ -1236,6 +1270,7 @@ def _thread_process_frame(i, f): for i, f in tqdm(enumerate(lights), total=n, desc=" Processing", unit="frame", disable=args.verbose): + _check_cancel(args) result = _process_single_frame( f.path, f.header, masters, args.debayer_method, args.white_balance, ca_correction=getattr(args, 'ca_correction', False), diff --git a/src/merge.py b/src/merge.py index ed17d80..6949425 100644 --- a/src/merge.py +++ b/src/merge.py @@ -36,6 +36,7 @@ import numpy as np from src.models import Config +from src.utils import embed_to_shape as _embed_to_shape from src.utils import get_logger, safe_print _log = get_logger() @@ -57,22 +58,6 @@ def _detect_stars(lum: np.ndarray) -> Optional[Any]: return None -def _embed_to_shape(arr: np.ndarray, H: int, W: int) -> np.ndarray: - """Place ``arr`` in the top-left of an (H, W[, C]) zero canvas (crop if - larger). Pixel coordinates are preserved, so a transform computed on the - embedded luminance applies directly to the embedded image.""" - if arr.shape[0] == H and arr.shape[1] == W: - return arr - if arr.ndim == 3: - out = np.zeros((H, W, arr.shape[2]), dtype=arr.dtype) - else: - out = np.zeros((H, W), dtype=arr.dtype) - h = min(H, arr.shape[0]) - w = min(W, arr.shape[1]) - out[:h, :w] = arr[:h, :w] - return out - - def _register_stack(new_lum: np.ndarray, prev_lum: np.ndarray, new_stars: Optional[Any]) -> Tuple[Optional[Any], Optional[Tuple[float, float]]]: diff --git a/src/models.py b/src/models.py index 2f38ad0..361f52a 100644 --- a/src/models.py +++ b/src/models.py @@ -6,6 +6,16 @@ from typing import Dict, List, Optional, Tuple +class RunCancelled(Exception): + """Raised at a cooperative checkpoint (``args._cancel_event`` set) to + unwind a run cleanly -- not a failure. Checked in the Phase 1 per-frame + loops (``frame_processor.execute_frame_processing``, the highest-value + spot: usually the longest-running phase) and between targets in + ``cli.process_directory`` (multi-session/hierarchical runs). Phases 2-4 + of a single target aren't interruptible yet -- once one starts it runs + to completion, same as before this existed.""" + + class Config: """Central configuration for magic numbers and thresholds.""" HOT_PIXEL_THRESHOLD = 12.0 diff --git a/src/originvision.py b/src/originvision.py index 9bf6ffa..762dfdd 100644 --- a/src/originvision.py +++ b/src/originvision.py @@ -64,7 +64,7 @@ def score_lights_with_originvision(lights: List[FrameInfo], args) -> None: model_path = _originvision_model(args) if model_path is None: return - workers = max(1, int(getattr(args, 'originvision_workers', 2))) + workers = max(1, int(getattr(args, 'originvision_workers', 8))) targets = [f for f in lights if f.accepted] if not targets: diff --git a/src/originvision_infer.py b/src/originvision_infer.py index 79d25f1..5b83ef9 100644 --- a/src/originvision_infer.py +++ b/src/originvision_infer.py @@ -197,12 +197,82 @@ def _prep_rgb(rgb: np.ndarray) -> Optional[np.ndarray]: return np.ascontiguousarray(arr[:, :, :3], dtype=np.float32) +def _downsample_if_large(arr: np.ndarray, max_long_side: int) -> np.ndarray: + """Shrink ``arr`` (aspect preserved, no crop) if its longer side exceeds + ``max_long_side``, same gaussian-prefiltered bilinear zoom as + ``_resize_center_crop`` -- just earlier, before the percentile stretch, + and without the crop. + + Measured on a full Origin sensor frame (1936x1096): the percentile + stretch + final resize scale with *input* pixel count even though only a + 256x256 crop is ever used, so at full resolution 85% of a + ``score_rgb`` call (1315 -> 196 ms) was spent processing pixels the model + never sees. Both backends call this identically (before the + native/onnxruntime dispatch below), so native/fallback parity holds and + this is a pure precomputation, not a behaviour fork. + + Not free of numerical effect -- a second resampling stage changes the + stretch's own percentile estimate and adds another antialiasing pass on + top of ``_resize_center_crop``'s. Validated against the un-downsampled + path the same way ``_resize_center_crop``'s own cv2->scipy swap was + (docstring above): category/defect/stray-light flags unchanged, quality + score within the same few-points-on-a-0-400-scale tolerance already + accepted there (session-relative only, never an absolute gate). + """ + h, w = arr.shape[:2] + long_side = max(h, w) + if long_side <= max_long_side: + return arr + scale = max_long_side / long_side + sigma = ((1.0 / scale) - 1.0) / 2.0 + f = arr + if sigma > 0.01: + f = ndimage.gaussian_filter(f, (sigma, sigma, 0) if f.ndim == 3 else (sigma, sigma)) + factor = (scale, scale, 1) if f.ndim == 3 else (scale, scale) + return np.ascontiguousarray(ndimage.zoom(f, factor, order=1, mode='reflect'), + dtype=np.float32) + + +# `_resize_center_crop` already needs 2x oversample margin on the shorter +# side to antialias well into `size`; capping the longer side at this factor +# leaves that margin on both axes while still discarding the bulk of a +# full-res frame's pixels before the expensive full-frame percentile scan. +# +# Opt-in (score_rgb's fast_preprocess=False by default), not a default-on +# speedup: measured on a real full-res frame (1370x2833, Fireworks Galaxy), +# the `reject`/`quality` heads are genuinely resolution-sensitive, not just +# resampling-noise-sensitive like _resize_center_crop's own cv2->scipy swap +# (which left them within ~0.002/a few points). Sweeping this factor on that +# same frame (ms/call, speedup, defect_probability delta, quality_score +# delta vs. no pre-downsample): +# factor=2 (512px): 400ms 7.04x defect +0.40 quality -160 +# factor=3 (768px): 413ms 6.82x defect +0.26 quality -92 +# factor=4 (1024px): 492ms 5.72x defect +0.13 quality -35 +# factor=6 (1536px): 822ms 3.42x defect +0.06 quality -22 +# factor=8 (2048px): 1494ms 1.88x defect +0.04 quality -3 +# defect_probability swinging by tenths (not thousandths) on one real frame +# at every tested factor -- including the mild ones -- is enough to flip +# is_defective on a frame that sits near 0.5, and that flag feeds +# auto_settings.py's defensive nudges (trail-reject on, stronger chroma +# denoise), not just a log line. Left off by default pending a decision on +# whether/how to expose it (a CLI flag, a specific factor) rather than +# shipping a silent accuracy/speed tradeoff. +_PREDOWNSAMPLE_FACTOR = 4 + + def score_rgb(rgb: np.ndarray, *, model_path: Optional[str] = None, - size: int = 256, shape_gate: bool = True) -> Optional[dict]: + size: int = 256, shape_gate: bool = True, + fast_preprocess: bool = False) -> Optional[dict]: """Score an ``(H, W, 3)`` RGB array (any range/dtype -- it's percentile- stretched here). Returns the result dict, or ``None`` on any failure (logged). Uses the native tract kernel when ``astro_native`` is built, otherwise the ``onnxruntime`` fallback. + + ``fast_preprocess`` (default off): pre-downsample large inputs before the + percentile stretch (see ``_downsample_if_large``'s docstring for the + measured speed-vs-accuracy tradeoff) -- real speedup, but the `reject`/ + `quality` heads shift more than this project's usual resampling + tolerance, so it's opt-in, not the default. """ mp = resolve_model_path(model_path) if mp is None: @@ -210,6 +280,8 @@ def score_rgb(rgb: np.ndarray, *, model_path: Optional[str] = None, arr = _prep_rgb(rgb) if arr is None: return None + if fast_preprocess: + arr = _downsample_if_large(arr, size * _PREDOWNSAMPLE_FACTOR) if _HAS_NATIVE_OV: try: diff --git a/src/pipeline.py b/src/pipeline.py index 30e5cd2..8d141df 100644 --- a/src/pipeline.py +++ b/src/pipeline.py @@ -1153,7 +1153,9 @@ def _psf_fallback_suffix() -> str: run_transient_detection( _transient_src, args.transient_detect, output_path, threshold=float(getattr(args, 'transient_threshold', 5.0)), - wcs=_wcs_for_transients) + wcs=_wcs_for_transients, + triage=bool(getattr(args, 'transient_triage', False)), + triage_model_path=getattr(args, 'transient_triage_model', None)) safe_print(f" Difference imaging: {time.time() - _td_start:.1f}s") except Exception as e: # Diagnostic add-on: a failure here must never cost the stack. diff --git a/src/transient_triage.py b/src/transient_triage.py new file mode 100644 index 0000000..6ce463d --- /dev/null +++ b/src/transient_triage.py @@ -0,0 +1,152 @@ +"""ZOGY candidate real/bogus triage (``--transient-triage``). + +``detect_transients`` (``src/difference_imaging.py``) returns every ``S_corr`` +peak above threshold with no filtering: cosmic rays, sub-pixel +registration-slip dipoles and hot pixels all surface as candidates alongside +genuine transients. This scores each candidate with a small CNN -- the same +role ZTF's BTSbot / Rubin's DIA triage play downstream of classical image +differencing -- and attaches a ``real_probability`` to it. Advisory only: it +never drops a candidate, exactly like ``originvision.py`` never touches +``FrameInfo.accepted``/``metrics['score']``. + +Native-only for now, deliberately -- unlike every other native kernel in this +project, there is no numpy/onnxruntime fallback yet (see the module docstring +in ``ext/astro_native/src/lib.rs``'s ``mod transient_triage``). Without +``astro_native`` built, ``--transient-triage`` self-disables with a warning. + +The bundled model (``src/data/transient_triage.onnx``) is trained entirely on +synthetic data (``tools/gen_transient_triage_data.py`` + +``tools/train_transient_triage.py``) -- no labelled real transients exist yet +-- so it is a first cut, not a production classifier. +""" +from __future__ import annotations + +import logging +import os +from typing import List, Optional, Sequence + +import numpy as np + +logger = logging.getLogger('originstack') + +try: + import astro_native as _native + _HAS_NATIVE_TRIAGE = hasattr(_native, 'transient_triage_score') +except Exception: # pragma: no cover + _native = None + _HAS_NATIVE_TRIAGE = False + +DEFAULT_STAMP_SIZE = 31 + + +def scoring_backend_available() -> bool: + """True when the native kernel is built. No fallback exists yet.""" + return _HAS_NATIVE_TRIAGE + + +def bundled_model_path() -> str: + """Path to the model shipped inside the package.""" + return os.path.join(os.path.dirname(os.path.abspath(__file__)), 'data', + 'transient_triage.onnx') + + +def resolve_model_path(explicit: Optional[str] = None) -> Optional[str]: + """Return the first usable model path: an explicit override, else the + bundled copy. ``None`` if neither exists on disk.""" + for cand in (explicit, bundled_model_path()): + if cand and os.path.isfile(cand): + return cand + return None + + +def _extract_stamp(arr: np.ndarray, y: float, x: float, size: int) -> np.ndarray: + """Fixed-size square cutout centred on ``(y, x)``, reflect-padded at the + frame border -- same boundary convention as this codebase's other + fixed-window extractions (``_resize_center_crop``, the native + gaussian/median kernels' mirror boundary).""" + h, w = arr.shape + half = size // 2 + cy, cx = int(round(y)), int(round(x)) + top, left = cy - half, cx - half + pad_top = max(0, -top) + pad_left = max(0, -left) + pad_bottom = max(0, (top + size) - h) + pad_right = max(0, (left + size) - w) + if pad_top or pad_left or pad_bottom or pad_right: + padded = np.pad(arr, ((pad_top, pad_bottom), (pad_left, pad_right)), + mode='reflect') + top += pad_top + left += pad_left + return padded[top:top + size, left:left + size] + return arr[top:top + size, left:left + size] + + +def _normalize(stamp: np.ndarray, sigma: float) -> np.ndarray: + s = sigma if (sigma is not None and np.isfinite(sigma) and sigma > 0) else 1.0 + # A candidate near the footprint edge can pull in NaN fill from outside + # the warped reference's coverage (difference_imaging.py's `valid` mask); + # zero is the right fill -- both epochs arrive background-subtracted, so + # zero already means "sky" everywhere else in this pipeline's ZOGY code. + clean = np.nan_to_num(stamp.astype(np.float32), nan=0.0, posinf=0.0, neginf=0.0) + return clean / np.float32(s) + + +def build_stamps(new_lum: np.ndarray, ref_lum: np.ndarray, + difference: np.ndarray, positions: Sequence, + sigma_new: float, sigma_ref: float, + sigma_diff: float, size: int = DEFAULT_STAMP_SIZE) -> np.ndarray: + """Build the ``(N, 3, size, size)`` NCHW input for ``transient_triage_score``. + + Channels are ``new``, ``ref``, ``difference`` (the standard real/bogus + "triplet"), each normalized by its own frame-level robust sigma so a + candidate's stamp is architecture-agnostic across sessions of different + noise level -- not the ``originvision`` percentile stretch, which is for + photographic display and would destroy the physical sigma units ZOGY's + own significance already relies on. + """ + n = len(positions) + out = np.zeros((n, 3, size, size), dtype=np.float32) + for i, (y, x) in enumerate(positions): + out[i, 0] = _normalize(_extract_stamp(new_lum, y, x, size), sigma_new) + out[i, 1] = _normalize(_extract_stamp(ref_lum, y, x, size), sigma_ref) + out[i, 2] = _normalize(_extract_stamp(difference, y, x, size), sigma_diff) + return out + + +def score_candidates(new_lum: np.ndarray, ref_lum: np.ndarray, + difference: np.ndarray, transients: Sequence, + sigma_new: float, sigma_ref: float, sigma_diff: float, + *, model_path: Optional[str] = None, + size: int = DEFAULT_STAMP_SIZE) -> List[Optional[float]]: + """Return one ``real_probability`` (or ``None``) per entry in + ``transients``, in the same order. ``None`` for every candidate -- logged + once, not per candidate -- when the native backend or the model file + isn't available.""" + if not transients: + return [] + + if not _HAS_NATIVE_TRIAGE: + logger.warning("transient triage requested but astro_native's " + "transient_triage_score is unavailable -- skipping " + "(no candidates will be scored)") + return [None] * len(transients) + + mp = resolve_model_path(model_path) + if mp is None: + logger.warning("transient triage requested but no model found " + "(bundled src/data/transient_triage.onnx missing) " + "-- skipping") + return [None] * len(transients) + + positions = [(t.y, t.x) for t in transients] + stamps = build_stamps(np.asarray(new_lum, dtype=np.float32), + np.asarray(ref_lum, dtype=np.float32), + np.asarray(difference, dtype=np.float32), + positions, sigma_new, sigma_ref, sigma_diff, size=size) + try: + probs = _native.transient_triage_score(stamps, mp, size) + except Exception as exc: + logger.warning(f"transient triage: native inference failed ({exc}) " + f"-- skipping") + return [None] * len(transients) + return [float(p) for p in probs] diff --git a/src/utils.py b/src/utils.py index aeff06f..1888b54 100644 --- a/src/utils.py +++ b/src/utils.py @@ -224,6 +224,31 @@ def format_time(seconds: float) -> str: return f"{hours}h {mins}m" +def embed_to_shape(arr, H: int, W: int): + """Place ``arr`` in the top-left of an ``(H, W[, C])`` zero canvas (crop + if larger). Pixel coordinates are preserved, so a transform computed on + the embedded array applies directly to it. + + Shared by ``merge.py`` (a previous stack's own shape rarely matches the + current run's) and ``difference_imaging.py`` (two independently-stacked + sessions of the same target routinely differ in pixel dimensions -- + different dither pattern, different Phase 3 common-crop -- even though + they're the same field). A no-op (identity, no copy) when the shape + already matches. + """ + import numpy as np + if arr.shape[0] == H and arr.shape[1] == W: + return arr + if arr.ndim == 3: + out = np.zeros((H, W, arr.shape[2]), dtype=arr.dtype) + else: + out = np.zeros((H, W), dtype=arr.dtype) + h = min(H, arr.shape[0]) + w = min(W, arr.shape[1]) + out[:h, :w] = arr[:h, :w] + return out + + def get_memory_usage_mb() -> float: """Get current process memory usage in MB.""" if HAS_PSUTIL: diff --git a/tests/test_cancel.py b/tests/test_cancel.py new file mode 100644 index 0000000..646fec0 --- /dev/null +++ b/tests/test_cancel.py @@ -0,0 +1,161 @@ +"""Tests for cooperative run cancellation (the desktop app's Cancel button): +``RunCancelled`` (src/models.py), ``frame_processor._check_cancel``, +``cli.process_directory``'s per-target guard, and ``RunManager``'s +cancel()/is_cancelling() plumbing (src/desktop_control.py). + +Cancellation is cooperative -- a ``threading.Event`` checked at a handful of +checkpoints, not a thread kill -- so these tests exercise the checkpoints +directly rather than timing a real multi-minute stacking run. +""" +from __future__ import annotations + +import argparse +import os +import tempfile +import threading +import unittest +from unittest.mock import patch + +from src.models import RunCancelled + + +class TestCheckCancel(unittest.TestCase): + def test_raises_when_event_set(self): + from src.frame_processor import _check_cancel + ns = argparse.Namespace(_cancel_event=threading.Event()) + ns._cancel_event.set() + with self.assertRaises(RunCancelled): + _check_cancel(ns) + + def test_no_op_when_event_unset(self): + from src.frame_processor import _check_cancel + ns = argparse.Namespace(_cancel_event=threading.Event()) + _check_cancel(ns) # must not raise + + def test_no_op_without_a_cancel_event_at_all(self): + """A plain CLI run never sets args._cancel_event -- must be a no-op, + not an AttributeError.""" + from src.frame_processor import _check_cancel + _check_cancel(argparse.Namespace()) + + +class TestExecuteFrameProcessingStopsEarly(unittest.TestCase): + def test_sequential_path_raises_before_processing_any_frame(self): + """A pre-cancelled event must stop the sequential dispatch path + (n < 4, no process pool) before it touches the first frame -- + cheap and deterministic, unlike timing a real multi-frame run.""" + from src.frame_processor import execute_frame_processing + from src.models import FrameInfo, ProcessingStats + + lights = [FrameInfo(path='does-not-exist.fits', type='light', header={})] + args = argparse.Namespace( + parallel=1, verbose=False, debayer_method='malvar', white_balance='grayworld', + ca_correction=False, cosmic_ray_rejection=False, advanced_metrics=True, + pre_gradient_removal=False, trail_reject=False, + _cancel_event=threading.Event()) + args._cancel_event.set() + + with patch('src.frame_processor._process_single_frame') as mock_proc: + with self.assertRaises(RunCancelled): + execute_frame_processing( + lights, {}, args, + mem_rgb=None, mem_lum=None, mm_rgb_path='', mm_lum_path='', + cached_lums=[None], rgb_shape=(1, 4, 4, 3), lum_shape=(1, 4, 4), + rejected_reasons={}, stats=ProcessingStats()) + mock_proc.assert_not_called() + + +class TestProcessDirectoryPerTargetGuard(unittest.TestCase): + def test_raises_before_the_first_target_when_precancelled(self): + from src.cli import process_directory + + with tempfile.TemporaryDirectory() as tmp: + for name in ('session_a', 'session_b'): + d = os.path.join(tmp, name) + os.makedirs(d) + open(os.path.join(d, 'light_000.fits'), 'w').close() + + args = argparse.Namespace(hierarchical=True, mosaic=False, + preset=None, _cancel_event=threading.Event()) + args._cancel_event.set() + + with patch('src.cli._want_combine_sessions', return_value=False), \ + patch('src.cli.discover_frames') as mock_discover: + with self.assertRaises(RunCancelled): + process_directory(tmp, os.path.join(tmp, 'out.fits'), args) + mock_discover.assert_not_called() + + +class TestRunManagerCancel(unittest.TestCase): + def test_cancel_is_a_noop_while_idle(self): + from src.desktop_control import RunManager + rm = RunManager() + rm.cancel() + self.assertFalse(rm._cancel_event.is_set()) + + def test_cancel_sets_the_event_while_running(self): + from src.desktop_control import RunManager + rm = RunManager() + rm.status = 'running' + rm.cancel() + self.assertTrue(rm._cancel_event.is_set()) + + def test_is_cancelling_reflects_both_status_and_event(self): + from src.desktop_control import RunManager + rm = RunManager() + self.assertFalse(rm.is_cancelling()) + rm.status = 'running' + self.assertFalse(rm.is_cancelling()) + rm._cancel_event.set() + self.assertTrue(rm.is_cancelling()) + rm.status = 'ok' + self.assertFalse(rm.is_cancelling(), "a finished run is not 'cancelling'") + + def test_start_gives_each_run_a_fresh_event(self): + """A cancel() from a previous run must not leak into the next one.""" + from src.desktop_control import RunManager + rm = RunManager() + rm._cancel_event.set() + with patch('src.cli.process_directory'): + rm.start({'directory': 'foo', 'output': 'bar.fits'}) + rm.thread.join(timeout=5) + self.assertEqual(rm.status, 'ok') + + def test_start_threads_its_cancel_event_onto_args(self): + """frame_processor._check_cancel / cli.process_directory's guard + read args._cancel_event -- RunManager._run must set it to the + SAME Event object cancel() sets.""" + from src.desktop_control import RunManager + rm = RunManager() + captured = {} + + def _capture(directory, output, args): + captured['ev'] = args._cancel_event + + with patch('src.cli.process_directory', side_effect=_capture): + rm.start({'directory': 'foo', 'output': 'bar.fits'}) + rm.thread.join(timeout=5) + self.assertIs(captured['ev'], rm._cancel_event) + + def test_pipeline_raising_run_cancelled_sets_status_cancelled_not_error(self): + from src.desktop_control import RunManager + rm = RunManager() + with patch('src.cli.process_directory', side_effect=RunCancelled('stop')): + result = rm.start({'directory': 'foo', 'output': 'bar.fits'}) + self.assertTrue(result['ok']) + rm.thread.join(timeout=5) + self.assertEqual(rm.status, 'cancelled') + + def test_run_finished_receives_cancelled_status(self): + from src.desktop_control import RunManager + rm = RunManager() + with patch('src.cli.process_directory', side_effect=RunCancelled('stop')), \ + patch('src.ui_events.get_ui_events') as mock_get_wv: + mock_wv = mock_get_wv.return_value + rm.start({'directory': 'foo', 'output': 'bar.fits'}) + rm.thread.join(timeout=5) + mock_wv.run_finished.assert_called_once_with('cancelled', None) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_difference_imaging.py b/tests/test_difference_imaging.py index 8ab953f..a785cc1 100644 --- a/tests/test_difference_imaging.py +++ b/tests/test_difference_imaging.py @@ -15,7 +15,9 @@ import unittest import numpy as np +import pytest +import src.transient_triage as _tt_mod from src.difference_imaging import ( Transient, _prepare_psf, @@ -435,9 +437,21 @@ def test_a_missing_reference_returns_none(self): self.assertIsNone(run_transient_detection( _rgb(new), f"{self.tmp.name}/nope.fits", self.out_path)) - def test_mismatched_frame_sizes_return_none(self): - new, ref = self._epochs() - self.assertIsNone(self._run(new[:200, :200], ref)) + def test_mismatched_frame_sizes_are_reconciled_not_rejected(self): + """Two independently-stacked sessions of the same target routinely + differ in pixel dimensions (different dither pattern, different + Phase 3 common-crop) even though they cover the same field -- this + used to be a hard failure. `_align_reference` now embeds the smaller + epoch onto the other's grid (`src.utils.embed_to_shape`, the same + trick `merge.py` uses for a previous stack's own mismatched shape) + instead of refusing, and a real transient inside the overlap is + still found afterward.""" + ty, tx = 80.0, 90.0 # inside the 200x200 crop below + new, ref = self._epochs(transient=(ty, tx, 14000.0)) + summary = self._run(new[:200, :200], ref) + self.assertIsNotNone(summary, "a shape mismatch must not be a hard failure") + best = max(summary['transients'], key=lambda t: t.significance) + self.assertLess(math.hypot(best.y - ty, best.x - tx), 3.0) def test_the_registration_sigma_is_measured_and_floored_not_hardcoded(self): """The old code returned a literal 0.3 and reported it as a measured @@ -493,6 +507,36 @@ def test_a_fully_covered_pair_reports_full_coverage(self): self.assertIsNotNone(summary) self.assertGreater(summary['covered_fraction'], 0.99) + def test_stars_beyond_a_smaller_references_extent_are_not_transients(self): + """A genuinely smaller reference (a different session's own Phase 3 + crop, not just a slice of the same array) gets zero-padded onto the + new stack's grid before registration (`_align_reference` / + `src.utils.embed_to_shape`). Stars in `new` beyond the reference's + real extent have nothing to subtract against there -- same failure + mode as a rotated reference's empty wedges, just from padding instead + of rotation -- and must not be reported as 'brightening' either.""" + # n_stars kept low relative to the small shape: _wide_field's + # rejection sampling (each star >20px from every other) needs real + # headroom -- 30 stars in this shape's margins is near the packing + # limit and made the sampling loop pathologically slow. + stars = _wide_field(n_stars=12, seed=9, shape=(140, 150)) + edge_stars = [(180.0, 210.0, 9000.0), (190.0, 30.0, 9000.0), + (30.0, 220.0, 9000.0)] + new = _render_field((220, 240), stars + edge_stars, fwhm=3.0, + sky=0.0, noise=1.0, seed=51) + # A genuinely smaller array -- the reference's own (unpadded) shape, + # not new[:140, :150] -- so embedding must zero-pad it, not just crop. + ref = _render_field((140, 150), stars, fwhm=3.0, sky=0.0, noise=1.0, seed=50) + + summary = self._run(new, ref) + + self.assertIsNotNone(summary, "a smaller reference must not be a hard failure") + self.assertLess(summary['covered_fraction'], 0.95, + "the padded exterior is not reference coverage") + self.assertEqual( + [t for t in summary['transients'] if t.kind == 'brightening'], [], + "stars beyond the reference's real extent are not transients") + def test_outputs_are_written_next_to_the_output_path(self): import os new, ref = self._epochs(transient=(110.0, 120.0, 14000.0)) @@ -503,6 +547,45 @@ def test_outputs_are_written_next_to_the_output_path(self): with self.subTest(suffix=suffix): self.assertTrue(os.path.exists(stem + suffix)) + def test_triage_disabled_leaves_real_probability_none_and_off_the_csv(self): + new, ref = self._epochs(transient=(110.0, 120.0, 14000.0)) + summary = self._run(new, ref, triage=False) + self.assertIsNotNone(summary) + self.assertTrue(all(t.real_probability is None for t in summary['transients'])) + with open(summary['catalog']) as fh: + header = fh.readline() + self.assertNotIn('real_probability', header) + + def test_triage_requested_but_unavailable_does_not_crash(self): + """--transient-triage without a native backend/model self-disables + with a warning (checked via the returned real_probability, not the + log) rather than raising -- mirrors --originvision's own gate.""" + import src.transient_triage as tt_mod + had = tt_mod._HAS_NATIVE_TRIAGE + tt_mod._HAS_NATIVE_TRIAGE = False + try: + new, ref = self._epochs(transient=(110.0, 120.0, 14000.0)) + summary = self._run(new, ref, triage=True) + finally: + tt_mod._HAS_NATIVE_TRIAGE = had + self.assertIsNotNone(summary) + self.assertTrue(all(t.real_probability is None for t in summary['transients'])) + + @pytest.mark.skipif( + not _tt_mod.scoring_backend_available() or _tt_mod.resolve_model_path(None) is None, + reason='native transient_triage_score / bundled model absent -- run ' + 'tools/gen_transient_triage_data.py + tools/train_transient_triage.py') + def test_triage_populates_real_probability_and_csv_column(self): + new, ref = self._epochs(transient=(110.0, 120.0, 14000.0)) + summary = self._run(new, ref, triage=True) + self.assertIsNotNone(summary) + self.assertTrue(summary['transients'], "expected at least the injected transient") + self.assertTrue(all(t.real_probability is not None for t in summary['transients'])) + self.assertTrue(all(0.0 <= t.real_probability <= 1.0 for t in summary['transients'])) + with open(summary['catalog']) as fh: + header = fh.readline() + self.assertIn('real_probability', header) + class TestErodeFootprint(unittest.TestCase): """``_erode`` shrinks the covered region away from *uncovered* pixels only. diff --git a/tests/test_native.py b/tests/test_native.py index 4a15d7c..d874693 100644 --- a/tests/test_native.py +++ b/tests/test_native.py @@ -2645,3 +2645,39 @@ def test_gpu_quality_pool_size_stays_within_core_budget(): assert fp._gpu_quality_pool_size(n_workers=20, cpu_count=16, n_frames=100) == 1 # Never more threads than there are frames to process. assert fp._gpu_quality_pool_size(n_workers=2, cpu_count=16, n_frames=3) == 3 + + +# --------------------------------------------------------------------------- +# transient_triage_score (src/transient_triage.py, --transient-triage) +# --------------------------------------------------------------------------- + +import src.transient_triage as _tt_mod # noqa: E402 + +_tt_model = _tt_mod.resolve_model_path(None) +_have_tt = hasattr(native, 'transient_triage_score') and _tt_model is not None +_tt_skip = pytest.mark.skipif( + not _have_tt, + reason='native transient_triage_score / bundled model absent -- run ' + 'tools/gen_transient_triage_data.py + tools/train_transient_triage.py') + + +@_tt_skip +def test_transient_triage_score_native_shape_and_range(): + rng = np.random.default_rng(3) + stamps = rng.normal(0, 1, (5, 3, 31, 31)).astype(np.float32) + probs = native.transient_triage_score(stamps, _tt_model, 31) + assert len(probs) == 5 + assert all(0.0 <= p <= 1.0 for p in probs) + + +@_tt_skip +def test_transient_triage_score_native_empty_batch(): + stamps = np.zeros((0, 3, 31, 31), dtype=np.float32) + assert native.transient_triage_score(stamps, _tt_model, 31) == [] + + +@_tt_skip +def test_transient_triage_score_native_rejects_wrong_shape(): + stamps = np.zeros((2, 3, 20, 20), dtype=np.float32) # size mismatch + with pytest.raises(ValueError): + native.transient_triage_score(stamps, _tt_model, 31) diff --git a/tests/test_originvision.py b/tests/test_originvision.py index f80c253..3b2b7a2 100644 --- a/tests/test_originvision.py +++ b/tests/test_originvision.py @@ -69,6 +69,27 @@ def test_resize_center_crop_upscale(self): out = infer_mod._resize_center_crop(img, 128) assert out.shape == (128, 128, 3) + def test_downsample_if_large_is_a_no_op_under_the_cap(self): + img = (np.random.default_rng(2).random((100, 150, 3)) * 255).astype(np.float32) + out = infer_mod._downsample_if_large(img, max_long_side=200) + assert out is img # identity, not just equal -- no copy when already small + + def test_downsample_if_large_preserves_aspect_and_caps_long_side(self): + img = (np.random.default_rng(3).random((1000, 2000, 3)) * 255).astype(np.float32) + out = infer_mod._downsample_if_large(img, max_long_side=500) + assert out.shape[1] == 500 # long side (width here) hits the cap exactly + assert abs(out.shape[0] / out.shape[1] - img.shape[0] / img.shape[1]) < 0.01 + assert out.dtype == np.float32 + + def test_downsample_if_large_never_crops(self): + """Unlike _resize_center_crop, this must keep the whole frame -- + cropping here, before the percentile stretch even runs, would change + which pixels the stretch is computed over.""" + img = (np.random.default_rng(4).random((300, 900, 3)) * 255).astype(np.float32) + out = infer_mod._downsample_if_large(img, max_long_side=300) + assert out.shape[0] == 100 # short side scaled down, not cropped away + assert out.shape[1] == 300 + class _FakeCatSession: """Minimal onnxruntime-style session with only a `category` head, so the @@ -171,6 +192,25 @@ def test_score_rgb_returns_expected_keys(self): assert r['category'] in ('galaxy', 'nebula', 'star_cluster', 'comet') assert 0.0 <= r['category_confidence'] <= 1.0 + def test_fast_preprocess_defaults_off_and_matches_explicit_false(self): + """The dormant fast_preprocess opt-in must never change behaviour + unless a caller explicitly asks for it -- default-arg omission and + an explicit False must be identical.""" + rgb = self._synth_rgb() + r_default = infer_mod.score_rgb(rgb) + r_explicit_false = infer_mod.score_rgb(rgb, fast_preprocess=False) + assert r_default == r_explicit_false + + def test_fast_preprocess_true_actually_changes_the_input(self): + """Opting in must take a measurably different (cheaper) path -- a + small synthetic frame is already near/under the downsample cap, so + scale up first to guarantee the cap actually bites.""" + big = np.tile(self._synth_rgb(), (3, 3, 1)) # 900x1200, well over 256*4 + r_full = infer_mod.score_rgb(big) + r_fast = infer_mod.score_rgb(big, fast_preprocess=True) + assert r_full is not None and r_fast is not None + assert r_full['quality_score'] != r_fast['quality_score'] + def test_untrained_and_unused_heads_not_surfaced(self): """The bundled v4 graph emits 8 outputs incl. `trailing` (untrained, excluded from `tasks`) and `background_grid` (trained but unused). @@ -495,9 +535,11 @@ def test_all_samples_failing_returns_none(self): class TestCliOriginvisionResolution: def _parse(self, tmp_path, *extra): + # No explicit --originvision -- it's on by default now (a single + # --no-originvision action, same no-positive-flag shape as --auto; + # see the flag's own definition in cli.py for why). from src import cli - return cli.parse_args(['-d', str(tmp_path), '-o', str(tmp_path / 'o.fits'), - '--originvision', *extra]) + return cli.parse_args(['-d', str(tmp_path), '-o', str(tmp_path / 'o.fits'), *extra]) @_real_infer def test_originvision_stays_enabled_with_bundled_model(self, tmp_path, monkeypatch): diff --git a/tools/gen_transient_triage_data.py b/tools/gen_transient_triage_data.py new file mode 100644 index 0000000..9945ff6 --- /dev/null +++ b/tools/gen_transient_triage_data.py @@ -0,0 +1,193 @@ +"""Synthetic training-data generator for the ZOGY transient-triage model +(``--transient-triage``, ``src/transient_triage.py``). + +No labelled real transients exist yet, so this bootstraps a training set the +same way this codebase already validates ZOGY itself +(``tests/test_difference_imaging.py``): synthetic star fields, rendered at two +different seeings, run through the *real* ``zogy()`` + ``detect_transients()`` +so the stamps a model trains on match what ``run_transient_detection`` actually +produces in production -- not a shortcut simulation of what a candidate stamp +"should" look like. + +Four scene kinds, chosen at random per pair: + +- ``real`` -- an extra star present only in the new epoch (genuine + brightening). Positive label. +- ``cosmic_ray`` -- a single-pixel spike added post-hoc to the new epoch + only, with no PSF. Negative label. +- ``dipole`` -- the new epoch's star field is rendered with a small + sub-pixel ``(dy, dx)`` offset from the reference, and + ``zogy()`` is deliberately given a smaller + ``astrometric_sigma`` than the true offset -- an + *undersuppressed* registration slip, i.e. a hard negative + of exactly the artefact ``astrometric_sigma`` exists to + catch. Negative label. +- ``hot_pixel`` -- a fixed-position single-pixel spike in the new epoch's + noise realization only (not present in the reference, not + aligned with any star). Negative label. + +For a ``real`` pair, every OTHER candidate the detector turns up (there can be +more than one, e.g. noise peaks) is also a hard negative -- only the injected +position is positive. + +Usage: + python tools/gen_transient_triage_data.py --n-pairs 4000 --out transient_triage_data.npz +""" +from __future__ import annotations + +import argparse +import math +import os +import sys + +import numpy as np +from scipy.signal import fftconvolve + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..')) + +from src.difference_imaging import ( # noqa: E402 + detect_transients, + estimate_background_sigma, + zogy, +) +from src.transient_triage import DEFAULT_STAMP_SIZE, build_stamps # noqa: E402 + +_MATCH_RADIUS_PX = 4.0 # a candidate this close to the injected star is the positive + + +def _gaussian_psf(size: int, fwhm: float) -> np.ndarray: + """Normalised Gaussian kernel, odd-sized and centred -- same construction + as tests/test_difference_imaging.py's own PSF helper, kept independent + here rather than importing from a test module.""" + if size % 2 == 0: + size += 1 + sigma = fwhm / (2.0 * math.sqrt(2.0 * math.log(2.0))) + c = size // 2 + yy, xx = np.mgrid[0:size, 0:size] + g = np.exp(-(((yy - c) ** 2 + (xx - c) ** 2) / (2.0 * sigma ** 2))) + return g / g.sum() + + +def _render_field(shape, stars, fwhm: float, sky: float, noise: float, + rng: np.random.Generator) -> np.ndarray: + """Star field convolved to a given seeing, with Gaussian read noise -- + same shape as tests/test_difference_imaging.py's ``_render_field``.""" + h, w = shape + img = np.zeros((h, w), dtype=np.float64) + for y, x, flux in stars: + iy, ix = int(round(y)), int(round(x)) + if 0 <= iy < h and 0 <= ix < w: + img[iy, ix] += flux + img = fftconvolve(img, _gaussian_psf(21, fwhm), mode='same') + return img + sky + rng.normal(0.0, noise, (h, w)) + + +def _random_star_field(rng: np.random.Generator, shape, n_stars: int): + h, w = shape + stars = [] + while len(stars) < n_stars: + y, x = rng.uniform(20, h - 20), rng.uniform(20, w - 20) + if all(math.hypot(y - sy, x - sx) > 15 for sy, sx, _ in stars): + stars.append((y, x, float(rng.uniform(2000, 12000)))) + return stars + + +def make_pair(rng: np.random.Generator, size: int = DEFAULT_STAMP_SIZE, + shape=(160, 180), n_stars: int = 30, max_candidates: int = 6): + """Build one synthetic (new, ref) pair, run it through the real ZOGY path, + and return ``(stamps, labels)`` for whatever candidates were detected -- + zero or more per pair, since a bogus-kind pair can turn up nothing and a + noisy one can turn up spurious hard negatives alongside the label.""" + kind = rng.choice(['real', 'cosmic_ray', 'dipole', 'hot_pixel']) + stars = _random_star_field(rng, shape, n_stars) + ref_fwhm, new_fwhm = rng.uniform(2.5, 4.5), rng.uniform(2.5, 4.5) + + ref = _render_field(shape, stars, fwhm=ref_fwhm, sky=0.0, noise=1.0, rng=rng) + + inject_yx = None + new_stars = list(stars) + astro_sigma = 0.3 + + if kind == 'real': + h, w = shape + iy, ix = rng.uniform(20, h - 20), rng.uniform(20, w - 20) + new_stars.append((iy, ix, float(rng.uniform(3000, 20000)))) + inject_yx = (iy, ix) + new = _render_field(shape, new_stars, fwhm=new_fwhm, sky=0.0, noise=1.0, rng=rng) + elif kind == 'dipole': + dy, dx = rng.uniform(0.4, 1.2) * rng.choice([-1, 1]), rng.uniform(0.4, 1.2) * rng.choice([-1, 1]) + shifted = [(y + dy, x + dx, f) for y, x, f in stars] + new = _render_field(shape, shifted, fwhm=new_fwhm, sky=0.0, noise=1.0, rng=rng) + # Undersuppressed on purpose: the true offset is ~0.4-1.2 px/axis, + # this is the floor ZOGY normally applies when nothing better is + # measured -- exactly the case that leaves a residual dipole. + astro_sigma = 0.3 + else: + new = _render_field(shape, new_stars, fwhm=new_fwhm, sky=0.0, noise=1.0, rng=rng) + h, w = shape + py, px = int(rng.uniform(15, h - 15)), int(rng.uniform(15, w - 15)) + spike = float(rng.uniform(4000, 15000)) + if kind == 'cosmic_ray': + new[py, px] += spike # no PSF -- a single raw pixel, unlike a star + else: # hot_pixel + new[py, px] += spike * 0.6 + new[py, px] += rng.normal(0.0, 1.0) + + psf_new, psf_ref = _gaussian_psf(21, new_fwhm), _gaussian_psf(21, ref_fwhm) + result = zogy(new, ref, psf_new, psf_ref, astrometric_sigma=(astro_sigma, astro_sigma)) + candidates = detect_transients(result.score_corr, threshold=5.0, + max_candidates=max_candidates) + if not candidates: + return np.zeros((0, 3, size, size), dtype=np.float32), np.zeros((0,), dtype=np.float32) + + labels = np.zeros(len(candidates), dtype=np.float32) + if inject_yx is not None: + iy, ix = inject_yx + for i, c in enumerate(candidates): + if math.hypot(c.y - iy, c.x - ix) <= _MATCH_RADIUS_PX: + labels[i] = 1.0 + + sigma_new = estimate_background_sigma(new) + sigma_ref = estimate_background_sigma(ref) + sigma_diff = estimate_background_sigma(result.difference) + stamps = build_stamps(new.astype(np.float32), ref.astype(np.float32), + result.difference, [(c.y, c.x) for c in candidates], + sigma_new, sigma_ref, sigma_diff, size=size) + return stamps, labels + + +def main(): + parser = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument('--n-pairs', type=int, default=4000, + help='Number of synthetic (new, ref) epoch pairs to generate (default: 4000)') + parser.add_argument('--size', type=int, default=DEFAULT_STAMP_SIZE, + help=f'Stamp size (default: {DEFAULT_STAMP_SIZE}, must match training/inference)') + parser.add_argument('--seed', type=int, default=0) + parser.add_argument('--out', default=None, + help='Output .npz path (default: tools/../transient_triage_data.npz)') + args = parser.parse_args() + + out = args.out or os.path.abspath(os.path.join( + os.path.dirname(__file__), '..', 'transient_triage_data.npz')) + + rng = np.random.default_rng(args.seed) + all_stamps, all_labels = [], [] + for i in range(args.n_pairs): + stamps, labels = make_pair(rng, size=args.size) + if len(labels): + all_stamps.append(stamps) + all_labels.append(labels) + if (i + 1) % 200 == 0: + print(f' {i + 1}/{args.n_pairs} pairs ' + f'({sum(len(l) for l in all_labels)} candidates so far)') + + X = np.concatenate(all_stamps, axis=0) if all_stamps else np.zeros((0, 3, args.size, args.size), dtype=np.float32) + y = np.concatenate(all_labels, axis=0) if all_labels else np.zeros((0,), dtype=np.float32) + np.savez(out, X=X, y=y, size=args.size) + n_pos = int(y.sum()) + print(f'Wrote {len(y)} labelled stamps ({n_pos} real, {len(y) - n_pos} bogus) to {out}') + + +if __name__ == '__main__': + main() diff --git a/tools/mine_real_transient_data.py b/tools/mine_real_transient_data.py new file mode 100644 index 0000000..07a991f --- /dev/null +++ b/tools/mine_real_transient_data.py @@ -0,0 +1,290 @@ +"""Mine REAL astrophotography sessions for transient-triage training data, +as a companion to the fully-synthetic ``tools/gen_transient_triage_data.py``. + +Real light frames of the same target on different nights give two things a +synthetic star field can't: + +- **Real negatives (bogus class)**: stack each session with this project's + own pipeline (a real ``originstack.py`` run, not a shortcut), then run the + same ``--transient-detect`` comparison this codebase ships between + consecutive sessions. Since no known real transient is expected in most + amateur fields, every candidate that survives is a genuine artifact -- + cosmic ray, registration-slip dipole, hot pixel -- that got past Phase 1, + the real production population rather than a synthetic guess at it. +- **Real positives (real-transient class)**: still can't get for free (no + labelled real transients exist), but injecting a synthetic point source + into a *copy* of one real stacked epoch before differencing rides on real + noise, real PSF and real artifacts -- a meaningful upgrade over a fully + synthetic star field for the "real" class too. + +Usage: + python tools/mine_real_transient_data.py \\ + --target-dir "G:\\astro\\Astrophotography\\Fireworks Galaxy" \\ + --work-dir transient_triage_real_work \\ + --out transient_triage_real_data.npz \\ + [--max-sessions N] [--n-inject-per-session 3] [--threshold 5.0] + +Each session subfolder under ``--target-dir`` is stacked once and cached in +``--work-dir`` (an existing ``.fits`` there is reused, not +re-stacked) -- a real multi-hundred-frame session can take minutes, so re-runs +while iterating on this script don't pay that cost twice. Never writes +anything back into ``--target-dir``. +""" +from __future__ import annotations + +import argparse +import glob +import math +import os +import subprocess +import sys +import time + +import numpy as np + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..')) + +from src.difference_imaging import _compare_epochs, estimate_background_sigma # noqa: E402 +from src.transient_triage import DEFAULT_STAMP_SIZE, build_stamps # noqa: E402 + +_MATCH_RADIUS_PX = 4.0 +_ORIGINSTACK_PY = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'originstack.py')) + + +def _gaussian_psf(size: int, fwhm: float) -> np.ndarray: + """Normalised Gaussian kernel -- same construction as + tools/gen_transient_triage_data.py's copy, kept independent per this + project's tools/ convention of freestanding scripts.""" + if size % 2 == 0: + size += 1 + sigma = fwhm / (2.0 * math.sqrt(2.0 * math.log(2.0))) + c = size // 2 + yy, xx = np.mgrid[0:size, 0:size] + g = np.exp(-(((yy - c) ** 2 + (xx - c) ** 2) / (2.0 * sigma ** 2))) + return g / g.sum() + + +def _inject_point_source(rgb: np.ndarray, y: float, x: float, flux: float, + size: int = 21, fwhm: float = 3.0) -> np.ndarray: + """Additively place a synthetic point source into a COPY of ``rgb``, + equally across channels. ``flux`` is total (the kernel already sums to + 1), matching the convention tools/gen_transient_triage_data.py's + ``_render_field`` uses (place flux at one pixel, then PSF-convolve).""" + out = rgb.copy() + patch = _gaussian_psf(size, fwhm) * flux + half = size // 2 + h, w = rgb.shape[:2] + y0, x0 = int(round(y)) - half, int(round(x)) - half + ys, ye = max(0, y0), min(h, y0 + size) + xs, xe = max(0, x0), min(w, x0 + size) + if ys >= ye or xs >= xe: + return out + py0, px0 = ys - y0, xs - x0 + py1, px1 = py0 + (ye - ys), px0 + (xe - xs) + sub_patch = patch[py0:py1, px0:px1] + if out.ndim == 3: + out[ys:ye, xs:xe, :] += sub_patch[:, :, None] + else: + out[ys:ye, xs:xe] += sub_patch + return out + + +def _load_rgb(path: str) -> np.ndarray: + from src.io_fits import load_fits + arr, _ = load_fits(path) + arr = np.asarray(arr) + if arr.ndim == 3 and arr.shape[0] in (3, 4) and arr.shape[0] < arr.shape[-1]: + arr = np.transpose(arr, (1, 2, 0)) + return arr + + +def discover_sessions(target_dir: str) -> list: + """Immediate subdirectories of ``target_dir`` that contain FITS light + frames, sorted by name -- session folder names are timestamped + (``Target_YYYY-MM-DD_HH-MM-SS``), so name order is chronological order.""" + sessions = [] + for entry in sorted(os.listdir(target_dir)): + d = os.path.join(target_dir, entry) + if not os.path.isdir(d): + continue + if glob.glob(os.path.join(d, '*.fit*')): + sessions.append(d) + return sessions + + +def stack_session(session_dir: str, out_fits: str) -> bool: + """Stack one session with the real pipeline, caching the result. Returns + True on success (including a cache hit).""" + if os.path.exists(out_fits): + print(f' (cached) {os.path.basename(out_fits)}') + return True + t0 = time.time() + proc = subprocess.run( + [sys.executable, _ORIGINSTACK_PY, '-d', session_dir, '-o', out_fits], + capture_output=True, text=True) + elapsed = time.time() - t0 + if proc.returncode != 0 or not os.path.exists(out_fits): + print(f' FAILED to stack {session_dir} ({elapsed:.0f}s):') + print(' ' + (proc.stderr or proc.stdout)[-2000:].replace('\n', '\n ')) + return False + print(f' stacked {os.path.basename(out_fits)} in {elapsed:.0f}s') + return True + + +def mine_negatives(new_path: str, ref_path: str, size: int, threshold: float): + """Every candidate from a real cross-session comparison is a hard + negative -- no known real transient is expected between two ordinary + nights of the same amateur target.""" + new_rgb, ref_rgb = _load_rgb(new_path), _load_rgb(ref_path) + comparison = _compare_epochs(new_rgb, ref_rgb, threshold=threshold) + if comparison is None or not comparison.transients: + return np.zeros((0, 3, size, size), dtype=np.float32), np.zeros((0,), dtype=np.float32) + + sigma_new = estimate_background_sigma(comparison.new_lum) + sigma_ref = estimate_background_sigma(comparison.ref_lum) + sigma_diff = estimate_background_sigma(comparison.difference) + positions = [(t.y, t.x) for t in comparison.transients] + stamps = build_stamps(comparison.new_lum.astype(np.float32), + comparison.ref_lum.astype(np.float32), + comparison.difference, positions, + sigma_new, sigma_ref, sigma_diff, size=size) + labels = np.zeros(len(positions), dtype=np.float32) # all bogus + return stamps, labels + + +def mine_positives(stack_path: str, size: int, threshold: float, + n_inject: int, rng: np.random.Generator): + """Inject synthetic point sources into a copy of a real stacked epoch, + difference against the untouched original, and label the recovered + injection sites real (everything else found is a hard negative).""" + base_rgb = _load_rgb(stack_path) + from src.difference_imaging import _to_luminance + base_sigma = estimate_background_sigma(_to_luminance(base_rgb)) + h, w = base_rgb.shape[:2] + + all_stamps, all_labels = [], [] + margin = 25 + for _ in range(n_inject): + y, x = rng.uniform(margin, h - margin), rng.uniform(margin, w - margin) + flux = float(rng.uniform(20, 60)) * base_sigma + injected = _inject_point_source(base_rgb, y, x, flux, + fwhm=float(rng.uniform(2.5, 4.0))) + comparison = _compare_epochs(injected, base_rgb, threshold=threshold) + if comparison is None or not comparison.transients: + continue + sigma_new = estimate_background_sigma(comparison.new_lum) + sigma_ref = estimate_background_sigma(comparison.ref_lum) + sigma_diff = estimate_background_sigma(comparison.difference) + positions = [(t.y, t.x) for t in comparison.transients] + stamps = build_stamps(comparison.new_lum.astype(np.float32), + comparison.ref_lum.astype(np.float32), + comparison.difference, positions, + sigma_new, sigma_ref, sigma_diff, size=size) + labels = np.array([1.0 if math.hypot(t.y - y, t.x - x) <= _MATCH_RADIUS_PX else 0.0 + for t in comparison.transients], dtype=np.float32) + all_stamps.append(stamps) + all_labels.append(labels) + + if not all_stamps: + return np.zeros((0, 3, size, size), dtype=np.float32), np.zeros((0,), dtype=np.float32) + return np.concatenate(all_stamps, axis=0), np.concatenate(all_labels, axis=0) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument('--target-dir', required=True, + help='Directory of session subfolders for ONE target, e.g. ' + '"G:\\astro\\Astrophotography\\Fireworks Galaxy"') + parser.add_argument('--work-dir', default=None, + help='Where stacked FITS + sidecars are cached (default: ' + 'tools/../transient_triage_real_work/)') + parser.add_argument('--out', default=None, + help='Output .npz (default: tools/../transient_triage_real_data.npz)') + parser.add_argument('--append-to', default=None, + help='Merge with an existing .npz (e.g. the synthetic one from ' + 'gen_transient_triage_data.py) instead of overwriting --out') + parser.add_argument('--max-sessions', type=int, default=None, + help='Only stack/use the first N sessions, by chronological name ' + 'order (default: all)') + parser.add_argument('--session-filter', default=None, metavar='SUBSTR', + help='Only use session folders whose name contains this substring ' + '(applied before --max-sessions) -- handy for picking a small ' + 'session to validate against before committing to a big one') + parser.add_argument('--n-inject-per-session', type=int, default=3) + parser.add_argument('--threshold', type=float, default=5.0) + parser.add_argument('--size', type=int, default=DEFAULT_STAMP_SIZE) + parser.add_argument('--seed', type=int, default=0) + args = parser.parse_args() + + target_name = os.path.basename(os.path.normpath(args.target_dir)) + work_dir = args.work_dir or os.path.abspath(os.path.join( + os.path.dirname(__file__), '..', 'transient_triage_real_work', target_name)) + os.makedirs(work_dir, exist_ok=True) + out_path = args.out or os.path.abspath(os.path.join( + os.path.dirname(__file__), '..', 'transient_triage_real_data.npz')) + + sessions = discover_sessions(args.target_dir) + if args.session_filter: + sessions = [s for s in sessions if args.session_filter in os.path.basename(s)] + if args.max_sessions: + sessions = sessions[:args.max_sessions] + if len(sessions) < 2: + raise SystemExit(f"only {len(sessions)} session(s) with FITS lights found under " + f"{args.target_dir} -- need at least 2 to compare epochs") + print(f'{target_name}: {len(sessions)} session(s)') + + stacked_paths = [] + for s in sessions: + out_fits = os.path.join(work_dir, os.path.basename(s) + '.fits') + if stack_session(s, out_fits): + stacked_paths.append(out_fits) + + if len(stacked_paths) < 2: + raise SystemExit(f"only {len(stacked_paths)} session(s) stacked successfully -- " + f"need at least 2") + + rng = np.random.default_rng(args.seed) + all_stamps, all_labels = [], [] + + print('Mining real negatives from consecutive session pairs...') + for i in range(len(stacked_paths) - 1): + stamps, labels = mine_negatives(stacked_paths[i + 1], stacked_paths[i], + args.size, args.threshold) + print(f' {os.path.basename(stacked_paths[i + 1])} vs ' + f'{os.path.basename(stacked_paths[i])}: {len(labels)} candidate(s)') + if len(labels): + all_stamps.append(stamps) + all_labels.append(labels) + + if args.n_inject_per_session > 0: + print('Mining real-image-injection positives...') + for p in stacked_paths: + stamps, labels = mine_positives(p, args.size, args.threshold, + args.n_inject_per_session, rng) + n_pos = int(labels.sum()) + print(f' {os.path.basename(p)}: {n_pos} real, {len(labels) - n_pos} bogus ' + f'(of {args.n_inject_per_session} injected)') + if len(labels): + all_stamps.append(stamps) + all_labels.append(labels) + + X = (np.concatenate(all_stamps, axis=0) if all_stamps + else np.zeros((0, 3, args.size, args.size), dtype=np.float32)) + y = np.concatenate(all_labels, axis=0) if all_labels else np.zeros((0,), dtype=np.float32) + + if args.append_to and os.path.exists(args.append_to): + prev = np.load(args.append_to) + if int(prev['size']) != args.size: + raise SystemExit(f"--append-to size {int(prev['size'])} != --size {args.size}") + X = np.concatenate([prev['X'], X], axis=0) + y = np.concatenate([prev['y'], y], axis=0) + out_path = args.append_to + + np.savez(out_path, X=X, y=y, size=args.size) + n_pos = int(y.sum()) + print(f'Wrote {len(y)} labelled stamps ({n_pos} real, {len(y) - n_pos} bogus) to {out_path}') + + +if __name__ == '__main__': + main() diff --git a/tools/train_transient_triage.py b/tools/train_transient_triage.py new file mode 100644 index 0000000..48fbb5e --- /dev/null +++ b/tools/train_transient_triage.py @@ -0,0 +1,159 @@ +"""Train the ZOGY transient-triage model from synthetic data +(tools/gen_transient_triage_data.py) and export it to ONNX +(src/data/transient_triage.onnx, consumed by +astro_native.transient_triage_score / src/transient_triage.py). + +``torch`` is a script-local optional dependency, not part of this project's +runtime dependencies -- model training happens outside the shipped package, +the same stance ``src/data/originvision.onnx`` itself was trained under (see +vendor/originvision/README.md). + +The model is deliberately small (a handful of conv layers) given the tiny +31x31x3 input and a synthetic-only training set -- there is no reason to +reach for originvision's 256x256-real-photograph-classifier capacity here. + +Usage: + pip install torch onnx + python tools/gen_transient_triage_data.py --n-pairs 4000 + python tools/train_transient_triage.py +""" +from __future__ import annotations + +import argparse +import os + +import numpy as np + +try: + import torch + import torch.nn as nn +except ImportError as exc: # pragma: no cover - environment-dependent + raise SystemExit( + "tools/train_transient_triage.py needs torch, which is not part of " + "this project's runtime dependencies (model training happens outside " + "the shipped package -- see vendor/originvision/README.md for the " + "same stance on originvision.onnx). Install it with: pip install torch" + ) from exc + + +class TriageNet(nn.Module): + """Small conv net -> one logit. The native kernel applies sigmoid itself, + so this exports a raw logit, matching astro_native's `compute()`.""" + + def __init__(self): + super().__init__() + self.features = nn.Sequential( + nn.Conv2d(3, 16, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), + nn.Conv2d(16, 32, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), + nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(inplace=True), + nn.AdaptiveAvgPool2d(1), + ) + self.fc = nn.Linear(64, 1) + + def forward(self, x): + x = self.features(x) + x = x.flatten(1) + return self.fc(x).squeeze(-1) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument('--data', default=None, + help='.npz from gen_transient_triage_data.py ' + '(default: tools/../transient_triage_data.npz)') + parser.add_argument('--epochs', type=int, default=20) + parser.add_argument('--batch-size', type=int, default=64) + parser.add_argument('--lr', type=float, default=1e-3) + parser.add_argument('--val-frac', type=float, default=0.15) + parser.add_argument('--seed', type=int, default=0) + parser.add_argument('--out', default=None, + help='Output ONNX path (default: src/data/transient_triage.onnx)') + args = parser.parse_args() + + data_path = args.data or os.path.abspath(os.path.join( + os.path.dirname(__file__), '..', 'transient_triage_data.npz')) + out_path = args.out or os.path.abspath(os.path.join( + os.path.dirname(__file__), '..', 'src', 'data', 'transient_triage.onnx')) + + npz = np.load(data_path) + X, y, size = npz['X'], npz['y'], int(npz['size']) + n = len(y) + if n < 50: + raise SystemExit(f"only {n} labelled stamps in {data_path} -- generate " + f"more with tools/gen_transient_triage_data.py first") + + rng = np.random.default_rng(args.seed) + perm = rng.permutation(n) + n_val = max(1, int(n * args.val_frac)) + val_idx, train_idx = perm[:n_val], perm[n_val:] + + torch.manual_seed(args.seed) + device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + model = TriageNet().to(device) + opt = torch.optim.Adam(model.parameters(), lr=args.lr) + loss_fn = nn.BCEWithLogitsLoss() + + X_t = torch.from_numpy(X).float() + y_t = torch.from_numpy(y).float() + + def _batches(idx): + order = idx.copy() + rng.shuffle(order) + for i in range(0, len(order), args.batch_size): + b = order[i:i + args.batch_size] + yield X_t[b].to(device), y_t[b].to(device) + + for epoch in range(args.epochs): + model.train() + total_loss = 0.0 + for xb, yb in _batches(train_idx): + opt.zero_grad() + loss = loss_fn(model(xb), yb) + loss.backward() + opt.step() + total_loss += loss.detach().item() * len(yb) + model.eval() + with torch.no_grad(): + val_pred = (torch.sigmoid(model(X_t[val_idx].to(device))) > 0.5).float() + val_acc = float((val_pred.cpu() == y_t[val_idx]).float().mean()) + n_train = max(1, len(train_idx)) + print(f'epoch {epoch + 1}/{args.epochs} ' + f'train_loss={total_loss / n_train:.4f} val_acc={val_acc:.3f}') + + model.eval() + model.to('cpu') # export from CPU regardless of training device + os.makedirs(os.path.dirname(out_path), exist_ok=True) + dummy = torch.zeros(1, 3, size, size) + torch.onnx.export( + model, dummy, out_path, + input_names=['stamps'], output_names=['logit'], + dynamic_axes={'stamps': {0: 'batch'}, 'logit': {0: 'batch'}}, + opset_version=13, + # The newer dynamo-based exporter (torch's default since 2.x) needs + # `onnxscript`, an extra dependency beyond torch itself; the legacy + # TorchScript-tracing exporter doesn't and is plenty for this model. + dynamo=False, + ) + + # Lightweight provenance metadata -- astro_native's kernel doesn't read + # it (it takes `size` as an explicit call argument and applies sigmoid + # itself), but it mirrors originvision.onnx's own metadata_props and + # matters for a future re-train/re-sync. Best-effort: onnx is not a + # project dependency either. + try: + import onnx + m = onnx.load(out_path) + for k, v in {'stamp_size': str(size), 'channels': 'new,ref,diff', + 'trained_on': 'synthetic'}.items(): + e = m.metadata_props.add() + e.key, e.value = k, v + onnx.save(m, out_path) + except ImportError: + print("(skipping ONNX metadata -- `pip install onnx` to include it)") + + print(f'Wrote {out_path}') + + +if __name__ == '__main__': + main()