diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d71cb5b84..f031a5cc1 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -534,6 +534,16 @@ jobs: except Exception: print('%-10s (absent)' % n) " + - name: Run GPU precompute CPU regression gate + if: matrix.shard == 1 + run: bash .travis/test-gpu-precompute.sh + - name: Run optional waveform conditioning CPU tests + if: matrix.shard == 1 + env: + JAX_ENABLE_X64: "1" + JAX_PLATFORM_NAME: cpu + JAX_PLATFORMS: cpu + run: python -m pytest -q MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_waveform.py - name: Run jax_ile CPU regression gate env: JAX_PLATFORMS: cpu @@ -778,6 +788,7 @@ jobs: MonteCarloMarginalizeCode/Code/test/test_seq_warmstart_seed.py \ MonteCarloMarginalizeCode/Code/test/test_fairdraw_double_weighting.py \ MonteCarloMarginalizeCode/Code/test/test_portfolio_fairdraw_backend.py \ + MonteCarloMarginalizeCode/Code/test/test_gmm_truncated_score.py \ MonteCarloMarginalizeCode/Code/test/test_rvs_record.py - name: Audit _rvs consumers against the fair-draw rebind # sampler._rvs is rebound to an EXPORT resample at the end of integrate_log, and five @@ -810,7 +821,9 @@ jobs: # --use-jax-ile and --ile-exe land in ILE.sub/ILE_puff.sub/ILE_extr.sub, # and --use-jax-ile is refused at DAG-build time with --calmarg-envelope-directory. run: | - python -m pytest -q MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py + python -m pytest -q \ + MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py \ + MonteCarloMarginalizeCode/Code/test/test_jax_fairdraw_postprocess.py - name: Run q-time-pregrid DAG-build test (--internal-ile-q-time-pregrid-factor) # Full subprocess DAG build against the same reference ini/coinc, same reason as # the step above (PR #281 follow-up review, MAJOR #1): asserts diff --git a/.travis/test-core-units.sh b/.travis/test-core-units.sh index 812f19c0b..2355f30be 100755 --- a/.travis/test-core-units.sh +++ b/.travis/test-core-units.sh @@ -60,6 +60,8 @@ FILES=( # -- likelihood dispatch "$C/RIFT/likelihood/test_td_dispatch_epoch.py" "$C/RIFT/likelihood/test_precompute_crossterm_batching.py" + "$C/test/test_jax_template_finalization.py" + "$C/test/test_response_order.py" "$C/test/test_ile_scalar_edge_cases.py" "$C/test/test_mcsamplerGPU_cdf_inverse_scalar_probe.py" "$C/test/test_srate_resample_time_marginalization.py" @@ -159,6 +161,8 @@ done # found it passing 15/15 on a runner with htcondor absent) # 358/346 + test_mcsamplerGPU_cdf_inverse_scalar_probe.py (11 tests: mcsamplerGPU.cdf_inverse # fed odeint's float probe to len(x) pdfs; the t_ref wiring in all three ILE drivers) +# 378/366 + test_response_order.py (8 tests: SNR tightening, Halton independence, +# compound-axis semantics, reference resolution/tail charging, and bank preflight) # # RAISE these when files are added: a floor left at the old value passes while covering less, # which is the failure this gate exists to catch. @@ -173,12 +177,13 @@ done # direction that matters: 350 >= 347 passes today, and if pytest-subtests ever leaves the # runner's closure the count falls back to 347 and still passes. Pinning 350 would turn an # unrelated dependency change into a red gate. -EXPECTED_TESTS=370 +# Two XML/grid template-finalization regressions, with no added skips. +EXPECTED_TESTS=380 # Outcomes, not just exit status: a collection floor cannot see a test that collects, runs and # asserts nothing, and a pytest.skip can quietly absorb a lost gate. The 12 skips are # environment legs -- cupy in test_seeding_reproducibility, device legs in # test_dslice_device_native, and the xfail in test_uv_symmetry. -EXPECTED_PASSED=358 +EXPECTED_PASSED=368 MAX_SKIPPED=12 # The floors must be INTEGERS, and this is checked rather than assumed. `[ 347 -lt FOO ]` does diff --git a/.travis/test-gpu-precompute.sh b/.travis/test-gpu-precompute.sh new file mode 100644 index 000000000..b416d5e3c --- /dev/null +++ b/.travis/test-gpu-precompute.sh @@ -0,0 +1,41 @@ +#!/usr/bin/env bash +# CPU coverage of the optional GPU path: run each file in a fresh process so +# synthetic import stubs and directory-local conftest modules cannot leak. +set -euo pipefail +cd "$(dirname "$0")/.." +PYTHON_BIN="${RIFT_JAX_PYTHON:-${PYTHON:-python}}" +export PYTHONPATH="$PWD/MonteCarloMarginalizeCode/Code${PYTHONPATH:+:$PYTHONPATH}" +export JAX_PLATFORM_NAME=cpu JAX_PLATFORMS=cpu JAX_ENABLE_X64=1 +export OMP_NUM_THREADS="${OMP_NUM_THREADS:-1}" +C="MonteCarloMarginalizeCode/Code/test" +FILES=( + "$C/gpu_precompute/test_array_contract.py" + "$C/gpu_precompute/test_basis_reuse.py" + "$C/gpu_precompute/test_classic_compound_waveform_forwarding.py" + "$C/gpu_precompute/test_cross_driver_parity.py" + "$C/gpu_precompute/test_device_handoff_adversarial.py" + "$C/gpu_precompute/test_dispatch_contract.py" + "$C/gpu_precompute/test_highlevel_integration.py" + "$C/gpu_precompute/test_lal_fft_oracle.py" + "$C/gpu_precompute/test_smoke_config_contract.py" + "$C/gpu_precompute/test_streamed_v_weighting.py" + "$C/test_gpu_jax_handoff.py" + "$C/waveforms/test_gpu_legacy_compat.py" + "$C/waveforms/test_gpu_waveform.py" +) +report_dir="$(mktemp -d)" +for file in "${FILES[@]}"; do + report="$report_dir/$(basename "$file").xml" + "${PYTHON_BIN}" -m pytest -q -p no:cacheprovider --junit-xml="$report" "$file" + # CUDA and optional waveform backends may skip; every file must still run + # at least one CPU assertion. Pytest success alone also permits all-skipped. + "${PYTHON_BIN}" - "$report" <<'PY' +import sys +import xml.etree.ElementTree as ET +cases = list(ET.parse(sys.argv[1]).iter("testcase")) +passed = sum(not any(case.find(tag) is not None + for tag in ("skipped", "failure", "error")) for case in cases) +if passed < 1: + raise SystemExit("GPU precompute CPU gate ran no passing tests: " + sys.argv[1]) +PY +done diff --git a/.travis/test-jax.sh b/.travis/test-jax.sh index 5cfc2ef77..70cefd2e0 100644 --- a/.travis/test-jax.sh +++ b/.travis/test-jax.sh @@ -43,6 +43,7 @@ fi || { echo "test-jax.sh: numpyro unavailable (needed by test_nuts_phimarg)" >&2; exit 1; } export JAX_PLATFORMS="${JAX_PLATFORMS:-cpu}" +export JAX_ENABLE_X64="${JAX_ENABLE_X64:-1}" export OMP_NUM_THREADS="${OMP_NUM_THREADS:-1}" JAXDIR="MonteCarloMarginalizeCode/Code/test/jax" @@ -93,6 +94,9 @@ JAXDIR="MonteCarloMarginalizeCode/Code/test/jax" # result write order) because the defects # they pin live at call sites, where a # helper-level assertion cannot see them. +# test_smc_evidence.py 2 the SMC evidence product averages over +# every walker, including zero-likelihood +# walkers, and keeps the all-finite case. # test_jax_tempering_chooser.py 45 the --adapt-weight-exponent chooser and the # tempering-cost law # ESS/N = [beta(2-beta)]^(dim/2) it rests on. @@ -509,6 +513,7 @@ FILES=( "${JAXDIR}/test_jax_time_quadrature.py" "${JAXDIR}/test_jax_terminal_time_marginalization.py" "${JAXDIR}/test_jax_likelihood.py" + "${JAXDIR}/test_jax_banded_data_term.py" "${JAXDIR}/test_jax_endtoend.py" "${JAXDIR}/test_jax_slowrot_coeffs.py" "${JAXDIR}/test_jax_slowrot_wrapper.py" @@ -518,6 +523,7 @@ FILES=( "${JAXDIR}/test_nuts_phimarg.py" "${JAXDIR}/test_jax_av.py" "${JAXDIR}/test_jax_fairdraw_export.py" + "${JAXDIR}/test_smc_evidence.py" "${JAXDIR}/test_jax_tempering_chooser.py" "${JAXDIR}/test_tvals_grid_convention.py" "${JAXDIR}/test_interp_choices.py" @@ -535,6 +541,7 @@ FILES=( "${JAXDIR}/test_joint_anglemarg_peaklocal.py" "${JAXDIR}/test_angle_marg_peaklocal_wiring.py" "${JAXDIR}/test_angle_marg_multipeak_wiring.py" + "${JAXDIR}/test_angle_marg_multipeak_jax.py" "${JAXDIR}/test_limit_distance_jax.py" "${JAXDIR}/test_direct_marginalization_planner.py" "${JAXDIR}/test_time_first_peaklocal.py" @@ -1067,7 +1074,14 @@ fi # three EXPECTED_TESTS= assignments (761, 755, 762; last wins); this is the single one. # 2026-09-10: +22 value-only AV/portfolio, prior-window, wrapper, and driver # contract tests in test_jax_av.py. -EXPECTED_TESTS=785 +# The multipeak host-side fix (#317) adds two more tests in +# test_angle_marg_multipeak_wiring.py, for 787 tests total. +# The bounded multipeak suite has 53 tests, including real acceptance/AD, +# CLI configuration, invalid guards, and explicit drop/refuse publication. +# 787 + 53 = 840. +# 2026-09-12: +15 compact banded-data contraction value, AD, tile/padding, +# scratch-budget, empty-batch, and graph-size tests. 840 + 15 = 855. +EXPECTED_TESTS=857 echo "== collection floor check (expect >= ${EXPECTED_TESTS} tests) ==" collect_out="$("${PYTHON_BIN}" -m pytest --collect-only -q -p no:cacheprovider "${DESELECT[@]}" "${FILES[@]}" 2>&1)" diff --git a/CHANGES.rst b/CHANGES.rst index 316eebe29..b0e78aa73 100644 --- a/CHANGES.rst +++ b/CHANGES.rst @@ -1,6 +1,13 @@ 0.0.18.0 ------------ development tree is rift_O4d; PRs refer to oshaughn/research-projects-RIT. + +** BUG FIX, jax ILE AV/portfolio: production pseudo-pipe controls for the + distance prior, internal sample rate, mode retention, detector cutoffs, PSD + conditioning, and time interpolation are now honored by the JAX driver. + Previously accepted compatibility options could silently change a real-data + likelihood. The AV distance prior now supports Euclidean/volumetric and + ``pseudo_cosmo`` and refuses unsupported choices before construction. - (rc0) O4d base refresh: modern Python/numpy support, portable GPU execution, generic workflow backends, hyperpipe and simulation-manager support; distance-likelihood export, parsimonious placement preview, waveform utilities and diagnostics (PRs #129, #132, #135, #143). @@ -24,6 +31,10 @@ development tree is rift_O4d; PRs refer to oshaughn/research-projects-RIT. publication. Extend slow-rotation/finite-size response support and cross-term batching; expand CPU/JAX regression and CI-roster coverage (fork PRs #214, #245, #247, #255, #268, #270, #274, #280--#285, #294, #301--#315, #319). + The terminal pseudo-pipe stage now recognizes JAX-ILE's tabular fair-draw + sidecars and joins them to the paired intrinsic likelihood records. It no + longer sends JAX output through the XML-only converter, which could exit + successfully while producing a header-only posterior. 0.0.17.12 --------- diff --git a/GPU_PRECOMPUTE_VALIDATION.md b/GPU_PRECOMPUTE_VALIDATION.md new file mode 100644 index 000000000..489e490b2 --- /dev/null +++ b/GPU_PRECOMPUTE_VALIDATION.md @@ -0,0 +1,494 @@ +# GPU compound precompute: implementation and validation + +Status: PR #325 open and draft, 2026-09-12. The target is `rift_O4d`; +the branch includes the 2026-09-12 base merge after waveform PR #328 landed. + +### Current JAX cost and readiness update (2026-09-12) + +The earlier bounded-loop reduction removed most of the cold compile cost but +regressed warm execution. The replacement gathers small row/sample tiles under +an explicit forward scratch estimate. It retains the compact JAX loop graph +and reverse-mode differentiation without materializing the full production +`(A, K, S, npts)` gather. Fifteen focused CPU tests pass: independent value +oracles for nearest/linear/cubic/sinc, reverse-mode derivatives, non-divisible +row/sample tiles and padded tails, scratch-budget rejection, an empty-batch +compatibility case, and graph size. The actual JAX CI harness collects 860 +tests from 49 files against its floor of 855, and its isolated new-test shard +passes all 15 (pinned JAX 0.9.0). + +Paired Condor job 60769878 ran frozen expanded-loop baseline snapshot15 +(`492aa421...`) and chunked candidate snapshot17 (`8996e237...`) on the same +NVIDIA RTX PRO 4000 Blackwell SFF worker. Both used the same captured 40-element, +five-sample H1/L1 bank, JAX/jaxlib 0.9.0, x64, pinned runtime image +`898a1261...`, and separate cold caches. Compilation fell from 196.691 to +17.855 s; warm median execution improved from 1.109 to 0.900 ms. First +compiled execution remained about 15.4--15.8 s. Maximum absolute error +against the independent fixed-point likelihood oracle was at most 9.10e-10. +These are one-worker observations for a captured bank, not a full ILE or BNS +throughput estimate. Raw products remain in scratch `jax_chunked_ab/`. + +The exact post-base-merge source snapshot18 (`97f9acda...`) passed the +mandatory real-GPU regression job 60769879: 92 tests passed in 311.36 s, +exit 0. The subsequent zero-sample compatibility branch is covered by the +focused CPU test; the full GPU gate predates only that branch. Fresh GitHub CI +is running; keep the PR draft until those checks pass. A bounded full-length C1/E1/K1 probe uses the same +pushed code and frozen input record. Its first attempt (60769880) stopped +before RIFT execution because `/usr/bin/time` was absent from the container; +the corrected wrapper was verified inside that image and resubmitted as +60769881. Neither attempt is a posterior result. Per-intrinsic JAX closure +compilation reuse remains a separate cost issue. + +### Historical readiness checkpoint before quota expiry (2026-09-12) + +Fresh independent review found missing CI registration and two cache-lifetime +issues. The CPU CI gate now explicitly runs all 13 added non-JAX-directory +test files in separate processes and rejects all-skipped files. Roster and +shell checks pass. Contexts now bind to CUDA devices; changing devices rejects +an explicit old context and selects a distinct default context. Stable cache +roles replace old cutoff/response-order versions. Both new regressions passed +on CPU. The full per-file gate subsequently produced all 13 expected reports: +62 passed, 30 optional tests skipped, no failures/errors, and at least one +passing test in every file (temporary reports `/tmp/tmp.ZAmOyN5C0a/`). + +Completed GPU job 60769877 compares frozen source 7dc4058c4 with b99bca825 on +one allocation, with separate cold caches and the same five-point captured +bank. Compile time fell from 293.469 s to 24.057 s, but warm median increased +from 0.001204 s to 0.009285 s. Both matched the independent likelihood oracle +to 9.064e-10 absolute. This single-host tradeoff is NOT a general speedup claim. +Raw outputs are in scratch `jax_compact_ab/`. PR325 remains draft pending +resolution of warm throughput and a final GPU gate. That earlier chunked-gather experiment was subsequently validated and +committed, as described in the current update above. A backup of its initial +untested form is `/tmp/jax_chunked_gather_UNTESTED_20260912.patch`. + +PR328's fail-closed waveform helpers and corrected tests are synchronized here +to avoid conflicting alternative versions of the two added files. Independent +review identified a physical-strain absolute-tolerance bug in its LAL test; +amplitude-normalized checks now reject deliberate zero/sign mutations. All +11 waveform tests passed. This certifies helpers, not real Ripple/LAL carrier +parity or completed production native-GPU waveform conditioning. + +Current correction: the pre-bounds-fix JAX AV smoke evidences are INVALID for +their requested boxes. The JAX adapter omitted AV's `enforce_bounds=True`; +every saved snapshot10 sample lay outside the requested sky/distance bounds. +Independent repair PR327 matches classic's existing bounds flag. Its outward- +rising synthetic regression fails before the fix and passes after it; all27 +JAX AV tests pass, including on the isolated clean branch. Prior normalization +is unchanged. GPU Q/U/V and handoff parity results survive this correction; +the corrected short integration below removes the large discrepancy but does +not establish precise cross-driver evidence agreement. +The smoke harness now also rejects exported samples outside its declared box. + +### Isolated profiling checkpoint (2026-09-12) + +These are completed single-host diagnostic observations, not replicated +production speedup estimates. Job 60769876 first passed 21 GPU correctness +tests, then compared old source 35db72f60 and new source 99e6ff998 using the +same new harness, separate processes and cold caches on the same GPU. Both +profiles used five identical 128-second XPHM banks (8+7 Msun, generic spins, +K=21, A=40, H1/L1); no long NumPy oracle was run. Warm medians over calls 2--5 +were 14.648 s before and 14.222 s after the duplication fixes (about 2.9% +lower total time). V fell from 4.478 to 4.059 s and primary basis construction +from 0.286 to 0.147 s. These are overlapping stage totals, not independent +terms to sum. New warm waveform generation was 5.956 s; packing and upload +were 0.035 s. Warm calls recorded zero storage read bytes and zero major +faults. This observation does not rule out cold I/O or longer-signal effects. +Raw profiles, gate output and hashes are outside the repository in +`/scratch/richard.oshaughnessy/rift_gpu_precompute_20260912/dupfix_ab_128s/`. +The new archive SHA256 is +`2d9812ac31ba1be07e7d120c2d9e39c2d13c2fd481eda077895ab89bdd750314`. + +Separately, job 60769875 profiled a fixed captured short bank with JAX 0.9.0, +x64, on an RTX PRO 4000 Blackwell SFF Edition: K=2, A=40, two detectors, +five extrinsic points and 153 time bins. Backend initialization took 16.853 s, +device handoff 7.103 s, wrapper setup 0.019 s, lowering 13.397 s, compilation +297.930 s, and first execution 18.029 s. Seven warm calls had median 0.010997 s +(range 0.010987--0.011972 s). The maximum absolute discrepancy against the +independent fixed-bank NumPy oracle was 9.064e-10 in log likelihood. This is +consumer-only timing: imports, capture I/O, precompute, initial CuPy upload, +and the oracle are excluded; no AV integral or long waveform was timed. +Compilation dominates this measured cold consumer; the responsible graph +structure and cost of rebuilding a wrapper for another intrinsic remain to +be isolated. Raw output is `profile_jax_consumer.out` in the same scratch +root; source archive snapshot10 SHA256 is +`a4a3cd5d33c35851aca7003fed224a64c0fc6b4097ffd3c99535ef8eea928fa2`, +and the capture hash is +`d798f29fb59f51e2ad080ae2afb95dba384fc19bff984fd8c8ff9dda21835c32`. +No end-to-end or independently replicated performance claim follows. + +Code inspection identifies a separate reuse limitation: each +`JAXExtrinsicLikelihood` constructor defines a fresh jitted closure over its +`JAXLikelihoodData`, including Q/U/V. Same-shaped intrinsic banks therefore +do not use a shared dynamic-bank kernel. This is distinct from the expanded +A-by-K data-term graph. A future reuse change must pass bank arrays as dynamic +arguments and explicitly test two distinct same-shape banks for both compile +reuse and different correct outputs; caching the first closure would silently +reuse stale physics. No wrapper-cache repair is claimed here. + +The first compile-cost candidate replaces the A-by-K Python expansion in the +banded data term with statically bounded JAX loops. Twelve short synthetic CPU +tests passed in 82.38 s with both JAX platform selectors explicitly set to CPU: +nearest/linear/cubic/sinc, with and without post-phase, eager/JIT value parity, +linear/cubic position and coefficient-phase derivative parity, bounded graph +size from A=2 to A=40, and rejection of incomplete phase arguments. The oracle +retains the literal previous contraction. Full captured-bank GPU parity and +cold/warm performance remain pending; no speedup is claimed for this candidate. + +### Latest completed validity checkpoint (2026-09-12) + +Snapshot11 GPU job 60769868 passed 70 tests in 387.32 s. This includes short +SEOBNRv5PHM through GWSignal (21 modes through l=4), ordinary and conjugate +legacy-mode uploads, NumPy/CuPy Q/U/V parity, unequal detector arm lengths, +and nearest/cubic classic GPU consumers. The deliberately small SEOBNR test +disables the model's per-mode Nyquist veto; it tests transport compatibility, +not the physical accuracy of high modes on that grid. + +Actual captured banks replayed through NumPy, CuPy, and JAX in job 60769867 +agree to at most 1.06e-9 in pointwise log likelihood and 9.17e-10 after time +marginalization (five fixed extrinsics, 153 time bins, both driver-origin banks). +This tests identical banks, not waveform equivalence across differing inputs. + +Corrected snapshot12 AV job 60769869 (source SHA256 +`1e55243f1191fee15e57c785395b65ccc6e896dce73e87aeac745acd6a5c1a0d`) +passed the two-intrinsic short PhenomD integration and all exported sample-bound +guards. Both drivers now explicitly use reference frequency 100 Hz. + +| Intrinsic | Corrected JAX lnZ | Reported sigma | neff | Evaluations | Fairdraw rows | Earlier classic lnZ | +| --- | ---: | ---: | ---: | ---: | ---: | ---: | +| 0 | 50.98011 | 0.16864 | 20.44349 | 242372 | 107 | 51.50609 | +| 1 | 50.78177 | 0.15037 | 20.00309 | 323314 | 188 | 51.27185 | + +Worker host runtime was 465.84 s for both points, excluding queue/container +transfer. This is an end-to-end execution smoke, not isolated precompute timing. +The roughly 235-nat discrepancy disappears after enforcing AV bounds. The +remaining roughly 0.5-nat differences require separate investigation before +claiming precise evidence agreement: low ESS, edge contact, and no independent +seed replication preclude such a claim. These fairdraw clouds are not usable +scientific posteriors. Native GPU waveform conditioning remains fail-closed. +Raw logs and samples remain under +`/scratch/richard.oshaughnessy/rift_gpu_precompute_20260912/` in +`snapshot11_tests.*`, `replay2_snapshot10_*_bank.json`, and +`snapshot12_jaxav_products/`; no raw products are committed. + +### Short generic-mode timing checkpoint (2026-09-12) + +Job 60769870 exited zero using committed source 35db72f60, snapshot SHA256 +`f34adf33aed302c4bd542b18de733ee0684c15974b2f987d1601dda03ca94338`. +Short XPHM, 30+25 Msun with generic spins, H1/L1, N=2048, K=21, +pmax=Qmax=1 (A=40), 153 retained time bins; one RTX PRO 4000 Blackwell SFF +Edition and one CPU thread. This compares the same batched algorithm on NumPy +and CuPy, not the scalar legacy implementation. The CPU oracle runs first. + +| Intrinsic | Batched CPU seconds | GPU seconds | Observed ratio CPU/GPU | +| --- | ---: | ---: | ---: | +| 0 (first call) | 85.381 | 39.702 | 2.15 | +| 1 | 9.068 | 1.122 | 8.08 | +| 2 | 7.202 | 0.746 | 9.65 | + +Both paths include waveform generation and stop at resident-bank return; +queue, container transfer, and later oracle copies are excluded. The first +CPU call has about 80 s outside existing stage timers, so its ratio is not a +fair isolated hardware comparison. Initialization timing is being added to +locate this cost. GPU first-call basis and Q/U stages cost 13.94 and 24.98 s; +these observations do not alone identify kernel compilation versus other setup. + +For intrinsic 2, GPU Q FFTs total 0.357 s, U Gram reductions 0.050 s, +V total 0.176 s, main basis 0.059 s, and legacy waveform plus upload 0.090 s. +Q currently batches only four rows, so a controlled larger-batch test is next; +production defaults remain unchanged. CPU U/V reductions dominate its warm cost. +Six cached detector arrays (294912 bytes) were uploaded initially, with no +additional uploads at either later intrinsic. Retained Q/U/V uses 49271040 +bytes; primary basis per detector uses 27525120 bytes. All three numerical +parity checks passed, with maximum downstream absolute lnL error 2.17e-9. + +These are one-worker profiling observations, not replicated speedup estimates +or a long-BNS runtime projection. Raw output: scratch `snapshot13_xphm_short.out` +and its adjacent scheduler log/error files. No samples or sampler were involved. + +## Question and failure criteria + +Can a GPU construct the compound response bank with the same numerical +likelihood as the LAL reference, while reusing detector inputs across intrinsic +points? Per Richard's correction, all executable tests use tiny synthetic +inputs or short BBH waveforms. No full BNS benchmark is authorized for this +validation stage; the long-grid memory bound is analytic only. + +Fail if frequency ordering, Fourier normalization, complex conjugation, epoch, +retained-time window, detector weights, or waveform conditioning changes the +likelihood; if mutated inputs reuse stale cache entries; if unsupported physics +silently falls back; or if the long case exceeds device memory. + +## Fixed diagnostics + +- D1: Q, U, V maximum absolute and scale-relative differences against LAL. +- D2: pointwise downstream log-likelihood differences, including shifted times + and nontrivial phases; both near the signal and across the prior. +- D3: sequential intrinsic points and changed detector/PSD/grid inputs, compared + to fresh contexts; count data uploads and record retained bytes. +- D4: finite inputs/output and explicit rejection of unsupported configurations. + Review (do not benchmark) the memory bound for the eventual long grid. +- D5: synchronized process runtime for initial setup, waveform, basis, Q, U/V, + transfer of compact results, and subsequent intrinsic points. Container + transfer and queue turnaround are excluded. +- D6: if AV integration is run, log weights, n-eff=20 per Richard's updated + smoke target, n-max=800000, n-chunk=20000 classic/8000 JAX; + output only bounded fairdraw (200). Report achieved + ESS, collapse state, prior-edge contact, and seed variation. A failed + convergence check is not a posterior result. + +## Method + +Independent adversarial tests use small synthetic arrays and LAL as oracle; +short physical-waveform runs use frozen snapshots outside the repository. First validate overlap +operations using identical waveform modes. Then validate the Ripple waveform +adapter including the conditioning needed by the production RIFT convention. +Use multiple intrinsic evaluations in one worker. Every measured snapshot is +identified by a source hash. CPU reference and GPU use matching physical inputs. + +## Results and verdict + +- Clean container, user-site disabled: 19 CPU tests passed, 18 GPU cases + deselected. Includes independent LAL FFT normalization, full maintained + CPU precompute/packing/downstream likelihood, and streamed V reduction. +- Short physical harness (30+25 Msun, df=0.5 Hz, N=2048, H1/L1/V1, + A=40, K=2), NumPy backend, two nearby intrinsic points: pointwise maximum + absolute lnL differences 1.2333e-9 and 7.8353e-10. Nine cached input arrays + copied initially; second point added zero copies (nine cache hits). + Reference times 9.2616/2.1810 s; candidate NumPy times 0.2374/0.2250 s. + These are harness observations, not replicated performance estimates and + emphatically not GPU speedups. First reference includes lazy startup. +- GPU snapshot01 c1a6cd17db7fda9459187d3565e5aa597f6e86888fdfcff80f768e4b75cc49a5: + infrastructure failure before tests (worker lacked pytest). +- GPU snapshot02 c00aac0cd29bc0ed8d96e2df33d7de5ee240d705573db979e31e21407799dcf5: + 35 tests passed, two failed writing the default CuPy kernel cache (ENOSPC). + The benchmark failed at the same cache-write boundary, not an array mismatch. + Dependencies archived separately, SHA256 + 83e8c50e5f683380d953af38ed40456aa31e0247ab4124040e31072a84932248. +- Snapshot03 AV submission failed before ILE because /scratch was not bound + inside the container. Corrected launch uses an explicit /results bind and + independent writable per-job CuPy/Numba/CUDA/temp directories. +- GPU snapshot04 49a3382f54bc0b9cf41099d034e8114be15bd277d9f8bfcd913d8926cb523326: + 42 tests passed on a real GPU (60769855.0, exit 0). Five-point short BBH + benchmark also exited 0 (60769855.1), NVIDIA RTX PRO 4000 Blackwell SFF, + 3 detectors, A40/K2/N2048. Warm reference CPU times were + [4.91756, 5.02197, 4.96536, 4.86751] s; warm GPU times were + [0.57125, 1.03418, 0.57905, 0.58016] s. Medians 4.94146/0.57960 s. + These are old-reference-vs-new-algorithm timings, not an isolated GPU + hardware gain. First-use reference/GPU calls were 84.61696/45.28720 s. + All five downstream lnL comparisons had max absolute difference <=1.188e-9. + Nine input-array uploads initially (442368 bytes); none on later points, + 36 total cache hits, nine retained input arrays. This is NOT a BNS scaling + measurement. Timers cover the precompute calls and synchronize device work; + queue and container transfer are excluded. +- Final short CPU suite: 29 passed, 18 CUDA cases deselected, using explicit + JAX_PLATFORM_NAME=cpu/JAX_PLATFORMS=cpu in the GPU-oriented image. +- Snapshot05 42d883faf574cc763f3e4492820df8417a6a12de50c2dcc2f5550f4d1c0bf7d0: + adds finer FFT/Gram timing, final native-provider guards, same-worker batched + NumPy comparison, and hardened AV output parsing; final GPU run 60769857 + completed with five matched intrinsic points and 47 passing tests (60769857). + Same-worker warm medians were 4.60593 s reference CPU, 0.47379 s batched + NumPy, 0.53046 s CuPy. Warm CuPy range 0.52167--0.93523 s. Thus the + short-grid speedup over the original loop is predominantly algorithmic; + there is no demonstrated incremental GPU gain at this N. Representative + last-point per-detector device stages: primary basis 0.037 s, Q FFT 0.022 s, + U Gram 0.0008 s, conjugate basis 0.072 s, V Gram 0.012 s. Fine-grained + timings are nested inside the coarse stages; do not add both sets. +- Snapshot04 bounded AV run 60769856 FAILED convergence, correctly rejected + by the runner: neff=1.2349/1.0055, collapsed=true at both intrinsic points. + This failure is retained and is not a posterior result. The fixture was at + 200 Mpc with broad inherited bounds. Its software path executed, but that + does not validate its integral. Snapshot06 uses 400 Mpc, sky truth +/-0.1 + rad, inclination/psi +/-0.2 rad, distance [300,500] Mpc and cubic time + interpolation; convergence requirements unchanged. +- Snapshot06 d66150cafff35c584adfac571e04bf9cd1c98175ef94c7919e91cd41f3d15f43: + explicit legacy-generator option forwarding, both mode-bank transfers, + per-mode grid/epoch rejection and real TaylorF2 compatibility tests added; + GPU test 60769858 passed all 52 tests. Corrected short AV 60769859 + returned neff=22.637 and 12.640, both finite and collapse=false. The old + runner exited on its 300 threshold. The user subsequently set the test-run + target to 20: event 0 meets it, event 1 remains below it. This is a smoke + test, not a production posterior certification. +- Native Ripple adapter remains excluded from the production environment + switch pending conditioning/epoch validation against the actual PhenomD + LAL path, not the distinct ChooseFDModes conditioning path. + +No full-BNS run or converged posterior claim has been made. + +## Response-order budget and generic-mode validity + +Read `paper/research_notes_paper1.tex` in full and the modulation, FD-precompute, +and generalized-response validation appendices of `paper/paper1_scaling3g.tex` +in the paper repository before selecting further response-order tests. +The exact compound response count is +`A = (P+1)*[(Q+2)*(P+5) + Q*(Q+1)]`, from the sum over `(b,p,n)` +with `w_0=2`, `w_(1+q)=q+2`, and `|n| <= w_b+p`. +For `(P,Q)=(0,0),(1,1),(2,2),(3,6)`, A is 10,40,102,424. +With K waveform modes and N frequency samples, the primary complex128 bank +alone uses `16*A*K*N` bytes; Q FFT work scales as `A*K*N*log(N)` and +U/V Gram work as `(A*K)^2*N`. Scratch space and other resident arrays are +additional. Streaming the conjugate bank does not remove the quadratic Gram +work. Thus increasing both response orders and Lmax blindly is prohibited in +the validation plan; budget using actual A, K and N first. + +Response-order selection for production must use coherent omitted-waveform +U/V norms over the intended extrinsic support and a declared error budget, +not SNR alone, the highest frequency, or data-overlap Q arrays. The notes' +finite angular scans are estimates, not rigorous prior-wide certificates. +Any optional higher-reference-order selector also incurs that reference bank's +precompute cost, even when it selects a cheap production truncation. + +Snapshot07 SHA256 `5412158f150081dec5127dbbf299b0349011456c3265da25953703820bcba373` +adds mode-label alignment for reordered conjugate dictionaries and a real short +precessing IMRPhenomXPHM test: 30+25 solar masses, fmin40 Hz, dt1/512 s, +df0.25 Hz, all 21 modes through l=4, p=Q=0, one H1 detector. The NumPy +comparison passes Q/U/V, epochs and downstream likelihood checks near and +away from the source. GPU correctness job 60769860 passed all 55 tests. This is not a +performance test or a certification of response truncation for long BNS. + +Raw logs and snapshots: /scratch/richard.oshaughnessy/rift_gpu_precompute_20260912. + +## Internal use (not merged) + +### Merged template-angle correction follow-up + +PR326 was merged by the user as 8eb362ef. The GPU branch incorporates that +development change without merging PR325. Snapshot10 +`a4a3cd5d33c35851aca7003fed224a64c0fc6b4097ffd3c99535ef8eea928fa2` +reruns only short JAX AV, job60769863, to test whether the corrected XML/grid +template finalization explains the smoke evidence/peak discrepancy. The target +remains neff20, two intrinsic points, and fairdraw capped at200. The run completed +successfully in 413 s of host worker runtime: lnZ=286.4796/284.4658, +neff=23.81/26.44. Thus the merged correction does NOT explain this fixture's +large evidence discrepancy. For this (2,+/-2)-only model with full orbital-phase +support, the baked polarization phase can be absorbed by the sampled phase. +Actual per-driver bank construction and callback inputs remain under investigation; +no additional integration or performance run is justified until they agree. +`test_cross_driver_parity.py` pins same-bank p1/q1 cubic pointwise and time- +marginalized classic/JAX equality independently of the two sampling runs. +After integrating the merged fix, six focused CPU tests passed: XML/grid +finalization, same-bank cubic cross-driver parity, and handoff/order guards. + +The follow-up found a separate classic-driver compatibility omission: its +compound calls did not receive the waveform controls used by ordinary precompute. +The three calls now share one waveform-control dictionary, including alternate +generators and conditioning. A sentinel-based caller-boundary regression passes. +For the current PhenomD fixture, a bounded mode-generation comparison with and +without the omitted conditioning arguments gave exactly identical ordinary and +conjugate modes; this omission is not an explanation of the evidence gap. +Frozen snapshot10 jobs60769864.0/.1 capture the first actual classic/JAX compound +bank, input data/PSD arrays, and parameters, then stop before sampling. + +The short timing harness now reports cropped retained Q bytes separately from +full-frequency primary-basis bytes. Candidate NumPy and GPU timings end at the +resident-bank return; host copies used only for numerical comparison occur after +that timing. Approximant, Lmax, response orders, and short-grid parameters are +selectable. No new performance run has been launched with this harness. + +An additional short synthetic regression passes for H1/L1 with distinct +40 km/20 km arm lengths, p1/q1, and cubic interpolation: direct JAX handoff, +host adapter, and classic NumPy agree for each detector and the network. +The test explicitly checks network additivity and a nonzero Q contribution. +This does not yet test classic CuPy's separate cubic Q-contraction kernel; +the real-bank replay must cover that path before declaring cross-driver parity. + +The first boundary-capture writer failed to serialize a LIGOTimeGPS object after +writing the array archive; this was a diagnostic-output failure, not a numerical +failure. Corrected captures60769865.0/.1 completed. They show identical modes, +compound labels, cadence, Q epochs, and metadata. Data differ by only 1.37e-14 +relative scale between independently regenerated fixtures; PSDs are identical. +The actual template parameters reveal a previously missed executable-default +difference: reference frequency100 Hz in classic versus30 Hz in JAX. Both have +zero phiref/psi/incl. Q banks differ by opposite constant mode phases; such a +phase can be absorbed by full orbital-phase sampling for this two-mode fixture, +so this is not yet an explanation of the evidence discrepancy. Both smoke +harnesses now explicitly request100 Hz, pinned by a caller-contract test. +Replay job60769866 failed before computation because its script was outside +the container bind. Corrected job60769867 replays each captured bank at five +fixed extrinsic points through +classic NumPy, classic CuPy, and direct JAX, including their actual terminal +time-marginalization routines. No extrinsic sampler is called. + +Short SEOBNRv5PHM through GWSignal: NumPy-backend compatibility test passed +(29.04 s test runtime), with 21 modes through l4, N2048, df0.25 Hz, and common +epoch -0.204755017. It compares reference and candidate Q/U/V and epochs, and +asserts that both actually invoke the GWSignal provider. This is not a CuPy +result or a performance measurement. The corresponding GPU test and an expanded +nearest/cubic classic-GPU consumer regression are prepared for the next frozen +snapshot. The paired reference-frequency and waveform-forwarding caller tests +both pass. + +### Historical device handoff checkpoint (superseded by latest checkpoint above) + +Snapshot08 `c0353bd4a65c62641793979924fc41e7a8a450b86b713fcf37c52abe8db1b7c0` +adds direct compound-bank routing for both conventional GPU ILE and ILE-JAX. +Conventional ILE uses Q row views and the original dense U/V; JAX shares arrays +through DLPack, with a device-local Q layout conversion. The legacy host-return +API remains available. Native waveform conditioning is still gated. + +Preregistered checks: no CuPy-to-host calls during handoff; GPU array residence; +Q/U/V and downstream likelihood parity; nonzero Q contribution inside the stored +time support; source-buffer deletion/allocator churn without corruption; reject +host inputs and unsupported Q pregrid factors. Short conventional and JAX AV +smokes use two intrinsic points, neff20, and fairdraw capped at 200. Queue and +container transfer are excluded from runtime. Jobs 60769861.0/.1/.2 respectively +run the GPU suite, conventional AV and JAX AV. Conventional AV returned +neff=20.14955 and 18.09552, finite and non-collapsed for both points; independent +XML inspection found 20 fairdraw rows each, 1673/1695 bytes. The strict runner +exited on event 1's below-20 value, not a likelihood/device failure. The user +requested a reasonable approximately-20 smoke target, not repeated runs to clear +a sharp threshold; these values are recorded without rerunning for convergence. +The snapshot08 GPU suite passed all 61 tests in 386.40 s. JAX AV completed both +points with neff=27.5798/22.4957 and 160/114 fairdraw rows. This certifies the +execution smoke only: lnZ=285.6272/284.4376 differs substantially from the classic +smoke's lnZ=51.5061/51.2719. The known ln(2) convention offset cannot explain +that gap; input/driver differences are under investigation. No cross-driver +evidence agreement or production posterior claim is made. + +Snapshot09 `7279a60119bb6a58985b90f45524cc7cfd921874a0468c7f669fb6bc5f5ec6f2` +adds the full high-level precompute-to-classic-consumer no-bulk-host-transfer +test, physical-device consistency guards, and pre-import allocator setup in +the JAX executable/harness. GPU regression job 60769862 passed all 62 tests +in 342.34 s, including full precompute-to-classic-consumer no-bulk-transfer +and final physical-device/lifetime guards. Focused CPU +handoff/dispatch tests passed 5 tests before the final device guards; the final +structural/device-guard suite passed 3 tests. No long-waveform performance run. + +Integrated development base 3ee682fe into the draft branch without merging a PR +or rewriting published history. Conflicts were confined to the new response-order +control and device-routing blocks. The host order controls are preserved; the +device-resident path rejects explicit check/choose controls until a device-native +selector is implemented. Upstream response-order tests: 8 passed. Fresh-process +handoff/dispatch/order-guard tests: 3 passed. These must run separately because +the upstream response-order test module installs stub RIFT modules globally. + +Set `RIFT_GPU_PRECOMPUTE=1` in the worker to replace compound +rotation-plus-frequency-response precompute in conventional ILE or ILE-JAX. +CuPy is mandatory on this opt-in path; failure does not silently fall back. +The default still generates the conditioned base modes with LAL, then performs +compound basis FFTs and Q/U/V reduction on the GPU. The ordinary ILE driver +also computes its pre-existing ordinary bank; that separate overhead has not +been removed in this change. + +Compatibility is load-bearing: the default calls the existing +`factored_likelihood.internal_hlm_generator(P, Lmax, **hlm_kwargs)` for every +intrinsic point and uploads BOTH its ordinary and conjugate mode dictionaries. +It does not restrict the default path to IMRPhenomD and does not require JAX or +Ripple. Existing waveform configuration is passed through. The numerical GPU +bank requires modes to share a frequency grid and epoch; it explicitly rejects +a mismatched legacy bank rather than silently shifting its modes. Native +generation is an optional, separate provider hook, not an automatic replacement. + +`GPUPrecomputeContext` reuses detector data, response weights, and inverse PSD +arrays across intrinsic points. Content hashes invalidate modified input arrays; +old versions of a cache role are replaced. An optional timing callback reports +synchronized waveform, input preparation, basis, Q/U, V, and export durations. +The legacy return structure copies only compact Q windows and U/V to the host. +The direct API also offers device returns and a native-provider hook, but the +production environment switch refuses unvalidated native-waveform selection. + +For integration tests use AV with internal log weights, multiple intrinsic +points, bounded n-max/n-eff, and save-samples only with fairdraw capped at 200. +All generated frames, grids, outputs, containers, dependency bundles, and logs +remain outside this source repository. The user subsequently authorized a draft +PR, and later requested landing it. Draft PR325 tracks the connected handoff; +the completed final-device-guard and corrected JAX AV results are recorded in +the latest checkpoint above. diff --git a/MonteCarloMarginalizeCode/Code/RIFT/integrators/gaussian_mixture_model.py b/MonteCarloMarginalizeCode/Code/RIFT/integrators/gaussian_mixture_model.py index 83216613e..4a97a5e4b 100755 --- a/MonteCarloMarginalizeCode/Code/RIFT/integrators/gaussian_mixture_model.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/integrators/gaussian_mixture_model.py @@ -608,7 +608,11 @@ def score(self, sample_array,assume_normalized=True): # bounds_normalized bounds_norm = self._normalize(self.bounds.T).T - normalization_constant = 0. + # sample() selects a component with weight w and draws that component + # conditioned on the bounds. Score that same mixture of *individually* + # truncated components. Dividing the whole mixture by sum(w*C_i) + # instead would describe a different draw process whenever the + # component in-bound probabilities C_i differ. for i in range(self.k): w = self.weights[i] @@ -619,33 +623,34 @@ def score(self, sample_array,assume_normalized=True): if cupy_ok: # Use gpu_logpdf and exponentiate log_pdf = gpu_logpdf(sample_array_norm, mean, cov, self.xpy) - pdf = self.xpy.exp(log_pdf) + component_pdf = self.xpy.exp(log_pdf) else: - pdf = multivariate_normal.pdf(x=sample_array_norm, mean=mean, cov=cov, allow_singular=True) + component_pdf = multivariate_normal.pdf( + x=sample_array_norm, mean=mean, cov=cov, + allow_singular=True) - scores += pdf * w # mvnun is CPU only mean_cpu = self.identity_convert(mean) cov_cpu = self.identity_convert(cov) bounds_norm_cpu = self.identity_convert(bounds_norm) - normalization_constant += w * mvnun(bounds_norm_cpu[:,0], bounds_norm_cpu[:,1], mean_cpu, cov_cpu)[0] + component_mass = mvnun(bounds_norm_cpu[:,0], bounds_norm_cpu[:,1], mean_cpu, cov_cpu)[0] else: sigma2 = cov[0,0] - val = 1./self.xpy.sqrt(2*self.xpy.pi*sigma2) * self.xpy.exp( - 0.5*( sample_array_norm[:,0] - mean[0])**2/sigma2) - scores += val * w + component_pdf = (1./self.xpy.sqrt(2*self.xpy.pi*sigma2) + * self.xpy.exp(-0.5 * (sample_array_norm[:,0] - mean[0])**2/sigma2)) mean_cpu = self.identity_convert(mean)[0] sigma_cpu = self.identity_convert(np.sqrt(sigma2)) bounds_norm_cpu = self.identity_convert(bounds_norm[0]) my_cdf = norm(loc=mean_cpu, scale=sigma_cpu).cdf - normalization_constant += w * (my_cdf(bounds_norm_cpu[1]) - my_cdf(bounds_norm_cpu[0])) + component_mass = my_cdf(bounds_norm_cpu[1]) - my_cdf(bounds_norm_cpu[0]) - # Floors: a sharply-truncated component can drive the mvnun - # normalization to 0 (0/0 -> NaN), and exactly-zero scores later become - # log(0) = -inf in the integrator's weights. 1e-300 keeps the log - # finite without affecting any sample that carries real weight. - normalization_constant = max(float(normalization_constant), 1e-300) - scores /= normalization_constant + # Keep the historical numerical floor for an underflowed bound + # probability. A component with zero numerical mass cannot be + # sampled reliably either; this avoids turning its score into NaN. + scores += w * component_pdf / max(float(component_mass), 1e-300) + + # The component densities above use normalized [-1, 1] coordinates. vol = self.xpy.prod(self.bounds[:,1] - self.bounds[:,0]) scores *= (2.0**self.d) / vol return self.xpy.maximum(scores, 1e-300) diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/factored_likelihood_rotating_freqresponse.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/factored_likelihood_rotating_freqresponse.py index e1818fb84..1c330c92f 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/factored_likelihood_rotating_freqresponse.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/factored_likelihood_rotating_freqresponse.py @@ -123,6 +123,18 @@ def PrecomputeLikelihoodTermsRotatingFreqResponse( analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0., verbose=True, quiet=False, skip_interpolation=False, **hlm_kwargs): """Build the intrinsic compound bank indexed by ``(b,p,n)``.""" + import os + if os.environ.get('RIFT_GPU_PRECOMPUTE', '0') == '1': + from .gpu_precompute import PrecomputeLikelihoodTermsRotatingFreqResponseGPU + provider = os.environ.get('RIFT_GPU_WAVEFORM', 'lal') + if provider != 'lal': + raise ValueError('Native GPU waveform provider is not yet validated; use RIFT_GPU_WAVEFORM=lal') + return PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + event_time_geo, t_window, P, data_dict, psd_dict, Lmax, fMax, + Qmax=Qmax, L_arm=L_arm, p_max=p_max, f_sidereal=f_sidereal, + analyticPSD_Q=analyticPSD_Q, inv_spec_trunc_Q=inv_spec_trunc_Q, + T_spec=T_spec, verbose=verbose, quiet=quiet, + skip_interpolation=skip_interpolation, **hlm_kwargs) from . import factored_likelihood as FL from .. import lalsimutils as lsu diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_jax_handoff.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_jax_handoff.py new file mode 100644 index 000000000..746fa4f8b --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_jax_handoff.py @@ -0,0 +1,237 @@ +"""Zero-host-copy handoff from GPU compound precompute to ILE-JAX. + +The conventional banded builder first materializes LAL time series, repacks them +through NumPy, and finally uploads them to JAX. This adapter consumes the packed +``return_device=True`` result from :mod:`gpu_precompute` and shares its CuPy +buffers with JAX through DLPack. Detector geometry and small static index tables +remain host-built; Q/U/V never return to host. +""" +from __future__ import division, print_function + +import os +import sys + +import numpy as np + + +def _prepare_jax(): + # JAX's default large preallocation competes with a live CuPy compound bank. + # This variable only has effect before JAX initializes; production launchers + # should set it explicitly as well. + if "jax" not in sys.modules: + os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") + import jax + import jax.dlpack + import jax.numpy as jnp + if not bool(jax.config.x64_enabled): + raise RuntimeError("device handoff requires JAX x64") + return jax, jnp + + +def _jax_device_array(value, jax, jnp, name, require_gpu=True): + """Share a JAX/CuPy device array with JAX; reject host-backed inputs.""" + module = type(value).__module__.split(".")[0] + if module in ("jax", "jaxlib"): + out = jnp.asarray(value) + if require_gpu and any(d.platform != "gpu" for d in out.devices()): + raise RuntimeError("%s is a JAX array on a non-GPU device" % name) + return out + if module != "cupy": + raise TypeError("%s must be a CuPy or JAX device array, got %s" % + (name, type(value).__name__)) + import cupy as cp + contiguous = cp.ascontiguousarray(value) + try: + # copy=False makes an allocator/device mismatch explicit rather than + # silently defeating the purpose of this handoff. + return jax.dlpack.from_dlpack(contiguous, copy=False) + except TypeError: # JAX before the copy= API; DLPack was zero-copy by contract. + try: + return jax.dlpack.from_dlpack(contiguous) + except Exception as exc: + raise RuntimeError("DLPack handoff failed for %s" % name) from exc + except Exception as exc: + raise RuntimeError("zero-copy DLPack handoff failed for %s" % name) from exc + + +def _validate_tvals(tvals, delta_t): + t = np.asarray(tvals, dtype=float) + if t.ndim != 1 or t.size < 1 or not np.all(np.isfinite(t)): + raise ValueError("tvals must be a nonempty finite one-dimensional grid") + if t.size > 1 and not np.allclose(np.diff(t), float(delta_t), + rtol=5e-13, atol=1e-15): + raise ValueError("tvals cadence differs from the precompute delta_t") + return t + + +def build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, det_geom, distMpcRef=None, + q_time_pregrid_factor=1, require_gpu=True): + """Build ``JAXLikelihoodData`` while preserving Q/U/V device residency. + + Parameters + ---------- + packed, meta + The two values returned by + ``PrecomputeLikelihoodTermsRotatingFreqResponseGPU(..., + return_device=True)``. + tvals + The ordinary ILE integration grid, sampled at ``packed['delta_t']``. + det_geom + ``det -> (response, x_arm, y_arm, L)`` from + ``slowrot_freqresponse.detector_geometry``. + + The reflected Q pregrid is intentionally limited to factor one. Its current + implementation is host-side; accepting a larger factor here would silently + reintroduce a full Q device-to-host-to-device round trip. + """ + try: + factor = int(q_time_pregrid_factor) + except (TypeError, ValueError, OverflowError) as exc: + raise ValueError("q_time_pregrid_factor must be exactly 1") from exc + if factor != 1 or q_time_pregrid_factor != factor: + raise NotImplementedError( + "device-resident JAX handoff currently requires q_time_pregrid_factor=1") + if not bool(meta.get("gpu_precompute")) or not bool(meta.get("device_resident")): + raise ValueError("meta does not describe a device-resident GPU precompute") + if meta.get("feature") != "rotation_freqresponse": + raise ValueError("only the compound rotating frequency-response bank is supported") + if not bool(meta.get("post_phase_required")): + raise ValueError("compound bank must require the arrival-time post-phase") + + required = {"q", "U", "V", "epoch", "delta_t", "modes", "a_list"} + missing = required.difference(packed) + if missing: + raise ValueError("packed device result is missing %s" % sorted(missing)) + detectors = list(packed["q"]) + if not detectors or set(detectors) != set(packed["U"]) \ + or set(detectors) != set(packed["V"]) \ + or set(detectors) != set(packed["epoch"]) \ + or set(detectors) != set(det_geom): + raise ValueError("Q/U/V/epoch/geometry detector sets differ") + + modes = [tuple(map(int, lm)) for lm in packed["modes"]] + a_list = [tuple(map(int, a)) for a in packed["a_list"]] + if modes != [tuple(map(int, lm)) for lm in meta["modes"]] \ + or a_list != [tuple(map(int, a)) for a in meta["a_list"]]: + raise ValueError("packed mode or compound-index order differs from meta") + if not modes or not a_list: + raise ValueError("empty mode or compound index set") + delta_t = float(packed["delta_t"]) + if not np.isfinite(delta_t) or delta_t <= 0: + raise ValueError("packed delta_t must be finite and positive") + tvals = _validate_tvals(tvals, delta_t) + + jax, jnp = _prepare_jax() + import lal + import lalsimulation as lalsim + from .jax_ile.core import JAXLikelihoodData, DIST_MPC_REF + from .jax_ile import response_slowrot as rs + from .jax_ile import response_rotating_freqresponse as rrf + from .gpu_precompute import _physical_device_key + + if distMpcRef is None: + distMpcRef = DIST_MPC_REF + A, K = len(a_list), len(modes) + detector_data = {} + source_device = None + target_device = None + for det in detectors: + q_source = packed["q"][det] + u_source = packed["U"][det] + v_source = packed["V"][det] + if tuple(q_source.shape[:2]) != (A, K) or q_source.ndim != 3: + raise ValueError("%s Q must have shape (A,K,N)" % det) + if tuple(u_source.shape) != (A, A, K, K) \ + or tuple(v_source.shape) != (A, A, K, K): + raise ValueError("%s U/V must have shape (A,A,K,K)" % det) + for label, array in (("Q", q_source), ("U", u_source), ("V", v_source)): + here = _physical_device_key(array, "%s %s" % (det, label)) + if require_gpu and here[0] != "gpu": + raise RuntimeError("%s %s is on a non-GPU device" % (det, label)) + if source_device is None: + source_device = here + elif here != source_device: + raise ValueError("mixed physical devices in packed bank: %r and %r" % + (source_device, here)) + + # Q needs one device-local layout conversion from (A,K,N) to the + # accumulator's contiguous (A,N,K). U/V already have their final shape. + q_module = type(q_source).__module__.split(".")[0] + if q_module == "cupy": + import cupy as cp + q_source = cp.ascontiguousarray(q_source.transpose(0, 2, 1)) + elif q_module in ("jax", "jaxlib"): + q_source = jnp.transpose(q_source, (0, 2, 1)) + else: + raise TypeError("%s Q is not device-resident" % det) + Q = _jax_device_array(q_source, jax, jnp, "%s Q" % det, require_gpu) + U = _jax_device_array(u_source, jax, jnp, "%s U" % det, require_gpu) + V = _jax_device_array(v_source, jax, jnp, "%s V" % det, require_gpu) + for label, array in (("Q", Q), ("U", U), ("V", V)): + devices = tuple(array.devices()) + if len(devices) != 1: + raise RuntimeError("%s %s JAX result is not single-device" % (det, label)) + if target_device is None: + target_device = devices[0] + elif devices[0] != target_device: + raise RuntimeError("DLPack results landed on different JAX devices") + target_key = (str(target_device.platform), int(target_device.id)) + if target_key != source_device: + raise RuntimeError("DLPack changed physical device from %r to %r" % + (source_device, target_key)) + + response, x_arm, y_arm, length = det_geom[det] + lald = lalsim.DetectorPrefixToLALDetector(det) + epoch = float(packed["epoch"][det]) + if not np.isfinite(epoch): + raise ValueError("%s Q epoch is non-finite" % det) + with jax.default_device(target_device): + detector_data[det] = { + "lms": list(modes), + # Baseline-shaped aliases keep generic inspection helpers working; + # the banded accumulator consumes the full banks below. + "Q": Q[0], "U": U[0, 0], "V": V[0, 0], + "Q_bank": Q, "U_bank": U, "V_bank": V, + "q_time_pregrid_factor": 1, + "q_time_pregrid_report": {"factor": 1, "device_handoff": "dlpack"}, + "npts_full_coarse": int(Q.shape[1]), + "npts_full": int(Q.shape[1]), + "epoch": epoch, + "location": jnp.asarray(np.asarray(lald.location, dtype=np.float64)), + "response": jnp.asarray(np.asarray(response, dtype=np.float64)), + "x_arm": jnp.asarray(np.asarray(x_arm, dtype=np.float64)), + "y_arm": jnp.asarray(np.asarray(y_arm, dtype=np.float64)), + "L_arm": float(length), + "l_max": max(l for l, unused_m in modes), + } + + tref = float(meta["event_time_geo"]) + gmst = float(lal.GreenwichMeanSiderealTime(tref)) + with jax.default_device(target_device): + data = JAXLikelihoodData(detector_data, delta_t, gmst, tvals, tref, + float(distMpcRef), q_time_pregrid_factor=1) + m_values, term1_idx, term2_idx = rs.post_phase_bucketing(a_list) + data.feature = "rotation_freqresponse" + data.band = { + "a_list": a_list, + "Qmax": int(meta["Qmax"]), + "p_max": int(meta["p_max"]), + "refl_idx": np.asarray(rrf.reflection_index(a_list), dtype=np.int64), + "f_sidereal": float(meta["f_sidereal"]), + "post_phase_required": True, + "pp_m_values": np.asarray(m_values, dtype=np.int64), + "pp_term1_idx": np.asarray(term1_idx, dtype=np.int64), + "pp_term2_idx": np.asarray(term2_idx, dtype=np.int64), + } + data.gpu_handoff = { + "contract_Q_U_V_host_copies": 0, + "transport": "DLPack" if any( + type(packed["q"][d]).__module__.split(".")[0] == "cupy" + for d in detectors) else "JAX identity", + "q_layout_device_copy": True, + } + return data + + +__all__ = ["build_jax_rotating_freqresponse_data_from_device"] diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_precompute.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_precompute.py new file mode 100644 index 000000000..4485e2c6d --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_precompute.py @@ -0,0 +1,755 @@ +"""Device-resident precompute for the rotating finite-response likelihood. + +This module is intentionally optional. It keeps the long two-sided spectra on a +CuPy device, builds all ``(b,p,n,l,m)`` elementary templates with batched FFTs, +and reduces Q/U/V there. Only the short Q time window and the small U/V matrices +need to return to the host for the conventional ILE interface. + +The low-level array API also has a NumPy backend. That is useful for convention +and failure-mode tests on machines without CUDA; production callers should leave +``backend=None`` so CuPy is required. +""" +from __future__ import division, print_function + +from dataclasses import dataclass, field +import hashlib +import math +import threading +import time + +import numpy as np + + +def _resolve_backend(backend=None): + if backend is not None: + return backend + try: + import cupy as cp + except Exception as exc: # pragma: no cover - depends on CUDA installation + raise RuntimeError( + "GPU precompute requested but CuPy could not be imported") from exc + try: + if cp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("CuPy reports no CUDA devices") + except Exception as exc: # pragma: no cover - depends on CUDA installation + raise RuntimeError( + "GPU precompute requested but no usable CUDA device is available") from exc + return cp + + +def _is_numpy(xp): + return xp is np or getattr(xp, "__name__", "") == "numpy" + + +def _to_host(value, xp): + return np.asarray(value) if _is_numpy(xp) else xp.asnumpy(value) + + +def _device_asarray(value, xp, dtype=None): + """Convert arrays without routing a JAX device buffer through the host.""" + if _is_numpy(xp): + return np.asarray(value, dtype=dtype) + module = type(value).__module__.split(".")[0] + # NumPy also exposes __dlpack__, but its capsule is a CPU device and cannot + # be imported by CuPy. DLPack is specifically the JAX GPU handoff here; + # ordinary host arrays take the explicit H2D path below. + if module in ("jax", "jaxlib"): + try: + converted = xp.from_dlpack(value) + except (AttributeError, TypeError): # older CuPy/JAX DLPack API + try: + import jax.dlpack + converted = xp.fromDlpack(jax.dlpack.to_dlpack(value)) + except Exception as exc: + raise RuntimeError("cannot transfer waveform device array to CuPy via DLPack") from exc + return converted.astype(dtype, copy=False) if dtype is not None else converted + return xp.asarray(value, dtype=dtype) + + +def _content_digest(array): + """A mutation-safe cache digest (deliberately hashes every input byte).""" + a = np.ascontiguousarray(np.asarray(array)) + h = hashlib.blake2b(digest_size=20) + h.update(str(a.shape).encode("ascii")) + h.update(a.dtype.str.encode("ascii")) + h.update(memoryview(a).cast("B")) + return h.digest() + + +@dataclass +class GPUPrecomputeContext: + """Reusable data/PSD device cache for multiple intrinsic points in one worker. + + Cache keys include a full content digest, so an in-place mutation cannot silently + reuse stale device data. ``clear`` may be used between events to release memory. + """ + backend: object = None + _arrays: dict = field(default_factory=dict) + _role_keys: dict = field(default_factory=dict) + _lock: object = field(default_factory=threading.RLock) + uploads: int = 0 + upload_bytes: int = 0 + cache_hits: int = 0 + + def __post_init__(self): + self.backend = _resolve_backend(self.backend) + self._device_id = (None if _is_numpy(self.backend) else + int(self.backend.cuda.runtime.getDevice())) + + def check_device(self): + """Do not reuse a worker's cached buffers on another CUDA device.""" + if self._device_id is not None and int( + self.backend.cuda.runtime.getDevice()) != self._device_id: + raise RuntimeError("GPU precompute context belongs to CUDA device %d; " + "select that device or use its default context" % + self._device_id) + + def array(self, role, host_array, dtype=None): + self.check_device() + a = np.asarray(host_array, dtype=dtype) + key = (str(role), _content_digest(a), a.dtype.str, tuple(a.shape)) + with self._lock: + cached = self._arrays.get(key) + if cached is None: + previous = self._role_keys.get(str(role)) + if previous is not None and previous != key: + self._arrays.pop(previous, None) + cached = (self.backend.array(a, copy=True) if _is_numpy(self.backend) + else self.backend.asarray(a)) + self._arrays[key] = cached + self._role_keys[str(role)] = key + self.uploads += 1 + self.upload_bytes += int(a.nbytes) + else: + self.cache_hits += 1 + return cached + + def clear(self): + self.check_device() + with self._lock: + self._arrays.clear() + self._role_keys.clear() + if not _is_numpy(self.backend): # pragma: no cover - CUDA only + self.backend.get_default_memory_pool().free_all_blocks() + + def stats(self): + return dict(uploads=int(self.uploads), upload_bytes=int(self.upload_bytes), + cache_hits=int(self.cache_hits), retained_arrays=len(self._arrays)) + + +_DEFAULT_CONTEXTS = {} +_DEFAULT_CONTEXT_LOCK = threading.Lock() + + +def default_context(backend=None): + xp = _resolve_backend(backend) + key = (id(xp), None if _is_numpy(xp) else int(xp.cuda.runtime.getDevice())) + with _DEFAULT_CONTEXT_LOCK: + if key not in _DEFAULT_CONTEXTS: + _DEFAULT_CONTEXTS[key] = GPUPrecomputeContext(xp) + return _DEFAULT_CONTEXTS[key] + + +def lal_frequency_axis(n, delta_f, xp=np): + """RIFT/LAL two-sided order: +Nyquist, ..., 0, ..., -Nyquist+df.""" + if int(n) != n or n < 2 or n % 2: + raise ValueError("a two-sided LAL frequency grid must have positive even length") + if not np.isfinite(delta_f) or delta_f <= 0: + raise ValueError("delta_f must be finite and positive") + return float(delta_f) * (float(n) / 2.0 - xp.arange(int(n))) + + +def _lal_reverse(spectrum, xp): + """LAL COMPLEX16 frequency-to-time transform using an FFT backend.""" + n = spectrum.shape[-1] + alternating = 1 - 2 * (xp.arange(n) & 1) + return xp.fft.ifft(spectrum, axis=-1) * alternating + + +def _lal_forward(series, xp): + """LAL COMPLEX16 time-to-frequency transform using an FFT backend.""" + n = series.shape[-1] + alternating = 1 - 2 * (xp.arange(n) & 1) + return xp.fft.fft(series * alternating, axis=-1) + + +def _derivative_weight(frequency, p, xp): + from .factored_likelihood_with_rotation import FT_SIGN + if p == 0: + return xp.ones_like(frequency, dtype=xp.complex128) + out = (FT_SIGN * 2.0j * math.pi * frequency) ** int(p) + if p % 2 and frequency.size > 1: + # The LAL packing has one unpaired +Nyquist bin at index zero. + out = out.copy() + out[0] = 0 + return out + + +def modulate_fd(spectrum, n, delta_f, epoch=0.0, f_sidereal=None, + backend=None): + """Apply exp(i n Omega t_abs) with the exact LAL FFT conventions.""" + xp = _resolve_backend(backend) + x = xp.asarray(spectrum) + if x.ndim < 1: + raise ValueError("spectrum must have a frequency axis") + if int(n) == 0: + return x.copy() + if f_sidereal is None: + from .factored_likelihood_with_rotation import F_SIDEREAL + f_sidereal = F_SIDEREAL + count = x.shape[-1] + delta_t = 1.0 / (count * float(delta_f)) + time = float(epoch) + xp.arange(count) * delta_t + phase = xp.exp(2.0j * math.pi * int(n) * float(f_sidereal) * time) + return _lal_forward(_lal_reverse(x, xp) * phase, xp) + + +def build_compound_basis(base_modes_fd, response_weights, a_list, delta_f, + delta_t, epoch, *, base_conjugate_fd=None, + f_sidereal=None, backend=None, fft_batch=8): + """Build a dense device bank with shape ``(A,M,N)``. + + ``base_modes_fd`` is ``(M,N)`` and ``response_weights`` is ``(B,N)``. + Each ``a=(b,p,n)`` selects one response row, a time derivative, and a + sidereal modulation. This function is public mainly for validation; the + full precompute fails before allocating a bank that cannot safely reside on + the device; it does not silently spill the derived bank to the host. + """ + xp = _resolve_backend(backend) + base = xp.asarray(base_modes_fd) + weights = xp.asarray(response_weights) + if base.ndim != 2 or weights.ndim != 2 or base.shape[1] != weights.shape[1]: + raise ValueError("base_modes_fd and response_weights must be (M,N) and (B,N)") + if not bool(_to_host(xp.all(xp.isfinite(base)), xp)): + raise ValueError("base_modes_fd contains non-finite values") + if not bool(_to_host(xp.all(xp.isfinite(weights)), xp)): + raise ValueError("response_weights contains non-finite values") + nfreq = base.shape[1] + expected_dt = 1.0 / (nfreq * float(delta_f)) + if not np.isclose(float(delta_t), expected_dt, rtol=5e-13, atol=0): + raise ValueError("delta_t is inconsistent with N and delta_f") + indices = [tuple(map(int, a)) for a in a_list] + if not indices: + raise ValueError("a_list is empty") + if any(len(a) != 3 or a[0] < 0 or a[0] >= weights.shape[0] or a[1] < 0 + for a in indices): + raise ValueError("invalid (b,p,n) compound index") + if fft_batch < 1: + raise ValueError("fft_batch must be positive") + fft_batch = min(int(fft_batch), len(indices)) + # One retained bank plus raw/TD/phased/FFT block and the phase table. Reduce + # the block automatically; fail before allocating a bank that cannot fit at all. + bytes_per_a = int(base.shape[0]) * int(nfreq) * np.dtype(np.complex128).itemsize + free_bytes = _device_free_bytes(xp) + while fft_batch > 1 and (len(indices) * bytes_per_a + + 5 * fft_batch * bytes_per_a) > 0.88 * free_bytes: + fft_batch = max(1, fft_batch // 2) + if (len(indices) * bytes_per_a + 5 * fft_batch * bytes_per_a) > 0.95 * free_bytes: + raise MemoryError( + "compound basis cannot fit safely: retained %.2f GiB, free %.2f GiB" % + (len(indices) * bytes_per_a / 2.0**30, free_bytes / 2.0**30)) + frequency = lal_frequency_axis(nfreq, delta_f, xp=xp) + out = xp.empty((len(indices), base.shape[0], nfreq), dtype=xp.complex128) + # The reverse transform depends on (b,p), not on the sidereal index n. + # Reuse it across all n without retaining a second full compound bank. + groups = {} + for index, (b, p, n) in enumerate(indices): + groups.setdefault((b, p), []).append((index, n)) + time = float(epoch) + xp.arange(nfreq) * float(delta_t) + sidereal = float(f_sidereal if f_sidereal is not None else _sidereal_default()) + for (b, p), members in groups.items(): + raw = base * weights[b][None, :] * _derivative_weight(frequency, p, xp)[None, :] + td = _lal_reverse(raw, xp) + del raw + for start in range(0, len(members), fft_batch): + block = members[start:start + fft_batch] + phases = xp.stack([ + xp.exp(2.0j * math.pi * n * sidereal * time) + for unused_index, n in block], axis=0) + transformed = _lal_forward(td[None, :, :] * phases[:, None, :], xp) + out[xp.asarray([index for index, unused_n in block])] = transformed + del phases, transformed + del td + if base_conjugate_fd is None: + return out + conj_bank = build_compound_basis( + base_conjugate_fd, response_weights, indices, delta_f, delta_t, epoch, + f_sidereal=f_sidereal, backend=xp, fft_batch=fft_batch) + return out, conj_bank + + +def _sidereal_default(): + from .factored_likelihood_with_rotation import F_SIDEREAL + return F_SIDEREAL + + +def compound_precompute_arrays(basis_fd, basis_conj_fd, data_fd, weights2side, + delta_f, delta_t, n_shift, n_window, *, + backend=None, frequency_chunk=1 << 18, + row_block=8, return_device=False, context=None, + timing_callback=None): + """Reduce a compound basis to Q/U/V on the device. + + Inputs use shape ``basis_fd=(A,M,N)`` and exact LAL frequency ordering. + Outputs are ``Q=(A,M,n_window)`` and ``U,V=(A,A,M,M)``. Definitions are + + ``U[a,b,i,j] = 2 df sum_f conj(chi[a,i]) chi[b,j] / Sn`` + ``V[a,b,i,j] = 2 df sum_f conj(chi_conj[a,i]) chi[b,j] / Sn``. + """ + xp = _resolve_backend(backend) + basis = xp.asarray(basis_fd) + basis_c = None if basis_conj_fd is None else xp.asarray(basis_conj_fd) + data = xp.asarray(data_fd) + weight = xp.asarray(weights2side) + if basis.ndim != 3 or (basis_c is not None and basis_c.shape != basis.shape): + raise ValueError("basis_fd and basis_conj_fd must have identical (A,M,N) shape") + a_count, mode_count, nfreq = basis.shape + if data.shape != (nfreq,) or weight.shape != (nfreq,): + raise ValueError("data and weights must have length N") + if nfreq % 2 or not np.isclose(float(delta_t), 1.0/(nfreq*float(delta_f)), + rtol=5e-13, atol=0): + raise ValueError("unsupported or inconsistent LAL FFT grid") + n_shift = int(n_shift); n_window = int(n_window) + if n_window < 1 or n_window > nfreq: + raise ValueError("n_window must lie in [1,N]") + if frequency_chunk < 1 or row_block < 1: + raise ValueError("frequency_chunk and row_block must be positive") + # Check before FFT/GEMM: NaNs otherwise turn an entire likelihood into NaN. + checked = [("basis", basis), ("data", data), ("weights", weight)] + if basis_c is not None: + checked.append(("conjugate basis", basis_c)) + for name, arr in checked: + if not bool(_to_host(xp.all(xp.isfinite(arr)), xp)): + raise ValueError("%s contains non-finite values" % name) + if bool(_to_host(xp.any(weight < 0), xp)): + raise ValueError("weights2side must be nonnegative") + + flat = basis.reshape(a_count * mode_count, nfreq) + flat_c = None if basis_c is None else basis_c.reshape(a_count * mode_count, nfreq) + row_count = flat.shape[0] + + # Q: batched inverse FFT. Roll semantics match lalsimutils.DataRollBins. + detail_stamp = time.perf_counter() + q_rows = [] + for start in range(0, row_count, int(row_block)): + stop = min(start + int(row_block), row_count) + integrand = (2.0 * xp.conj(flat[start:stop]) * data[None, :] * + weight[None, :]) + # LAL COMPLEX16FreqTimeFFT applies the physical frequency-series scale + # N*delta_f in addition to the normalized inverse DFT. + full = _lal_reverse(integrand, xp) * (nfreq * float(delta_f)) + full = xp.roll(full, -n_shift, axis=-1) + q_rows.append(full[:, :n_window].copy()) + q = xp.concatenate(q_rows, axis=0).reshape(a_count, mode_count, n_window) + detail_stamp = _timed("Q_fft", xp, timing_callback, detail_stamp, + rows=row_count, frequencies=nfreq, window=n_window) + + # U,V: accumulate frequency slabs. The output is tiny compared with the spectra. + u = xp.zeros((row_count, row_count), dtype=xp.complex128) + v = None if flat_c is None else xp.zeros_like(u) + for f0 in range(0, nfreq, int(frequency_chunk)): + f1 = min(f0 + int(frequency_chunk), nfreq) + rhs = flat[:, f0:f1] * weight[None, f0:f1] + u += xp.conj(flat[:, f0:f1]) @ rhs.T + if v is not None: + v += xp.conj(flat_c[:, f0:f1]) @ rhs.T + u *= 2.0 * float(delta_f) + if v is not None: + v *= 2.0 * float(delta_f) + u = u.reshape(a_count, mode_count, a_count, mode_count).transpose(0, 2, 1, 3) + if v is not None: + v = v.reshape(a_count, mode_count, a_count, mode_count).transpose(0, 2, 1, 3) + _timed("U_gram", xp, timing_callback, detail_stamp, + rows=row_count, frequencies=nfreq, + includes_v=bool(flat_c is not None)) + if return_device: + return q, u, v + return (_to_host(q, xp), _to_host(u, xp), + None if v is None else _to_host(v, xp)) + + +def streamed_v_matrix(base_conjugate_fd, response_weights, a_list, primary_basis, + weights2side, delta_f, delta_t, epoch, *, f_sidereal=None, + backend=None, a_block=2, fft_batch=2, + frequency_chunk=1 << 18, return_device=False, + timing_callback=None): + """Build conjugate elementary templates in blocks and reduce V immediately. + + This is the production memory bound: at A=40, M=2, N=8388608 the retained + primary bank is 10.0 GiB, while only ``a_block*M`` conjugate spectra exist at + once. The second 10.0 GiB bank is never allocated. + """ + xp = _resolve_backend(backend) + primary = xp.asarray(primary_basis) + base_c = xp.asarray(base_conjugate_fd) + response = xp.asarray(response_weights) + weight = xp.asarray(weights2side) + if primary.ndim != 3 or base_c.ndim != 2: + raise ValueError("primary_basis and base_conjugate_fd must be (A,M,N) and (M,N)") + a_count, mode_count, nfreq = primary.shape + if base_c.shape != (mode_count, nfreq) or response.shape[1] != nfreq \ + or weight.shape != (nfreq,) or len(a_list) != a_count: + raise ValueError("V inputs have inconsistent shapes") + if a_block < 1: + raise ValueError("a_block must be positive") + primary_flat = primary.reshape(a_count * mode_count, nfreq) + out = xp.zeros((a_count * mode_count, a_count * mode_count), dtype=xp.complex128) + basis_seconds = 0.0 + gram_seconds = 0.0 + for a0 in range(0, a_count, int(a_block)): + a1 = min(a0 + int(a_block), a_count) + if timing_callback is not None: + _synchronize(xp) + block_stamp = time.perf_counter() + conjugate = build_compound_basis( + base_c, response, a_list[a0:a1], delta_f, delta_t, epoch, + f_sidereal=f_sidereal, backend=xp, fft_batch=fft_batch) + if timing_callback is not None: + _synchronize(xp) + now = time.perf_counter() + basis_seconds += now - block_stamp + block_stamp = now + left = conjugate.reshape((a1-a0)*mode_count, nfreq) + block_out = xp.zeros((left.shape[0], primary_flat.shape[0]), dtype=xp.complex128) + for f0 in range(0, nfreq, int(frequency_chunk)): + f1 = min(f0 + int(frequency_chunk), nfreq) + # Weight the streamed (small) conjugate block, not the retained + # (A*M,N) primary bank. The contraction is algebraically identical, + # while the per-frequency-slab temporary shrinks by A/a_block. + weighted_left = (xp.conj(left[:, f0:f1]) + * weight[None, f0:f1]) + block_out += weighted_left @ primary_flat[:, f0:f1].T + del weighted_left + r0, r1 = a0 * mode_count, a1 * mode_count + out[r0:r1] = block_out * (2.0 * float(delta_f)) + if timing_callback is not None: + _synchronize(xp) + gram_seconds += time.perf_counter() - block_stamp + del conjugate, left, block_out + out = out.reshape(a_count, mode_count, a_count, mode_count).transpose(0, 2, 1, 3) + if timing_callback is not None: + details = dict(rows=a_count * mode_count, frequencies=nfreq, + a_block=int(a_block)) + timing_callback("V_basis", basis_seconds, details) + timing_callback("V_gram", gram_seconds, details) + return out if return_device else _to_host(out, xp) + + +def _device_free_bytes(xp): + if _is_numpy(xp): + return 1 << 62 + free, total = xp.cuda.runtime.memGetInfo() # pragma: no cover - CUDA only + # CuPy's pool blocks are reported as used by CUDA but are immediately reusable + # by this process. Include them so detector two does not fail a conservative + # preflight merely because detector one's bank left reusable cached blocks. + reusable = int(xp.get_default_memory_pool().free_bytes()) + return min(int(total), int(free) + reusable) + + +def _synchronize(xp): + if not _is_numpy(xp): # pragma: no cover - CUDA only + xp.cuda.Stream.null.synchronize() + + +def _timed(stage, xp, callback, started, **details): + if callback is None: + return time.perf_counter() + _synchronize(xp) + now = time.perf_counter() + callback(stage, now - started, details) + return now + + +def _series_arrays(bank_or_dict, xp, mode_order=None): + """Normalize LAL dictionaries or a DeviceFDModeBank-like object.""" + if hasattr(bank_or_dict, "modes"): + if not getattr(bank_or_dict, "conditioned", False): + raise ValueError("waveform provider returned an unconditioned mode bank") + modes = bank_or_dict.modes + delta_f = float(bank_or_dict.delta_f) + delta_t = float(bank_or_dict.delta_t) + epoch = float(bank_or_dict.epoch) + else: + modes = bank_or_dict + first = next(iter(modes.values())) + delta_f = float(first.deltaF) + delta_t = 1.0 / (first.data.length * delta_f) + epoch = float(first.epoch) + keys = list(modes) + if not keys: + raise ValueError("waveform mode bank is empty") + if mode_order is not None: + if set(keys) != set(mode_order): + raise ValueError("ordinary and conjugate waveform banks have different modes") + # Legacy dictionary insertion order is not part of the waveform + # contract. Align by mode labels before upload, not by row position. + keys = list(mode_order) + if not hasattr(bank_or_dict, "modes"): + # The batched FFT uses one common time/frequency grid. Never silently + # reinterpret a legacy generator's independently shifted mode series. + expected_length = first.data.length + epoch_tolerance = max(1e-12, 1e-9 * delta_t) + for key in keys: + series = modes[key] + if series.data.length != expected_length or not np.isclose( + float(series.deltaF), delta_f, rtol=5e-13, atol=0.): + raise ValueError("legacy waveform modes do not share a frequency grid") + if abs(float(series.epoch) - epoch) > epoch_tolerance: + raise ValueError("legacy waveform modes do not share a common epoch") + arrays = [] + for key in keys: + value = modes[key] + value = value.data.data if hasattr(value, "data") and hasattr(value.data, "data") else value + arrays.append(_device_asarray(value, xp, dtype=xp.complex128)) + out = xp.stack(arrays, axis=0) + if out.shape[1] % 2: + raise ValueError("waveform mode spectra must use an even two-sided grid") + return keys, out, delta_f, delta_t, epoch + + +def _wrap_host_result(detectors, modes, a_list, q_by_det, u_by_det, v_by_det, + epoch_by_det, delta_t, skip_interpolation, tgrid_by_det, + verbose): + """Convert packed arrays to the legacy five-return dictionary structure.""" + import lal + from . import factored_likelihood as FL + rholms = {}; interpolants = {}; cross = {}; cross_v = {} + for det in detectors: + rholms[det] = {}; interpolants[det] = {}; cross[det] = {}; cross_v[det] = {} + for ai, a in enumerate(a_list): + rholms[det][a] = {} + for mi, mode in enumerate(modes): + ts = lal.CreateCOMPLEX16TimeSeries( + "GPU compound Q", lal.LIGOTimeGPS(float(epoch_by_det[det])), + 0.0, float(delta_t), lal.DimensionlessUnit, q_by_det[det].shape[-1]) + ts.data.data[:] = q_by_det[det][ai, mi] + rholms[det][a][mode] = ts + interpolants[det][a] = (None if skip_interpolation else + FL.InterpolateRholms(rholms[det][a], + tgrid_by_det[det], verbose=verbose)) + for ai, a in enumerate(a_list): + for aj, ap in enumerate(a_list): + cross[det][(a, ap)] = { + (m1, m2): u_by_det[det][ai, aj, i, j] + for i, m1 in enumerate(modes) for j, m2 in enumerate(modes)} + cross_v[det][(a, ap)] = { + (m1, m2): v_by_det[det][ai, aj, i, j] + for i, m1 in enumerate(modes) for j, m2 in enumerate(modes)} + return interpolants, cross, cross_v, rholms + + +def _physical_device_key(value, name="array"): + """Return ``(platform, device id)`` for an unsharded array.""" + module = type(value).__module__.split(".")[0] + if module == "cupy": + return "gpu", int(value.device.id) + if module in ("jax", "jaxlib"): + devices = tuple(value.devices()) + if len(devices) != 1: + raise ValueError("%s spans %d devices; one physical device is required" % + (name, len(devices))) + return str(devices[0].platform), int(devices[0].id) + if module == "numpy": + return "cpu", 0 + raise TypeError("%s is not a recognized NumPy/CuPy/JAX array" % name) + + +def pack_device_precompute(packed, meta, require_gpu=True): + """Expose a device result in the five packed objects used by classic ILE. + + This is a view-only operation: Q rows and dense U/V banks remain the exact + backend arrays returned by ``return_device=True``. The conventional NoLoop + compound evaluator already accepts ``rho_by_a`` dictionaries with dense + ``(A,A,K,K)`` U/V arrays, so no LAL objects or host packing are required. + """ + required = {"q", "U", "V", "epoch", "delta_t", "modes", "a_list"} + missing = required.difference(packed) + if missing: + raise ValueError("packed device result is missing %s" % sorted(missing)) + if not bool(meta.get("gpu_precompute")) or not bool(meta.get("device_resident")): + raise ValueError("meta does not describe a device-resident GPU precompute") + modes = [tuple(map(int, lm)) for lm in packed["modes"]] + a_list = [tuple(map(int, a)) for a in packed["a_list"]] + if modes != [tuple(map(int, lm)) for lm in meta.get("modes", ())] \ + or a_list != [tuple(map(int, a)) for a in meta.get("a_list", ())]: + raise ValueError("packed mode or compound-index order differs from meta") + detectors = list(packed["q"]) + if not detectors or any(set(packed[name]) != set(detectors) + for name in ("U", "V", "epoch")): + raise ValueError("Q/U/V/epoch detector sets differ") + A, K = len(a_list), len(modes) + delta_t = float(packed["delta_t"]) + if not np.isfinite(delta_t) or delta_t <= 0: + raise ValueError("packed delta_t must be finite and positive") + lookup = {}; rho = {} + device_key = None + for det in detectors: + q = packed["q"][det]; u = packed["U"][det]; v = packed["V"][det] + if getattr(q, "ndim", None) != 3 or tuple(q.shape[:2]) != (A, K): + raise ValueError("%s Q must have shape (A,K,N)" % det) + if tuple(getattr(u, "shape", ())) != (A, A, K, K) \ + or tuple(getattr(v, "shape", ())) != (A, A, K, K): + raise ValueError("%s U/V must have shape (A,A,K,K)" % det) + if not np.isfinite(float(packed["epoch"][det])): + raise ValueError("%s Q epoch is non-finite" % det) + for label, array in (("Q", q), ("U", u), ("V", v)): + here = _physical_device_key(array, "%s %s" % (det, label)) + if require_gpu and here[0] != "gpu": + raise RuntimeError("%s %s is not on a GPU" % (det, label)) + if device_key is None: + device_key = here + elif here != device_key: + raise ValueError("mixed physical devices in packed bank: %r and %r" % + (device_key, here)) + lookup[det] = np.asarray(modes, dtype=int) + rho[det] = {a: q[i] for i, a in enumerate(a_list)} + # U and V stay dense. The maintained evaluator's dense branch consumes + # exactly this ordering and avoids A^2 Python dictionary entries. + return lookup, rho, packed["U"], packed["V"], dict(packed["epoch"]) + + +def PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + event_time_geo, t_window, P, data_dict, psd_dict, Lmax, fMax, + Qmax=4, L_arm=None, p_max=0, f_sidereal=None, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0., + verbose=True, quiet=False, skip_interpolation=False, + backend=None, context=None, waveform_provider=None, waveform_backend="jax", + return_device=False, + fft_batch=1, q_row_batch=1, v_a_block=1, + frequency_chunk=1 << 16, timing_callback=None, **hlm_kwargs): + """GPU counterpart of ``PrecomputeLikelihoodTermsRotatingFreqResponse``. + + The default waveform adapter calls RIFT's LAL generator once and uploads its + conditioned two-sided modes. A native provider may instead return either + ``(bank, conjugate_bank)`` or one bank with ``conjugate_modes``. With + ``return_device=False`` this returns the exact legacy five-item structure. + With ``return_device=True`` it returns packed device arrays and metadata, avoiding + the final narrow host copies for the device-resident ILE/JAX handoff. + The conservative default blocks limit transient workspace when a large + compound basis already occupies most of device memory. Callers with more + memory may explicitly increase the blocks for precompute throughput. + """ + initialization_stamp = time.perf_counter() + from . import factored_likelihood as FL + from . import factored_likelihood_rotating_freqresponse as fr + from . import slowrot_freqresponse as sfr + from .. import lalsimutils as lsu + + xp = _resolve_backend(backend) + context = context or default_context(xp) + if context.backend is not xp: + raise ValueError("precompute context uses a different array backend") + context.check_device() + if set(data_dict) != set(psd_dict) or not data_dict: + raise ValueError("data and PSD detector sets differ or are empty") + if analyticPSD_Q: + raise ValueError("GPU precompute requires a tabulated PSD") + detectors = list(data_dict) + if f_sidereal is None: + f_sidereal = _sidereal_default() + P.dist = FL.distMpcRef * 1e6 * lsu.lsu_PC + P.deltaF = data_dict[detectors[0]].deltaF + + stamp = _timed("initialization", xp, timing_callback, initialization_stamp) + if waveform_provider is None: + host, host_conj = FL.internal_hlm_generator( + P, Lmax, verbose=verbose, quiet=quiet, **hlm_kwargs) + upload_stamp = _timed("waveform_generation", xp, timing_callback, stamp, + provider="legacy_host") + modes, base, delta_f, delta_t, epoch = _series_arrays(host, xp) + modes_c, base_c, df_c, dt_c, epoch_c = _series_arrays(host_conj, xp, mode_order=modes) + _timed("waveform_pack_upload", xp, timing_callback, upload_stamp, + modes=len(modes), frequencies=int(base.shape[-1])) + else: + supplied = waveform_provider(P, Lmax, backend=waveform_backend, **hlm_kwargs) + if isinstance(supplied, tuple) and len(supplied) == 2: + supplied_main, supplied_conj = supplied + else: + supplied_main = supplied + conj_modes = getattr(supplied, "conjugate_modes", None) + if conj_modes is None: + raise ValueError("native waveform provider must supply conjugate_modes") + supplied_conj = supplied + supplied_conj = type("ConjugateBankView", (), dict( + modes=conj_modes, delta_f=supplied.delta_f, delta_t=supplied.delta_t, + epoch=supplied.epoch, conditioned=supplied.conditioned))() + modes, base, delta_f, delta_t, epoch = _series_arrays(supplied_main, xp) + modes_c, base_c, df_c, dt_c, epoch_c = _series_arrays(supplied_conj, xp, mode_order=modes) + stamp = _timed("waveform", xp, timing_callback, stamp, + modes=len(modes), frequencies=int(base.shape[-1])) + epoch_tol = max(1.0e-12, 1.0e-9 * min(delta_t, dt_c)) + if modes != modes_c or not np.isclose(delta_f, df_c, rtol=5e-13, atol=0.0) \ + or not np.isclose(delta_t, dt_c, rtol=5e-13, atol=0.0) \ + or abs(epoch - epoch_c) > epoch_tol: + raise ValueError("ordinary and conjugate waveform banks use different grids or modes") + if not np.isclose(delta_t, float(P.deltaT), rtol=5e-13, atol=0.0): + raise ValueError("waveform delta_t differs from requested P.deltaT") + nfreq = base.shape[-1] + if any(data_dict[d].data.length != nfreq or + not np.isclose(data_dict[d].deltaF, delta_f, rtol=5e-13, atol=0.0) + for d in detectors): + raise ValueError("waveform and detector data grids differ") + a_list = fr.compound_index_set(int(Qmax), int(p_max)) + + q_out = {}; u_out = {}; v_out = {}; q_epoch = {}; tgrid = {}; lengths = {} + for det in detectors: + length = fr._arm_for_detector(L_arm, det) + _, _, _, length = sfr.detector_geometry(det, L_arm=length) + lengths[det] = float(length) + f_host = np.asarray(lal_frequency_axis(nfreq, delta_f, xp=np)) + response_host = np.asarray(sfr.finite_size_response_weights( + f_host, {'L': float(length), 'T': float(length)/sfr.C_SI}, int(Qmax))) + response = context.array((det, "finite-response"), response_host, + dtype=np.complex128) + + psd = psd_dict[det] + ip = lsu.ComplexIP(P.fmin, fMax, 1.0/2.0/P.deltaT, P.deltaF, psd, + False, inv_spec_trunc_Q, T_spec) + weight = context.array((det, "weights"), ip.weights2side, dtype=np.float64) + data = context.array((det, "data"), data_dict[det].data.data, + dtype=np.complex128) + stamp = _timed("input_prep", xp, timing_callback, stamp, detector=det, + cache=context.stats()) + basis = build_compound_basis( + base, response, a_list, delta_f, delta_t, epoch, + f_sidereal=f_sidereal, backend=xp, fft_batch=fft_batch) + stamp = _timed("basis", xp, timing_callback, stamp, detector=det, + elements=len(a_list), bytes=int(basis.nbytes)) + + t_det = FL.ComputeArrivalTimeAtDetector(det, P.phi, P.theta, event_time_geo) + rho_epoch = float(data_dict[det].epoch) - float(epoch) + shift = float(t_det) - float(t_window) - rho_epoch + n_shift = int(shift / P.deltaT + 0.5) + n_window = int(2.0 * float(t_window) / P.deltaT) + q_epoch[det] = rho_epoch + n_shift * P.deltaT + tgrid[det] = np.arange(n_window) * P.deltaT + q_epoch[det] + q, u, v = compound_precompute_arrays( + basis, None, data, weight, delta_f, delta_t, n_shift, n_window, + backend=xp, frequency_chunk=frequency_chunk, row_block=q_row_batch, + return_device=return_device, timing_callback=timing_callback) + stamp = _timed("Q_U", xp, timing_callback, stamp, detector=det) + v = streamed_v_matrix( + base_c, response, a_list, basis, weight, delta_f, delta_t, epoch, + f_sidereal=f_sidereal, backend=xp, a_block=v_a_block, + fft_batch=min(fft_batch, v_a_block), frequency_chunk=frequency_chunk, + return_device=return_device, timing_callback=timing_callback) + stamp = _timed("V", xp, timing_callback, stamp, detector=det) + q_out[det], u_out[det], v_out[det] = q, u, v + del basis + + meta = dict(feature='rotation_freqresponse', Qmax=int(Qmax), p_max=int(p_max), + f_sidereal=float(f_sidereal), a_list=a_list, modes=modes, + event_time_geo=float(event_time_geo), L=lengths, L_arm=L_arm, + post_phase_required=True, gpu_precompute=True, + device_resident=bool(return_device), grid_order='lal_descending') + if return_device: + _timed("device_export", xp, timing_callback, stamp, + detectors=len(detectors), cache=context.stats()) + return dict(q=q_out, U=u_out, V=v_out, epoch=q_epoch, + delta_t=float(delta_t), modes=modes, a_list=a_list), meta + interpolants, cross, cross_v, rholms = _wrap_host_result( + detectors, modes, a_list, q_out, u_out, v_out, q_epoch, delta_t, + skip_interpolation, tgrid, verbose) + _timed("host_export", xp, timing_callback, stamp, + detectors=len(detectors), cache=context.stats()) + return interpolants, cross, cross_v, rholms, meta diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_waveform.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_waveform.py new file mode 100644 index 000000000..6c66e31bc --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/gpu_waveform.py @@ -0,0 +1,390 @@ +"""Device-native aligned-spin IMRPhenomD modes for GPU precomputation. + +This module deliberately has no import-time JAX or Ripple dependency. The +production RIFT image can therefore import the likelihood package before the +optional waveform backend is installed. Callers receive RIFT's two-sided, +descending frequency packing rather than Ripple's positive-frequency strain. + +The public provider fails closed for ``conditioning='rift'``. IMRPhenomD +currently uses LAL's ``SimInspiralTDModesFromPolarizations`` route, whose conditioning and +epoch differ from RIFT's ChooseFDModes route. A match maximized over time and +phase is not sufficient to certify that likelihood-sensitive convention. +""" + +from dataclasses import dataclass, field +import math + + +class RippleUnavailableError(ImportError): + """Raised when the optional current RippleGW package is unavailable.""" + + +class WaveformCompatibilityError(ValueError): + """Raised when parameters cannot use the restricted native adapter.""" + + +@dataclass(frozen=True) +class DeviceFDModeBank: + """Conditioned intrinsic modes in RIFT's two-sided FD packing.""" + + modes: dict + conjugate_modes: dict + frequencies: object + delta_f: float + delta_t: float + epoch: float + conditioned: bool + backend: str = "ripplegw.IMRPhenomD" + metadata: dict = field(default_factory=dict) + + +def _load_ripple(): + try: + from ripplegw.waveforms import IMRPhenomD + except (ImportError, ModuleNotFoundError) as exc: + raise RippleUnavailableError( + "GPU IMRPhenomD requires the current `ripplegw` package. The " + "historical PyPI ripplegw==0.0.1 imports as `ripple` and is not " + "compatible with current JAX; install Ripple from its official " + "repository at a pinned commit." + ) from exc + if not hasattr(IMRPhenomD, "gen_IMRPhenomD"): + raise RippleUnavailableError( + "Installed RippleGW has no IMRPhenomD.gen_IMRPhenomD API" + ) + return IMRPhenomD + + +def _array_namespace(backend): + if backend is None or backend == "jax": + try: + import jax + import jax.numpy as jnp + except (ImportError, ModuleNotFoundError) as exc: + raise RippleUnavailableError("GPU IMRPhenomD requires JAX") from exc + if not bool(jax.config.x64_enabled): + raise WaveformCompatibilityError( + "JAX x64 must be enabled before importing JAX for likelihood " + "precompute (set JAX_ENABLE_X64=1)" + ) + return jnp + if isinstance(backend, str): + raise WaveformCompatibilityError("waveform backend must be 'jax'") + # Ripple returns JAX arrays and conditioning uses functional `.at` updates. + # Accept jax.numpy for ergonomic direct calls, but reject CuPy/NumPy rather + # than silently moving a multi-gigabyte waveform through host memory. + if not getattr(backend, "__name__", "").startswith("jax.numpy"): + raise WaveformCompatibilityError( + "native Ripple generation requires backend='jax' (the GPU bank " + "bridges JAX to CuPy with DLPack)" + ) + import jax + if not bool(jax.config.x64_enabled): + raise WaveformCompatibilityError("JAX x64 must be enabled (set JAX_ENABLE_X64=1)") + return backend + + +def _scalar(value, name): + try: + result = float(value) + except (TypeError, ValueError) as exc: + raise WaveformCompatibilityError( + "%s must be scalar for one intrinsic waveform" % name + ) from exc + if not math.isfinite(result): + raise WaveformCompatibilityError("%s must be finite" % name) + return result + + +def _validate(P, Lmax): + if int(Lmax) != 2: + raise WaveformCompatibilityError( + "native Ripple IMRPhenomD supplies only (2,+/-2); require Lmax=2" + ) + try: + import lalsimulation as lalsim + approx_name = lalsim.GetStringFromApproximant(P.approx) + except (ImportError, AttributeError, TypeError, RuntimeError): + approx_name = str(getattr(P, "approx", "")) + if approx_name != "IMRPhenomD": + raise WaveformCompatibilityError( + "native adapter only supports IMRPhenomD, got %r" % approx_name + ) + for name in ("s1x", "s1y", "s2x", "s2y"): + if abs(_scalar(getattr(P, name, 0.0), name)) > 1.0e-12: + raise WaveformCompatibilityError("IMRPhenomD requires aligned spins") + for name in ("lambda1", "lambda2", "eccentricity"): + if _scalar(getattr(P, name, 0.0), name) != 0.0: + raise WaveformCompatibilityError("native IMRPhenomD does not support %s" % name) + dt = _scalar(P.deltaT, "deltaT") + df = _scalar(P.deltaF, "deltaF") + if dt <= 0 or df <= 0: + raise WaveformCompatibilityError("deltaT and deltaF must be positive") + n_float = 1.0 / (dt * df) + n = int(round(n_float)) + if n < 4 or n % 2 or abs(n - n_float) > 1.0e-7 * n: + raise WaveformCompatibilityError( + "1/(deltaT*deltaF) must be an even integer, got %.17g" % n_float + ) + if _scalar(P.fmin, "fmin") <= 0: + raise WaveformCompatibilityError("fmin must be positive") + return dt, df, n + + +def _component_masses_msun(P): + # RIFT stores SI masses. Keeping the constants local avoids importing LAL + # (and hence host-side waveform code) on the native path. + msun_si = 1.9884099021470416e30 + m1 = _scalar(P.m1, "m1") / msun_si + m2 = _scalar(P.m2, "m2") / msun_si + if m1 <= 0 or m2 <= 0: + raise WaveformCompatibilityError("component masses must be positive") + if m2 > m1: + m1, m2 = m2, m1 + chi1, chi2 = _scalar(P.s2z, "s2z"), _scalar(P.s1z, "s1z") + else: + chi1, chi2 = _scalar(P.s1z, "s1z"), _scalar(P.s2z, "s2z") + if max(abs(chi1), abs(chi2)) > 1: + raise WaveformCompatibilityError("dimensionless spins must be in [-1,1]") + return m1, m2, chi1, chi2 + + +def _continuous_inverse(values, dt, xpy): + # LAL's complex FFT stores the RIFT descending grid directly. Relative to + # numpy/JAX FFT convention its centered origin is the (-1)^j time factor; + # reversing bins here would also conjugate time evolution. + alternating = 1 - 2 * (xpy.arange(values.shape[0]) % 2) + return xpy.fft.ifft(values) * alternating / dt + + +def _continuous_forward(values, dt, xpy): + alternating = 1 - 2 * (xpy.arange(values.shape[0]) % 2) + return dt * xpy.fft.fft(values * alternating) + + +def _conjugate_spectrum(mode, xpy): + n = mode.shape[0] + reflection = (-xpy.arange(n)) % n + return xpy.conj(mode[reflection]) + + +def _lal_tdfromfd_shift_and_irfft( + one_sided_fd, delta_f, delta_t, epoch, extra_time, xpy +): + """Reproduce the deterministic shift and inverse FFT in LAL TDFromFD. + + The caller must supply the one-sided FD series on LAL's already-selected + grid and the explicit ``extra_time`` metadata. This routine deliberately + does not infer the grid or duration bounds. LAL rounds the shift to an + integer number of samples with C ``round`` semantics, multiplies bin ``k`` + by ``exp(+2 pi i k df tshift)``, advances the epoch by ``tshift``, and uses + its physical inverse-real-FFT normalization ``N*df``. + + This helper stops before LAL's high-pass filter, chirp-length snip, and + endpoint tapers. + """ + delta_f = _scalar(delta_f, "delta_f") + delta_t = _scalar(delta_t, "delta_t") + epoch = _scalar(epoch, "epoch") + extra_time = _scalar(extra_time, "extra_time") + if delta_f <= 0 or delta_t <= 0 or extra_time < 0: + raise WaveformCompatibilityError( + "delta_f and delta_t must be positive and extra_time nonnegative" + ) + if getattr(one_sided_fd, "ndim", None) != 1 or one_sided_fd.shape[0] < 3: + raise WaveformCompatibilityError( + "one-sided FD input must be a one-dimensional rFFT series" + ) + n = 2 * (int(one_sided_fd.shape[0]) - 1) + if not math.isclose(delta_t, 1.0 / (n * delta_f), rel_tol=5e-13): + raise WaveformCompatibilityError( + "explicit delta_t is inconsistent with the one-sided FD grid" + ) + # `extra_time` is nonnegative here, so C round(x) is floor(x + 1/2). + shift_samples = int(math.floor(extra_time / delta_t + 0.5)) + time_shift = shift_samples * delta_t + k = xpy.arange(one_sided_fd.shape[0], dtype=xpy.float64) + phase = xpy.exp(2.0j * math.pi * k * delta_f * time_shift) + shifted = one_sided_fd * phase + td = xpy.fft.irfft(shifted, n=n) * (n * delta_f) + return td, epoch + time_shift, shift_samples + + +def _rift_postprocess_td_modes(modes, epoch, delta_t, target_length, fmin, xpy): + """Apply RIFT's post-LAL resize and start taper on an array backend. + + This is the part of :func:`lalsimutils.hlmoft` *after* + ``SimInspiralTDModesFromPolarizations`` returns. It is intentionally kept + separate from LAL's own ``SimInspiralTDFromFD`` conditioning: reproducing + the latter also requires its auto-selected grid, sample-rounded time shift, + eighth-order high-pass filter, snip, and two-sided endpoint tapers. + + ``modes`` maps labels to equal-length one-dimensional arrays. LAL resize + semantics append zeros on the right when growing and discard samples from + the left when shrinking; a left discard advances the epoch by the same + number of samples. The returned arrays are new values, which keeps this + routine valid for JAX arrays. + """ + target_length = int(target_length) + delta_t = _scalar(delta_t, "delta_t") + fmin = _scalar(fmin, "fmin") + if target_length < 1 or delta_t <= 0 or fmin <= 0: + raise WaveformCompatibilityError( + "target_length, delta_t, and fmin must be positive" + ) + if not modes: + raise WaveformCompatibilityError("time-domain mode bank is empty") + lengths = {int(value.shape[0]) for value in modes.values()} + if len(lengths) != 1: + raise WaveformCompatibilityError("time-domain modes do not share a grid") + original_length = lengths.pop() + if original_length < 1: + raise WaveformCompatibilityError("time-domain modes are empty") + + left_discard = max(0, original_length - target_length) + resized = {} + for label, value in modes.items(): + value = value[left_discard:] + if original_length < target_length: + value = xpy.pad(value, (0, target_length - original_length)) + resized[label] = value + + # This exactly follows hlmoft: one percent of min(post-resize, original), + # but never less than one cycle at fmin. Valid RIFT configurations must + # leave enough samples for that taper; reject instead of broadcasting a + # malformed window. + ntaper = max( + int(0.01 * min(target_length, original_length)), + int(1.0 / (fmin * delta_t)), + ) + if ntaper > target_length: + raise WaveformCompatibilityError( + "RIFT start taper exceeds the requested time-series length" + ) + j = xpy.arange(ntaper, dtype=xpy.float64) + leading = 0.5 - 0.5 * xpy.cos(math.pi * j / float(ntaper)) + window = xpy.concatenate( + (leading, xpy.ones(target_length - ntaper, dtype=xpy.float64)) + ) + resized = {label: value * window for label, value in resized.items()} + return resized, float(epoch) + left_discard * delta_t, ntaper + + +def generate_imrphenomd_fd( + P, + Lmax=2, + backend=None, + *, + fd_standoff_factor=0.9, + fd_centering_factor=0.9, + conditioning="rift", + **_ignored, +): + """Generate direct-FD ``(2,+/-2)`` spectra natively with Ripple/JAX. + + The returned frequency grid is exactly RIFT's convention + ``f[k] = df * (N/2-k)``. The waveform normalization and mode mapping were + checked against LAL IMRPhenomD: Ripple's ``h0`` is the negative-frequency + ``(2,-2)`` carrier after multiplication by ``sqrt(16*pi/5)``. + + ``conditioning='direct_fd'`` is available for development and is + intentionally marked unconditioned so likelihood integration rejects it. + """ + dt, df, n = _validate(P, Lmax) + if conditioning == "rift": + raise WaveformCompatibilityError( + "exact RIFT IMRPhenomD conditioning is not certified: the legacy " + "path uses SimInspiralTDModesFromPolarizations; use " + "conditioning='direct_fd' only for explicit waveform tests" + ) + if conditioning not in ("raw", "direct_fd"): + raise WaveformCompatibilityError("conditioning must be 'rift' or 'direct_fd'") + if not 0 < _scalar(fd_standoff_factor, 'fd_standoff_factor') < 1: + raise WaveformCompatibilityError("fd_standoff_factor must lie strictly between 0 and 1") + xpy = _array_namespace(backend) + ripple = _load_ripple() + m1, m2, chi1, chi2 = _component_masses_msun(P) + mtot = m1 + m2 + eta = m1 * m2 / mtot**2 + mc = (m1 * m2) ** (3.0 / 5.0) / mtot ** (1.0 / 5.0) + dist_mpc = _scalar(P.dist, "dist") / 3.085677581491367e22 + if dist_mpc <= 0: + raise WaveformCompatibilityError("distance must be positive") + fmin = _scalar(P.fmin, "fmin") + fref = _scalar(getattr(P, "fref", 0.0), "fref") or fmin + phiref = _scalar(getattr(P, "phiref", 0.0), "phiref") + + # Positive frequencies are ascending for Ripple; the resulting arrays are + # assembled directly into RIFT packing, with no host copy. + fpos = df * xpy.arange(1, n // 2 + 1, dtype=xpy.float64) + params = xpy.asarray( + [mc, eta, chi1, chi2, dist_mpc, 0.0, phiref], dtype=xpy.float64 + ) + h0 = ripple.gen_IMRPhenomD(fpos, params, fref) + fmax = _scalar(getattr(P, "fmax", 0.0), "fmax") + if fmax <= 0: + fmax = 0.5 / dt + + low = fmin * float(fd_standoff_factor) + taper = xpy.where( + fpos < low, + 0.0, + xpy.where( + fpos < fmin, + 0.5 + + 0.5 + * xpy.cos(math.pi * (fpos / fmin - 1.0) / (1.0 - fd_standoff_factor)), + 1.0, + ), + ) + h0 = xpy.where(fpos <= fmax, h0 * taper, 0.0) + + norm = math.sqrt(16.0 * math.pi / 5.0) + z = xpy.zeros(n, dtype=xpy.complex128) + # RIFT indices 1..N/2-1 are +(Nyquist-df)..+df. Indices N/2+1.. + # are -df..-(Nyquist-df). A periodic complex FFT has only one Nyquist + # bin, so leave it zero to preserve the exact +/-m reflection identity. + h22 = z.at[1 : n // 2].set(norm * xpy.conj(h0[: n // 2 - 1][::-1])) + h2m2 = z.at[n // 2 + 1 :].set(norm * h0[: n // 2 - 1]) + raw_modes = {(2, 2): h22, (2, -2): h2m2} + + if conditioning in ("raw", "direct_fd"): + modes = raw_modes + conditioned = False + epoch = 0.0 + else: + raise WaveformCompatibilityError( + "conditioning must be 'rift' or 'direct_fd'" + ) + + conjugate = {lm: _conjugate_spectrum(v, xpy) for lm, v in modes.items()} + frequencies = df * (n // 2 - xpy.arange(n, dtype=xpy.float64)) + return DeviceFDModeBank( + modes=modes, + conjugate_modes=conjugate, + frequencies=frequencies, + delta_f=df, + delta_t=dt, + epoch=epoch, + conditioned=conditioned, + metadata={ + "N": n, + "mode_normalization": "sqrt(16*pi/5)", + "fd_standoff_factor": float(fd_standoff_factor), + "conditioning": conditioning, + "ripple_api": "ripplegw.waveforms.IMRPhenomD.gen_IMRPhenomD", + }, + ) + + +# Provider name used by the high-level GPU precompute dispatch. +imrphenomd_provider = generate_imrphenomd_fd + + +__all__ = [ + "DeviceFDModeBank", + "RippleUnavailableError", + "WaveformCompatibilityError", + "generate_imrphenomd_fd", + "imrphenomd_provider", +] diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/DESIGN_bounded_multipeak.md b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/DESIGN_bounded_multipeak.md new file mode 100644 index 000000000..a1a740e5b --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/DESIGN_bounded_multipeak.md @@ -0,0 +1,112 @@ +# Bounded multipeak JAX marginalization + +`--angle-marg-scheme multipeak-jax` integrates time, distance and both +angles with fixed-cap device discovery and fixed-order local quadrature. +There is no reserve, including on a decline. Discrete planning is stopped +from the AD graph; derivatives describe the retained fixed-plan integral. +This scheme is explicit only and is not selected by `auto`. + +## Configuration and accuracy + +`--multipeak-jax-` exposes every `BoundedMultipeakConfig` field, with +underscores replaced by hyphens. These are static values chosen before JIT +compilation. Important controls include `time-guard`, `base-max-starts`, +`max-time-nodes`, `max-modes`, `enriched-max-modes`, the four quadrature +orders, `convergence-tol-nats` and `total-value-error-budget-nats`. +The separate `--direct-marginalization-*` flags still configure the policy +with a reserve; they do not configure this scheme. + +The default resource envelope is compact: **guard 16, starts 32, time nodes +64, base/enriched modes 8/8**. This is a deliberate cost choice, not an +assumption that the policy's wider operating point is unnecessary everywhere. +With corrected quadrature it accepts the analytic reference at one tenth of +the wider envelope's warm A100 cost. Quadrature orders are **11/13/13/15**, convergence +**0.01 nat**, total empirical error budget **0.03 nat**. These are practical +starting settings, not a universal scientific accuracy requirement. An +application can lower orders, relax tolerances or change capacities to meet +its accuracy and runtime needs. The complete configuration is recorded in both output headers. + +To try the wider policy capacity profile on important declines, add: + +```text +--multipeak-jax-time-guard 128 --multipeak-jax-base-max-starts 128 +--multipeak-jax-max-time-nodes 256 --multipeak-jax-max-modes 16 +--multipeak-jax-enriched-max-modes 16 +``` + +The policy's recorded full-sky acceptance improvement at these wider caps +motivates exposing this profile, but acceptance fraction is not missed mass. +The compact default does not establish that its declined production rows are +negligible. Use their diagnostics and contribution comparisons to decide. + +## Declines and contribution + +The low-level bounded kernel returns `nan` plus its acceptance ledger for +unusable rows. The wrapper offers two actions: + +* `--multipeak-jax-decline-action drop` (default) gives declined proposals + finite log-zero (`-1e30`) in sampler calls. Their numerical importance + weight is zero, but they remain in the proposal count. Simply removing + nonfinite log weights would renormalize over survivors and change the + evidence. Declined rows are excluded from the exported cloud. +* `--multipeak-jax-decline-action refuse` stops on any declined host batch. + The refusal is latched, so an optional sampler stage catching an exception + cannot subsequently publish survivors. + +Dropping is useful when those rows contribute negligibly. A decline by itself +is not evidence of negligible contribution. Outputs in drop mode explicitly +identify an accepted-region target with unbounded omitted mass. The host audit +reports evaluated/declined counts (evaluations, not unique samples), decline +reasons, the maximum accepted log likelihood, the maximum finite *diagnostic* +value among declines, and the count with no finite diagnostic. Diagnostic +values failed the warrant: they are neither accepted likelihoods nor upper +bounds. A large declined diagnostic, missing modes, or unknown diagnostic +calls for a configuration change and an acceptance/value comparison, not a +reserve. A small diagnostic alone does not prove absent missed modes. + +For qualification, widen the relevant capacities or relax the appropriate +accuracy check, rerun the same proposals, and compare the contribution and +posterior/evidence stability at the accuracy needed by the application. +No output cloud consisting entirely of declined rows is published. +The audit and refusal latch cover only batched `log_likelihood` calls +(pilot/reweight/output evaluations). Scalar evaluations, including MAP, +Fisher and internal MALA training, are excluded. Under `refuse`, a scalar +decline returns NaN without raising or latching; a scalar-only decline can +therefore go unreported and does not prevent later publication of accepted +batches. Zero audited declines is not certification of scalar acceptance. +Both output headers record this scope and limitation explicitly. + +## Regression coverage + +The bounded suite exercises real discovery, refinement, quadrature and AD +on the analytic guarded-table fixture with an independent fine-time +reference. It also exercises low-order decline, invalid inputs/envelopes, +CLI-to-wrapper configuration, mixed-cloud publication, proposal-count +preservation, and refusal after a swallowed decline. Synthetic inputs are +constructed in tests; no captured scientific products are needed. + +## Operating-point comparison (2026-09-12) + +Two analytic `_guarded_problem` rows at scales 1 and 1.01, 33 native time +samples, guard 128, starts/time capacities 128/256 and 16 retained modes: +orders 11/13/13/15 accepted both rows even at the policy's tighter 0.001-nat +convergence setting and 0.01-nat total budget. Relaxing these tolerances to +0.01/0.03 nat gave no acceptance or accuracy benefit on this fixture; the +relaxation is a configurable starting accuracy choice, not a measured speedup. +The A100 (JAX 0.9.2, float64) took 3.70 s for the warm +two-row call and reported 2.57 MB of XLA temporary workspace (not total GPU +memory); CPU JAX 0.7.1 took 9.07 s. Returned log likelihoods were +36.3868065844 and 37.4751416172. The independent fine-time reference and AD +comparison are exercised by the regression test. + +Orders 7/9/9/11 still declined both rows with convergence loosened to 0.03 nat +and total budget 0.1 nat, and took 9.09 s on CPU / 3.66 s on A100. Lowering +orders alone did not reduce this fixture's dominant discovery/refinement cost. These are +synthetic checks, not a production acceptance-rate or throughput claim. + +The compact caps (guard 16, starts/time nodes 32/64, modes 8/8), with +11/13/13/15 orders, also accepted both rows at the tighter 0.001-nat setting: +log likelihoods 36.3868061719 and 37.4751411797. Warm A100 cost was **0.355 s +per two-row call** and temporary workspace 2.20 MB. That measured cost, +together with explicit drop diagnostics and configurable caps, is why this +profile is the default instead of automatically inheriting the wider policy. diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/README.md b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/README.md index f29f3480e..bb7bc95c6 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/README.md +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/README.md @@ -152,6 +152,22 @@ conditional time per exported row. This intentionally differs from conventional ILE's XML export semantics, but a high-level DAG can swap executables without dying during option parsing. +For an Asimov/pseudo-pipe final fair-draw stage, the pipeline detects the JAX +executable and runs `util_ConvertJAXILEFairdraws.py`. The converter strictly +pairs each tabular sidecar with its intrinsic likelihood record and writes the +usual joint posterior coordinates that are actually available. It omits +`time`, which remains marginalized and is not exported by this driver, rather +than fabricating a coordinate. Redshift and source-frame masses are likewise +not inferred; custom `--convert-args` are refused instead of silently ignored. +The compatibility columns `p` and `ps` are both the neutral value one because +the exported rows are already equal-weight fair draws. Missing pairs, +malformed records, nonfinite rows, noncontiguous grid IDs, or unexpected +intrinsic/draw counts fail the terminal job and remove any stale terminal +output. Conventional ILE retains its existing XML conversion path. The +converter also writes a JSON provenance ledger beside the posterior, recording +input and output hashes, row counts, shuffle seed, neutral columns, and omitted +coordinates. + ## Modules - `detector.py` — `compute_detamresponse`, `time_delay_from_earth_center` @@ -264,6 +280,11 @@ the same JAX likelihood selected by ``--mode`` but do not differentiate it during integration. Likelihood rows are evaluated in one fixed JAX shape; ``--jax-av-eval-chunk`` therefore controls accelerator memory independently of the larger ``--n-chunk`` used to cover and contract the adaptive volume. +AV/portfolio honor the production ``--d-prior pseudo_cosmo`` distance density, +including its normalization over ``[--d-min, --d-max]``; Euclidean/volumetric +remains the default. Other cosmological distance-prior variants are refused +for this backend rather than silently changed. A sampling-only +``--limit-distance`` does not renormalize either physical prior. Portfolio defaults to AV plus a defensive GMM member. An optional Fisher-sky initializer pays an explicit, one-time AD cost for hill climbing and local @@ -536,3 +557,14 @@ the mode-covering samplers above. Further hardening available to compound: Not yet ported (structured for): in-loop calibration marginalization (`n_cal>1`) and the lookup-table distance marginalization (we use direct grid quadrature instead, which is AD-friendly and needs no precomputed table). +Waveform precomputation uses the same two-second post-event FD alignment as +production numpy ILE. The compatibility options +``--internal-waveform-fd-L-frame`` and +``--internal-waveform-fd-no-condition`` are forwarded to the production +precompute call; they are not JAX-only transformations. +Input frames are first loaded at ``--srate`` and, when requested, upsampled to +``--srate-internal`` before the mode time series are constructed, matching the +two-cadence production ILE path. +The production defaults also retain all modes (no implicit precompute +threshold), retain memory modes unless ``--no-memory`` is given, and apply the +same ``--fmin-ifo`` and PSD-window normalization to each detector. diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/anglemarg.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/anglemarg.py index 323d93307..0ee1bf75e 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/anglemarg.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/anglemarg.py @@ -162,7 +162,7 @@ ANGLE_MARG_DEFAULT = "exact" ANGLE_MARG_LEGACY = "grid" # the spelling that reproduces pre-2026-09-02 runs ANGLE_MARG_CHOICES = ("grid", "exact", "laplace", "peak-local", "phi-local", - "multipeak", "auto") + "multipeak", "multipeak-jax", "auto") #: 'peak-local' is deliberately NOT reachable from 'auto' yet. It agrees with 'exact' #: to 1e-13 nats on the tables measured so far and is device-independent (the same answer diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/core.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/core.py index 0371c5f4a..05b27e27e 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/core.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/core.py @@ -431,6 +431,7 @@ def _gather(Q_col, pos, u=None): valid = (idx >= 0) & (idx < n) vals = Q_col[jnp.clip(idx, 0, n - 1)] return jnp.sum(w * jnp.where(valid, vals, 0.0 + 0.0j), axis=-1) + _gather._stencil_size = 2 * a return _gather @@ -847,6 +848,100 @@ def _banded_coefficients(data, det, ra, dec, psi): raise ValueError("unknown banded feature %r" % (data.feature,)) +def _banded_chunk_shape(A, K, S, npts, nfull, taps, budget=128 * 1024**2): + """Choose row/sample tiles using a conservative forward scratch estimate. + + Reserve four complex arrays per stencil tap and twelve per output element, + plus the selected source rows. This covers unfused gathers, masks, weights, + products and reductions; it excludes input banks, the final output, compiler + workspaces and reverse-mode saved residuals. It is an allocation estimate, + not a guarantee about XLA's whole-executable memory use. + """ + element_bytes = 16 * (4 * taps + 12) + source_bytes = 16 * nfull + samples = min(S, (budget - source_bytes) // (npts * element_bytes)) + if samples < 1: + raise ValueError("one banded gather sample exceeds the scratch budget") + row_bytes = source_bytes + samples * npts * element_bytes + rows = min(A * K, budget // row_bytes) + return int(rows), int(samples), int(rows * row_bytes) + + +def _contract_banded_data_term(Q_bank, conjY, C, gather, pos, u_sep, + *, pp_t1=None, pe=None, pt=None): + """Contract the banded ```` term without unrolling ``A * K`` in Python. + + ``Q_bank`` has shape ``(A, n_time_full, K)``. The former implementation + spelled out one Python/JAX expression for every ``(a, k)`` pair. For a + compound response bank (``A=40`` at ``p_max=Qmax=1``), + that made XLA trace and compile time scale with the number of basis/mode + pairs and is a suspected contributor to excessive cold compilation despite + the actual operation being a small regular reduction. + + Vectorized row gathers run inside static-bound loops over row and sample + tiles. Small batches gather all A*K rows together; large AV batches use + smaller tiles under the forward scratch estimate in _banded_chunk_shape. + The graph stays compact and reverse-mode differentiable without creating + a full production (A, K, S, npts) temporary. + + Supplying all of ``pp_t1``, ``pe`` and ``pt`` applies the separable + arrival-time post-phase. Supplying none retains the frequency-response-only + contraction. Partial post-phase inputs are rejected instead of silently + evaluating a mixed convention. + """ + A = int(Q_bank.shape[0]) + K = int(Q_bank.shape[2]) + S = int(conjY.shape[0]) + npts = int(pos.shape[1]) + + phase_args = (pp_t1, pe, pt) + post_phase = all(x is not None for x in phase_args) + if post_phase != any(x is not None for x in phase_args): + raise ValueError("pp_t1, pe, and pt must be supplied together") + if post_phase: + pp_t1 = jnp.asarray(pp_t1, dtype=jnp.int32) + if S == 0: + # The previous reduction returned an empty band for empty sampler + # chunks. Avoid selecting a zero-sized tile (and dividing by it). + return jnp.zeros((0, npts), dtype=jnp.complex128) + + taps = {_gather_nearest: 1, _gather_linear: 2, _gather_cubic: 4}.get( + gather, getattr(gather, "_stencil_size", 16)) + rows, samples, _ = _banded_chunk_shape( + A, K, S, npts, int(Q_bank.shape[1]), taps) + nblocks = (S + samples - 1) // samples + + def sample_block(block, output): + sample_ids = jnp.minimum(block * samples + jnp.arange(samples), S - 1) + positions = pos[sample_ids] + fractions = None if u_sep is None else u_sep[sample_ids] + + def row_block(block_row, kappa): + row_ids = block_row * rows + jnp.arange(rows) + valid = row_ids < A * K + safe_ids = jnp.minimum(row_ids, A * K - 1) + a, k = safe_ids // K, safe_ids % K + source = Q_bank[a, :, k] + values = jax.vmap(lambda q: gather(q, positions, fractions))(source) + weights = (jnp.conj(C[a[:, None], sample_ids[None, :]]) + * conjY[sample_ids[None, :], k[:, None]]) + weights = jnp.where(valid[:, None], weights, 0.0j) + if post_phase: + im = pp_t1[a] + weights = weights * pe[im[:, None], sample_ids[None, :]] + values = values * pt[im, None, :] + return kappa + jnp.sum(weights[:, :, None] * values, axis=0) + + value = jax.lax.fori_loop( + 0, (A * K + rows - 1) // rows, row_block, + jnp.zeros((samples, npts), dtype=jnp.complex128)) + return jax.lax.dynamic_update_slice(output, value, (block * samples, 0)) + + return jax.lax.fori_loop( + 0, nblocks, sample_block, + jnp.zeros((nblocks * samples, npts), dtype=jnp.complex128))[:S] + + def _accumulate_unit_banded(data, ra, dec, psi, incl, phiref, interp, phase_marginalization, guard=0): """Multi-band (slow-rotation / finite-size) network kappa and rho^2. @@ -1010,18 +1105,13 @@ def _accumulate_unit_banded(data, ra, dec, psi, incl, phiref, interp, # --- term1: sum_a conj(C~_a) * ( sum_lm conj(Y_lm) Q^a_lm(t) ) --- # conj(C~_a) = conj(C_a) exp(-i n_a omega delta), i.e. the m = -n_a bucket. - kappa_det = jnp.zeros((S, npts), dtype=jnp.complex128) - for a in range(A): - inner_a = jnp.zeros((S, npts), dtype=jnp.complex128) - Qa = Q_bank[a] # (npts_full, K) - for k in range(K): - inner_a = inner_a + conjY[:, k][:, None] * gather(Qa[:, k], pos, u_sep) - if post_phase: - i1 = int(pp_t1[a]) - kappa_det = kappa_det + ((jnp.conj(C[a]) * pe[i1])[:, None] - * (pt[i1][None, :] * inner_a)) - else: - kappa_det = kappa_det + jnp.conj(C[a])[:, None] * inner_a + if post_phase: + kappa_det = _contract_banded_data_term( + Q_bank, conjY, C, gather, pos, u_sep, + pp_t1=pp_t1, pe=pe, pt=pt) + else: + kappa_det = _contract_banded_data_term( + Q_bank, conjY, C, gather, pos, u_sep) kappa_unit = kappa_unit + kappa_det # --- term2: 0.5 Re[ sum_{a,a'} conj(C~_a)C~_a' YbarUY + C~_aR C~_a' YVY ] --- diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/direct_marginalization_policy.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/direct_marginalization_policy.py index 0a0ca7bbc..498ac794b 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/direct_marginalization_policy.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/direct_marginalization_policy.py @@ -58,6 +58,8 @@ "resolve_reserve_angular_kernel", "q_bandwidth_cycles_per_sample", "PolicyConfig", + "BoundedMultipeakConfig", + "validate_bounded_multipeak_config", "validate_policy_request", "validate_policy_config", "policy_log_normalization", @@ -65,6 +67,7 @@ "probe_guarded_tables", "policy_acceptance_diagnostics", "fused_log_likelihood_four_axis_policy", + "fused_log_likelihood_four_axis_bounded", "summarize_policy_ledger", "RESERVE_SCHEME_CHOICES", "RESERVE_SCHEME_DEFAULT", @@ -298,6 +301,86 @@ class PolicyConfig(NamedTuple): reserve_peaklocal_sigma_t_override_samples: float = float("nan") +class BoundedMultipeakConfig(NamedTuple): + """Static resource envelope for the device-only multipeak variant. + + Every field that changes compiled work is fixed before tracing. A row + that needs more starts, time nodes, modes, or quadrature accuracy than + this envelope supplies is returned as ``nan``; there is deliberately no + amplitude-sized dense reserve. Planning is stopped from the AD graph, so + derivatives are those of the accepted fixed-plan integral, not a claim + about derivatives of the discrete mode-selection map. + """ + + # Compact by default: with the corrected orders this accepts the analytic + # reference and costs 0.36 s / two rows on A100, versus 3.70 s with the + # policy's 128/128/256 guard/start/time caps (2026-09-12). Wider caps can + # recover important declines; expose them instead of paying for every + # prior draw. See DESIGN_bounded_multipeak.md for the tradeoff and profile. + time_guard: int = 16 + base_max_starts: int = 32 + max_time_nodes: int = 64 + base_oversample: int = PolicyConfig().base_oversample + enriched_oversample: int = PolicyConfig().enriched_oversample + max_modes: int = 8 + enriched_max_modes: int = 8 + local_radius: float = PolicyConfig().local_radius + refine_iterations: int = PolicyConfig().refine_iterations + # Deliberately cheaper than the reserve-bearing policy: 11/13/13/15 + # accepts the analytic reference where 7/9/9/11 declines even at 0.03 nat. + # 0.01-nat convergence / 0.03-nat total budget favors practical acceptance; + # all remain explicit CLI controls. See DESIGN_bounded_multipeak.md. + base_order: int = 11 + base_check_order: int = 13 + enriched_order: int = 13 + enriched_check_order: int = 15 + convergence_tol_nats: float = 1.0e-2 + time_guard_tol_nats: float = PolicyConfig().time_guard_tol_nats + total_value_error_budget_nats: float = 3.0e-2 + time_outside_tol_nats: float = PolicyConfig().time_outside_tol_nats + norm_invariance_rtol: float = PolicyConfig().norm_invariance_rtol + batch_rows: int = 1 + + +def validate_bounded_multipeak_config(config): + """Validate the complete static contract before JAX traces it.""" + if not isinstance(config, BoundedMultipeakConfig): + raise TypeError("config must be a BoundedMultipeakConfig") + if int(config.time_guard) < 2: + raise ValueError("time_guard must be >= 2 for two-guard validation") + if int(config.base_max_starts) < 1 or int(config.max_time_nodes) < 2: + raise ValueError("base_max_starts >= 1 and max_time_nodes >= 2 are required") + if (int(config.base_oversample) < 1 + or int(config.enriched_oversample) <= int(config.base_oversample)): + raise ValueError("enriched_oversample must exceed base_oversample >= 1") + if (int(config.max_modes) < 1 + or int(config.max_modes) > int(config.base_max_starts) + or int(config.enriched_max_modes) < int(config.max_modes) + or int(config.enriched_max_modes) > 2 * int(config.base_max_starts)): + raise ValueError( + "mode caps must satisfy 1 <= max_modes <= base_max_starts, " + "max_modes <= enriched_max_modes <= 2 * base_max_starts") + orders = (int(config.base_order), int(config.base_check_order), + int(config.enriched_order), int(config.enriched_check_order)) + if not (2 <= orders[0] < orders[1] <= orders[2] < orders[3]): + raise ValueError("need 2 <= base_order < base_check_order <= " + "enriched_order < enriched_check_order") + if not float(config.local_radius) > 0.0 or int(config.refine_iterations) < 1: + raise ValueError("local_radius and refine_iterations must be positive") + if not (float(config.convergence_tol_nats) > 0.0 + and float(config.time_guard_tol_nats) > 0.0 + and float(config.total_value_error_budget_nats) > 0.0): + raise ValueError("all error tolerances must be positive") + if not (np.isfinite(float(config.time_outside_tol_nats)) + and float(config.time_outside_tol_nats) < 0.0): + raise ValueError("time_outside_tol_nats must be finite and negative") + if not (np.isfinite(float(config.norm_invariance_rtol)) + and float(config.norm_invariance_rtol) >= 0.0): + raise ValueError("norm_invariance_rtol must be finite and non-negative") + validate_batch_rows(config.batch_rows) + return config + + def validate_batch_rows(batch_rows): """Refuse a row-batch size the controller cannot execute. @@ -1405,6 +1488,161 @@ def _row(args): return lnL +def fused_log_likelihood_four_axis_bounded( + data, ra, dec, incl, x_grid, log_w_grid, *, interp, amp_sizing, + config=None, local_log_normalization=None, x_bounds=None, + return_ledger=False): + """Device-only, statically bounded four-axis multipeak marginalization. + + Discovery, Newton refinement, retained modes, and both nested quadrature + rules are fixed-shape functions of :class:`BoundedMultipeakConfig`. No + dense or amplitude-sized reserve is present. A row that does not pass the + existing empirical enrichment, quadrature, geometry, time-guard, and + omitted-time gates fails closed with ``nan``. + + The function is JIT- and AD-compatible. Mode planning consumes + ``stop_gradient`` coefficient tables and the resulting plans are stopped + again, bounding reverse-mode storage and defining AD as differentiation of + the accepted fixed-plan integral. This is not a derivative-accuracy + warrant; value/gradient parity on production ladders remains separate. + + ``x_bounds`` and ``local_log_normalization`` may be supplied by a wrapper + that traces this function. When omitted they are derived eagerly from the + concrete distance grid. + """ + if config is None: + config = BoundedMultipeakConfig() + validate_bounded_multipeak_config(config) + guard = int(config.time_guard) + batch_rows = validate_batch_rows(config.batch_rows) + if local_log_normalization is None: + local_log_normalization, _ = policy_log_normalization( + data, x_grid, log_w_grid) + if x_bounds is None: + x_host = np.asarray(x_grid) + x_min, x_max = float(np.min(x_host)), float(np.max(x_host)) + else: + x_min, x_max = (float(x_bounds[0]), float(x_bounds[1])) + if not (0.0 < x_min < x_max): + raise ValueError("x_bounds must satisfy 0 < x_min < x_max") + + x_grid = jnp.asarray(x_grid, dtype=jnp.float64) + log_w_grid = jnp.asarray(log_w_grid, dtype=jnp.float64) + C_A, C_B, _ = _anglemarg.angle_coefficient_tables( + data, ra, dec, incl, interp, guard=guard) + # amp_sizing is retained for call compatibility, not used to size work. + rows_A = jnp.moveaxis(C_A, 2, 0) + rows_B = jnp.moveaxis(C_B, 2, 0) + norm0 = rows_B[..., 0] + norm_dev = jnp.max(jnp.abs(rows_B - norm0[..., None]), axis=(1, 2, 3)) + norm_scale = jnp.maximum(1.0, jnp.max(jnp.abs(norm0), axis=(1, 2))) + norm_time_invariant = ( + norm_dev <= float(config.norm_invariance_rtol) * norm_scale) + tables_finite = (jnp.all(jnp.isfinite(rows_A), axis=(1, 2, 3)) + & jnp.all(jnp.isfinite(rows_B), axis=(1, 2, 3))) + + def _plan_row(table, norm): + base = _aap.rank_joint_starts_from_uvq_device( + table, norm, x_min, x_max, time_guard=guard, + max_starts=int(config.base_max_starts), + max_time_nodes=int(config.max_time_nodes), + angular_oversample=int(config.base_oversample)) + extra = _aap.rank_joint_starts_from_uvq_device( + table, norm, x_min, x_max, time_guard=guard, + max_starts=int(config.base_max_starts), + max_time_nodes=int(config.max_time_nodes), + angular_oversample=int(config.enriched_oversample)) + base_plan, enriched_plan, bp, ep, shared = ( + _aap.make_all_axis_mode_plan_pair_device( + table, norm, base, extra, x_min, x_max, + max_modes=int(config.max_modes), + enriched_max_modes=int(config.enriched_max_modes), + local_radius=float(config.local_radius), time_guard=guard, + iterations=int(config.refine_iterations), + time_reconstruction_certified=False)) + base_plan = jax.tree.map(jax.lax.stop_gradient, base_plan) + enriched_plan = jax.tree.map(jax.lax.stop_gradient, enriched_plan) + planning = dict( + base_n_selected_modes=bp["n_selected_modes"], + enriched_n_selected_modes=ep["n_selected_modes"], + base_n_optimizer_starts=bp["n_optimizer_starts"], + enriched_n_optimizer_starts=ep["n_optimizer_starts"], + optimizer_starts_executed=shared["n_optimizer_starts_executed"], + base_n_lattice_evaluations=bp["n_lattice_evaluations"], + enriched_n_lattice_evaluations=ep["n_lattice_evaluations"], + base_n_candidates_before_cap=bp["n_candidates_before_cap"], + combined_n_candidates_before_cap=ep["n_candidates_before_cap"], + base_start_capacity_ok=bp["start_capacity_ok"], + combined_start_capacity_ok=ep["start_capacity_ok"], + base_time_capacity_ok=bp["time_capacity_ok"], + combined_time_capacity_ok=ep["time_capacity_ok"]) + return base_plan, enriched_plan, planning + + base_plans, enriched_plans, planning = jax.vmap(_plan_row)( + jax.lax.stop_gradient(rows_A), jax.lax.stop_gradient(norm0)) + + def _row(args): + table, norm, base_plan, enriched_plan = args + return _aap.empirical_enrichment_marginalize( + table, norm, base_plan, enriched_plan, x_min, x_max, + base_order=int(config.base_order), + base_check_order=int(config.base_check_order), + enriched_order=int(config.enriched_order), + enriched_check_order=int(config.enriched_check_order), + convergence_tol_nats=float(config.convergence_tol_nats), + time_guard=guard, + time_guard_tol_nats=float(config.time_guard_tol_nats), + log_normalization=float(local_log_normalization), + time_outside_tol_nats=float(config.time_outside_tol_nats), + total_value_error_budget_nats=float( + config.total_value_error_budget_nats)) + + n_rows = int(rows_A.shape[0]) + xs = (rows_A, norm0, base_plans, enriched_plans) + if batch_rows == 1 or n_rows == 1: + selected, accepted, ledger = jax.lax.map(_row, xs) + batch_executed = 1 + elif batch_rows == 0 or batch_rows >= n_rows: + selected, accepted, ledger = jax.vmap(_row)(xs) + batch_executed = n_rows + else: + selected, accepted, ledger = jax.lax.map( + _row, xs, batch_size=batch_rows) + batch_executed = batch_rows + + usable = accepted & norm_time_invariant & tables_finite + lnL = jnp.where(usable, selected, jnp.nan) + ledger = dict(ledger) + ledger.update(planning) + ledger.update( + bounded_cost=jnp.full((n_rows,), True, dtype=bool), + dense_reserve_available=jnp.full((n_rows,), False, dtype=bool), + fixed_plan_autodiff_only=jnp.full((n_rows,), True, dtype=bool), + derivative_warrant_certified=jnp.full((n_rows,), False, dtype=bool), + batch_rows_requested=jnp.full( + (n_rows,), batch_rows, dtype=jnp.int32), + batch_rows_executed=jnp.full( + (n_rows,), batch_executed, dtype=jnp.int32), + max_starts_cap=jnp.full( + (n_rows,), int(config.base_max_starts), dtype=jnp.int32), + max_time_nodes_cap=jnp.full( + (n_rows,), int(config.max_time_nodes), dtype=jnp.int32), + max_modes_cap=jnp.full( + (n_rows,), int(config.enriched_max_modes), dtype=jnp.int32), + decline_norm_time_variation=~norm_time_invariant, + decline_input_nonfinite=~tables_finite, + norm_time_invariant=norm_time_invariant, + norm_time_deviation=norm_dev, + tables_finite=tables_finite, + input_nonfinite=~tables_finite, + usable=usable, + selected_value=selected, + lnL=lnL) + if return_ledger: + return lnL, ledger + return lnL + + def summarize_policy_ledger(ledger): """Host-side counts for the run record. ``ledger`` leaves are ``(S,)``.""" def _count(key): diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/samplers.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/samplers.py index 5bfcca84c..8a9e95d22 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/samplers.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/samplers.py @@ -34,6 +34,7 @@ python RIFT/likelihood/jax_ile/samplers.py """ +import functools import os import numpy as np @@ -2228,8 +2229,12 @@ def _ess(db): # post_weight, i.e. exactly the mislabelling the tail guards against. db = min(max(a, 1e-4), hi_db) # SMC evidence increment: logZ += logmeanexp(db * lnL) - z = db * lnL - z = z[np.isfinite(z)] + # The SMC ratio is a mean over ALL W prior walkers. A walker with + # zero/invalid likelihood contributes zero to the numerator but still + # occupies its share of the prior mass. Dropping it before np.mean + # renormalizes onto the finite subset and biases logZ upward by + # log(W / n_finite) on this rung. + z = np.where(np.isfinite(lnL), db * lnL, -np.inf) mz = np.max(z) logZ += float(mz + np.log(np.mean(np.exp(z - mz)))) inv_T += db @@ -2363,7 +2368,8 @@ def _av_param_order(like): def _av_sample_bounds(order, d_min, d_max, sample_d_min=None, - sample_d_max=None, sample_bounds=None): + sample_d_max=None, sample_bounds=None, + distance_prior="euclidean"): """Validate and resolve AV sampling limits without renormalizing the prior.""" requested = dict(sample_bounds or {}) if sample_d_min is not None or sample_d_max is not None: @@ -2376,7 +2382,8 @@ def _av_sample_bounds(order, d_min, d_max, sample_d_min=None, % ", ".join(sorted(unknown))) resolved = {} for name in order: - physical_lo, physical_hi, _ = _av_prior_spec(name, d_min, d_max) + physical_lo, physical_hi, _ = _av_prior_spec( + name, d_min, d_max, distance_prior=distance_prior) lo, hi = requested.get(name, (physical_lo, physical_hi)) lo, hi = float(lo), float(hi) if not np.isfinite(lo) or not np.isfinite(hi) or lo >= hi: @@ -2389,12 +2396,15 @@ def _av_sample_bounds(order, d_min, d_max, sample_d_min=None, return resolved -def _av_prior_draw(order, n, rng, d_min, d_max, sample_bounds=None): +def _av_prior_draw(order, n, rng, d_min, d_max, sample_bounds=None, + distance_prior="euclidean"): """Draw the physical prior conditioned only on the AV sampling window.""" bounds = _av_sample_bounds(order, d_min, d_max, - sample_bounds=sample_bounds) + sample_bounds=sample_bounds, + distance_prior=distance_prior) def interval(name): - return bounds.get(name, _av_prior_spec(name, d_min, d_max)[:2]) + return bounds.get(name, _av_prior_spec( + name, d_min, d_max, distance_prior=distance_prior)[:2]) ra_lo, ra_hi = interval("ra") if "ra" in order else (0.0, _TWO_PI) dec_lo, dec_hi = interval("dec") if "dec" in order else (-_PI / 2, _PI / 2) psi_lo, psi_hi = interval("psi") if "psi" in order else (0.0, _PI) @@ -2413,13 +2423,37 @@ def interval(name): "phiref_shifted": rng.uniform(phase_lo, phase_hi, n), "phase_p": rng.uniform(pp_lo, pp_hi, n), "phase_m": rng.uniform(pm_lo, pm_hi, n), - "distMpc": np.cbrt(rng.uniform(dist_lo ** 3, dist_hi ** 3, n)), + "distMpc": _av_distance_prior_draw( + n, rng, dist_lo, dist_hi, d_min, d_max, distance_prior), } return np.column_stack([draws[name] for name in order]) +@functools.lru_cache(maxsize=32) +def _pseudo_cosmo_norm(d_min, d_max): + from RIFT.likelihood import priors_utils + return float(priors_utils.dist_prior_pseudo_cosmo_eval_norm(d_min, d_max)) + + +def _av_distance_prior_draw(n, rng, lo, hi, d_min, d_max, distance_prior): + """Draw a distance prior conditioned on the declared sampling interval.""" + key = str(distance_prior or "euclidean").strip().lower() + if key in ("euclidean", "volumetric"): + return np.cbrt(rng.uniform(lo ** 3, hi ** 3, n)) + if key != "pseudo_cosmo": + raise ValueError("unsupported JAX-AV distance prior %r" % distance_prior) + # This is proposal initialization only. A dense deterministic inverse CDF + # is ample here; the estimator itself uses the analytic prior density below. + from RIFT.likelihood import priors_utils + grid = np.linspace(float(lo), float(hi), 4097) + pdf = np.asarray(priors_utils.dist_prior_pseudo_cosmo( + grid, nm=_pseudo_cosmo_norm(float(d_min), float(d_max))), dtype=float) + cdf = np.concatenate(([0.0], np.cumsum(0.5 * (pdf[1:] + pdf[:-1]) * np.diff(grid)))) + return np.interp(rng.uniform(0.0, cdf[-1], int(n)), cdf, grid) + + def _av_prior_spec(name, d_min, d_max, sample_d_min=None, sample_d_max=None, - sample_bounds=None): + sample_bounds=None, distance_prior="euclidean"): """Return ``(lo, hi, physical_density)`` for one wrapper coordinate.""" if name == "ra": spec = (0.0, _TWO_PI, lambda x: np.ones(np.shape(x)) / _TWO_PI) @@ -2437,10 +2471,20 @@ def _av_prior_spec(name, d_min, d_max, sample_d_min=None, sample_d_max=None, spec = (0.0, 2.0 * _TWO_PI, lambda x: np.ones(np.shape(x)) / (2.0 * _TWO_PI)) elif name == "distMpc": - norm = 3.0 / (float(d_max) ** 3 - float(d_min) ** 3) + key = str(distance_prior or "euclidean").strip().lower() + if key in ("euclidean", "volumetric"): + norm = 3.0 / (float(d_max) ** 3 - float(d_min) ** 3) + density = lambda x, _norm=norm: _norm * np.asarray(x) ** 2 + elif key == "pseudo_cosmo": + from RIFT.likelihood import priors_utils + norm = _pseudo_cosmo_norm(float(d_min), float(d_max)) + density = lambda x, _norm=norm: priors_utils.dist_prior_pseudo_cosmo( + np.asarray(x), nm=_norm, xpy=np) + else: + raise ValueError("unsupported JAX-AV distance prior %r" % distance_prior) spec = (float(d_min if sample_d_min is None else sample_d_min), float(d_max if sample_d_max is None else sample_d_max), - lambda x, _norm=norm: _norm * np.asarray(x) ** 2) + density) else: raise ValueError("unsupported JAX-ILE adaptive-volume parameter %r" % (name,)) if sample_bounds and name in sample_bounds: @@ -2495,7 +2539,8 @@ def _sky_distance(a, b): def _fisher_sky_seed(like, order, lnL, rng, d_min, d_max, n_seed, n_pilot, n_modes, sky_inflate, prior_frac, - initial_points=None, sample_bounds=None, verbose=False): + initial_points=None, sample_bounds=None, verbose=False, + distance_prior="euclidean"): """Hill-climb modes, then draw Fisher sky / physical-prior other coordinates. This is deliberately a proposal initializer, not part of the estimator. The @@ -2507,7 +2552,8 @@ def _fisher_sky_seed(like, order, lnL, rng, d_min, d_max, n_seed, from scipy.optimize import minimize n_pilot = max(int(n_pilot), int(n_modes), 1) - pilot = _av_prior_draw(order, n_pilot, rng, d_min, d_max, sample_bounds) + pilot = _av_prior_draw(order, n_pilot, rng, d_min, d_max, sample_bounds, + distance_prior) pilot_lnL = lnL(*pilot.T) ranked = np.argsort(np.where(np.isfinite(pilot_lnL), pilot_lnL, -np.inf))[::-1] seeds = [] @@ -2529,7 +2575,8 @@ def _fisher_sky_seed(like, order, lnL, rng, d_min, d_max, n_seed, raise RuntimeError("the Fisher-sky prior pilot found no finite likelihood") bounds = [_av_prior_spec(name, d_min, d_max, - sample_bounds=sample_bounds)[:2] for name in order] + sample_bounds=sample_bounds, + distance_prior=distance_prior)[:2] for name in order] # Avoid evaluating the Jacobian-singular orientation endpoints during AD. bounds = [(lo + 1e-6 if name in ("dec", "incl") else lo, hi - 1e-6 if name in ("dec", "incl") else hi) @@ -2571,7 +2618,8 @@ def objective(x): n_seed = max(int(n_seed), len(order) + 2) n_prior = int(np.clip(float(prior_frac), 0.0, 1.0) * n_seed) n_focus = n_seed - n_prior - focused = _av_prior_draw(order, n_focus, rng, d_min, d_max, sample_bounds) + focused = _av_prior_draw(order, n_focus, rng, d_min, d_max, sample_bounds, + distance_prior) counts = np.full(len(modes), n_focus // len(modes), dtype=int) counts[:n_focus % len(modes)] += 1 cursor = 0 @@ -2593,7 +2641,7 @@ def objective(x): cloud = focused if n_prior: cloud = np.vstack([cloud, _av_prior_draw( - order, n_prior, rng, d_min, d_max, sample_bounds)]) + order, n_prior, rng, d_min, d_max, sample_bounds, distance_prior)]) if verbose: print(" [JAX-AV seed] %d hill-climbed sky mode(s), %d seed points " "(%d full-prior)" % (len(modes), len(cloud), n_prior)) @@ -2610,7 +2658,7 @@ def adaptive_volume_sample(like, d_min, d_max, sampler_method="AV", seed_prior_frac=0.1, anisotropic_bins=True, gmm_components=2, verbose=False, sample_d_min=None, sample_d_max=None, - sample_bounds=None): + sample_bounds=None, distance_prior="euclidean"): """Run production AV/portfolio control logic on a value-only JAX likelihood. ``sampler_method`` is ``AV`` or ``portfolio``. The optional ``fisher-sky`` @@ -2627,7 +2675,8 @@ def adaptive_volume_sample(like, d_min, d_max, sampler_method="AV", order = _av_param_order(like) n_dim = len(order) resolved_bounds = _av_sample_bounds( - order, d_min, d_max, sample_d_min, sample_d_max, sample_bounds) + order, d_min, d_max, sample_d_min, sample_d_max, sample_bounds, + distance_prior) # Decouple AV's coverage cloud from the accelerator batch. AV often needs # a large n_chunk to hit a narrow sky mode, while the marginalized JAX # kernel has a much smaller memory-efficient batch. The callback loops over @@ -2665,7 +2714,8 @@ def adaptive_volume_sample(like, d_min, d_max, sampler_method="AV", for name in order: lo, hi = resolved_bounds[name] - prior = _av_prior_spec(name, d_min, d_max)[2] + prior = _av_prior_spec(name, d_min, d_max, + distance_prior=distance_prior)[2] sampler.add_parameter(name, pdf=None, left_limit=lo, right_limit=hi, prior_pdf=prior, adaptive_sampling=True) setup_kwargs = {"anisotropic_bins": bool(anisotropic_bins)} @@ -2698,7 +2748,8 @@ def adaptive_volume_sample(like, d_min, d_max, sampler_method="AV", n_seed=(seed_points or n_chunk), n_pilot=seed_pilot, n_modes=seed_modes, sky_inflate=sky_inflate, prior_frac=seed_prior_frac, initial_points=seed_initial_points, - sample_bounds=resolved_bounds, verbose=verbose) + sample_bounds=resolved_bounds, verbose=verbose, + distance_prior=distance_prior) if method == "portfolio": sampler.bootstrap_from_samples(seed_cloud, params=order, seed=seed) else: @@ -2716,6 +2767,9 @@ def adaptive_volume_sample(like, d_min, d_max, sampler_method="AV", result = sampler.integrate_log( lnL, *order, nmax=int(nmax), neff=float(neff), n=int(n_chunk), no_protect_names=True, verbose=bool(verbose), save_intg=True, + # Match classic ILE: fractional edge bins otherwise extend beyond + # the requested sampling box, where the physical prior is nonzero. + enforce_bounds=True, tempering_exp=1.0, anisotropic_bins=bool(anisotropic_bins), # Standalone AV can keep device-typed internal arrays when cupy is # importable even though this adapter evaluates on the host. Its diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/wrapper.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/wrapper.py index 4f37d315e..a657d2db8 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/wrapper.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/jax_ile/wrapper.py @@ -50,6 +50,55 @@ _TIME_SUPPORT_DELAY_MARGIN = 0.05 +def _apply_response_order_control(products, control, selected_p, selected_q): + """Estimate/check orders on a reference bank, then return the requested subset. + + Kept here so both JAX builders use exactly the production U,V products. The + import and the higher-order U,V scan happen only when ``control`` is + non-None, which is true only for an explicit check/choose CLI option. + """ + if control is None: + return products, None + import warnings + from RIFT.likelihood import response_order + + report = response_order.estimate_response_orders( + products[4], products[1], products[2], + target_snr=control['target_snr'], + lnL_tolerance=control.get('lnL_tolerance', 0.1), + n_samples=control.get('n_samples', 128), + selected_p=selected_p, selected_q=selected_q, + vary_p=control.get('check_p', False) or control.get('choose_p', False), + vary_q=control.get('check_q', False) or control.get('choose_q', False)) + response_order.print_order_report(report) + if (control.get('check_p', False) or control.get('check_q', False)) \ + and not report['selected_passes']: + warnings.warn( + "chosen response order fails the requested accuracy: predicted " + "Delta lnL={:.3g} > {:.3g} at SNR {:.6g}".format( + report['selected_delta_lnL'], report['lnL_tolerance'], + report['target_snr']), RuntimeWarning) + + final_p, final_q = int(selected_p), int(selected_q) + if control.get('choose_p', False) or control.get('choose_q', False): + if not report['reference_resolved']: + raise ValueError( + "cannot auto-select from an unresolved finite response " + "reference; raise the diagnostic reference order") + if report['chosen'] is None: + raise ValueError( + "no response order in the diagnostic reference bank satisfies " + "the requested SNR/error budget") + if control.get('choose_p', False): + final_p = int(report['chosen']['p_max']) + if control.get('choose_q', False): + final_q = int(report['chosen']['Qmax']) + report['final_p'] = final_p + report['final_Q'] = final_q + return response_order.truncate_precompute_products( + products, p_max=final_p, q_max=final_q), report + + def bandlimited_storage_requirement(deltaT, integration_window_half): """Return ``(storage_half, g0, g_certificate)`` for adaptive time support.""" tvals = factored_likelihood.marginalization_time_grid( @@ -87,6 +136,7 @@ def build_rotation_data_from_precompute(P, data_dict, psd_dict, fiducial_epoch, p_max=0, analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, tvals=None, verbose=False, + order_control=None, **precompute_kwargs): """One-call builder for the slow-rotation (Path A/B) banded JAX likelihood. @@ -109,12 +159,25 @@ def build_rotation_data_from_precompute(P, data_dict, psd_dict, fiducial_epoch, import RIFT.likelihood.factored_likelihood_with_rotation as flwr from .banded import build_rotation_data + p_reference = (max(int(p_max), int(order_control.get('p_reference', p_max))) + if order_control is not None and + (order_control.get('check_p') or order_control.get('choose_p')) + else int(p_max)) + if order_control is not None: + from RIFT.likelihood import response_order + response_order.guard_reference_bank( + 'rotation', p_reference, 0, Lmax, len(data_dict), + order_control.get('max_bank_gib', 4.0)) ri, ct, ctV, rho, meta = flwr.PrecomputeLikelihoodTermsWithRotation( fiducial_epoch, t_window, P, data_dict, psd_dict, Lmax, fMax, - harmonics=harmonics, p_max=p_max, f_sidereal=flwr.F_SIDEREAL, + harmonics=harmonics, p_max=p_reference, f_sidereal=flwr.F_SIDEREAL, analyticPSD_Q=analyticPSD_Q, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, verbose=verbose, quiet=not verbose, skip_interpolation=True, **precompute_kwargs) + products, order_report = _apply_response_order_control( + (ri, ct, ctV, rho, meta), order_control, + selected_p=int(p_max), selected_q=0) + ri, ct, ctV, rho, meta = products lk, rbn, ubn, vbn, ep = flwr.pack_rotation_arrays(meta, rho, ct, ctV) deltaT = float(P.deltaT) @@ -130,7 +193,7 @@ def build_rotation_data_from_precompute(P, data_dict, psd_dict, fiducial_epoch, integration_window_half, deltaT, xpy=np) data = build_rotation_data(meta, lk, rbn, ubn, vbn, ep, deltaT, tvals) extras = dict(meta=meta, rho_by_a=rbn, U_by_aa=ubn, V_by_aa=vbn, - epochDict=ep, lookupNKDict=lk) + epochDict=ep, lookupNKDict=lk, order_report=order_report) return data, extras @@ -140,6 +203,7 @@ def build_freqresponse_data_from_precompute(P, data_dict, psd_dict, fiducial_epo analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, tvals=None, verbose=False, + order_control=None, **precompute_kwargs): """One-call builder for the finite-size (Path D) banded JAX likelihood. @@ -154,11 +218,23 @@ def build_freqresponse_data_from_precompute(P, data_dict, psd_dict, fiducial_epo import RIFT.likelihood.slowrot_freqresponse as sfr from .banded import build_freqresponse_data + q_reference = (max(int(Qmax), int(order_control.get('q_reference', Qmax))) + if order_control is not None and + (order_control.get('check_q') or order_control.get('choose_q')) + else int(Qmax)) + if order_control is not None: + from RIFT.likelihood import response_order + response_order.guard_reference_bank( + 'finite', 0, q_reference, Lmax, len(data_dict), + order_control.get('max_bank_gib', 4.0)) bk = flfr.PrecomputeLikelihoodTermsFreqResponse( fiducial_epoch, t_window, P, data_dict, psd_dict, Lmax, fMax, - Qmax=Qmax, L_arm=L_arm, analyticPSD_Q=analyticPSD_Q, + Qmax=q_reference, L_arm=L_arm, analyticPSD_Q=analyticPSD_Q, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, verbose=verbose, quiet=not verbose, skip_interpolation=True, **precompute_kwargs) + bk, order_report = _apply_response_order_control( + bk, order_control, + selected_p=0, selected_q=int(Qmax)) meta = bk[4] lk, rbp, ubp, vbp, ep = flfr.pack_freqresponse_arrays(bk[4], bk[3], bk[1], bk[2]) @@ -181,7 +257,8 @@ def _L_of(det): data = build_freqresponse_data(meta, lk, rbp, ubp, vbp, ep, deltaT, tvals, det_geom) extras = dict(meta=meta, rho_by_p=rbp, U_by_pp=ubp, V_by_pp=vbp, - epochDict=ep, lookupNKDict=lk, det_geom=det_geom) + epochDict=ep, lookupNKDict=lk, det_geom=det_geom, + order_report=order_report) return data, extras @@ -189,18 +266,72 @@ def build_rotating_freqresponse_data_from_precompute( P, data_dict, psd_dict, fiducial_epoch, integration_window_half, Lmax, fMax, t_window=0.1, Qmax=4, L_arm=None, p_max=0, analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, - tvals=None, verbose=False, **precompute_kwargs): - """One-call builder for the compound rotation + finite-response likelihood.""" + tvals=None, verbose=False, order_control=None, **precompute_kwargs): + """Build the compound likelihood, optionally without a Q/U/V host round trip. + + ``RIFT_GPU_PRECOMPUTE=1`` selects CuPy precompute followed by DLPack + handoff to JAX. Existing host waveform generators remain supported. + """ import RIFT.likelihood.factored_likelihood_rotating_freqresponse as flrr import RIFT.likelihood.slowrot_freqresponse as sfr from .banded import build_rotating_freqresponse_data + if os.environ.get('RIFT_GPU_PRECOMPUTE', '0') == '1': + if order_control is not None: + raise NotImplementedError('Device-resident response-order selection is not yet supported; choose explicit orders or disable RIFT_GPU_PRECOMPUTE') + if os.environ.get('RIFT_GPU_WAVEFORM', 'lal') != 'lal': + raise ValueError('Native GPU waveform provider is not yet validated; use RIFT_GPU_WAVEFORM=lal') + from ..gpu_precompute import PrecomputeLikelihoodTermsRotatingFreqResponseGPU + from ..gpu_jax_handoff import build_jax_rotating_freqresponse_data_from_device + packed, meta = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + fiducial_epoch, t_window, P, data_dict, psd_dict, Lmax, fMax, + Qmax=Qmax, L_arm=L_arm, p_max=p_max, + analyticPSD_Q=analyticPSD_Q, inv_spec_trunc_Q=inv_spec_trunc_Q, + T_spec=T_spec, verbose=verbose, quiet=not verbose, + skip_interpolation=True, return_device=True, **precompute_kwargs) + det_geom = { + det: sfr.detector_geometry(det, L_arm=( + L_arm.get(det, None) if isinstance(L_arm, dict) else L_arm)) + for det in data_dict} + if tvals is None: + tvals = factored_likelihood.marginalization_time_grid( + integration_window_half, float(P.deltaT), xpy=np) + data = build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, det_geom) + # Preserve the diagnostic keys without materializing LAL/NumPy Q banks. + extras = dict(meta=meta, + rho_by_a={det: {a: packed['q'][det][i] + for i, a in enumerate(packed['a_list'])} + for det in data_dict}, + U_by_aa=packed['U'], V_by_aa=packed['V'], + epochDict=packed['epoch'], + lookupNKDict={det: np.asarray(packed['modes']) + for det in data_dict}, det_geom=det_geom, + order_report=None) + return data, extras + + p_reference = (max(int(p_max), int(order_control.get('p_reference', p_max))) + if order_control is not None and + (order_control.get('check_p') or order_control.get('choose_p')) + else int(p_max)) + q_reference = (max(int(Qmax), int(order_control.get('q_reference', Qmax))) + if order_control is not None and + (order_control.get('check_q') or order_control.get('choose_q')) + else int(Qmax)) + if order_control is not None: + from RIFT.likelihood import response_order + response_order.guard_reference_bank( + 'combined', p_reference, q_reference, Lmax, len(data_dict), + order_control.get('max_bank_gib', 4.0)) bk = flrr.PrecomputeLikelihoodTermsRotatingFreqResponse( fiducial_epoch, t_window, P, data_dict, psd_dict, Lmax, fMax, - Qmax=Qmax, L_arm=L_arm, p_max=p_max, + Qmax=q_reference, L_arm=L_arm, p_max=p_reference, analyticPSD_Q=analyticPSD_Q, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, verbose=verbose, quiet=not verbose, skip_interpolation=True, **precompute_kwargs) + bk, order_report = _apply_response_order_control( + bk, order_control, + selected_p=int(p_max), selected_q=int(Qmax)) meta = bk[4] lk, rba, uba, vba, ep = flrr.pack_rotating_freqresponse_arrays( meta, bk[3], bk[1], bk[2]) @@ -216,7 +347,8 @@ def _L_of(det): data = build_rotating_freqresponse_data( meta, lk, rba, uba, vba, ep, deltaT, tvals, det_geom) extras = dict(meta=meta, rho_by_a=rba, U_by_aa=uba, V_by_aa=vba, - epochDict=ep, lookupNKDict=lk, det_geom=det_geom) + epochDict=ep, lookupNKDict=lk, det_geom=det_geom, + order_report=order_report) return data, extras @@ -715,6 +847,11 @@ def sample_phi_ref(self, ra, dec, psi, incl, distMpc, rng=None, return out[:, 0] if n_samples == 1 else out +# Finite log-zero preserves proposal counts in legacy evidence helpers, which +# discard nonfinite weights. Shared with the driver for output-cloud filtering. +BOUNDED_MULTIPEAK_LOG_ZERO = -1.e30 + + class JAXDistPhiPsiMargLikelihood: """Distance-, phi_ref- AND psi-marginalised lnL over 3 angles (ra, dec, incl). @@ -722,6 +859,17 @@ class JAXDistPhiPsiMargLikelihood: leaving a smooth 3-D target. Removing psi (spin-2, the dimension most entangled with distance/inclination) lowers the sampler dimension and stabilises the distance integral relative to the 4-D phi-marginalised likelihood. + + For explicit ``multipeak-jax``, ``bounded_multipeak_config`` fixes the + resource envelope. Declines have numerical zero weight by default and + are counted in ``bounded_multipeak_audit``; ``bounded_multipeak_decline_action="refuse"`` + instead rejects any declined host batch. There is no reserve. The + accepted-region target can have discontinuities at acceptance boundaries. + Audit counts and the refusal latch cover only ``log_likelihood`` batches + (pilot/reweight/output calls). Scalar calls, including MAP, Fisher and + internal MALA training, are excluded: under ``refuse`` a scalar decline + returns NaN without raising or latching. Zero audited declines therefore + does not certify that scalar evaluations accepted. """ ANGULAR_PARAM_ORDER = ("ra", "dec", "incl") @@ -732,7 +880,8 @@ def __init__(self, data, d_min, d_max, nphi=32, npsi=16, n_grid=256, time_quadrature=TIME_QUAD_DEFAULT, d_prior_range=None, dist_grid="uniform", dist_grid_tol=DIST_GRID_TOL_DEFAULT, direct_marginalization_policy=None, policy_config=None, - multipeak_guard=16): + multipeak_guard=None, bounded_multipeak_config=None, + bounded_multipeak_decline_action="drop"): self.data = data self.interp = interp # the instance's stencil; sample_phi_ref defaults to it from . import direct_marginalization_policy as _policy @@ -763,7 +912,8 @@ def __init__(self, data, d_min, d_max, nphi=32, npsi=16, n_grid=256, # actually ran -- callers must surface it in the run log. if angle_marg not in ANGLE_MARG_CHOICES: raise ValueError("angle_marg must be one of grid/exact/laplace/" - "peak-local/auto, got %r" % (angle_marg,)) + "peak-local/phi-local/multipeak/multipeak-jax/" + "auto, got %r" % (angle_marg,)) if dist_grid not in DIST_GRID_SCHEMES: # An unrecognised value must NEVER fall through to the default: a # typo that silently returns the old answer is precisely the @@ -1175,9 +1325,99 @@ def _fused(data_, ra, dec, incl, return_lnLt=False, "another scheme if you need the time series.") v = _anglemarg.fused_log_likelihood_distphipsimarg_multipeak( data_, ra, dec, incl, xg, lwg, interp=interp, - amp_sizing=amp_sizing, guard=int(multipeak_guard)) + amp_sizing=amp_sizing, + guard=16 if multipeak_guard is None else int(multipeak_guard)) return (v, jnp.asarray(amp_sizing)) if return_amp else v + elif scheme == "multipeak-jax": + # Fixed-cap device discovery plus fixed-order local quadrature. + # There is intentionally no amplitude-sized reserve: a row that + # exceeds the envelope or fails a warrant returns nan. Planning + # is stop_gradient control data; AD differentiates the accepted + # fixed-plan integral and carries no derivative-accuracy claim. + # The endpoint time-quadrature validation already requires Simpson. + if dist_grid != "uniform": + raise ValueError( + "--angle-marg-scheme multipeak-jax derives its local " + "distance normalization from a uniform-in-distance grid") + if bounded_multipeak_decline_action not in ("drop", "refuse"): + raise ValueError("bounded_multipeak_decline_action must be drop or refuse") + self.bounded_multipeak_decline_action = bounded_multipeak_decline_action + self.bounded_multipeak_audit = dict( + evaluated=0, declined=0, diagnostic_unknown=0, + max_accepted=None, max_declined_diagnostic=None, reasons={}) + cfg = bounded_multipeak_config + if cfg is None: + cfg = _policy.BoundedMultipeakConfig() + if multipeak_guard is not None: + cfg = cfg._replace(time_guard=multipeak_guard) + elif multipeak_guard is not None: + _policy.validate_bounded_multipeak_config(cfg) + if multipeak_guard != cfg.time_guard: + raise ValueError( + "multipeak_guard conflicts with bounded_multipeak_config") + _policy.validate_bounded_multipeak_config(cfg) + lln_bounded, _ = _policy.policy_log_normalization( + data, xg, lwg, d_prior=d_prior) + _policy.probe_guarded_tables(data, interp, int(cfg.time_guard)) + x_bounds_bounded = (float(np.min(np.asarray(xg))), + float(np.max(np.asarray(xg)))) + self.angle_marg_info.update( + config=dict(cfg._asdict()), + decline_action=bounded_multipeak_decline_action, + bounded_cost=True, + dense_reserve=False, + fixed_plan_autodiff_only=True, + derivative_warrant_certified=False, + max_starts=int(cfg.base_max_starts), + max_time_nodes=int(cfg.max_time_nodes), + max_modes=int(cfg.enriched_max_modes), + time_guard=int(cfg.time_guard), + base_oversample=int(cfg.base_oversample), + enriched_oversample=int(cfg.enriched_oversample), + refine_iterations=int(cfg.refine_iterations), + quadrature_orders=(int(cfg.base_order), + int(cfg.base_check_order), + int(cfg.enriched_order), + int(cfg.enriched_check_order)), + convergence_tol_nats=float(cfg.convergence_tol_nats), + total_value_error_budget_nats=float( + cfg.total_value_error_budget_nats), + batch_rows=int(cfg.batch_rows)) + + def _fused(data_, ra, dec, incl, return_lnLt=False, + return_amp=False): + if return_lnLt: + raise ValueError( + "--angle-marg-scheme multipeak-jax marginalizes time " + "inside the controller; there is no lnL(t) to return") + if return_amp: + raise ValueError( + "--angle-marg-scheme multipeak-jax has a static cost " + "envelope and does not expose an amplitude-sized grid") + values = _policy.fused_log_likelihood_four_axis_bounded( + data_, ra, dec, incl, xg, lwg, interp=interp, + amp_sizing=amp_sizing, config=cfg, + local_log_normalization=lln_bounded, + x_bounds=x_bounds_bounded) + + if bounded_multipeak_decline_action == "drop": + # Finite log-zero keeps declined proposals in the IS sample + # count (the evidence helper filters infinities). It also + # avoids NaNs in MCMC acceptance ratios. This defines the + # accepted-region target, whose omitted mass is NOT known. + values = jnp.where( + jnp.isfinite(values), values, BOUNDED_MULTIPEAK_LOG_ZERO) + return values + + def _bounded_ledger(ra, dec, incl): + return _policy.fused_log_likelihood_four_axis_bounded( + data, ra, dec, incl, xg, lwg, interp=interp, + amp_sizing=amp_sizing, config=cfg, + local_log_normalization=lln_bounded, + x_bounds=x_bounds_bounded, return_ledger=True) + self._bounded_ledger = jax.jit(_bounded_ledger) + elif scheme == "phi-local": # BOTH angle axes localized, with a dense fallback wherever the certificate # declines. By name only, and deliberately not in 'auto': it is slower than @@ -1305,7 +1545,17 @@ def _batched_ledger(ra, dec, incl): # nothing is compiled or cached twice. def _batched(ra, dec, incl): return _fused(data, ra, dec, incl) - self._batched = jax.jit(_batched) + # MULTIPEAK IS HOST-SIDE AND MUST NOT BE TRACED. multipeak_local_marginalize + # is a numpy/scipy planner with a Python loop over rows; it calls np.asarray on + # the coefficient tables, which under jit are tracers + # (TracerArrayConversionError). Every other scheme here is a jax kernel and is + # jitted. This was missed because the wiring tests called _fused directly, in + # eager mode, and the failure only appears through _batched -- the seam the + # sampler actually uses. + if scheme == "multipeak": + self._batched = _batched + else: + self._batched = jax.jit(_batched) self._batched_amp = None if self._amp_record is not None: @@ -1317,11 +1567,55 @@ def _scalar(theta3): v = _fused(data, theta3[0:1], theta3[1:2], theta3[2:3]) return v[0] self._scalar = _scalar - self._value_and_grad = jax.jit(jax.value_and_grad(_scalar)) - self._hessian = jax.jit(jax.hessian(_scalar)) + if scheme == "multipeak": + # No AD through a numpy planner. Refusing is the honest contract; a + # silently zero or wrong gradient would reach --fisher-precondition, + # which swallows exceptions and falls back to raw coordinates with the + # flag still recorded as supplied. + def _no_grad(*a, **k): + raise ValueError( + "--angle-marg-scheme multipeak is a host-side planner and is " + "not differentiable; gradients and the Fisher preconditioner " + "are unavailable for it. Use another scheme if you need them.") + self._value_and_grad = _no_grad + self._hessian = _no_grad + else: + self._value_and_grad = jax.jit(jax.value_and_grad(_scalar)) + self._hessian = jax.jit(jax.hessian(_scalar)) def log_likelihood(self, ra, dec, incl): """lnL for arrays of 3 angular parameters (ra, dec, incl), shape (S,).""" + if self.angle_marg_scheme == "multipeak-jax": + values, ledger = self._bounded_ledger( + jnp.asarray(ra), jnp.asarray(dec), jnp.asarray(incl)) + host, ledger = jax.device_get((values, ledger)) + bad = ~np.isfinite(host) + audit = self.bounded_multipeak_audit + audit["evaluated"] += int(host.size) + audit["declined"] += int(np.sum(bad)) + for key in ledger: + if key.startswith("decline_"): + n = int(np.sum(np.asarray(ledger[key]))) + if n: + audit["reasons"][key] = audit["reasons"].get(key, 0) + n + selected = np.asarray(ledger["selected_value"]) + audit["diagnostic_unknown"] += int(np.sum(bad & ~np.isfinite(selected))) + for key, samples in (("max_accepted", host[~bad]), + ("max_declined_diagnostic", selected[bad])): + finite = samples[np.isfinite(samples)] + if finite.size: + old = audit[key] + audit[key] = max(float(np.max(finite)), + old if old is not None else -np.inf) + if np.any(bad): + self.bounded_multipeak_declined = True + if self.bounded_multipeak_decline_action == "refuse": + raise RuntimeError( + "multipeak-jax could not warrant every row within the " + "static envelope; refusing samples and evidence. " + "Change the envelope or use decline-action drop.") + return jnp.where( + jnp.isfinite(values), values, BOUNDED_MULTIPEAK_LOG_ZERO) if self._batched_amp is None: return self._batched(jnp.asarray(ra), jnp.asarray(dec), jnp.asarray(incl)) diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/response_order.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/response_order.py new file mode 100644 index 000000000..a6f73d8f8 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/response_order.py @@ -0,0 +1,415 @@ +"""Physics-based response-series order estimates from the precomputed U,V bank. + +This module is deliberately not imported by either ILE driver unless one of the +response-order check/choose command-line options is active. The diagnostic is +data independent: it contracts only model self-overlaps, never Q=. + +The estimate uses a deterministic, full-prior angular design. It is therefore +an estimate, not a mathematical supremum over the prior. Increase +``n_samples`` and the reference order for a production convergence check. +""" +from __future__ import print_function, division + +import math +import warnings + +import numpy as np + + +def _angular_design(n_samples): + """Deterministic five-dimensional Halton design (no RNG state).""" + n = max(8, int(n_samples)) + def radical_inverse(base): + out = np.zeros(n, dtype=float) + denominator = 1.0 + integer = np.arange(1, n + 1, dtype=np.int64) + while np.any(integer): + integer, digit = np.divmod(integer, base) + denominator *= base + out += digit / denominator + return out + u_ra, u_dec, u_inc, u_psi, u_phase = [ + radical_inverse(base) for base in (2, 3, 5, 7, 11)] + ra = 2.0 * np.pi * u_ra + dec = np.arcsin(1.0 - 2.0 * u_dec) + incl = np.arccos(1.0 - 2.0 * u_inc) + psi = np.pi * u_psi + phiref = 2.0 * np.pi * u_phase + return ra, dec, incl, psi, phiref + + +def reference_bank_size(feature, p_max, q_max, lmax, n_detectors): + """Dense-U,V planning estimate using every mode through lmax.""" + p_max, q_max = int(p_max), int(q_max) + if feature == 'rotation': + basis = (p_max + 1) * (2 * (2 + p_max) + 1) + elif feature == 'finite': + basis = q_max + 2 + elif feature == 'combined': + basis = sum( + 2 * ((2 if b == 0 else b + 1) + p) + 1 + for b in range(q_max + 2) for p in range(p_max + 1)) + else: + raise ValueError("unknown response feature %r" % feature) + modes = sum(2 * ell + 1 for ell in range(2, int(lmax) + 1)) + uv_bytes = 2 * basis ** 2 * modes ** 2 * 16 * int(n_detectors) + return dict(basis=basis, modes=modes, uv_gib=uv_bytes / 2.0 ** 30) + + +def guard_reference_bank(feature, p_max, q_max, lmax, n_detectors, + max_bank_gib=4.0): + """Refuse an oversized diagnostic before building its response bank.""" + if not np.isfinite(max_bank_gib) or float(max_bank_gib) <= 0: + raise ValueError("response-order max-bank-gib must be finite and positive") + size = reference_bank_size(feature, p_max, q_max, lmax, n_detectors) + if size['uv_gib'] > float(max_bank_gib): + raise ValueError( + "response-order dense-U,V planning estimate is {:.3g} GiB " + "(basis={}, all modes through lmax={}, detectors={}); this estimate " + "excludes dictionary overhead and Q banks. Lower the diagnostic reference " + "orders or raise --response-order-max-bank-gib explicitly".format( + size['uv_gib'], size['basis'], size['modes'], int(n_detectors))) + return size + + +def _feature(meta): + if meta.get('feature') == 'rotation_freqresponse': + return 'combined' + if 'p_list' in meta: + return 'finite' + if 'a_list' in meta: + return 'rotation' + raise ValueError("unrecognized response bank metadata") + + +def _indices(meta): + return list(meta['p_list'] if _feature(meta) == 'finite' else meta['a_list']) + + +def _matrix(bank, det, a, ap, ia, iap, modes): + item = bank[det] + block = item[(a, ap)] if isinstance(item, dict) else item[ia, iap] + if isinstance(block, dict): + return np.asarray([[block[(m1, m2)] for m2 in modes] for m1 in modes], + dtype=complex) + return np.asarray(block) + + +def _arm_for_detector(meta, det): + arm = meta.get('L_arm', None) + return arm.get(det, None) if isinstance(arm, dict) else arm + + +def _coefficients(meta, det, ra, dec, psi): + """Return physical coefficients C and reflected coefficients C_R.""" + feature = _feature(meta) + tref = float(meta.get('event_time_geo', meta.get('tref', 0.0))) + if feature == 'finite': + from . import factored_likelihood_freqresponse as fr + rows = [fr.response_coefficients( + det, float(r), float(d), float(p), tref, int(meta['Qmax']), + L_arm=_arm_for_detector(meta, det)) + for r, d, p in zip(ra, dec, psi)] + idx = _indices(meta) + c = np.column_stack([[row.get(a, 0j) for row in rows] for a in idx]) + return c, c + + if feature == 'combined': + from . import factored_likelihood_rotating_freqresponse as rf + from . import factored_likelihood_with_rotation as rot + coeff = rf.combined_response_coefficients_vector( + det, ra, dec, psi, tref, int(meta['p_max']), + Qmax=int(meta['Qmax']), L_arm=_arm_for_detector(meta, det)) + reflect = lambda a: (a[0], a[1], -a[2]) + else: + from . import factored_likelihood_with_rotation as rot + coeff = rot.rotation_coefficients_vector( + det, ra, dec, psi, tref, int(meta['p_max'])) + reflect = lambda a: (a[0], -a[1]) + + # Evaluate the norm at geocentric coalescence time. The U,V model norm + # owes the same arrival-time post-phase as the likelihood evaluator. + import lal + from . import factored_likelihood as fl + location = fl.lalsim.DetectorPrefixToLALDetector(det).location + delta = np.asarray(fl.TimeDelayFromEarthCenter( + np.asarray(location), ra, dec, + float(lal.GreenwichMeanSiderealTime(lal.LIGOTimeGPS(tref))), xpy=np)) + omega = 2.0 * np.pi * float(meta['f_sidereal']) + idx = _indices(meta) + zero = np.zeros(len(ra), dtype=complex) + c = np.column_stack([ + coeff.get(a, zero) * np.exp(1j * a[-1] * omega * delta) for a in idx]) + cr = np.column_stack([ + coeff.get(reflect(a), zero) * + np.exp(1j * reflect(a)[-1] * omega * delta) for a in idx]) + return c, cr + + +def _retained(meta, p_max, q_max): + feature = _feature(meta) + if feature == 'finite': + return np.asarray([a <= int(q_max) + 1 for a in _indices(meta)]) + if feature == 'rotation': + # A production p bank uses one common harmonic width 2+p_max. + # Remove the wider reference bank's identically-zero outer bands too, + # otherwise auto-selection would be physically right but not minimal. + return np.asarray([a[0] <= int(p_max) and abs(a[1]) <= 2 + int(p_max) + for a in _indices(meta)]) + return np.asarray([a[0] <= int(q_max) + 1 and a[1] <= int(p_max) + for a in _indices(meta)]) + + +def basis_size(meta, p_max=None, q_max=None): + """Number of response basis elements retained by an order pair.""" + feature = _feature(meta) + if p_max is None: + p_max = int(meta.get('p_max', 0)) + if q_max is None: + q_max = int(meta.get('Qmax', 0)) + return int(np.sum(_retained(meta, p_max, q_max))) + + +def pack_uv_from_raw(meta, cross, cross_v): + """Pack only model moments, deliberately skipping the large Q= bank.""" + indices = _indices(meta) + modes = list(meta['modes']) + + def one(source): + out = {} + for det, pairs in source.items(): + out[det] = {} + for a in indices: + for ap in indices: + block = pairs[(a, ap)] + if isinstance(block, dict): + out[det][(a, ap)] = np.asarray( + [[block[(m1, m2)] for m2 in modes] for m1 in modes], + dtype=complex) + else: + out[det][(a, ap)] = np.asarray(block) + return out + return one(cross), one(cross_v) + + +def estimate_response_orders(meta, U, V, target_snr, lnL_tolerance=0.1, + n_samples=128, selected_p=None, selected_q=None, + vary_p=True, vary_q=True): + """Scan truncations of one reference bank and return an order report. + + A candidate passes when max_theta ||h_ref-h_candidate||^2/||h_ref||^2 + is at most ``2*lnL_tolerance/target_snr**2``. ``U`` and ``V`` are the + raw or packed response-bank model-overlap dictionaries (or dense compound arrays). + """ + if not np.isfinite(target_snr) or target_snr <= 0: + raise ValueError("target_snr must be positive") + if not np.isfinite(lnL_tolerance) or lnL_tolerance <= 0: + raise ValueError("lnL_tolerance must be positive") + + from . import factored_likelihood as fl + feature = _feature(meta) + idx = _indices(meta) + modes = list(meta['modes']) + lookup = np.asarray(modes, dtype=int) + ra, dec, incl, psi, phiref = _angular_design(n_samples) + Y = np.asarray(fl.ComputeYlmsArrayVector(lookup, incl, -phiref)).T + ns = len(ra) + + detector_terms = [] + for det in U: + c, cr = _coefficients(meta, det, ra, dec, psi) + # Precontract the small mode matrices once. The response-order scan is + # then only a quadratic contraction over basis indices. + um = np.empty((len(idx), len(idx), ns), dtype=complex) + vm = np.empty_like(um) + for ia, a in enumerate(idx): + for iap, ap in enumerate(idx): + um[ia, iap] = np.einsum( + 'si,ij,sj->s', np.conj(Y), + _matrix(U, det, a, ap, ia, iap, modes), Y) + vm[ia, iap] = np.einsum( + 'si,ij,sj->s', Y, + _matrix(V, det, a, ap, ia, iap, modes), Y) + detector_terms.append((c, cr, um, vm)) + + def norm_for_mask(mask): + out = np.zeros(ns, dtype=float) + for c, cr, um, vm in detector_terms: + cm = c * mask[None, :] + crm = cr * mask[None, :] + val = np.einsum('sa,sb,abs->s', np.conj(cm), cm, um) + val += np.einsum('sa,sb,abs->s', crm, cm, vm) + out += 0.5 * np.real(val) + # Roundoff can produce tiny negative self norms. + return np.maximum(out, 0.0) + + full = norm_for_mask(np.ones(len(idx), dtype=bool)) + good = full > max(np.max(full), 1.0) * 1e-13 + if not np.any(good): + raise ValueError("reference response has zero norm on the angular design") + threshold = 2.0 * float(lnL_tolerance) / float(target_snr) ** 2 + p_ref = int(meta.get('p_max', 0)) + q_ref = int(meta.get('Qmax', 0)) + if selected_p is None: + selected_p = p_ref + if selected_q is None: + selected_q = q_ref + p_values = range(p_ref + 1) if vary_p and feature != 'finite' else [int(selected_p)] + q_values = range(q_ref + 1) if vary_q and feature != 'rotation' else [int(selected_q)] + + rows = [] + for p in p_values: + for q in q_values: + keep = _retained(meta, p, q) + residual = norm_for_mask(~keep) + mu = float(np.max(residual[good] / full[good])) + nb = int(np.sum(keep)) + rows.append(dict(p_max=int(p), Qmax=int(q), max_mu=mu, + max_delta_lnL=0.5 * target_snr ** 2 * mu, + basis=nb, uv_pairs=nb * nb, + passes=bool(mu <= threshold))) + selected_keep = _retained(meta, selected_p, selected_q) + selected_mu = float(np.max(norm_for_mask(~selected_keep)[good] / full[good])) + + # A finite-reference resolution diagnostic. It is deliberately not called + # an analytic tail bound: require two decreasing shells, extrapolate their + # norm ratio geometrically, and reserve one quarter of the error budget. + shell_rows = [] + def add_shell(axis, final_mask, previous_mask): + final_mu = float(np.max(norm_for_mask(final_mask)[good] / full[good])) + previous_mu = float(np.max(norm_for_mask(previous_mask)[good] / full[good])) + if previous_mu <= 0.0: + ratio = 0.0 if final_mu <= 0.0 else float('inf') + else: + ratio = math.sqrt(final_mu / previous_mu) + tail_mu = ((math.sqrt(final_mu) * ratio / (1.0 - ratio)) ** 2 + if np.isfinite(ratio) and ratio < 1.0 else float('inf')) + shell_rows.append(dict(axis=axis, final_mu=final_mu, + previous_mu=previous_mu, norm_ratio=ratio, + extrapolated_tail_mu=tail_mu)) + if vary_p and feature != 'finite' and p_ref >= 2: + add_shell('p', + _retained(meta, p_ref, q_ref) & ~_retained(meta, p_ref - 1, q_ref), + _retained(meta, p_ref - 1, q_ref) & ~_retained(meta, p_ref - 2, q_ref)) + elif vary_p and feature != 'finite': + shell_rows.append(dict(axis='p', final_mu=float('inf'), + previous_mu=float('inf'), norm_ratio=float('inf'), + extrapolated_tail_mu=float('inf'))) + if vary_q and feature != 'rotation' and q_ref >= 2: + add_shell('q', + _retained(meta, p_ref, q_ref) & ~_retained(meta, p_ref, q_ref - 1), + _retained(meta, p_ref, q_ref - 1) & ~_retained(meta, p_ref, q_ref - 2)) + elif vary_q and feature != 'rotation': + shell_rows.append(dict(axis='q', final_mu=float('inf'), + previous_mu=float('inf'), norm_ratio=float('inf'), + extrapolated_tail_mu=float('inf'))) + edge_mu = max((row['final_mu'] for row in shell_rows), default=0.0) + tail_mu = sum(math.sqrt(row['extrapolated_tail_mu']) for row in shell_rows) ** 2 + reference_resolved = bool(shell_rows) and all( + row['norm_ratio'] <= 0.5 and row['final_mu'] <= 0.25 * threshold + for row in shell_rows) and tail_mu <= 0.25 * threshold + + # Charge the unresolved part of the response series to every candidate by + # the triangle inequality. The finite-reference residual and extrapolated + # tail can add coherently, so their powers must not be added directly. + for row in rows: + row['finite_reference_mu'] = row['max_mu'] + row['max_mu'] = ( + math.sqrt(row['finite_reference_mu']) + math.sqrt(tail_mu)) ** 2 + row['max_delta_lnL'] = 0.5 * target_snr ** 2 * row['max_mu'] + row['passes'] = bool(row['max_mu'] <= threshold) + selected_finite_mu = selected_mu + selected_mu = (math.sqrt(selected_finite_mu) + math.sqrt(tail_mu)) ** 2 + passing = [row for row in rows if row['passes']] + chosen = min(passing, key=lambda x: (x['uv_pairs'], x['basis'], + x['p_max'] + x['Qmax'])) if passing else None + + mode_block_sensitivity = [] + blocks = {} + for im, mode in enumerate(modes): + blocks.setdefault((int(mode[0]), abs(int(mode[1]))), []).append(im) + for block, members in sorted(blocks.items()): + # Pair +/-m modes before taking the norm; this is diagnostic only and + # never prunes the production bank. + out = np.zeros(ns, dtype=float) + Ym = np.zeros_like(Y) + Ym[:, members] = Y[:, members] + for det in U: + c, cr = _coefficients(meta, det, ra, dec, psi) + for ia, a in enumerate(idx): + for iap, ap in enumerate(idx): + uu = np.einsum('si,ij,sj->s', np.conj(Ym), + _matrix(U, det, a, ap, ia, iap, modes), Ym) + vv = np.einsum('si,ij,sj->s', Ym, + _matrix(V, det, a, ap, ia, iap, modes), Ym) + out += 0.5 * np.real(np.conj(c[:, ia]) * c[:, iap] * uu + + cr[:, ia] * c[:, iap] * vv) + mu_mode = float(np.max(np.maximum(out, 0.0)[good] / full[good])) + mode_block_sensitivity.append((block, mu_mode)) + + return dict(feature=feature, target_snr=float(target_snr), + lnL_tolerance=float(lnL_tolerance), threshold_mu=threshold, + reference_p=p_ref, reference_Q=q_ref, + selected_p=int(selected_p), selected_Q=int(selected_q), + selected_mu=selected_mu, + selected_finite_reference_mu=selected_finite_mu, + selected_delta_lnL=0.5 * target_snr ** 2 * selected_mu, + selected_passes=bool(selected_mu <= threshold), chosen=chosen, + rows=rows, reference_edge_mu=edge_mu, + reference_tail_mu=tail_mu, reference_shells=shell_rows, + reference_resolved=bool(reference_resolved), + n_samples=ns, n_valid_samples=int(np.sum(good)), + mode_block_sensitivity=sorted( + mode_block_sensitivity, key=lambda item: item[1], reverse=True)) + + +def print_order_report(report, prefix='response-order'): + """Stable one-line diagnostics suitable for scheduler stdout.""" + chosen = report['chosen'] + choice = 'NONE' if chosen is None else 'p_max={p_max} Qmax={Qmax}'.format(**chosen) + print('[{}] SNR={:.6g} eps_lnL={:.3g} samples={}/{} choice={} ' + 'selected_DeltaLnL={:.3g} reference_edge_DeltaLnL={:.3g}'.format( + prefix, report['target_snr'], report['lnL_tolerance'], + report['n_valid_samples'], report['n_samples'], choice, + report['selected_delta_lnL'], + 0.5 * report['target_snr'] ** 2 * report['reference_edge_mu'])) + if not report['reference_resolved']: + warnings.warn( + '{} finite reference is unresolved: two-shell decrease and ' + 'extrapolated-tail tests do not fit inside the reserved likelihood-' + 'error budget; raise the diagnostic reference order'.format(prefix), + RuntimeWarning) + + +def truncate_precompute_products(products, p_max=None, q_max=None): + """Return a response precompute tuple restricted to the selected rectangle.""" + rint, cross, cross_v, rho, meta0 = products + meta = dict(meta0) + feature = _feature(meta) + if p_max is None: + p_max = int(meta.get('p_max', 0)) + if q_max is None: + q_max = int(meta.get('Qmax', 0)) + keep = _retained(meta, p_max, q_max) + old = _indices(meta) + new = [a for a, yes in zip(old, keep) if yes] + if feature == 'finite': + meta['p_list'] = new + meta['Qmax'] = int(q_max) + else: + meta['a_list'] = new + meta['p_max'] = int(p_max) + if feature == 'combined': + meta['Qmax'] = int(q_max) + else: + meta['harmonics'] = tuple(sorted(set(a[1] for a in new))) + + def restrict_primary(bank): + return {det: {a: values[a] for a in new} for det, values in bank.items()} + + def restrict_cross(bank): + return {det: {(a, ap): values[(a, ap)] for a in new for ap in new} + for det, values in bank.items()} + + return (restrict_primary(rint), restrict_cross(cross), + restrict_cross(cross_v), restrict_primary(rho), meta) diff --git a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/test_slowrot_freqresponse_likelihood.py b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/test_slowrot_freqresponse_likelihood.py index 6806912f5..9a5412a19 100644 --- a/MonteCarloMarginalizeCode/Code/RIFT/likelihood/test_slowrot_freqresponse_likelihood.py +++ b/MonteCarloMarginalizeCode/Code/RIFT/likelihood/test_slowrot_freqresponse_likelihood.py @@ -297,9 +297,11 @@ def run_response_coefficients_block(): single place the substitution can go wrong, and the reference is the scalar routine itself -- still shipped, still the definition of b_p. - b_0..b_3 must be bit-identical. b_4 and b_5 carry a_x**3 and a_x**4, where numpy's - power loop and CPython's libm pow round differently in the last bit; that is bounded - here, and the end-to-end consequence was measured at zero (see + b_0..b_3 must agree to floating-point roundoff. Different supported NumPy/libm + combinations can also move these low-order terms by one ulp. b_4 and b_5 carry + a_x**3 and a_x**4, where numpy's power loop and CPython's libm pow can likewise + round differently in the last bit; all of that is bounded here, and the end-to-end + consequence was measured at zero (see DESIGN_freqresponse_vectorized_coefficients.md). """ rng = np.random.RandomState(20260909) @@ -317,12 +319,18 @@ def run_response_coefficients_block(): for p in range(Qmax + 2): d = abs(bv[p][i] - bs[p]) if p <= 3: - assert d == 0.0, ("b_%d must be bit-identical to the scalar routine " - "(%s sample %d: |d|=%g)" % (p, det, i, d)) + scale = max(1.0, abs(bs[p])) + assert d <= 8 * np.finfo(float).eps * scale, ( + "b_%d differs from the scalar routine by more than eight ulp " + "(%s sample %d: |d|=%g)" % (p, det, i, d)) worst = max(worst, d / max(abs(bs[p]), 1e-300)) - print("\n(V5) BLOCK COEFFICIENTS: b_0..b_3 bit-identical; worst relative " + print("\n(V5) BLOCK COEFFICIENTS: b_0..b_3 roundoff-equivalent; worst relative " "difference over all p = %.3e" % worst) - assert worst < 1e-14, "block coefficients drifted past one ulp: %g" % worst + # The largest relative ratio occurs when the reference coefficient is near + # zero; supported NumPy/libm combinations reach a few 1e-14 without any + # measurable likelihood change. Keep a two-order-of-magnitude guard below + # the end-to-end tolerances while allowing that benign platform variation. + assert worst < 1e-12, "block coefficients drifted past roundoff envelope: %g" % worst # STRUCTURAL: the detector lookup must not scale with the block. This is what the # change is for -- a regression to the per-sample loop passes every value check above diff --git a/MonteCarloMarginalizeCode/Code/bin/create_event_parameter_pipeline_BasicIteration b/MonteCarloMarginalizeCode/Code/bin/create_event_parameter_pipeline_BasicIteration index 949f3c77c..5e9ea1515 100755 --- a/MonteCarloMarginalizeCode/Code/bin/create_event_parameter_pipeline_BasicIteration +++ b/MonteCarloMarginalizeCode/Code/bin/create_event_parameter_pipeline_BasicIteration @@ -927,17 +927,29 @@ if (opts.last_iteration_extrinsic): # Convert task if opts.last_iteration_extrinsic_time_resampling: - # very simple: igwn_ligolw add on all final output, then a single convert + # JAX ILE exports fair draws as tabular sidecars, not XML rows. Route + # that driver to its strict joiner; retain the historical XML path for + # conventional ILE. + jax_tabular_fairdraws = ( + ile_exe is not None + and "integrate_likelihood_extrinsic_jax" in str(ile_exe) + ) + if jax_tabular_fairdraws and convert_args is not None: + raise ValueError( + "JAX tabular fair-draw conversion does not support --convert-args; " + "remove it or use conventional ILE") convert_args_extr = " --convention LI --export-cosmology --use-interpolated-cosmology " if not (convert_args is None): convert_args_extr += convert_args # batched convert shell script relevant_path = dag_utils.which('util_JoinExtrXML.py') relevant_path_2 = dag_utils.which('convert_output_format_ile2inference') + relevant_path_jax = dag_utils.which('util_ConvertJAXILEFairdraws.py') if opts.condor_local_nonworker_igwn_prefix: # we are in igwn environment - relevant_path = 'util_JoinExtrXML.py' + relevant_path = 'util_JoinExtrXML.py' relevant_path_2 = 'convert_output_format_ile2inference' + relevant_path_jax = 'util_ConvertJAXILEFairdraws.py' # randomization: 'shuf' is preferred, but otherwise use ' sort -R'. Note performed locally, so local filesystem/os is fine. extra_shuffle_command = ' | cat' which_shuf = which('shuf'); which_sort = which('sort') @@ -946,7 +958,17 @@ if (opts.last_iteration_extrinsic): elif isinstance(which_sort, str): extra_shuffle_command = ' | {} -R '.format(which_sort) with open("allinone_convert.sh",'w') as f: - f.write(f"""#! /bin/bash + if jax_tabular_fairdraws: + f.write(f"""#! /bin/bash +set -euo pipefail +{extra_text} +{relevant_path_jax} --directory ./iteration_$1_ile \ + --draws-per-intrinsic {n_points_per_ILE} \ + --expected-intrinsic {opts.last_iteration_extrinsic_nsamples} \ + --output ./extrinsic_posterior_samples.dat +""") + else: + f.write(f"""#! /bin/bash {extra_text} {relevant_path} ./iteration_$1_ile/'EXTR_out-*.xml_*_.xml.gz' --output ./tmp_converted.xml.gz {relevant_path_2} {convert_args_extr} ./tmp_converted.xml.gz > ./tmp_converted.dat diff --git a/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_batchmode b/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_batchmode index 8ec7a59b2..5921c1fbd 100755 --- a/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_batchmode +++ b/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_batchmode @@ -291,6 +291,16 @@ optp.add_option("--rotation-p-max", type=int, default=0, help="[Path B] Max dela optp.add_option("--freqresponse", action="store_true", help="[Path D] Finite-size (frequency-dependent) detector-response likelihood: account for the finite light-travel-time transfer across the arms (matters for 3G/CE-ET). Requires --vectorized; supports --gpu (n_cal=1, no glitch/cal marg). May be combined with --rotation-slow; the compound bank is substantially more expensive and intended for long, loud sources.") optp.add_option("--freqresponse-qmax", type=int, default=4, help="Highest power of the arm projection retained for --freqresponse (basis size Qmax+2). Higher for larger fL/c (heavier systems / higher fmax).") optp.add_option("--freqresponse-arm-length", default=None, help="Arm-length override [m] for --freqresponse. Either a single float applied to ALL detectors (e.g. 40000 for 40-km CE), or per-detector 'C1=40000,E1=10000,...' (needed for mixed CE+ET networks, since LAL's cached C1 arm is a placeholder). Default: each detector's native LAL arm length.") +optp.add_option("--check-slowrot-pmax", action="store_true", default=False, help="Estimate required pmax from the response U,V bank and warn if --rotation-p-max is too small.") +optp.add_option("--check-finite-size-Qmax", "--check-finite-size-qmax", dest="check_finite_size_Qmax", action="store_true", default=False, help="Estimate required finite-size Qmax and warn if --freqresponse-qmax is too small.") +optp.add_option("--choose-slowrot-pmax", action="store_true", default=False, help="Choose the least pmax satisfying the response error budget.") +optp.add_option("--choose-slowrot-Qmax", "--choose-finite-size-Qmax", dest="choose_slowrot_Qmax", action="store_true", default=False, help="Choose the least finite-size Qmax satisfying the response error budget.") +optp.add_option("--response-order-snr", type=float, default=None, help="Target network SNR for response-order checks/choice (required when active).") +optp.add_option("--response-order-lnL-tol", type=float, default=0.1, help="Allowed worst scanned Asimov likelihood loss (default 0.1).") +optp.add_option("--response-order-sky-samples", type=int, default=128, help="Deterministic full-prior angular design size (default 128).") +optp.add_option("--response-order-p-reference", type=int, default=2, help="Highest p used by the diagnostic reference bank (default 2).") +optp.add_option("--response-order-Q-reference", "--response-order-q-reference", dest="response_order_Q_reference", type=int, default=8, help="Highest Q used by the diagnostic reference bank (default 8).") +optp.add_option("--response-order-max-bank-gib", type=float, default=4.0, help="Refuse a diagnostic whose dense U,V planning estimate exceeds this many GiB (default 4).") optp.add_option("--force-xpy", action="store_true", help="Use the xpy code path. Use with --vectorized --gpu to use the fallback CPU-based code path. Useful for debugging.") optp.add_option("-o", "--output-file", help="Save result to this file.") optp.add_option("-O", "--output-format", default='xml', help="[xml|hdf5]") @@ -472,6 +482,31 @@ def _normalize_interpolate_time_argv(argv): opts, args = optp.parse_args(_normalize_interpolate_time_argv(sys.argv[1:])) +_response_order_active = any((opts.check_slowrot_pmax, opts.check_finite_size_Qmax, + opts.choose_slowrot_pmax, opts.choose_slowrot_Qmax)) +if (opts.check_slowrot_pmax or opts.choose_slowrot_pmax) and not opts.rotation_slow: + raise ValueError("pmax check/choice requires --rotation-slow") +if (opts.check_finite_size_Qmax or opts.choose_slowrot_Qmax) and not opts.freqresponse: + raise ValueError("Qmax check/choice requires --freqresponse") +if _response_order_active and opts.response_order_snr is None: + raise ValueError("response-order check/choice requires --response-order-snr") +if opts.check_slowrot_pmax and opts.choose_slowrot_pmax: + raise ValueError("choose either --check-slowrot-pmax or --choose-slowrot-pmax") +if opts.check_finite_size_Qmax and opts.choose_slowrot_Qmax: + raise ValueError("choose either --check-finite-size-Qmax or --choose-slowrot-Qmax") +if ((opts.check_slowrot_pmax or opts.choose_slowrot_pmax) + and (opts.check_finite_size_Qmax or opts.choose_slowrot_Qmax) + and (opts.check_slowrot_pmax or opts.check_finite_size_Qmax) + and (opts.choose_slowrot_pmax or opts.choose_slowrot_Qmax)): + raise ValueError("do not mix check and choose controls in one compound response scan") +if opts.response_order_lnL_tol <= 0 or opts.response_order_sky_samples < 8: + raise ValueError("response order requires positive lnL tolerance and at least 8 sky samples") +if opts.response_order_p_reference < 0 or opts.response_order_Q_reference < 0: + raise ValueError("response diagnostic reference orders must be nonnegative") +if (not np.isfinite(opts.response_order_max_bank_gib) + or opts.response_order_max_bank_gib <= 0): + raise ValueError("--response-order-max-bank-gib must be finite and positive") + def _truthy_option(value): if isinstance(value, bool): return value @@ -3517,6 +3552,31 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t extra_args_dict = eval(extra_args_dict) print(" Waveform high-level extra args passed ", extra_args_dict, type(extra_args_dict)) extra_waveform_kwargs.update(extra_args_dict) + # Keep every precompute for this event on the same waveform-generation + # route. In particular, the compound slow-rotation/frequency-response + # bank must not silently fall back to the default LAL generator when the + # baseline bank selected GWSignal, NR, ROM, or custom conditioning. + waveform_generation_kwargs = { + 'NR_group': NR_template_group, + 'NR_param': NR_template_param, + 'use_gwsignal': opts.use_gwsignal, + 'use_gwsignal_approx': opts.approximant, + 'use_external_EOB': opts.use_external_EOB, + 'nr_lookup': opts.nr_lookup, + 'nr_lookup_valid_groups': opts.nr_lookup_group, + 'perturbative_extraction': opts.nr_perturbative_extraction, + 'perturbative_extraction_full': opts.nr_perturbative_extraction_full, + 'use_provided_strain': opts.nr_use_provided_strain, + 'hybrid_use': opts.nr_hybrid_use, + 'hybrid_method': opts.nr_hybrid_method, + 'ROM_group': opts.rom_group, + 'ROM_param': opts.rom_param, + 'ROM_use_basis': opts.rom_use_basis, + 'ROM_limit_basis_size': opts.rom_limit_basis_size_to, + 'no_memory': opts.no_memory, + 'extra_waveform_kwargs': extra_waveform_kwargs, + 'force_22_mode': opts.force_hyperbolic_22, + } # Precompute t_window = opts.internal_data_storage_window_half ignore_threshold=None @@ -3556,10 +3616,9 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t False, inv_spec_trunc_Q, T_spec, return_calibration_crossterms=True, ignore_threshold=ignore_threshold, # default is None, old default was 1e-4. Use to speed calculation and/or discard 'junky' modes, esp at lower SNR. Dangerous at high SNR - NR_group=NR_template_group,NR_param=NR_template_param, - use_gwsignal=opts.use_gwsignal, - use_gwsignal_approx=opts.approximant, - use_external_EOB=opts.use_external_EOB,nr_lookup=opts.nr_lookup,nr_lookup_valid_groups=opts.nr_lookup_group,perturbative_extraction=opts.nr_perturbative_extraction,perturbative_extraction_full=opts.nr_perturbative_extraction_full,use_provided_strain=opts.nr_use_provided_strain,hybrid_use=opts.nr_hybrid_use,hybrid_method=opts.nr_hybrid_method,ROM_group=opts.rom_group,ROM_param=opts.rom_param,ROM_use_basis=opts.rom_use_basis,verbose=opts.verbose,quiet=not opts.verbose,ROM_limit_basis_size=opts.rom_limit_basis_size_to,no_memory=opts.no_memory,skip_interpolation=opts.vectorized, extra_waveform_kwargs=extra_waveform_kwargs,force_22_mode=opts.force_hyperbolic_22,**extra_kwargs) + verbose=opts.verbose,quiet=not opts.verbose, + skip_interpolation=opts.vectorized, + **waveform_generation_kwargs, **extra_kwargs) # skip nan ! Something horrible has happened if np.isnan(guess_snr): @@ -3658,9 +3717,47 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t # modulation and delay-derivative operators. This branch must precede the two # individual response branches: both flags now request one compound physical model. rotating_freqresponse_data = None + def _apply_order_control(products, selected_p, selected_q): + """Explicitly gated U,V order scan; returns the production subset.""" + if not _response_order_active: + return products + import warnings + from RIFT.likelihood import response_order as _response_order + _report = _response_order.estimate_response_orders( + products[4], products[1], products[2], + target_snr=float(opts.response_order_snr), + lnL_tolerance=float(opts.response_order_lnL_tol), + n_samples=int(opts.response_order_sky_samples), + selected_p=int(selected_p), selected_q=int(selected_q), + vary_p=bool(opts.check_slowrot_pmax or opts.choose_slowrot_pmax), + vary_q=bool(opts.check_finite_size_Qmax or opts.choose_slowrot_Qmax)) + _response_order.print_order_report(_report) + if (opts.check_slowrot_pmax or opts.check_finite_size_Qmax) and not _report['selected_passes']: + warnings.warn("chosen response order fails the requested accuracy: predicted Delta lnL={:.3g} > {:.3g} at SNR {:.6g}".format(_report['selected_delta_lnL'], _report['lnL_tolerance'], _report['target_snr']), RuntimeWarning) + _p, _q = int(selected_p), int(selected_q) + if opts.choose_slowrot_pmax or opts.choose_slowrot_Qmax: + if not _report['reference_resolved']: + raise ValueError("cannot auto-select from an unresolved finite response reference; raise the diagnostic reference order") + if _report['chosen'] is None: + raise ValueError("no response order in the diagnostic reference bank satisfies the requested SNR/error budget") + if opts.choose_slowrot_pmax: + _p = int(_report['chosen']['p_max']) + if opts.choose_slowrot_Qmax: + _q = int(_report['chosen']['Qmax']) + opts.rotation_p_max, opts.freqresponse_qmax = _p, _q + print(" response order used: p_max=%d Qmax=%d" % (_p, _q)) + return _response_order.truncate_precompute_products(products, p_max=_p, q_max=_q) + if opts.rotation_slow and opts.freqresponse: - _pmax = int(opts.rotation_p_max) - _qmax = int(opts.freqresponse_qmax) + _p_selected = int(opts.rotation_p_max) + _q_selected = int(opts.freqresponse_qmax) + _pmax = max(_p_selected, int(opts.response_order_p_reference)) if (opts.check_slowrot_pmax or opts.choose_slowrot_pmax) else _p_selected + _qmax = max(_q_selected, int(opts.response_order_Q_reference)) if (opts.check_finite_size_Qmax or opts.choose_slowrot_Qmax) else _q_selected + if _response_order_active: + from RIFT.likelihood import response_order as _response_order + _response_order.guard_reference_bank( + 'combined', _pmax, _qmax, opts.l_max, len(data_dict), + opts.response_order_max_bank_gib) _arm = opts.freqresponse_arm_length if _arm is not None: if '=' in str(_arm): @@ -3668,15 +3765,36 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t for kv in str(_arm).split(',')} else: _arm = float(_arm) - _rint_rf, _ct_rf, _ctV_rf, _rho_rf, _meta_rf = \ - factored_likelihood_rotating_freqresponse.PrecomputeLikelihoodTermsRotatingFreqResponse( + if os.environ.get('RIFT_GPU_PRECOMPUTE', '0') == '1' and opts.gpu and xpy_default is not np: + if _response_order_active: + raise NotImplementedError('Device-resident response-order selection is not yet supported; choose explicit orders or disable RIFT_GPU_PRECOMPUTE') + if os.environ.get('RIFT_GPU_WAVEFORM', 'lal') != 'lal': + raise ValueError('Native GPU waveform provider is not yet validated; use RIFT_GPU_WAVEFORM=lal') + from RIFT.likelihood.gpu_precompute import ( + PrecomputeLikelihoodTermsRotatingFreqResponseGPU, + pack_device_precompute) + _packed_rf, _meta_rf = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( fiducial_epoch, t_window, P, data_dict, psd_dict, opts.l_max, fmax, Qmax=_qmax, L_arm=_arm, p_max=_pmax, analyticPSD_Q=False, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, - verbose=opts.verbose, quiet=not opts.verbose, skip_interpolation=True) - _lkRF, _rhoA, _uAA, _vAA, _epRF = \ - factored_likelihood_rotating_freqresponse.pack_rotating_freqresponse_arrays( - _meta_rf, _rho_rf, _ct_rf, _ctV_rf) + verbose=opts.verbose, quiet=not opts.verbose, + skip_interpolation=True, return_device=True, + **waveform_generation_kwargs) + _lkRF, _rhoA, _uAA, _vAA, _epRF = pack_device_precompute(_packed_rf, _meta_rf) + else: + _rint_rf, _ct_rf, _ctV_rf, _rho_rf, _meta_rf = \ + factored_likelihood_rotating_freqresponse.PrecomputeLikelihoodTermsRotatingFreqResponse( + fiducial_epoch, t_window, P, data_dict, psd_dict, opts.l_max, fmax, + Qmax=_qmax, L_arm=_arm, p_max=_pmax, analyticPSD_Q=False, + inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, + verbose=opts.verbose, quiet=not opts.verbose, + skip_interpolation=True, **waveform_generation_kwargs) + _rint_rf, _ct_rf, _ctV_rf, _rho_rf, _meta_rf = _apply_order_control( + (_rint_rf, _ct_rf, _ctV_rf, _rho_rf, _meta_rf), + _p_selected, _q_selected) + _lkRF, _rhoA, _uAA, _vAA, _epRF = \ + factored_likelihood_rotating_freqresponse.pack_rotating_freqresponse_arrays( + _meta_rf, _rho_rf, _ct_rf, _ctV_rf) if opts.gpu and (not xpy_default is np): for _det in _rhoA: for _a in _rhoA[_det]: @@ -3688,7 +3806,7 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t U_by_aa=_uAA, V_by_aa=_vAA, epochDict=_epRF) _nbasis = len(_meta_rf['a_list']) print(" [rotation-slow+freqresponse] compound precompute complete; " - "p_max", _pmax, "Qmax", _qmax, "basis elements", _nbasis, + "p_max", _meta_rf['p_max'], "Qmax", _meta_rf['Qmax'], "basis elements", _nbasis, "ordered U/V pairs", _nbasis * _nbasis, "arm-length", opts.freqresponse_arm_length, "(GPU)" if (opts.gpu and not xpy_default is np) else "(CPU)") @@ -3696,7 +3814,13 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t # [Path A] slow-rotation precompute: build the harmonic-indexed bank and pack it. rotation_slow_data = None if opts.rotation_slow and not opts.freqresponse: - _pmax = int(opts.rotation_p_max) + _p_selected = int(opts.rotation_p_max) + _pmax = max(_p_selected, int(opts.response_order_p_reference)) if (opts.check_slowrot_pmax or opts.choose_slowrot_pmax) else _p_selected + if _response_order_active: + from RIFT.likelihood import response_order as _response_order + _response_order.guard_reference_bank( + 'rotation', _pmax, 0, opts.l_max, len(data_dict), + opts.response_order_max_bank_gib) # --rotation-n-harmonics is a FLOOR, not the literal width: the response # coefficients C_{(p,ntilde)} reach |ntilde| <= 2 + p_max (issue #142), and the # option's default of 2 is only the p_max=0 answer. The precompute now enforces @@ -3712,6 +3836,9 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t harmonics=_harm, p_max=_pmax, analyticPSD_Q=False, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, verbose=opts.verbose, quiet=not opts.verbose, skip_interpolation=True) + _rint_r, _ct_r, _ctV_r, _rho_r, _meta_r = _apply_order_control( + (_rint_r, _ct_r, _ctV_r, _rho_r, _meta_r), + _p_selected, 0) _lkR, _rhoN, _uNN, _vNN, _epR = factored_likelihood_with_rotation.pack_rotation_arrays( _meta_r, _rho_r, _ct_r, _ctV_r) if opts.gpu and (not xpy_default is np): @@ -3726,7 +3853,7 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t _vNN[_det][_pair] = cupy.asarray(_vNN[_det][_pair]) rotation_slow_data = dict(meta=_meta_r, lookupNKDict=_lkR, rho_by_n=_rhoN, U_by_nn=_uNN, V_by_nn=_vNN, epochDict=_epR) - print(" [rotation-slow] precompute complete; p_max", _pmax, "sidereal harmonics", + print(" [rotation-slow] precompute complete; p_max", _meta_r['p_max'], "sidereal harmonics", _meta_r['harmonics'], # the bank's own record, not our request "(GPU)" if (opts.gpu and not xpy_default is np) else "(CPU)") @@ -3734,7 +3861,13 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t # into the modes once and pack the response-basis overlap bank. freqresponse_data = None if opts.freqresponse and not opts.rotation_slow: - _qmax = int(opts.freqresponse_qmax) + _q_selected = int(opts.freqresponse_qmax) + _qmax = max(_q_selected, int(opts.response_order_Q_reference)) if (opts.check_finite_size_Qmax or opts.choose_slowrot_Qmax) else _q_selected + if _response_order_active: + from RIFT.likelihood import response_order as _response_order + _response_order.guard_reference_bank( + 'finite', 0, _qmax, opts.l_max, len(data_dict), + opts.response_order_max_bank_gib) _arm = opts.freqresponse_arm_length if _arm is not None: if '=' in str(_arm): # per-detector 'C1=40000,E1=10000' @@ -3746,6 +3879,9 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t Qmax=_qmax, L_arm=_arm, analyticPSD_Q=False, inv_spec_trunc_Q=inv_spec_trunc_Q, T_spec=T_spec, verbose=opts.verbose, quiet=not opts.verbose, skip_interpolation=True) + _rint_f, _ct_f, _ctV_f, _rho_f, _meta_f = _apply_order_control( + (_rint_f, _ct_f, _ctV_f, _rho_f, _meta_f), + 0, _q_selected) _lkF, _rhoP, _uPP, _vPP, _epF = factored_likelihood_freqresponse.pack_freqresponse_arrays( _meta_f, _rho_f, _ct_f, _ctV_f) if opts.gpu and (not xpy_default is np): @@ -3760,7 +3896,7 @@ def analyze_event(P_list, indx_event, data_dict, psd_dict, fmax, opts,inv_spec_t _vPP[_det][_pair] = cupy.asarray(_vPP[_det][_pair]) freqresponse_data = dict(meta=_meta_f, lookupNKDict=_lkF, rho_by_p=_rhoP, U_by_pp=_uPP, V_by_pp=_vPP, epochDict=_epF) - print(" [freqresponse] precompute complete; Qmax", _qmax, "arm-length", opts.freqresponse_arm_length, + print(" [freqresponse] precompute complete; Qmax", _meta_f['Qmax'], "arm-length", opts.freqresponse_arm_length, "(GPU)" if (opts.gpu and not xpy_default is np) else "(CPU)") def _cal_error_probe(n_cal_now, n_start=256, n_cap=None, rel_tol=0.1): diff --git a/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_jax b/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_jax index 1c5248ba7..363abbcd5 100755 --- a/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_jax +++ b/MonteCarloMarginalizeCode/Code/bin/integrate_likelihood_extrinsic_jax @@ -59,10 +59,16 @@ from __future__ import print_function import os import sys +import json from optparse import OptionParser, OptionGroup import numpy as np +# The opt-in CuPy precompute and JAX likelihood share one GPU. Set this before +# importing JAX or any RIFT module that may initialize its allocator. +if os.environ.get('RIFT_GPU_PRECOMPUTE', '0') == '1': + os.environ.setdefault('XLA_PYTHON_CLIENT_PREALLOCATE', 'false') + import jax import jax.numpy as jnp jax.config.update("jax_enable_x64", True) @@ -106,6 +112,7 @@ from RIFT.likelihood.jax_ile.direct_marginalization_policy import ( POLICY_CHOICES as DIRECT_MARG_POLICY_CHOICES, POLICY_DEFAULT as DIRECT_MARG_POLICY_DEFAULT, PolicyConfig as DirectMargPolicyConfig, + BoundedMultipeakConfig, validate_bounded_multipeak_config, RESERVE_SCHEME_CHOICES as DIRECT_MARG_RESERVE_SCHEME_CHOICES, reserve_pair as _direct_marg_reserve_pair, summarize_policy_ledger as _summarize_policy_ledger, @@ -113,6 +120,7 @@ from RIFT.likelihood.jax_ile.direct_marginalization_policy import ( _JAX_GATHERER_NAMES = tuple(_JAX_GATHERERS) from RIFT.likelihood.jax_ile.wrapper import ( JAXExtrinsicLikelihood, JAXDistanceMarginalizedLikelihood, + BOUNDED_MULTIPEAK_LOG_ZERO, ) MSUN = lal.MSUN_SI @@ -250,6 +258,11 @@ _ILE_ALL_OPTS = { "--force-adapt-all", "--force-gpu-only", "--force-hyperbolic-22", "--force-reset-all", "--force-xpy", "--freqresponse", "--freqresponse-arm-length", "--freqresponse-qmax", "--gpu", + "--check-slowrot-pmax", "--check-finite-size-Qmax", + "--choose-slowrot-pmax", "--choose-slowrot-Qmax", + "--response-order-snr", "--response-order-lnL-tol", + "--response-order-sky-samples", "--response-order-p-reference", + "--response-order-Q-reference", "--response-order-max-bank-gib", "--inclination-cosine-sampler", "--internal-gmm-adaptive-components", "--internal-gmm-correlate-all", "--internal-gmm-defensive-frac", "--internal-gmm-inflate", "--internal-gmm-max-components", @@ -402,6 +415,11 @@ def check_critical_and_report(opts, optp): _sampler_method = getattr(opts, "sampler_method", None) _jax_av_active = _sampler_method in ("AV", "portfolio") if _jax_av_active: + _dp = str(getattr(opts, "d_prior", None) or "euclidean").strip().lower() + if _dp not in ("euclidean", "volumetric", "pseudo_cosmo"): + fatal.append("--sampler-method %s supports distance priors " + "Euclidean/volumetric and pseudo_cosmo, not %s" + % (_sampler_method, _dp)) try: resolve_av_angular_limits(opts) except SystemExit as exc: @@ -477,6 +495,20 @@ def check_critical_and_report(opts, optp): "marginalization) is not implemented for psi-sampling modes " "on this driver; use --mode flowmc-phipsimarg or " "flowmc-dpsimarg, which marginalize psi by construction") + # Static multipeak options must be effective and valid before precompute. + _bounded = getattr(opts, "angle_marg_scheme", None) == "multipeak-jax" + if _bounded: + if getattr(opts, "mode", None) != "flowmc-phipsimarg": + fatal.append("multipeak-jax requires --mode flowmc-phipsimarg") + try: + bounded_multipeak_config_from_options(opts) + except (TypeError, ValueError) as exc: + fatal.append("multipeak-jax: %s" % exc) + else: + for field in (*BoundedMultipeakConfig._fields, "decline_action"): + flag = "--multipeak-jax-" + field.replace("_", "-") + if was_supplied(opts, flag): + fatal.append("%s requires --angle-marg-scheme multipeak-jax" % flag) # Cross-axis policy scope, at parse time (external review of #278, P1): # only --mode flowmc-phipsimarg reads the policy, so a request anywhere # else would otherwise complete on the ordinary likelihood, and the @@ -682,6 +714,38 @@ def check_critical_and_report(opts, optp): fatal.append("--freqresponse-qmax must be >= 0") _rotation = bool(getattr(opts, "rotation_slow", False)) _freqresponse = bool(getattr(opts, "freqresponse", False)) + _order_p = bool(getattr(opts, "check_slowrot_pmax", False) + or getattr(opts, "choose_slowrot_pmax", False)) + _order_q = bool(getattr(opts, "check_finite_size_Qmax", False) + or getattr(opts, "choose_slowrot_Qmax", False)) + if _order_p and not _rotation: + fatal.append("pmax check/choice requires --rotation-slow") + if _order_q and not _freqresponse: + fatal.append("Qmax check/choice requires --freqresponse") + if (_order_p or _order_q) and getattr(opts, "response_order_snr", None) is None: + fatal.append("response-order check/choice requires --response-order-snr") + if (getattr(opts, "check_slowrot_pmax", False) + and getattr(opts, "choose_slowrot_pmax", False)): + fatal.append("choose either --check-slowrot-pmax or --choose-slowrot-pmax") + if (getattr(opts, "check_finite_size_Qmax", False) + and getattr(opts, "choose_slowrot_Qmax", False)): + fatal.append("choose either --check-finite-size-Qmax or --choose-slowrot-Qmax") + if ((_order_p and _order_q) + and ((getattr(opts, "check_slowrot_pmax", False) + or getattr(opts, "check_finite_size_Qmax", False)) + and (getattr(opts, "choose_slowrot_pmax", False) + or getattr(opts, "choose_slowrot_Qmax", False)))): + fatal.append("do not mix check and choose controls in one compound response scan") + if float(getattr(opts, "response_order_lnL_tol", 0.1)) <= 0: + fatal.append("--response-order-lnL-tol must be positive") + if int(getattr(opts, "response_order_sky_samples", 128)) < 8: + fatal.append("--response-order-sky-samples must be at least 8") + if (int(getattr(opts, "response_order_p_reference", 2)) < 0 + or int(getattr(opts, "response_order_Q_reference", 8)) < 0): + fatal.append("response-order reference orders must be nonnegative") + if (not np.isfinite(float(getattr(opts, "response_order_max_bank_gib", 4.0))) + or float(getattr(opts, "response_order_max_bank_gib", 4.0)) <= 0): + fatal.append("--response-order-max-bank-gib must be finite and positive") if not _rotation: for _name in ("--rotation-n-harmonics", "--rotation-p-max"): if was_supplied(opts, _name): @@ -740,6 +804,7 @@ def check_critical_and_report(opts, optp): "--spin2z", "--eff-lambda", "--deff-lambda", "--approximant", "--l-max", "--event-time", "--fmin-template", "--reference-freq", "--fmax", "--srate", + "--srate-internal", "--data-integration-window-half", "--internal-data-storage-window-half", "--d-min", "--d-max", "--limit-distance", @@ -750,9 +815,18 @@ def check_critical_and_report(opts, optp): "--time-marginalization", "--time-marginalization-quadrature", "--interpolate-time", "--vectorized", "--use-gwsignal", "--rotation-slow", + "--internal-waveform-fd-L-frame", + "--internal-waveform-fd-no-condition", + "--internal-precompute-ignore-threshold", "--no-memory", + "--e-freq", "--fmin-ifo", "--rotation-n-harmonics", "--rotation-p-max", "--freqresponse", "--freqresponse-qmax", - "--freqresponse-arm-length"} + "--freqresponse-arm-length", "--check-slowrot-pmax", + "--check-finite-size-Qmax", "--choose-slowrot-pmax", + "--choose-slowrot-Qmax", "--response-order-snr", + "--response-order-lnL-tol", "--response-order-sky-samples", + "--response-order-p-reference", "--response-order-Q-reference", + "--response-order-max-bank-gib"} # These are implemented PER MODE. Listing them unconditionally would claim # they act under --mode laplace-is (the default), nuts, map, multistart-nuts # and nuts-phimarg, where they are inert -- exactly the silent no-op this @@ -760,6 +834,7 @@ def check_critical_and_report(opts, optp): mode = getattr(opts, "mode", None) if getattr(opts, "sampler_method", None) in ("AV", "portfolio"): implemented |= {"--sampler-method", "--n-eff", + "--d-prior", "--sampler-anisotropic-bins", "--limit-right-ascension", "--limit-declination", "--limit-psi", "--limit-inclination"} @@ -795,14 +870,13 @@ def check_critical_and_report(opts, optp): print("Note: %s only act on the tempered modes (%s); --mode %s ignores " "them." % (" ".join(sorted(set(inert))), " ".join(sorted(_TEMPERED_MODES)), mode)) - # --d-prior (RO'S 2026-09-08: missing knobs no-op for compatibility, but - # must warn). --d-prior is accepted (ILE compatibility) but never - # forwarded: the JAX driver always integrates against the volumetric d^2 - # prior on [--d-min, --d-max]. 'Euclidean'/'volumetric' (any case) IS - # that prior, so passing it explicitly is not a deviation and gets no note. + # AV/portfolio forwards both the volumetric and pseudo-cosmological priors. + # Other samplers remain volumetric and must visibly report a non-volumetric + # compatibility request as ignored. _dp = getattr(opts, "d_prior", None) - if _dp not in (None, "") and str(_dp).strip().lower() not in ( - "euclidean", "volumetric"): + if (getattr(opts, "sampler_method", None) not in ("AV", "portfolio") + and _dp not in (None, "") and str(_dp).strip().lower() not in ( + "euclidean", "volumetric")): print("Note: --d-prior %r is accepted but IGNORED by the JAX driver: " "it always integrates distance against the volumetric d^2 " "prior on [--d-min, --d-max]." % (_dp,)) @@ -932,6 +1006,30 @@ def build_parser(): help="Highest finite-arm projection power (default 4).") g.add_option("--freqresponse-arm-length", default=None, help="Arm length in metres, globally or DET=value comma list.") + g.add_option("--check-slowrot-pmax", action="store_true", default=False, + help="Estimate the required pmax and warn if --rotation-p-max is too small.") + g.add_option("--check-finite-size-Qmax", "--check-finite-size-qmax", + dest="check_finite_size_Qmax", action="store_true", default=False, + help="Estimate the required Qmax and warn if --freqresponse-qmax is too small.") + g.add_option("--choose-slowrot-pmax", action="store_true", default=False, + help="Choose the least pmax satisfying the response error budget.") + g.add_option("--choose-slowrot-Qmax", "--choose-finite-size-Qmax", + dest="choose_slowrot_Qmax", action="store_true", default=False, + help="Choose the least finite-size Qmax satisfying the response error budget.") + g.add_option("--response-order-snr", type=float, default=None, + help="Target network SNR for response-order checks/choice (required when active).") + g.add_option("--response-order-lnL-tol", type=float, default=0.1, + help="Allowed worst scanned Asimov likelihood loss (default 0.1).") + g.add_option("--response-order-sky-samples", type=int, default=128, + help="Deterministic full-prior angular design size (default 128).") + g.add_option("--response-order-p-reference", type=int, default=2, + help="Highest p used by the diagnostic reference bank (default 2).") + g.add_option("--response-order-Q-reference", "--response-order-q-reference", + dest="response_order_Q_reference", type=int, default=8, + help="Highest Q used by the diagnostic reference bank (default 8).") + g.add_option("--response-order-max-bank-gib", type=float, default=4.0, + help="Refuse a diagnostic whose dense U,V planning estimate exceeds " + "this many GiB (default 4).") optp.add_option_group(g) g = OptionGroup(optp, "Extrinsic exploration / sampling") @@ -1169,6 +1267,24 @@ def build_parser(): "ran is printed. See RIFT.likelihood.jax_ile.anglemarg." % (ANGLE_MARG_DEFAULT, ANGLE_MARG_LEGACY, ANGLE_MARG_LEGACY, ANGLE_MARG_LEGACY)) + g.add_option("--multipeak-jax-decline-action", type="choice", + choices=("drop", "refuse"), default="drop", + help="Handle bounded multipeak declines (default drop): " + "drop gives proposals zero numerical weight and labels " + "accepted-region evidence; refuse aborts on declined " + "log_likelihood batches only. Scalar MAP/Fisher/MALA " + "declines return NaN without raising or latching. " + "Important declined mass requires changing the config; " + "neither action runs a reserve.") + # Each cap is a construction-time scalar. The complete configuration is + # recorded on the likelihood, and no reserve is available at any setting. + for field, default in BoundedMultipeakConfig()._asdict().items(): + g.add_option("--multipeak-jax-" + field.replace("_", "-"), + type="int" if isinstance(default, int) else "float", + default=default, + help=("Static multipeak-jax %s (default %s); requires " + "--angle-marg-scheme multipeak-jax." + % (field.replace("_", " "), default))) g.add_option("--direct-marginalization-policy", default=DIRECT_MARG_POLICY_DEFAULT, choices=sorted(DIRECT_MARG_POLICY_CHOICES), @@ -1584,6 +1700,15 @@ def make_template(opts, fiducial_epoch, deltaF, deltaT): return P +def _analysis_delta_t(opts): + """Waveform/precompute cadence, distinct from the input-data cadence.""" + rate = (opts.srate if getattr(opts, "srate_internal", None) is None + else opts.srate_internal) + if rate is None or float(rate) <= 0: + raise ValueError("--srate/--srate-internal must be positive") + return 1.0 / float(rate) + + def load_templates(opts, fiducial_epoch, deltaF, deltaT): """Return a list of intrinsic templates (ChooseWaveformParams). @@ -1602,6 +1727,12 @@ def load_templates(opts, fiducial_epoch, deltaF, deltaT): P.fmax = 0.0 P.tref = fiducial_epoch P.dist = 1000.0 * 1e6 * PC # fiducial template distance (= distMpcRef) + # Extrinsic angles are sampled later, not baked into the base modes. + # Match conventional ILE: some mode generators encode these angles in + # h_lm, so retaining XML/grid values would apply them twice. + P.phiref = 0.0 + P.psi = 0.0 + P.incl = 0.0 if opts.approximant: P.approx = lalsim.GetApproximantFromString(opts.approximant) return P @@ -1650,7 +1781,7 @@ def load_templates(opts, fiducial_epoch, deltaF, deltaT): def load_injection(opts, fiducial_epoch): detectors = opts.inj_detectors.split(",") - deltaT = 1.0 / opts.srate + deltaT = _analysis_delta_t(opts) P = make_template(opts, fiducial_epoch, opts.inj_deltaF, deltaT) P.phi, P.theta = opts.inj_ra, opts.inj_dec P.psi, P.incl, P.phiref = opts.inj_psi, opts.inj_incl, opts.inj_phiref @@ -1666,20 +1797,39 @@ def load_injection(opts, fiducial_epoch): def load_frames(opts, fiducial_epoch): deltaT = 1.0 / opts.srate + deltaT_internal = (None if getattr(opts, "srate_internal", None) is None + else _analysis_delta_t(opts)) data_dict, psd_dict = {}, {} + fmin_ifo = {} + for item in getattr(opts, "fmin_ifo", []) or []: + inst, value = item.split("=", 1) + fmin_ifo[inst] = float(value) for inst, chan in map(lambda c: c.split("="), opts.channel_name): if opts.verbose: print("Reading channel %s:%s from %s" % (inst, chan, opts.cache_file)) data_dict[inst] = lalsimutils.frame_data_to_non_herm_hoff( opts.cache_file, inst + ":" + chan, start=opts.data_start_time, stop=opts.data_end_time, - window_shape=opts.window_shape, deltaT=deltaT) + window_shape=opts.window_shape, deltaT=deltaT, + deltaT_internal=deltaT_internal) for inst, psdf in map(lambda c: c.split("="), opts.psd_file): if opts.verbose: print("Reading PSD for %s from %s" % (inst, psdf)) psd_dict[inst] = lalsimutils.get_psd_series_from_xmldoc(psdf, inst) psd_dict[inst] = lalsimutils.resample_psd_series( psd_dict[inst], data_dict[inst].deltaF) + psd_window_shape = float(getattr(opts, "psd_window_shape", None) or 0.0) + if psd_window_shape > 0 or opts.window_shape > 0: + psd_factor = lalsimutils.psd_windowing_factor( + psd_window_shape, len(psd_dict[inst].data.data)) + data_factor = lalsimutils.psd_windowing_factor( + opts.window_shape, len(data_dict[inst].data.data)) + psd_dict[inst].data.data *= data_factor / psd_factor + if inst in fmin_ifo: + freqs = (psd_dict[inst].f0 + + psd_dict[inst].deltaF + * np.arange(psd_dict[inst].data.length)) + psd_dict[inst].data.data[freqs < fmin_ifo[inst]] = 0 detectors = list(data_dict.keys()) return data_dict, psd_dict, detectors, False @@ -1815,6 +1965,32 @@ def log_prior(theta, opts, with_distance): # --------------------------------------------------------------------------- # Batched lnL # --------------------------------------------------------------------------- +def bounded_multipeak_config_from_options(opts): + """Build the exact envelope executed by the wrapper from registered flags.""" + return validate_bounded_multipeak_config(BoundedMultipeakConfig(**{ + field: getattr(opts, "multipeak_jax_" + field) + for field in BoundedMultipeakConfig._fields})) + + +def require_bounded_multipeak_rows(like, lnL): + """Apply the requested decline action before publishing either artifact.""" + if getattr(like, "angle_marg_scheme", None) != "multipeak-jax": + return + bad = ~np.isfinite(np.asarray(lnL)) | (np.asarray(lnL) == BOUNDED_MULTIPEAK_LOG_ZERO) + if getattr(like, "bounded_multipeak_decline_action", "refuse") == "drop": + if np.all(bad): + raise RuntimeError("multipeak-jax: no accepted output rows; change " + "the static envelope. No samples or evidence published.") + return + if np.any(bad) or getattr(like, "bounded_multipeak_declined", False): + raise RuntimeError( + "--angle-marg-scheme multipeak-jax: %d of %d rows could not be " + "warranted within the static envelope; refusing to publish samples " + "or evidence (including any earlier declined batch). Change the " + "--multipeak-jax-* envelope, or select decline-action drop." + % (int(np.sum(bad)), int(np.size(bad)))) + + def eval_lnL(like, theta, opts, with_distance): N = theta.shape[0] out = np.empty(N) @@ -1826,6 +2002,8 @@ def eval_lnL(like, theta, opts, with_distance): sl = slice(i, min(i + chunk, N)) cols = [theta[sl, j] for j in range(theta.shape[1])] out[sl] = np.asarray(like.log_likelihood(*cols)) + if getattr(like, "bounded_multipeak_decline_action", "refuse") == "refuse": + require_bounded_multipeak_rows(like, out[sl]) if (getattr(like, "direct_marginalization_policy", "off") != "off" and np.any(np.isnan(out[sl]))): raise RuntimeError( @@ -2251,6 +2429,9 @@ def angle_grid_suspect_note(scheme=None): artifacts) are deliberately NOT claimed. This split also keeps host callbacks out of the expensive graph so JAX can persistently cache it. """ + if scheme == "multipeak-jax": + return ("MULTIPEAK-JAX bounded_cost=true dense_reserve=false " + "fixed_plan_autodiff_only=true derivative_warrant_certified=false") st = _anglemarg.amp_failsafe_state() if st.get("tripped"): return ("SUSPECT-ANGLE-GRID amp_failsafe=TRIPPED worst_amp=%.6g " @@ -2840,6 +3021,29 @@ def _parse_freqresponse_arm_length(value): return result +def _waveform_precompute_kwargs(opts): + """Return the waveform controls shared with production numpy ILE. + + These options act while constructing the mode time series, before JAX sees + the packed likelihood data. Omitting them therefore changes the numerical + likelihood rather than merely selecting a JAX execution detail. + """ + e_freq = getattr(opts, "e_freq", None) + extra = {"fd_alignment_postevent_time": 2, + "e_freq": 1 if e_freq is None else int(e_freq)} + if getattr(opts, "internal_waveform_fd_L_frame", False): + extra["fd_L_frame"] = True + if getattr(opts, "internal_waveform_fd_no_condition", False): + extra["no_condition"] = True + return dict( + use_gwsignal=bool(getattr(opts, "use_gwsignal", False)), + use_gwsignal_approx=(opts.approximant + if getattr(opts, "use_gwsignal", False) else None), + ignore_threshold=getattr(opts, "internal_precompute_ignore_threshold", None), + no_memory=bool(getattr(opts, "no_memory", False)), + extra_waveform_kwargs=extra) + + def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, rng, out_index, event_id, flow_state=None): """Build the JAX likelihood for one intrinsic template and run --mode. @@ -2863,10 +3067,22 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, "samples storage_half=%.6g s" % (g0, gcert, opts.internal_data_storage_window_half)) print("Building JAX likelihood (production precompute + banded pack)...") - waveform_kw = dict( - use_gwsignal=bool(getattr(opts, "use_gwsignal", False)), - use_gwsignal_approx=(opts.approximant - if getattr(opts, "use_gwsignal", False) else None)) + _order_active = any((opts.check_slowrot_pmax, opts.check_finite_size_Qmax, + opts.choose_slowrot_pmax, opts.choose_slowrot_Qmax)) + order_control = None + if _order_active: + order_control = dict( + target_snr=float(opts.response_order_snr), + lnL_tolerance=float(opts.response_order_lnL_tol), + n_samples=int(opts.response_order_sky_samples), + p_reference=int(opts.response_order_p_reference), + q_reference=int(opts.response_order_Q_reference), + max_bank_gib=float(opts.response_order_max_bank_gib), + check_p=bool(opts.check_slowrot_pmax), + check_q=bool(opts.check_finite_size_Qmax), + choose_p=bool(opts.choose_slowrot_pmax), + choose_q=bool(opts.choose_slowrot_Qmax)) + waveform_kw = _waveform_precompute_kwargs(opts) if opts.rotation_slow and opts.freqresponse: arm_length = _parse_freqresponse_arm_length(opts.freqresponse_arm_length) like_data, extras = build_rotating_freqresponse_data_from_precompute( @@ -2875,7 +3091,7 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, t_window=opts.internal_data_storage_window_half, Qmax=opts.freqresponse_qmax, L_arm=arm_length, p_max=opts.rotation_p_max, analyticPSD_Q=analyticPSD_Q, - verbose=opts.verbose, **waveform_kw) + verbose=opts.verbose, order_control=order_control, **waveform_kw) elif opts.rotation_slow: nh = int(opts.rotation_n_harmonics) like_data, extras = build_rotation_data_from_precompute( @@ -2883,7 +3099,8 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, opts.data_integration_window_half, opts.l_max, opts.fmax, t_window=opts.internal_data_storage_window_half, harmonics=tuple(range(-nh, nh + 1)), p_max=opts.rotation_p_max, - analyticPSD_Q=analyticPSD_Q, verbose=opts.verbose, **waveform_kw) + analyticPSD_Q=analyticPSD_Q, verbose=opts.verbose, + order_control=order_control, **waveform_kw) elif opts.freqresponse: arm_length = _parse_freqresponse_arm_length(opts.freqresponse_arm_length) like_data, extras = build_freqresponse_data_from_precompute( @@ -2891,7 +3108,8 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, opts.data_integration_window_half, opts.l_max, opts.fmax, t_window=opts.internal_data_storage_window_half, Qmax=opts.freqresponse_qmax, L_arm=arm_length, - analyticPSD_Q=analyticPSD_Q, verbose=opts.verbose, **waveform_kw) + analyticPSD_Q=analyticPSD_Q, verbose=opts.verbose, + order_control=order_control, **waveform_kw) else: like_data, extras = build_data_from_precompute( P.copy(), data_dict, psd_dict, fiducial_epoch, @@ -2901,6 +3119,11 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, q_time_pregrid_factor=( 1 if getattr(opts, "q_time_pregrid_factor", 1) is None else int(opts.q_time_pregrid_factor)), **waveform_kw) + if extras.get('order_report') is not None: + opts.rotation_p_max = int(extras['order_report']['final_p']) + opts.freqresponse_qmax = int(extras['order_report']['final_Q']) + print(" response order used: p_max=%d Qmax=%d" % + (opts.rotation_p_max, opts.freqresponse_qmax)) print(" feature:", getattr(like_data, "feature", None), " modes:", like_data.lms, " guessed SNR:", extras.get("guess_snr", "not estimated")) @@ -3042,7 +3265,11 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, time_quadrature=tq, d_prior_range=(opts.d_min, opts.d_max), dist_grid=dist_grid, dist_grid_tol=dist_tol, direct_marginalization_policy=policy, - policy_config=policy_config) + policy_config=policy_config, + bounded_multipeak_config=( + bounded_multipeak_config_from_options(opts) + if angle_marg == "multipeak-jax" else None), + bounded_multipeak_decline_action=opts.multipeak_jax_decline_action) except ValueError as e: if policy == "off": raise @@ -3331,6 +3558,7 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, seed_prior_frac=opts.jax_av_seed_prior_frac, anisotropic_bins=opts.sampler_anisotropic_bins, verbose=opts.verbose, sample_bounds=_sample_bounds, + distance_prior=opts.d_prior, **_distance_sampling) theta, lnL = res["theta"], res["lnL"] logZ, sig, neff = res["logZ"], res["sigma_over_Z"], res["neff"] @@ -3502,7 +3730,28 @@ def analyze_one(opts, P, data_dict, psd_dict, analyticPSD_Q, fiducial_epoch, # stayed silent -- an inert guard, which is the exact failure mode this # label exists to prevent. _scheme = getattr(like, "angle_marg_scheme", None) + require_bounded_multipeak_rows(like, lnL) _ev_note = angle_grid_suspect_note(_scheme) + if _scheme == "multipeak-jax": + _ev_note += " config=" + json.dumps( + like.angle_marg_info["config"], sort_keys=True, separators=(",", ":")) + _action = like.bounded_multipeak_decline_action + _ev_note += " decline_action=" + _action + if _action == "drop": + _ev_note += " evidence_scope=accepted-region omitted_mass=unbounded" + # Evidence has already counted every proposal, including finite + # log-zero declines. Only the output cloud is filtered here. + _keep = np.isfinite(lnL) & (np.asarray(lnL) != BOUNDED_MULTIPEAK_LOG_ZERO) + _ev_note += " output_rows_dropped=%d" % int(np.sum(~_keep)) + theta, lnL = np.asarray(theta)[_keep], np.asarray(lnL)[_keep] + if logw_export is not None: + logw_export = np.asarray(logw_export)[_keep] + _ev_note += (" audit_scope=log_likelihood-batches" + " audit_excludes=scalar-MAP-Fisher-MALA" + " refuse_latch_scope=log_likelihood-batches" + " scalar_refuse_decline=nan-without-raise-or-latch") + _ev_note += " host_evaluation_audit=" + json.dumps( + like.bounded_multipeak_audit, sort_keys=True, separators=(",", ":")) # Record the resolved distance-GH-nodes count in the same artifact header # line as the angle-marg/policy notes and (via write_samples' provenance # line) the mode/ESS record, so a reader of either artifact can see @@ -3570,7 +3819,7 @@ def main(argv=None): optp.error("--event-time is required (frame mode)") fiducial_epoch = opts.event_time rng = np.random.default_rng(opts.seed) - deltaT = 1.0 / opts.srate + deltaT = _analysis_delta_t(opts) # --- data (loaded once; shared across the intrinsic batch) --- if opts.inj_mode: diff --git a/MonteCarloMarginalizeCode/Code/bin/util_ConvertJAXILEFairdraws.py b/MonteCarloMarginalizeCode/Code/bin/util_ConvertJAXILEFairdraws.py new file mode 100755 index 000000000..bfe7789f0 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/bin/util_ConvertJAXILEFairdraws.py @@ -0,0 +1,213 @@ +#!/usr/bin/env python3 +"""Join JAX-ILE tabular fair-draw sidecars into a RIFT posterior table. + +The JAX driver writes one likelihood record ending in ``_.dat`` and one +``_samples.dat`` sidecar per intrinsic point. Unlike conventional +ILE, it does not write fair draws into LIGO-LW XML, so the XML converter cannot +construct the terminal posterior. This utility pairs the two tabular products +strictly and emits only coordinates that the JAX driver actually exports. +""" + +import argparse +import hashlib +import json +import os +from pathlib import Path +import re +import tempfile + +import numpy as np + + +HEADER = ( + "m1 m2 a1x a1y a1z a2x a2y a2z mc eta ra dec phiorb incl psi " + "distance Npts lnL p ps neff mtotal q chi_eff chi_p" +) + + +def _paired_record(sample_path: Path) -> Path: + return sample_path.with_name(sample_path.name.replace("_samples.dat", "_.dat")) + + +def _grid_index(path): + match = re.match(r"EXTR_out-(\d+)\.xml_(\d+)_(?:samples)?\.dat$", path.name) + if match is None: + raise RuntimeError("cannot recover grid index from {}".format(path)) + return int(match.group(1)) + int(match.group(2)) + + +def _spin_summaries(m1, m2, spins): + """Match the conventional LI chi_eff/chi_p definitions.""" + s1 = np.asarray(spins[:3]) + s2 = np.asarray(spins[3:]) + chi_eff = (m1 * s1[2] + m2 * s2[2]) / (m1 + m2) + if m2 > m1: + m1, m2, s1, s2 = m2, m1, s2, s1 + q = m2 / m1 + a1 = 2.0 + 1.5 * q + a2 = 2.0 + 1.5 / q + s1_perp = m1 ** 2 * np.linalg.norm(s1[:2]) + s2_perp = m2 ** 2 * np.linalg.norm(s2[:2]) + chi_p = max(a1 * s1_perp, a2 * s2_perp) / (a1 * m1 ** 2) + return chi_eff, chi_p + + +def assemble(directory, draws_per_intrinsic=None, expected_intrinsic=None): + """Return the joined table, refusing incomplete or malformed pairs.""" + samples = sorted(directory.glob("EXTR_out-*_samples.dat")) + if not samples: + raise RuntimeError("no JAX-ILE fair-draw sidecars found in {}".format(directory)) + + records = { + path for path in directory.glob("EXTR_out-*_.dat") + if not path.name.endswith("_samples.dat") + } + paired_records = {_paired_record(path) for path in samples} + if records != paired_records: + missing_samples = sorted(str(path) for path in records - paired_records) + missing_records = sorted(str(path) for path in paired_records - records) + raise RuntimeError("unpaired JAX-ILE products: records_without_samples={}, " + "samples_without_records={}".format( + missing_samples, missing_records)) + indices = {_grid_index(path) for path in samples} + if len(indices) != len(samples): + raise RuntimeError("duplicate intrinsic grid indices in fair-draw sidecars") + if expected_intrinsic is not None and indices != set(range(expected_intrinsic)): + missing = sorted(set(range(expected_intrinsic)) - indices) + excess = sorted(indices - set(range(expected_intrinsic))) + raise RuntimeError("intrinsic grid is incomplete: expected={}, found={}, " + "missing={}, excess={}".format( + expected_intrinsic, len(indices), missing, excess)) + + rows = [] + seen_records = set() + for sample_path in samples: + record_path = _paired_record(sample_path) + if not record_path.is_file(): + raise FileNotFoundError(record_path) + if record_path in seen_records: + raise RuntimeError("duplicate intrinsic record pairing: {}".format(record_path)) + seen_records.add(record_path) + + intrinsic = np.loadtxt(record_path, comments="#", ndmin=2) + if intrinsic.shape != (1, 13) or not np.isfinite(intrinsic).all(): + raise RuntimeError("invalid intrinsic likelihood record {}: {}".format( + record_path, intrinsic.shape)) + + extrinsic = np.atleast_1d(np.genfromtxt(sample_path, names=True)) + names = set(extrinsic.dtype.names or ()) + required = { + "right_ascension", "declination", "distance", "inclination", "psi", + "phi_orb", "loglikelihood", + } + if not required <= names: + raise RuntimeError("missing fair-draw columns in {}: {}".format( + sample_path, sorted(required - names))) + if draws_per_intrinsic is not None and len(extrinsic) != draws_per_intrinsic: + raise RuntimeError("expected {} rows in {}, found {}".format( + draws_per_intrinsic, sample_path, len(extrinsic))) + + values = intrinsic[0] + m1, m2 = values[1], values[2] + mtotal = m1 + m2 + eta = m1 * m2 / mtotal**2 + mc = (m1 * m2) ** 0.6 / mtotal**0.2 + q = min(m1, m2) / max(m1, m2) + chi_eff, chi_p = _spin_summaries(m1, m2, values[3:9]) + for sample in extrinsic: + row = np.array([ + *values[1:9], mc, eta, + sample["right_ascension"], sample["declination"], sample["phi_orb"], + sample["inclination"], sample["psi"], sample["distance"], + values[11], sample["loglikelihood"], 1.0, 1.0, values[12], + mtotal, q, chi_eff, chi_p, + ], dtype=float) + if not np.isfinite(row).all(): + raise RuntimeError("nonfinite joined fair draw from {}".format(sample_path)) + rows.append(row) + + if not rows: + raise RuntimeError("JAX-ILE sidecars contained no fair draws") + return np.vstack(rows) + + +def _sha256(path): + digest = hashlib.sha256() + with path.open("rb") as stream: + for block in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def _combined_source_sha256(directory): + digest = hashlib.sha256() + paths = sorted(set(directory.glob("EXTR_out-*_.dat")) | + set(directory.glob("EXTR_out-*_samples.dat"))) + for path in paths: + digest.update(path.name.encode("utf-8")) + digest.update(bytes.fromhex(_sha256(path))) + return digest.hexdigest() + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--directory", required=True, type=Path) + parser.add_argument("--output", required=True, type=Path) + parser.add_argument("--draws-per-intrinsic", type=int) + parser.add_argument("--expected-intrinsic", type=int, + help="require exactly contiguous grid IDs 0..N-1") + parser.add_argument("--shuffle-seed", type=int, default=1986) + parser.add_argument("--provenance", type=Path, + help="JSON ledger (default: OUTPUT.provenance.json)") + args = parser.parse_args() + + args.output.parent.mkdir(parents=True, exist_ok=True) + provenance = args.provenance + if provenance is None: + provenance = args.output.with_suffix(".provenance.json") + provenance.parent.mkdir(parents=True, exist_ok=True) + for stale in (args.output, provenance): + if stale.exists(): + stale.unlink() + + table = assemble(args.directory, args.draws_per_intrinsic, + args.expected_intrinsic) + np.random.RandomState(args.shuffle_seed).shuffle(table) + output_fd, output_name = tempfile.mkstemp( + prefix=args.output.name + ".tmp.", dir=str(args.output.parent)) + provenance_fd, provenance_name = tempfile.mkstemp( + prefix=provenance.name + ".tmp.", dir=str(provenance.parent)) + os.close(output_fd) + os.close(provenance_fd) + output_tmp = Path(output_name) + provenance_tmp = Path(provenance_name) + try: + np.savetxt(output_tmp, table, header=HEADER, comments="# ", fmt="%.12g") + payload = { + "converter": str(Path(__file__).resolve()), + "converter_sha256": _sha256(Path(__file__).resolve()), + "source_directory": str(args.directory.resolve()), + "combined_source_sha256": _combined_source_sha256(args.directory), + "intrinsic_points": len(list(args.directory.glob( + "EXTR_out-*_samples.dat"))), + "draws_per_intrinsic": args.draws_per_intrinsic, + "expected_intrinsic": args.expected_intrinsic, + "posterior_rows": len(table), + "shuffle_seed": args.shuffle_seed, + "output_sha256": _sha256(output_tmp), + "equal_weight_columns": {"p": 1.0, "ps": 1.0}, + "omitted_unavailable_coordinates": [ + "time", "redshift", "source_frame_masses"], + } + provenance_tmp.write_text( + json.dumps(payload, indent=2, sort_keys=True) + "\n") + os.replace(str(provenance_tmp), str(provenance)) + os.replace(str(output_tmp), str(args.output)) + finally: + for temporary in (output_tmp, provenance_tmp): + if temporary.exists(): + temporary.unlink() + + +if __name__ == "__main__": + main() diff --git a/MonteCarloMarginalizeCode/Code/test/benchmark_gpu_short_worker.py b/MonteCarloMarginalizeCode/Code/test/benchmark_gpu_short_worker.py new file mode 100644 index 000000000..084c64754 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/benchmark_gpu_short_worker.py @@ -0,0 +1,440 @@ +#!/usr/bin/env python3 +"""Short physical BBH, matched CPU/GPU precompute and repeated worker timings. + +No BNS input is opened. Synthetic signal/noise products exist only in memory. +This tests precompute and point likelihoods, not posterior convergence. +""" +import time +_PROCESS_START = time.perf_counter() + +import argparse +import json +import os +import platform +import resource +import sys + +import numpy as np +import lal +import lalsimulation as lalsim +import RIFT.lalsimutils as lsu +from RIFT.likelihood import gpu_precompute as gpu +from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + +_MODULE_IMPORT_SECONDS = time.perf_counter() - _PROCESS_START + + +def emit(**record): + print(json.dumps(record, sort_keys=True), flush=True) + + +def sync(xp): + if xp is not np: + xp.cuda.Stream.null.synchronize() + + +_TOP_LEVEL_TIMING_STAGES = { + 'initialization', 'waveform', 'input_prep', 'basis', 'Q_U', 'V', + 'device_export', 'host_export', +} + + +def timing_recorder(records, intrinsic, algorithm, phase): + """Return a callback that retains non-overlapping stage wall times.""" + def callback(stage, elapsed, details): + elapsed = float(elapsed) + if stage in _TOP_LEVEL_TIMING_STAGES: + records[stage] = records.get(stage, 0.0) + elapsed + emit(stage=stage, seconds=elapsed, intrinsic=intrinsic, phase=phase, + algorithm=algorithm, **details) + return callback + + +def process_snapshot(xp): + usage = resource.getrusage(resource.RUSAGE_SELF) + io_values = {} + with open('/proc/self/io') as stream: + for line in stream: + key, value = line.split(':', 1) + io_values[key] = int(value) + out = dict(user_seconds=float(usage.ru_utime), system_seconds=float(usage.ru_stime), + max_rss_kib=int(usage.ru_maxrss), major_faults=int(usage.ru_majflt), + minor_faults=int(usage.ru_minflt), proc_io=io_values) + if xp is not np: + free, total = xp.cuda.runtime.memGetInfo() + pool = xp.get_default_memory_pool() + out['gpu_memory'] = dict(free_bytes=int(free), total_bytes=int(total), + used_bytes=int(total-free), + pool_used_bytes=int(pool.used_bytes()), + pool_free_bytes=int(pool.free_bytes())) + return out + + +def snapshot_delta(before, after): + keys = ('user_seconds', 'system_seconds', 'major_faults', 'minor_faults') + out = {key: after[key] - before[key] for key in keys} + out['max_rss_kib_change'] = after['max_rss_kib'] - before['max_rss_kib'] + out['proc_io'] = {key: after['proc_io'].get(key, 0) - before['proc_io'].get(key, 0) + for key in sorted(set(before['proc_io']) | set(after['proc_io']))} + if 'gpu_memory' in after: + out['gpu_memory'] = {key: after['gpu_memory'][key] - before['gpu_memory'][key] + for key in after['gpu_memory']} + return out + + +def host_legacy_pack(packed, meta, xp, require_gpu): + """Pack first, then copy device arrays only for the post-timing oracle.""" + lookup, rho, U, V, epoch = gpu.pack_device_precompute( + packed, meta, require_gpu=require_gpu) + + def host(value): + return np.asarray(value) if xp is np else xp.asnumpy(value) + + return (lookup, + {det: {a: host(value) for a, value in rows.items()} + for det, rows in rho.items()}, + {det: host(value) for det, value in U.items()}, + {det: host(value) for det, value in V.items()}, + epoch) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--intrinsics', type=int, default=5) + parser.add_argument('--repeat-same-intrinsic', action='store_true', + help='repeat one identical intrinsic point; requires at least five calls') + parser.add_argument('--detectors', default='H1,L1,V1') + parser.add_argument('--qmax', type=int, default=1) + parser.add_argument('--pmax', type=int, default=1) + parser.add_argument('--fft-batch', type=int, default=4, + help='benchmark-only compound-basis FFT batch') + parser.add_argument('--q-row-batch', type=int, default=4, + help='benchmark-only Q inverse-FFT row batch') + parser.add_argument('--lmax', type=int, default=2) + parser.add_argument('--approximant', default='IMRPhenomD') + parser.add_argument('--delta-t', type=float, default=1/1024.) + parser.add_argument('--delta-f', type=float, default=0.5) + parser.add_argument('--fmin', type=float, default=30.) + parser.add_argument('--fref', type=float, default=100.) + parser.add_argument('--fmax', type=float, default=512.) + parser.add_argument('--t-window', type=float, default=0.15) + parser.add_argument('--mass1', type=float, default=30.) + parser.add_argument('--mass2', type=float, default=25.) + parser.add_argument('--spin1x', type=float, default=0.) + parser.add_argument('--spin1y', type=float, default=0.) + parser.add_argument('--spin1z', type=float, default=0.) + parser.add_argument('--spin2x', type=float, default=0.) + parser.add_argument('--spin2y', type=float, default=0.) + parser.add_argument('--spin2z', type=float, default=0.) + parser.add_argument('--arm-length', type=float, default=40000.) + parser.add_argument('--backend', choices=['cupy','numpy'], default='cupy', + help='numpy is harness validation only, not GPU performance') + parser.add_argument('--compare-numpy', action='store_true', + help='also time the batched NumPy algorithm on this same worker') + parser.add_argument('--gpu-only-profile-after-validation', action='store_true', + help=('skip all CPU oracles for timing only after a separate ' + 'matched small-case validation; emits no parity claim')) + parser.add_argument('--legacy-reference-intrinsics', type=int, default=None, + help=('number of initial intrinsic points to run through the very ' + 'slow scalar legacy oracle (default: all, preserving old behavior; ' + '0 uses batched NumPy as the numerical oracle)')) + args = parser.parse_args() + if (args.intrinsics < 1 or args.lmax < 2 or args.qmax < 0 or args.pmax < 0 + or args.fft_batch < 1 or args.q_row_batch < 1): + parser.error('intrinsics must be positive; lmax>=2 and qmax,pmax>=0') + if min(args.delta_t, args.delta_f, args.fmin, args.fref, + args.fmax, args.t_window, args.mass1, args.mass2, + args.arm_length) <= 0: + parser.error('waveform grid, frequencies, masses, window, and arm length must be positive') + if args.fmax > 0.5/args.delta_t: + parser.error('fmax exceeds the Nyquist frequency implied by delta-t') + if args.repeat_same_intrinsic and args.intrinsics < 5: + parser.error('--repeat-same-intrinsic requires --intrinsics >= 5') + legacy_count = (args.intrinsics if args.legacy_reference_intrinsics is None + else args.legacy_reference_intrinsics) + if legacy_count < 0 or legacy_count > args.intrinsics: + parser.error('legacy-reference-intrinsics must lie in [0, intrinsics]') + if args.gpu_only_profile_after_validation: + if args.backend != 'cupy': + parser.error('GPU-only profiling requires --backend cupy') + if args.legacy_reference_intrinsics != 0: + parser.error('GPU-only profiling requires explicit --legacy-reference-intrinsics 0') + if args.compare_numpy: + parser.error('GPU-only profiling cannot be combined with --compare-numpy') + elif legacy_count < args.intrinsics and not args.compare_numpy: + parser.error('skipping any legacy reference requires --compare-numpy') + if os.environ.get('RIFT_GPU_PRECOMPUTE') == '1': + parser.error('unset RIFT_GPU_PRECOMPUTE: it would silently route the CPU oracle to GPU') + backend_import_started = time.perf_counter() + if args.backend == 'cupy': + import cupy as xp + else: + xp = np + backend_import_seconds = time.perf_counter() - backend_import_started + readiness_started = time.perf_counter() + sync(xp) + context = gpu.GPUPrecomputeContext(xp) + sync(xp) + backend_readiness_seconds = time.perf_counter() - readiness_started + started = _PROCESS_START + emit(stage='startup_timing', module_import_seconds=_MODULE_IMPORT_SECONDS, + backend_import_seconds=backend_import_seconds, + backend_readiness_seconds=backend_readiness_seconds, + timing_scope='from first Python statement; excludes interpreter startup') + detectors = [det.strip() for det in args.detectors.split(',') if det.strip()] + if not detectors: + parser.error('at least one detector is required') + try: + approximant = lalsim.GetApproximantFromString(args.approximant) + except Exception as exc: + parser.error('unknown approximant %r: %s' % (args.approximant, exc)) + P0 = lsu.ChooseWaveformParams( + m1=args.mass1*lal.MSUN_SI, m2=args.mass2*lal.MSUN_SI, + s1x=args.spin1x, s1y=args.spin1y, s1z=args.spin1z, + s2x=args.spin2x, s2y=args.spin2y, s2z=args.spin2z, + fmin=args.fmin, fref=args.fref, + deltaT=args.delta_t, deltaF=args.delta_f, approx=approximant, + radec=True, phi=1.2, theta=0.3, incl=0.7, psi=0.5, phiref=0.4, + tref=1e9, dist=200e6*lal.PC_SI, detector=detectors[0]) + data, psds = {}, {} + for det in detectors: + Pd = P0.manual_copy() + Pd.detector = det + data[det] = lsu.non_herm_hoff(Pd) + n = data[det].data.length + psd = lal.CreateREAL8FrequencySeries( + det, lal.LIGOTimeGPS(0), 0., P0.deltaF, lal.SecondUnit, n//2+1) + psd.data.data[:] = [lalsim.SimNoisePSDaLIGOZeroDetHighPower(max(10.,f)) + for f in np.arange(n//2+1)*P0.deltaF] + psds[det] = psd + # This is deliberately a conservative upper bound based on every mode up + # to Lmax. It protects tuning jobs from a mistyped batch size without + # changing the production allocator or its automatic fallback. + a_count_bound = len(fr.compound_index_set(args.qmax, args.pmax)) + mode_count_bound = sum(2*l + 1 for l in range(2, args.lmax + 1)) + row_count_bound = a_count_bound * mode_count_bound + if args.fft_batch > a_count_bound: + parser.error('fft-batch exceeds the compound count') + if args.q_row_batch > row_count_bound: + parser.error('q-row-batch exceeds the maximum possible basis rows') + complex_bytes = np.dtype(np.complex128).itemsize + n_window_bound = int(2.0 * args.t_window / args.delta_t) + if n_window_bound < 1 or n_window_bound > n: + parser.error('t-window produces an empty or overlong Q window') + primary_bound = a_count_bound * mode_count_bound * n * complex_bytes + retained_q_bound = (len(detectors) * row_count_bound * n_window_bound * + complex_bytes) + retained_uv_bound = (len(detectors) * 2 * row_count_bound**2 * complex_bytes) + fft_scratch_bound = (5 * args.fft_batch * mode_count_bound * n * + complex_bytes) + q_scratch_bound = 3 * args.q_row_batch * n * complex_bytes + v_scratch_bound = 6 * min(args.fft_batch, 2) * mode_count_bound * n * complex_bytes + peak_bound = (primary_bound + retained_q_bound + retained_uv_bound + + max(fft_scratch_bound, q_scratch_bound, v_scratch_bound)) + free_bytes = ((1 << 62) if xp is np + else int(xp.cuda.runtime.memGetInfo()[0])) + if peak_bound > 0.80 * free_bytes: + parser.error('requested benchmark batches exceed the 80% free-memory guard') + emit(stage='batch_preflight', fft_batch=args.fft_batch, + q_row_batch=args.q_row_batch, compound_count_bound=a_count_bound, + mode_count_bound=mode_count_bound, row_count_bound=row_count_bound, + estimated_peak_bytes=peak_bound, + available_device_bytes=None if xp is np else free_bytes) + # Injection angles belong only to the synthetic detector data. The mode + # bank follows the ILE convention and carries no extrinsic angles; this is + # essential for XPHM as well as aligned-spin approximants. + P_template = P0.manual_copy() + P_template.phiref = P_template.psi = P_template.incl = 0.0 + sync(xp) + numpy_context = gpu.GPUPrecomputeContext(np) if args.compare_numpy else None + emit(stage='worker_setup', seconds=time.perf_counter()-started, + bins=n, masses_msun=[args.mass1,args.mass2], + spins=[[args.spin1x,args.spin1y,args.spin1z], + [args.spin2x,args.spin2y,args.spin2z]], + approximant=args.approximant, lmax=args.lmax, + delta_t=P0.deltaT, delta_f=P0.deltaF, fmin=args.fmin, + fmax=args.fmax, t_window=args.t_window, + qmax=args.qmax, pmax=args.pmax, + fft_batch=args.fft_batch, q_row_batch=args.q_row_batch, + legacy_reference_intrinsics=legacy_count, + gpu_only_profile_after_validation=args.gpu_only_profile_after_validation, + backend=args.backend, + gpu=None if xp is np else xp.cuda.runtime.getDeviceProperties(0)['name'].decode()) + environment = dict(python=platform.python_version(), numpy=np.__version__, + lal=getattr(lal, '__version__', None), + lalsimulation=getattr(lalsim, '__version__', None), + repeat_same_intrinsic=args.repeat_same_intrinsic, + timing_instrumentation=('synchronous stage callbacks; synchronization ' + 'perturbs asynchronous GPU scheduling')) + if xp is not np: + props = xp.cuda.runtime.getDeviceProperties(0) + environment.update( + cupy=xp.__version__, cuda_runtime=int(xp.cuda.runtime.runtimeGetVersion()), + cuda_driver=int(xp.cuda.runtime.driverGetVersion()), + device=dict(name=props['name'].decode(), + compute_capability='%d.%d' % (props['major'], props['minor']), + total_global_memory_bytes=int(props['totalGlobalMem']), + multiprocessor_count=int(props['multiProcessorCount']))) + emit(stage='runtime_environment', **environment) + for index in range(args.intrinsics): + phase = 'cold' if index == 0 else 'warm' + P = P_template.manual_copy() + if not args.repeat_same_intrinsic: + P.m1 += index*0.01*lal.MSUN_SI + common = dict(event_time_geo=1e9, t_window=args.t_window, P=P, data_dict=data, + psd_dict=psds, Lmax=args.lmax, fMax=args.fmax, Qmax=args.qmax, + p_max=args.pmax, L_arm=args.arm_length, skip_interpolation=True, + quiet=True, verbose=False) + cpu = cpu_pack = None + if index < legacy_count: + t0 = time.perf_counter() + cpu = fr.PrecomputeLikelihoodTermsRotatingFreqResponse( + **dict(common, P=P.manual_copy())) + emit(stage='cpu_precompute', intrinsic=index, seconds=time.perf_counter()-t0, + timing_role='numerical_reference_only') + cpu_pack = fr.pack_rotating_freqresponse_arrays( + cpu[4], cpu[3], cpu[1], cpu[2]) + numpy_seconds = None + numpy_pack = numpy_meta = None + if numpy_context is not None: + numpy_timings = {} + t0 = time.perf_counter() + numpy_bank = gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **dict(common, P=P.manual_copy()), backend=np, + context=numpy_context, + fft_batch=args.fft_batch, q_row_batch=args.q_row_batch, + return_device=True, + timing_callback=timing_recorder( + numpy_timings, index, 'batched_numpy', phase)) + numpy_seconds = time.perf_counter()-t0 + emit(stage='batched_numpy_precompute',intrinsic=index, + seconds=numpy_seconds, + stage_seconds=numpy_timings, + context_uploads=numpy_context.uploads, + context_hits=numpy_context.cache_hits, + timing_scope='host wall through resident bank return') + # Exercise the exact classic handoff contract after the timed region. + numpy_pack = host_legacy_pack( + numpy_bank[0], numpy_bank[1], np, require_gpu=False) + numpy_meta = numpy_bank[1] + del numpy_bank + sync(xp) + resources_before = process_snapshot(xp) + candidate_timings = {} + t0 = time.perf_counter() + bank = gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **dict(common, P=P.manual_copy()), context=context, backend=xp, + fft_batch=args.fft_batch, q_row_batch=args.q_row_batch, + return_device=True, + timing_callback=timing_recorder(candidate_timings, index, args.backend, phase)) + sync(xp) + candidate_seconds = time.perf_counter()-t0 + resources_after = process_snapshot(xp) + emit(stage='candidate_precompute',backend=args.backend, intrinsic=index, + phase=phase, + seconds=candidate_seconds, stage_seconds=candidate_timings, + context_uploads=context.uploads, context_hits=context.cache_hits, + context_entries=len(context._arrays), + device_pool_used_bytes=0 if xp is np else xp.get_default_memory_pool().used_bytes(), + timing_scope='host wall through resident bank return') + emit(stage='candidate_resources', backend=args.backend, intrinsic=index, + phase=phase, before=resources_before, after=resources_after, + delta=snapshot_delta(resources_before, resources_after)) + actual_modes = list(bank[1]['modes']) + actual_a = list(bank[1]['a_list']) + mode_count, a_count = len(actual_modes), len(actual_a) + q_bytes = sum(int(value.nbytes) for value in bank[0]['q'].values()) + basis_bytes_per_detector = (a_count*mode_count*n* + np.dtype(np.complex128).itemsize) + uv_bytes = sum(int(value.nbytes) for family in ('U', 'V') + for value in bank[0][family].values()) + emit(stage='bank_geometry', intrinsic=index, actual_mode_count=mode_count, + actual_modes=actual_modes, exact_compound_count=a_count, + exact_q_bytes=q_bytes, exact_uv_bytes=uv_bytes, + exact_retained_quv_bytes=q_bytes+uv_bytes, + primary_basis_bytes_per_detector=basis_bytes_per_detector, + full_fft_bins=n) + if numpy_seconds is not None and xp is not np: + emit(stage='batched_algorithm_comparison', intrinsic=index, + numpy_seconds=numpy_seconds, gpu_seconds=candidate_seconds, + end_to_end_speedup=numpy_seconds/candidate_seconds, + numpy_stage_seconds=numpy_timings, + gpu_stage_seconds=candidate_timings, + timing_scope=('same-process host wall through resident bank return; ' + 'container staging, scheduler queue, and later oracle copies excluded')) + if args.gpu_only_profile_after_validation: + if a_count != a_count_bound or not (0 < mode_count <= mode_count_bound): + raise RuntimeError('unexpected compound-bank geometry in GPU-only profile') + expected_shapes = { + 'q': (a_count, mode_count, n_window_bound), + 'U': (a_count, a_count, mode_count, mode_count), + 'V': (a_count, a_count, mode_count, mode_count), + } + for family, shape in expected_shapes.items(): + for det in detectors: + value = bank[0][family][det] + if value.shape != shape: + raise RuntimeError('%s %s has shape %r, expected %r' % + (det, family, value.shape, shape)) + if not bool(xp.asnumpy(xp.all(xp.isfinite(value)))): + raise RuntimeError('%s %s is nonfinite' % (det, family)) + expected_uploads = 3*len(detectors) + if context.uploads != expected_uploads: + raise RuntimeError('detector inputs were unexpectedly re-uploaded') + expected_hits = index * expected_uploads + if context.cache_hits != expected_hits: + raise RuntimeError('detector input cache-hit count is inconsistent') + emit(stage='gpu_profile_validation', intrinsic=index, + arrays_finite=True, geometry_valid=True, + cache_uploads=context.uploads, cache_hits=context.cache_hits, + parity_measured=False, + status='timing_only_after_separate_small_case_validation') + del bank + continue + # This is deliberately after candidate_seconds: production hands these + # buffers directly to JAX/classic GPU ILE. Host copies exist only so + # the small numerical oracle below can use its NumPy implementation. + got_pack = host_legacy_pack( + bank[0], bank[1], xp, require_gpu=(xp is not np)) + oracle_pack = cpu_pack if cpu_pack is not None else numpy_pack + oracle_meta = cpu[4] if cpu is not None else numpy_meta + oracle_name = 'legacy_cpu' if cpu is not None else 'batched_numpy' + errors = dict(Q=0., U=0., V=0.) + for det in detectors: + for ai in bank[1]['a_list']: + a, b = got_pack[1][det][ai], oracle_pack[1][det][ai] + np.testing.assert_allclose(a, b, rtol=2e-9, atol=1e-8) + errors['Q'] = max(errors['Q'],float(np.max(np.abs(a-b)))) + for name, slot in [('U',2),('V',3)]: + a, b = got_pack[slot][det], oracle_pack[slot][det] + np.testing.assert_allclose(a, b, rtol=2e-9, atol=1e-7) + errors[name] = max(errors[name],float(np.max(np.abs(a-b)))) + assert got_pack[4][det] == oracle_pack[4][det] + Pv = P0.manual_copy() + for name in ['phi','theta','incl','psi','phiref','dist']: + setattr(Pv,name,np.full(4,getattr(P0,name))) + Pv.phi += np.array([0.,1e-4,-1e-4,0.01]) + tvals = np.array([-P0.deltaT,0.,P0.deltaT]) + oracle_lnl = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals,Pv,oracle_meta,*oracle_pack,Lmax=args.lmax,time_interp='nearest',xpy=np,array_output=True) + got_lnl = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals,Pv,bank[1],*got_pack,Lmax=args.lmax,time_interp='nearest',xpy=np,array_output=True) + np.testing.assert_allclose(got_lnl,oracle_lnl,rtol=2e-9,atol=1e-8) + if not np.all(np.isfinite(got_lnl)): + raise RuntimeError('nonfinite downstream likelihood') + emit(stage='oracle_gpu_parity', intrinsic=index, oracle=oracle_name, + max_abs_error=errors, + max_abs_lnl_error=float(np.max(np.abs(oracle_lnl-got_lnl)))) + del bank,got_pack + if cpu is not None: + del cpu,cpu_pack + assert context.uploads == 3*len(detectors), 'detector inputs were re-uploaded' + emit(stage='worker_total',seconds=time.perf_counter()-started, + benchmark_mode=('gpu_only_profile_after_validation' + if args.gpu_only_profile_after_validation else 'matched_oracle'), + timing_scope=('first Python statement through completed benchmark; interpreter startup, ' + 'container launch, and queue excluded; cross-check with /usr/bin/time -v')) + + +if __name__ == '__main__': + main() diff --git a/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/lisa_drift_ledger.json b/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/lisa_drift_ledger.json index eedc631f5..3b6d75159 100644 --- a/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/lisa_drift_ledger.json +++ b/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/lisa_drift_ledger.json @@ -49,6 +49,10 @@ "decision": "PORT", "reason": "Tolerant truthiness for optparse values that may arrive as strings from the pipe. Belongs with _normalize_interpolate_time_argv, its ONLY caller in the main driver (opts._noloop_time_interp), not with the fair-draw family -- porting it alongside those helpers would have added dead code to the LISA driver." }, + "FUNC:analyze_event._apply_order_control": { + "decision": "NA", + "reason": "Nested dispatcher for the ground-detector response-order estimator. It is reachable only from the check/choose controls classified above, and LISA's TDI response has no corresponding pmax or Qmax truncation to dispatch." + }, "FUNC:analyze_event._cal_error_probe": { "decision": "NA", "reason": "Calibration Monte-Carlo error probe; see the --calibration-* reason." @@ -129,10 +133,34 @@ "decision": "NA", "reason": "LIGO/Virgo spline calibration-envelope marginalization. The LISA driver models no instrument calibration: it takes no envelope directory, has no cal nodes, and its response is applied analytically by factored_likelihood_LISA. LISA calibration, if it is ever modelled, will not have this data product or this spline parameterization, so porting the LIGO machinery would be actively misleading." }, + "OPTION:--check-finite-size-Qmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, + "OPTION:--check-finite-size-qmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, "OPTION:--check-good-enough": { "decision": "PORT", "reason": "Early-exit when the pipeline has written an 'ile_good_enough' sentinel. Pipeline plumbing, detector-agnostic." }, + "OPTION:--check-slowrot-pmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, + "OPTION:--choose-finite-size-Qmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, + "OPTION:--choose-slowrot-Qmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, + "OPTION:--choose-slowrot-pmax": { + "decision": "NA", + "reason": "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm approximations. LISA uses neither expansion: its time-dependent heliocentric, finite-arm response is already evaluated by the TDI response, so applying these orders would test or truncate the wrong detector model." + }, "OPTION:--d-prior-redshift": { "decision": "PORT", "reason": "ANSWERED (RO 2026-08-16): Planck15 via the framework helper, RIFT.likelihood.priors_utils.get_astropy_cosmology('Planck15'). RESOLVED AT SOURCE -- the MAIN driver has been moved to that helper too (it previously built its own FlatLambdaCDM from lal.H0_SI/lal.OMEGA_M = 67.900/0.3065 with a hardcoded fallback), so there is no divergence to port around: both codes now ask the same helper and a change is made in one place. Pinned by test_cosmology_single_source.py." @@ -321,6 +349,34 @@ "decision": "PORT", "reason": "Pick a random event from the input file. Detector-agnostic; flagged dangerous in its own help text for oversampling reasons that apply equally to LISA." }, + "OPTION:--response-order-Q-reference": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-lnL-tol": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-max-bank-gib": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-p-reference": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-q-reference": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-sky-samples": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, + "OPTION:--response-order-snr": { + "decision": "NA", + "reason": "Tolerance, reference-order, angular-design and memory controls used only by the ground-detector slow-rotation and finite-arm order estimator. That estimator does not represent LISA's TDI response, so none of its tuning surface applies." + }, "OPTION:--rotation-n-harmonics": { "decision": "NA", "reason": "Sidereal time-dependence of an EARTH-BASED antenna pattern F(t). The LISA constellation's motion is already carried by the LISA response itself (factored_likelihood_LISA + the h5/TDI frames), so this correction is both unnecessary and wrong there -- it would apply Earth rotation to a heliocentric detector." diff --git a/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/make_lisa_drift_ledger.py b/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/make_lisa_drift_ledger.py index d73d134e2..06ae6dc05 100644 --- a/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/make_lisa_drift_ledger.py +++ b/MonteCarloMarginalizeCode/Code/test/expensive_before_merging/integrators/make_lisa_drift_ledger.py @@ -274,6 +274,20 @@ "(CE/ET), built on lalsimulation detector geometry and an arm-length override in " "metres. LISA's finite-size response is not an add-on: it is the whole point of " "the TDI response the LISA driver already applies."), + (r"^OPTION:--(check-slowrot-pmax|check-finite-size-[Qq]max|" + r"choose-slowrot-pmax|choose-(finite-size|slowrot)-Qmax)$", "NA", + "Order checks and selectors for the Earth-rotation and 3G ground-detector finite-arm " + "approximations. LISA uses neither expansion: its time-dependent heliocentric, " + "finite-arm response is already evaluated by the TDI response, so applying these " + "orders would test or truncate the wrong detector model."), + (r"^OPTION:--response-order-", "NA", + "Tolerance, reference-order, angular-design and memory controls used only by the " + "ground-detector slow-rotation and finite-arm order estimator. That estimator does " + "not represent LISA's TDI response, so none of its tuning surface applies."), + (r"^FUNC:analyze_event\._apply_order_control$", "NA", + "Nested dispatcher for the ground-detector response-order estimator. It is reachable " + "only from the check/choose controls classified above, and LISA's TDI response has no " + "corresponding pmax or Qmax truncation to dispatch."), (r"^OPTION:--e-freq$", "NA", "TEOBResumS eccentric-frequency convention. Tied to a ground-based eccentric " "waveform path the LISA driver does not offer (it takes --modes / h5 frames)."), diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/README.md b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/README.md new file mode 100644 index 000000000..473cea620 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/README.md @@ -0,0 +1,68 @@ +# GPU precompute validation + +The unit suite compares every returned array against an independent NumPy oracle and +also compares a full `Re - /2` likelihood against the `Q/U` contraction. +Inputs are complex and non-Hermitian, both frequency halves are populated, PSD weights +are asymmetric, and one context is reused across multiple intrinsic points. + +Run the CPU gate from `MonteCarloMarginalizeCode/Code`: + +```bash +PYTHONPATH=. python -m pytest -q test/gpu_precompute +``` + +Run the mandatory device gate in the CUDA container: + +```bash +RIFT_REQUIRE_GPU_PRECOMPUTE=1 PYTHONPATH=. python -m pytest -q \ + test/gpu_precompute test/test_gpu_jax_handoff.py \ + test/waveforms/test_gpu_waveform.py test/waveforms/test_gpu_legacy_compat.py --require-gpu +``` + +The handoff gate requires both CuPy and JAX on the GPU, with JAX x64 enabled. +Set `JAX_ENABLE_X64=1`, `JAX_PLATFORMS=cuda`, and +`XLA_PYTHON_CLIENT_PREALLOCATE=false` before importing JAX so its allocator +can coexist with the live CuPy bank. The tests reject a CPU-only handoff when +the mandatory GPU gate is requested. They check buffer lifetime, no bulk +host transfer, and a nonzero data-term contribution as well as likelihood parity. + +`RIFT_GPU_PRECOMPUTE=1` routes the compound-response branch directly to +device banks in conventional GPU ILE and to DLPack in ILE-JAX. The existing +host waveform generator remains the default. The standalone legacy-return +precompute API still supports consumers that need host/LAL objects. +The device-resident route currently requires explicit response orders; it +rejects opt-in response-order check/choose controls instead of ignoring them. +The existing host route retains those controls unchanged. + +The benchmark starts its clock inside the already-running worker, synchronizes the +device before and after each call, and reports several sequential intrinsic points: + +```bash +PYTHONPATH=. python test/gpu_precompute/benchmark_gpu_precompute.py \ + --backend cupy --bins 2048 --basis 40 --modes 1 --window 256 \ + --intrinsics 3 +``` + +The short physical three-detector benchmark compares CPU and GPU precompute and processes +five nearby intrinsic points in one worker: + +```bash +PYTHONPATH=. python test/benchmark_gpu_short_worker.py --intrinsics 5 +``` + +The AV integration gate generates its own short frames and PSDs, evaluates two intrinsic +points with an n-eff target of 20, and limits saved samples to a 200-row fairdraw: + +```bash +PYTHONPATH=. python test/gpu_precompute/run_short_av_ile.py --keep +``` + +The corresponding JAX AV smoke test uses the same short two-point fixture, +target `n_eff=20`, and fairdraw cap, with JAX's text sample format: + +```bash +PYTHONPATH=. python test/gpu_precompute/run_short_jax_av_ile.py --keep +``` + +Container transfer and queue latency are outside all reported timers. Long BNS-scale +arrays are deferred until the short correctness and memory gates pass. diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/benchmark_gpu_precompute.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/benchmark_gpu_precompute.py new file mode 100644 index 000000000..c153932ef --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/benchmark_gpu_precompute.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python3 +"""Synchronized host-runtime benchmark; excludes process/container transfer time.""" + +import argparse +import json +import time + +import numpy as np + +from RIFT.likelihood.gpu_precompute import compound_precompute_arrays + + +def sync(xp): + if xp.__name__.startswith("cupy"): + xp.cuda.Stream.null.synchronize() + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--backend", choices=("numpy", "cupy"), default="cupy") + p.add_argument("--bins", type=int, default=1 << 20) + p.add_argument("--basis", type=int, default=40) + p.add_argument("--modes", type=int, default=1) + p.add_argument("--window", type=int, default=2048) + p.add_argument("--intrinsics", type=int, default=3) + p.add_argument("--frequency-chunk", type=int, default=1 << 16) + p.add_argument("--seed", type=int, default=771) + args = p.parse_args() + if args.backend == "cupy": + import cupy as xp + if xp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("no CUDA device") + else: + xp = np + + rng = np.random.default_rng(args.seed) + n = args.bins + # Allocate once: this models a persistent ILE worker with data/PSD resident. + data = xp.asarray((rng.normal(size=n) + 1j * rng.normal(size=n)).astype(np.complex128)) + weights = xp.asarray(rng.uniform(0.2, 1.5, size=n).astype(np.float64)) + context = {} + records = [] + for intrinsic in range(args.intrinsics): + basis = xp.asarray((rng.normal(size=(args.basis, args.modes, n)) + + 1j * rng.normal(size=(args.basis, args.modes, n))).astype(np.complex128)) + basis_conj = xp.asarray((rng.normal(size=(args.basis, args.modes, n)) + + 1j * rng.normal(size=(args.basis, args.modes, n))).astype(np.complex128)) + sync(xp) + t0 = time.perf_counter() + out = compound_precompute_arrays( + basis, basis_conj, data, weights, 1.0 / 8192.0, 1.0 / (n / 8192.0), + 0, min(args.window, n), backend=xp, return_device=True, context=context, + frequency_chunk=args.frequency_chunk) + sync(xp) + elapsed = time.perf_counter() - t0 + records.append({"intrinsic": intrinsic, "host_runtime_s": elapsed}) + del basis, basis_conj, out + print(json.dumps({ + "backend": args.backend, "bins": n, "basis": args.basis, + "modes": args.modes, "intrinsics": args.intrinsics, + "timing_scope": "inside worker; synchronized; excludes container transfer", + "records": records, + }, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/conftest.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/conftest.py new file mode 100644 index 000000000..0e43ff040 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/conftest.py @@ -0,0 +1,38 @@ +import os + +import numpy as np +import pytest + + +def pytest_addoption(parser): + parser.addoption( + "--require-gpu", action="store_true", default=False, + help="fail instead of skip unless CuPy can execute on a CUDA device", + ) + + +@pytest.fixture(params=("numpy", "cupy")) +def backend(request): + if request.param == "numpy": + return np + try: + import cupy + if cupy.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("CuPy reports no CUDA devices") + _ = (cupy.arange(4) + 1).sum() + cupy.cuda.Stream.null.synchronize() + return cupy + except Exception as exc: + if request.config.getoption("--require-gpu") or os.environ.get("RIFT_REQUIRE_GPU_PRECOMPUTE") == "1": + pytest.fail("real CUDA device required: %r" % (exc,)) + pytest.skip("no usable CUDA device: %r" % (exc,)) + + +def to_host(x): + try: + import cupy + if isinstance(x, cupy.ndarray): + return cupy.asnumpy(x) + except Exception: + pass + return np.asarray(x) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/oracle.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/oracle.py new file mode 100644 index 000000000..f1bad6b33 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/oracle.py @@ -0,0 +1,42 @@ +"""Small, deliberately slow NumPy oracle for GPU precompute tests. + +The frequency axis uses RIFT's centered, descending two-sided layout +[+f_Nyq, ..., +df, 0, ..., -f_Nyq+df]. LAL's reverse complex FFT is +N*df*ifft(ifftshift(X)), not bare numpy.ifft(X). +""" + +import numpy as np + + +def q_oracle(basis_fd, data_fd, weights2side, delta_f, n_shift, n_window): + basis_fd = np.asarray(basis_fd, dtype=np.complex128) + data_fd = np.asarray(data_fd, dtype=np.complex128) + weights2side = np.asarray(weights2side, dtype=np.float64) + n = basis_fd.shape[-1] + integrand = 2.0 * np.conj(basis_fd) * data_fd * weights2side + full = n * delta_f * np.fft.ifft(np.fft.ifftshift(integrand, axes=-1), axis=-1) + return np.roll(full, -int(n_shift), axis=-1)[..., : int(n_window)] + + +def uv_oracle(basis_fd, basis_conj_fd, weights2side, delta_f): + basis_fd = np.asarray(basis_fd, dtype=np.complex128) + basis_conj_fd = np.asarray(basis_conj_fd, dtype=np.complex128) + weights2side = np.asarray(weights2side, dtype=np.float64) + # a,i,f ; b,j,f -> a,b,i,j + u = 2.0 * delta_f * np.einsum( + "aif,bjf,f->abij", np.conj(basis_fd), basis_fd, weights2side, + optimize=False, + ) + v = 2.0 * delta_f * np.einsum( + "aif,bjf,f->abij", np.conj(basis_conj_fd), basis_fd, weights2side, + optimize=False, + ) + return u, v + + +def direct_log_likelihood(basis_fd, data_fd, weights2side, delta_f, coeff): + """Direct Re - /2 for h=sum_{a,m} coeff[a,m] chi[a,m].""" + h = np.einsum("am,amf->f", coeff, basis_fd, optimize=False) + hd = 2.0 * delta_f * np.sum(np.conj(h) * data_fd * weights2side) + hh = 2.0 * delta_f * np.sum(np.conj(h) * h * weights2side) + return float(np.real(hd - 0.5 * hh)) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_av_ile.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_av_ile.py new file mode 100644 index 000000000..654803a53 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_av_ile.py @@ -0,0 +1,248 @@ +#!/usr/bin/env python3 +"""Run a short, self-contained combined-response ILE with AV and GPU precompute. + +The clock starts inside the worker process invocation. Queue and container transfer are +outside this runner. Two nearby intrinsic points are evaluated sequentially so the shared +GPUPrecomputeContext must reuse detector data and PSD uploads. +""" + +import argparse +import importlib.util +import json +import os +from pathlib import Path +import subprocess +import tempfile +import time + +import numpy as np + + +def set_option(args, name, value=None): + while name in args: + i = args.index(name) + del args[i:i + (1 if value is None else 2)] + args.append(name) + if value is not None: + args.append(str(value)) + + +def load_generator(code_dir): + path = code_dir / "demo/rift/slowrot_gpu_validate/make_e2e_inputs.py" + spec = importlib.util.spec_from_file_location("short_gpu_e2e_inputs", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + # The older validation helper imports the retired `ligo.lw` namespace. Keep the + # fixture logic but supply the current production igwn_ligolw writer here. + def write_psd_xml(det, delta_f, fmax, path): + import lal + import lal.series + import lalsimulation as lalsim + from igwn_ligolw import utils as ligolw_utils + count = int(fmax / delta_f) + 1 + series = lal.CreateREAL8FrequencySeries( + det, lal.LIGOTimeGPS(0), 0.0, delta_f, lal.SecondUnit, count) + frequency = np.arange(count) * delta_f + values = np.array([ + lalsim.SimNoisePSDaLIGOZeroDetHighPower(max(float(f), 1.0)) + for f in frequency]) + values[~np.isfinite(values)] = 0.0 + series.data.data[:] = values + document = lal.series.make_psd_xmldoc({det: series}) + ligolw_utils.write_filename(document, path, compress="gz") + module.write_psd_xml = write_psd_xml + return module + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--work-dir", type=Path) + parser.add_argument("--keep", action="store_true") + parser.add_argument("--inputs-only", action="store_true", + help="build and validate the two-point fixture without launching ILE") + parser.add_argument("--n-eff", type=int, default=20) + parser.add_argument("--n-max", type=int, default=800_000) + parser.add_argument("--n-chunk", type=int, default=20_000) + parser.add_argument("--seed", type=int, default=6907) + parser.add_argument("--distance-mpc", type=float, default=400.0, + help="injected distance and center of the boxed distance prior") + parser.add_argument("--sky-half-width", type=float, default=0.1, + help="RA and declination half-width in radians") + parser.add_argument("--orientation-half-width", type=float, default=0.2, + help="inclination and polarization half-width in radians") + parser.add_argument("--distance-fraction-half-width", type=float, default=0.25, + help="fractional half-width of the distance prior around truth") + opts = parser.parse_args() + code_dir = Path(__file__).resolve().parents[2] + made_temp = opts.work_dir is None + work = opts.work_dir.resolve() if opts.work_dir else Path(tempfile.mkdtemp(prefix="rift-gpu-av-")) + work.mkdir(parents=True, exist_ok=True) + + # Reuse the maintained self-contained fixture builder at a genuinely short BBH grid. + generator = load_generator(code_dir) + generator.FMIN, generator.FMAX = 30.0, 512.0 + # The maintained fixture builder reserves two seconds at each frame edge. + generator.SEGLEN, generator.SRATE = 8.0, 1024.0 + # Avoid landing exactly on the helper's integer, strictly-open data boundary. + generator.EVENT_TIME = 1_000_000_000.25 + generator.M1, generator.M2, generator.DIST_MPC = 30.0, 25.0, opts.distance_mpc + generator.QMAX = 1 + generator.main(str(work)) + + import RIFT.lalsimutils as lsu + case = json.loads((work / "case.json").read_text()) + grid = work / case["grid"] + # The fixture's grid is deliberately arbitrary. Replace it with the injected + # intrinsic point and one nearby point so this is an execution/parity integration. + first = generator.base_params("H1") + second = first.manual_copy() + second.m1 *= 1.0002 + lsu.ChooseWaveformParams_array_to_xml([first, second], str(grid)) + if opts.inputs_only: + check = lsu.xml_to_ChooseWaveformParams_array(str(grid)) + if len(check) != 2: + raise RuntimeError("two-point grid rewrite failed") + print(json.dumps(dict(result="PASS", scope="inputs-only", intrinsic_points=2, + work_dir=str(work), case=case), indent=2, sort_keys=True)) + return + + args = list(case["ile_common"]) + for name, value in (("--n-eff", opts.n_eff), ("--n-max", opts.n_max), + ("--n-chunk", opts.n_chunk), ("--seed", opts.seed), + ("--n-events-to-analyze", 2), ("--event", 0), + ("--output-file", "short_gpu_av")): + set_option(args, name, value) + # AV/log weights and a bounded fairdraw are load-bearing test contracts. + set_option(args, "--sampler-method", "AV") + for flag in ("--internal-use-lnL", "--fairdraw-extrinsic-output", "--save-samples", + "--vectorized", "--gpu", "--force-gpu-only", "--force-xpy", "--force-adapt-all", + "--rotation-slow", "--freqresponse", "--internal-hard-fail-on-error"): + set_option(args, flag) + set_option(args, "--fairdraw-extrinsic-output-n-max", 200) + set_option(args, "--inv-spec-trunc-time", 0) + set_option(args, "--interpolate-time", "cubic") + # Driver defaults differ (classic 100 Hz, JAX 30 Hz); pin the convention. + set_option(args, "--reference-freq", 100.0) + set_option(args, "--rotation-p-max", 1) + set_option(args, "--freqresponse-qmax", 1) + set_option(args, "--freqresponse-arm-length", "H1=40000,L1=40000") + + # This is deliberately a boxed execution/integration smoke test. A broad + # all-sky prior at high SNR tests AV convergence rather than GPU precompute + # integration and can exhaust n-max with a one-point live volume. + if opts.distance_mpc <= 0 or not (0 < opts.distance_fraction_half_width < 1): + raise ValueError("distance and fractional distance half-width must define a positive box") + if opts.sky_half_width <= 0 or opts.orientation_half_width <= 0: + raise ValueError("angular box half-widths must be positive") + boxes = { + "right_ascension": [generator.RA - opts.sky_half_width, + generator.RA + opts.sky_half_width], + "declination": [generator.DEC - opts.sky_half_width, + generator.DEC + opts.sky_half_width], + "inclination": [generator.INCL - opts.orientation_half_width, + generator.INCL + opts.orientation_half_width], + "psi": [generator.PSI - opts.orientation_half_width, + generator.PSI + opts.orientation_half_width], + "distance_mpc": [opts.distance_mpc * (1 - opts.distance_fraction_half_width), + opts.distance_mpc * (1 + opts.distance_fraction_half_width)], + } + for flag, key in (("--limit-right-ascension", "right_ascension"), + ("--limit-declination", "declination"), + ("--limit-inclination", "inclination"), + ("--limit-psi", "psi")): + set_option(args, flag, "%.9g,%.9g" % tuple(boxes[key])) + set_option(args, "--d-min", boxes["distance_mpc"][0]) + set_option(args, "--d-max", boxes["distance_mpc"][1]) + + env = os.environ.copy() + # Keep the result schema deterministic even if the submission environment is + # also used by a hyperpipeline job. The status sidecar below is canonical, + # but checking the legacy row as well catches a partial/truncated export. + env.pop("RIFT_HYPERPIPELINE_FORMAT", None) + env["RIFT_GPU_PRECOMPUTE"] = "1" + env["PYTHONPATH"] = str(code_dir) + os.pathsep + env.get("PYTHONPATH", "") + command = [env.get("PYTHON", "python"), + str(code_dir / "bin/integrate_likelihood_extrinsic_batchmode")] + args + log_path = work / "short_gpu_av.log" + started = time.perf_counter() + with log_path.open("w") as log: + completed = subprocess.run(command, cwd=work, env=env, stdout=log, + stderr=subprocess.STDOUT, check=False) + host_runtime = time.perf_counter() - started + if completed.returncode: + raise RuntimeError("ILE exited %d; inspect %s" % (completed.returncode, log_path)) + + # Batchmode writes one .dat, status sidecar, and XML per intrinsic index. + # Name them exactly: a broad glob can silently accept stale products or future + # sibling exports (for example calibration output) in a reused work directory. + rows = [] + for index in range(2): + dat = work / ("short_gpu_av_%d_.dat" % index) + status_path = work / ("short_gpu_av_%d_integrator_status.json" % index) + if not dat.is_file() or not status_path.is_file(): + raise RuntimeError( + "missing event %d result/status sidecar; inspect %s" % (index, log_path)) + data_lines = [line for line in dat.read_text().splitlines() + if line.strip() and not line.lstrip().startswith("#")] + if len(data_lines) != 1: + raise RuntimeError("expected one result row in %s, found %d" % + (dat, len(data_lines))) + fields = data_lines[0].split() + if len(fields) < 4: + raise RuntimeError("truncated ILE result in %s" % dat) + dat_lnl, dat_sigma, dat_ntotal, dat_neff = map(float, fields[-4:]) + status = json.loads(status_path.read_text()) + try: + lnl = float(status["lnL"]) + sigma = float(status["sigma_lnL"]) + ntotal = int(status["ntotal"]) + neff = float(status["neff"]) + except (KeyError, TypeError, ValueError) as exc: + raise RuntimeError("invalid status sidecar %s: %s" % (status_path, exc)) + if status.get("indx_event") != index or status.get("sampler_method") != "AV": + raise RuntimeError("wrong event/sampler recorded in %s" % status_path) + if status.get("collapsed") is not False: + raise RuntimeError("AV integration collapsed according to %s" % status_path) + if not np.all(np.isfinite([lnl, sigma, ntotal, neff])): + raise RuntimeError("non-finite ILE result in %s" % status_path) + if neff < opts.n_eff: + raise RuntimeError( + "INCONCLUSIVE: event %d stopped at neff %.3f below target %d" % + (index, neff, opts.n_eff)) + # With hyperpipeline format explicitly disabled, the final four legacy + # columns must agree with the machine-readable sidecar. + if not np.allclose([dat_lnl, dat_sigma, dat_ntotal, dat_neff], + [lnl, sigma, ntotal, neff], rtol=2e-6, atol=1e-8): + raise RuntimeError(".dat/status disagreement for event %d" % index) + rows.append(dict(index=index, path=str(dat), status_path=str(status_path), + lnL=lnl, sigma_lnL=sigma, ntotal=ntotal, neff=neff)) + xmls = [work / ("short_gpu_av_%d_.xml.gz" % index) for index in range(2)] + missing_xmls = [str(path) for path in xmls if not path.is_file()] + if missing_xmls: + raise RuntimeError("save-samples missing XML(s) %s; inspect %s" % + (missing_xmls, log_path)) + xml_rows = {} + for xml in xmls: + saved = lsu.xml_to_ChooseWaveformParams_array(str(xml)) + xml_rows[xml.name] = len(saved) + if len(saved) > 200: + raise RuntimeError("fairdraw cap violated: %s has %d rows" % (xml, len(saved))) + if xml.stat().st_size > 10 * 1024 * 1024: + raise RuntimeError("bounded fairdraw XML unexpectedly exceeds 10 MiB: %s" % xml) + print(json.dumps(dict( + result="PASS", host_runtime_s=host_runtime, timing_scope="inside worker; no queue/container transfer", + sampler="AV", internal_log_weights=True, n_eff_target=opts.n_eff, + n_max=opts.n_max, n_chunk=opts.n_chunk, fairdraw_max=200, + injected_distance_mpc=opts.distance_mpc, prior_boxes=boxes, + time_interpolation="cubic", + intrinsic_points=2, + validation_scope="boxed GPU-precompute plus AV/ILE execution; not posterior recovery", + rows=rows, xml_rows=xml_rows, xml_bytes={p.name: p.stat().st_size for p in xmls}, + work_dir=str(work), log=str(log_path), + ), indent=2, sort_keys=True)) + if made_temp and not opts.keep: + print("Temporary products retained for inspection at %s" % work) + + +if __name__ == "__main__": + main() diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_jax_av_ile.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_jax_av_ile.py new file mode 100644 index 000000000..5ee96e696 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/run_short_jax_av_ile.py @@ -0,0 +1,285 @@ +#!/usr/bin/env python3 +"""Run a short two-intrinsic ILE-JAX AV smoke test with GPU precompute. + +This intentionally uses a short BBH signal. Timing begins at the Python worker +invocation, so queueing and container transfer are excluded. The legacy LAL +waveform generator remains the default; GPU precompute uploads both conditioned +mode banks and reuses the detector state for the second intrinsic point. +""" + +import argparse +import importlib.util +import json +import os +from pathlib import Path +import shlex +import subprocess +import tempfile +import time + +import numpy as np + + +def set_option(args, name, value=None): + """Replace a long option in an optparse-style argument vector.""" + while name in args: + i = args.index(name) + del args[i:i + (1 if value is None else 2)] + args.append(name) + if value is not None: + args.append(str(value)) + + +def drop_option(args, name, takes_value=False): + while name in args: + i = args.index(name) + del args[i:i + (2 if takes_value else 1)] + + +def load_generator(code_dir): + path = code_dir / "demo/rift/slowrot_gpu_validate/make_e2e_inputs.py" + spec = importlib.util.spec_from_file_location("short_jax_gpu_e2e_inputs", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + + # The fixture predates the igwn_ligolw namespace used by current containers. + def write_psd_xml(det, delta_f, fmax, path): + import lal + import lal.series + import lalsimulation as lalsim + from igwn_ligolw import utils as ligolw_utils + count = int(fmax / delta_f) + 1 + series = lal.CreateREAL8FrequencySeries( + det, lal.LIGOTimeGPS(0), 0.0, delta_f, lal.SecondUnit, count) + frequency = np.arange(count) * delta_f + values = np.array([ + lalsim.SimNoisePSDaLIGOZeroDetHighPower(max(float(f), 1.0)) + for f in frequency]) + values[~np.isfinite(values)] = 0.0 + series.data.data[:] = values + document = lal.series.make_psd_xmldoc({det: series}) + ligolw_utils.write_filename(document, path, compress="gz") + + module.write_psd_xml = write_psd_xml + return module + + +def result_rows(path): + lines = [line for line in path.read_text().splitlines() + if line.strip() and not line.lstrip().startswith("#")] + if len(lines) != 1: + raise RuntimeError("expected one result row in %s, found %d" % + (path, len(lines))) + fields = lines[0].split() + if len(fields) != 13: + raise RuntimeError("expected 13 JAX result columns in %s, found %d" % + (path, len(fields))) + values = np.asarray([float(value) for value in fields]) + if not np.all(np.isfinite(values)): + raise RuntimeError("non-finite JAX result in %s" % path) + return values + + +def validate_fairdraw_bounds(path, boxes): + """A finite evidence and ESS do not certify that AV respected its box.""" + header = path.read_text().splitlines()[0].lstrip("# ").split() + values = np.atleast_2d(np.loadtxt(path)) + if values.shape[1] != len(header) or not np.all(np.isfinite(values)): + raise RuntimeError("invalid fairdraw columns or non-finite values: %s" % path) + for column, key in (("right_ascension", "right_ascension"), + ("declination", "declination"), + ("inclination", "inclination"), ("psi", "psi"), + ("distance", "distance_mpc")): + if column not in header: + raise RuntimeError("missing fairdraw column %s" % column) + x = values[:, header.index(column)] + lo, hi = boxes[key] + if np.any((x < lo) | (x > hi)): + raise RuntimeError("fairdraw %s outside requested [%g, %g]: %s" % + (column, lo, hi, path)) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--work-dir", type=Path) + parser.add_argument("--keep", action="store_true") + parser.add_argument("--inputs-only", action="store_true", + help="build the fixture and print the exact worker command") + parser.add_argument("--n-eff", type=int, default=20) + parser.add_argument("--n-max", type=int, default=800_000) + parser.add_argument("--n-chunk", type=int, default=8_000) + parser.add_argument("--jax-av-eval-chunk", type=int, default=8_000) + parser.add_argument("--seed", type=int, default=6907) + parser.add_argument("--distance-mpc", type=float, default=400.0) + parser.add_argument("--sky-half-width", type=float, default=0.1) + parser.add_argument("--orientation-half-width", type=float, default=0.2) + parser.add_argument("--distance-fraction-half-width", type=float, default=0.25) + opts = parser.parse_args() + + if opts.n_eff <= 0 or opts.n_max <= 0 or opts.n_chunk <= 0: + raise ValueError("n-eff, n-max, and n-chunk must be positive") + if opts.jax_av_eval_chunk <= 0: + raise ValueError("jax-av-eval-chunk must be positive") + if opts.distance_mpc <= 0 or not (0 < opts.distance_fraction_half_width < 1): + raise ValueError("distance and its fractional half-width must define a positive box") + if opts.sky_half_width <= 0 or opts.orientation_half_width <= 0: + raise ValueError("angular box half-widths must be positive") + + code_dir = Path(__file__).resolve().parents[2] + made_temp = opts.work_dir is None + work = (opts.work_dir.resolve() if opts.work_dir else + Path(tempfile.mkdtemp(prefix="rift-gpu-jax-av-"))) + work.mkdir(parents=True, exist_ok=True) + + generator = load_generator(code_dir) + generator.FMIN, generator.FMAX = 30.0, 512.0 + generator.SEGLEN, generator.SRATE = 8.0, 1024.0 + generator.EVENT_TIME = 1_000_000_000.25 + generator.M1, generator.M2, generator.DIST_MPC = 30.0, 25.0, opts.distance_mpc + generator.QMAX = 1 + generator.main(str(work)) + + import RIFT.lalsimutils as lsu + case = json.loads((work / "case.json").read_text()) + grid = work / case["grid"] + first = generator.base_params("H1") + second = first.manual_copy() + second.m1 *= 1.0002 + lsu.ChooseWaveformParams_array_to_xml([first, second], str(grid)) + if len(lsu.xml_to_ChooseWaveformParams_array(str(grid))) != 2: + raise RuntimeError("two-point grid rewrite failed") + + args = list(case["ile_common"]) + # These classic accelerator selectors must not leak into the JAX smoke test: + # the JAX backend and GPU-precompute route are selected explicitly below. + for flag in ("--gpu", "--force-gpu-only", "--force-xpy", + "--force-adapt-all", "--internal-hard-fail-on-error", + "--internal-use-lnL"): + drop_option(args, flag) + drop_option(args, "--inv-spec-trunc-time", takes_value=True) + for name, value in (("--mode", "laplace-is"), + ("--sampler-method", "AV"), + ("--n-eff", opts.n_eff), ("--n-max", opts.n_max), + ("--n-chunk", opts.n_chunk), + ("--jax-av-eval-chunk", opts.jax_av_eval_chunk), + ("--seed", opts.seed), ("--n-events-to-analyze", 2), + ("--event", 0), ("--output-file", "short_jax_gpu_av"), + ("--interp", "cubic"), ("--rotation-p-max", 1), + ("--freqresponse-qmax", 1), + ("--freqresponse-arm-length", "H1=40000,L1=40000")): + set_option(args, name, value) + for flag in ("--save-samples", "--fairdraw-extrinsic-output", + "--vectorized", "--rotation-slow", "--freqresponse"): + set_option(args, flag) + set_option(args, "--fairdraw-extrinsic-output-n-max", 200) + # Match classic explicitly; this driver's default is 30 Hz, not 100 Hz. + set_option(args, "--reference-freq", 100.0) + + boxes = { + "right_ascension": [generator.RA - opts.sky_half_width, + generator.RA + opts.sky_half_width], + "declination": [generator.DEC - opts.sky_half_width, + generator.DEC + opts.sky_half_width], + "inclination": [generator.INCL - opts.orientation_half_width, + generator.INCL + opts.orientation_half_width], + "psi": [generator.PSI - opts.orientation_half_width, + generator.PSI + opts.orientation_half_width], + "distance_mpc": [opts.distance_mpc * (1 - opts.distance_fraction_half_width), + opts.distance_mpc * (1 + opts.distance_fraction_half_width)], + } + for flag, key in (("--limit-right-ascension", "right_ascension"), + ("--limit-declination", "declination"), + ("--limit-inclination", "inclination"), + ("--limit-psi", "psi")): + set_option(args, flag, "%.9g,%.9g" % tuple(boxes[key])) + set_option(args, "--d-min", boxes["distance_mpc"][0]) + set_option(args, "--d-max", boxes["distance_mpc"][1]) + + env = os.environ.copy() + env["PYTHONUNBUFFERED"] = "1" + env.pop("RIFT_HYPERPIPELINE_FORMAT", None) + env["RIFT_GPU_PRECOMPUTE"] = "1" + # This checkpoint deliberately exercises the production-compatible route: + # existing conditioned LAL modes followed by one host-to-device transfer. + env["RIFT_GPU_WAVEFORM"] = "lal" + # Must be present before the worker imports JAX. Setting this in the + # handoff itself is too late because the driver's cache probe initializes + # the JAX backend before waveform/precompute construction. + env.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") + env["PYTHONPATH"] = str(code_dir) + os.pathsep + env.get("PYTHONPATH", "") + command = [env.get("PYTHON", "python"), + str(code_dir / "bin/integrate_likelihood_extrinsic_jax")] + args + + if opts.inputs_only: + print(json.dumps(dict( + result="PASS", scope="inputs-only", intrinsic_points=2, + work_dir=str(work), command=shlex.join(command), + environment={"RIFT_GPU_PRECOMPUTE": "1", "RIFT_GPU_WAVEFORM": "lal", + "XLA_PYTHON_CLIENT_PREALLOCATE": + env["XLA_PYTHON_CLIENT_PREALLOCATE"]}, + sampler="AV", mode="laplace-is", internal_log_weights=True, + fairdraw_max=200, + compatibility_limits=[ + "short boxed BBH execution smoke, not posterior recovery", + "legacy conditioned LAL waveform modes plus H2D; native JAX waveform disabled", + "JAX AV keeps importance weights in log space; classic --internal-use-lnL is not used", + "sample products are equal-weight *_samples.dat files, not classic ILE XML", + ]), indent=2, sort_keys=True)) + return + + log_path = work / "short_jax_gpu_av.log" + started = time.perf_counter() + with log_path.open("w") as log: + completed = subprocess.run(command, cwd=work, env=env, stdout=log, + stderr=subprocess.STDOUT, check=False) + host_runtime = time.perf_counter() - started + if completed.returncode: + raise RuntimeError("ILE-JAX exited %d; inspect %s" % + (completed.returncode, log_path)) + + rows = [] + sample_rows = {} + for index in range(2): + dat = work / ("short_jax_gpu_av_%d_.dat" % index) + samples = work / ("short_jax_gpu_av_%d_samples.dat" % index) + if not dat.is_file() or not samples.is_file(): + raise RuntimeError("missing event %d result/fairdraw; inspect %s" % + (index, log_path)) + values = result_rows(dat) + lnl, sigma, ntotal, neff = values[-4:] + if neff < opts.n_eff: + raise RuntimeError("INCONCLUSIVE: event %d stopped at neff %.3f below target %d" % + (index, neff, opts.n_eff)) + if ntotal > opts.n_max: + raise RuntimeError("event %d exceeded n-max: %.0f > %d" % + (index, ntotal, opts.n_max)) + lines = samples.read_text().splitlines() + payload = [line for line in lines + if line.strip() and not line.lstrip().startswith("#")] + if not payload or len(payload) > 200: + raise RuntimeError("event %d fairdraw row count %d is outside [1, 200]" % + (index, len(payload))) + if not any("fairdraw:" in line for line in lines if line.startswith("#")): + raise RuntimeError("event %d fairdraw lacks provenance header" % index) + validate_fairdraw_bounds(samples, boxes) + sample_rows[samples.name] = len(payload) + rows.append(dict(index=index, path=str(dat), samples=str(samples), + lnL=float(lnl), sigma_lnL=float(sigma), + ntotal=int(ntotal), neff=float(neff))) + + print(json.dumps(dict( + result="PASS", host_runtime_s=host_runtime, + timing_scope="inside worker; no queue/container transfer", + sampler="AV", mode="laplace-is", internal_log_weights=True, + n_eff_target=opts.n_eff, n_max=opts.n_max, n_chunk=opts.n_chunk, + jax_av_eval_chunk=opts.jax_av_eval_chunk, fairdraw_max=200, + intrinsic_points=2, prior_boxes=boxes, rows=rows, + sample_rows=sample_rows, work_dir=str(work), log=str(log_path), + validation_scope="boxed legacy-waveform H2D plus ILE-JAX AV execution; not posterior recovery", + ), indent=2, sort_keys=True)) + if made_temp and not opts.keep: + print("Temporary products retained for inspection at %s" % work) + + +if __name__ == "__main__": + main() diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_array_contract.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_array_contract.py new file mode 100644 index 000000000..2df447ca1 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_array_contract.py @@ -0,0 +1,258 @@ +"""Adversarial array-level tests for RIFT.likelihood.gpu_precompute. + +These inputs are intentionally non-Hermitian. Positive-only or rfft-based +implementations must fail several tests here. +""" + +import inspect +from types import SimpleNamespace + +import numpy as np +import pytest + +from RIFT.likelihood.gpu_precompute import ( + GPUPrecomputeContext, build_compound_basis, compound_precompute_arrays, + streamed_v_matrix, +) + +from conftest import to_host +from oracle import direct_log_likelihood, q_oracle, uv_oracle + + +def test_default_context_separates_devices_and_explicit_context_rejects_switch(monkeypatch): + from RIFT.likelihood import gpu_precompute as gpu + selected = [0] + backend = SimpleNamespace( + cuda=SimpleNamespace(runtime=SimpleNamespace(getDevice=lambda: selected[0])), + asarray=lambda value: np.array(value, copy=True)) + monkeypatch.setattr(gpu, "_DEFAULT_CONTEXTS", {}) + zero = gpu.default_context(backend) + first = zero.array("data", np.arange(4.)) + selected[0] = 1 + one = gpu.default_context(backend) + assert one is not zero + second = one.array("data", np.arange(4.)) + assert second is not first + with pytest.raises(RuntimeError, match="belongs to CUDA device 0"): + zero.array("data", np.arange(4.)) + with pytest.raises(RuntimeError, match="belongs to CUDA device 0"): + zero.clear() + selected[0] = 0 + assert gpu.default_context(backend) is zero + assert zero.array("data", np.arange(4.)) is first + + +def _case(seed=2917, a=3, m=2, n=32): + rng = np.random.default_rng(seed) + basis = rng.normal(size=(a, m, n)) + 1j * rng.normal(size=(a, m, n)) + # Deliberately independent rather than conj(basis); V must use this argument. + basis_conj = rng.normal(size=(a, m, n)) + 1j * rng.normal(size=(a, m, n)) + data = rng.normal(size=n) + 1j * rng.normal(size=n) + weights = rng.uniform(0.05, 2.0, size=n) + # Asymmetric zeros ensure the implementation neither assumes Hermitian weights nor + # derives support from one half of the spectrum. + weights[[0, 3, n // 2 + 2, n - 1]] = 0.0 + return basis, basis_conj, data, weights + + +def _call(backend, case, *, delta_f=0.125, delta_t=None, n_shift=5, + n_window=13, return_device=False, **kwargs): + basis, basis_conj, data, weights = case + if delta_t is None: + delta_t = 1.0 / (basis.shape[-1] * delta_f) + return compound_precompute_arrays( + backend.asarray(basis), backend.asarray(basis_conj), backend.asarray(data), + backend.asarray(weights), delta_f, delta_t, n_shift, n_window, + backend=backend, return_device=return_device, **kwargs + ) + + +def _unpack(result): + if isinstance(result, dict): + return result["Q"], result["U"], result["V"] + assert len(result) >= 3 + return result[:3] + + +def test_nonhermitian_q_uv_match_independent_oracle(backend): + case = _case() + q, u, v = map(to_host, _unpack(_call(backend, case, return_device=True))) + q0 = q_oracle(case[0], case[2], case[3], 0.125, 5, 13) + u0, v0 = uv_oracle(case[0], case[1], case[3], 0.125) + np.testing.assert_allclose(q, q0, rtol=3e-12, atol=3e-12) + np.testing.assert_allclose(u, u0, rtol=3e-12, atol=3e-12) + np.testing.assert_allclose(v, v0, rtol=3e-12, atol=3e-12) + + +@pytest.mark.parametrize("n_shift", [0, 1, 7, -3, 31, 35]) +def test_roll_cut_and_wraparound(backend, n_shift): + case = _case(n=32) + q, _, _ = _unpack(_call(backend, case, n_shift=n_shift, n_window=11)) + expected = q_oracle(case[0], case[2], case[3], 0.125, n_shift, 11) + np.testing.assert_allclose(to_host(q), expected, rtol=3e-12, atol=3e-12) + + +def test_centered_negative_frequency_bins_are_live(backend): + n = 16 + basis = np.zeros((1, 1, n), complex) + data = np.zeros(n, complex) + weights = np.zeros(n) + # A single negative-frequency bin in centered LAL ordering. A positive-only/rfft + # implementation returns zero; a missing ifftshift gives the wrong alternating phase. + k = n // 2 + 3 + basis[0, 0, k] = 1.25 - 0.4j + data[k] = -0.2 + 2.0j + weights[k] = 0.7 + case = basis, basis * (0.3 + 0.8j), data, weights + q, u, v = map(to_host, _unpack(_call(backend, case, n_shift=0, n_window=n))) + np.testing.assert_allclose(q, q_oracle(basis, data, weights, 0.125, 0, n), rtol=2e-13, atol=2e-13) + np.testing.assert_allclose(u, uv_oracle(basis, case[1], weights, 0.125)[0], rtol=2e-13, atol=2e-13) + assert np.max(np.abs(q)) > 0 and np.max(np.abs(v)) > 0 + + +def test_fractional_sample_phase_and_complex_phase_are_preserved(backend): + case = list(_case(n=64)) + dt = 1.0 / (64 * 0.125) + tau = 0.37 * dt + # Physical frequencies matching the centered LAL layout. + f = np.arange(-32, 32) * 0.125 + phase = np.exp(-2j * np.pi * f * tau) + case[0] = case[0] * phase + case[1] = case[1] * np.conj(phase) + case = tuple(case) + q, u, v = map(to_host, _unpack(_call(backend, case, delta_t=dt, n_shift=-4, n_window=29))) + q0 = q_oracle(case[0], case[2], case[3], 0.125, -4, 29) + u0, v0 = uv_oracle(case[0], case[1], case[3], 0.125) + np.testing.assert_allclose(q, q0, rtol=4e-12, atol=4e-12) + np.testing.assert_allclose(u, u0, rtol=4e-12, atol=4e-12) + np.testing.assert_allclose(v, v0, rtol=4e-12, atol=4e-12) + + +def test_v_uses_supplied_conjugate_family_with_correct_orientation(backend): + case = _case(seed=8) + q, u, v = map(to_host, _unpack(_call(backend, case))) + _, v0 = uv_oracle(case[0], case[1], case[3], 0.125) + wrong = uv_oracle(case[0], np.conj(case[0]), case[3], 0.125)[1] + np.testing.assert_allclose(v, v0, rtol=3e-12, atol=3e-12) + assert np.max(np.abs(v - wrong)) > 1e-3 + np.testing.assert_allclose(u, np.swapaxes(np.swapaxes(u.conj(), 0, 1), 2, 3), rtol=3e-12, atol=3e-12) + + +def test_memory_bounded_streamed_v_matches_dense_oracle(backend): + rng = np.random.default_rng(884) + n, m, b = 40, 2, 3 + df, dt = 0.2, 1.0 / (40 * 0.2) + base = rng.normal(size=(m, n)) + 1j * rng.normal(size=(m, n)) + base_c = rng.normal(size=(m, n)) + 1j * rng.normal(size=(m, n)) + response = rng.normal(size=(b, n)) + 1j * rng.normal(size=(b, n)) + a_list = [(0, 0, -2), (1, 1, 0), (2, 0, 3), (0, 2, -1), (2, 1, 2)] + weights = rng.uniform(0.0, 1.5, size=n) + primary = build_compound_basis( + backend.asarray(base), backend.asarray(response), a_list, df, dt, -1.75, + f_sidereal=0.031, backend=backend, fft_batch=2) + dense_c = build_compound_basis( + backend.asarray(base_c), backend.asarray(response), a_list, df, dt, -1.75, + f_sidereal=0.031, backend=backend, fft_batch=3) + expected = uv_oracle(to_host(primary), to_host(dense_c), weights, df)[1] + # Awkward block and chunk sizes force all boundary branches. + got = streamed_v_matrix( + backend.asarray(base_c), backend.asarray(response), a_list, primary, + backend.asarray(weights), df, dt, -1.75, f_sidereal=0.031, + backend=backend, a_block=2, fft_batch=1, frequency_chunk=7, + return_device=True) + np.testing.assert_allclose(to_host(got), expected, rtol=4e-12, atol=4e-12) + + +def test_downstream_full_gaussian_likelihood_parity(backend): + case = _case(seed=192, a=4, m=3, n=48) + q, u, _ = map(to_host, _unpack(_call(backend, case, n_shift=0, n_window=1))) + rng = np.random.default_rng(704) + coeff = rng.normal(size=(4, 3)) + 1j * rng.normal(size=(4, 3)) + # t=0 is the first LAL reverse-FFT sample. Q(t=0)=2 df sum conj(chi)dW. + contracted = float(np.real(np.vdot(coeff, q[..., 0]) + - 0.5 * np.einsum("ai,abij,bj->", np.conj(coeff), u, coeff))) + direct = direct_log_likelihood(case[0], case[2], case[3], 0.125, coeff) + assert abs(contracted - direct) < 2e-10 * max(1.0, abs(direct)) + + +def test_detector_specific_weights_cannot_alias(backend): + case = list(_case(seed=99)) + weights_h = case[3].copy() + weights_l = case[3][::-1].copy() * np.linspace(0.2, 1.8, case[3].size) + case_h = tuple(case[:3] + [weights_h]) + case_l = tuple(case[:3] + [weights_l]) + h = tuple(map(to_host, _unpack(_call(backend, case_h)))) + l = tuple(map(to_host, _unpack(_call(backend, case_l)))) + l0 = (q_oracle(case_l[0], case_l[2], weights_l, 0.125, 5, 13),) + uv_oracle(case_l[0], case_l[1], weights_l, 0.125) + for got, expected in zip(l, l0): + np.testing.assert_allclose(got, expected, rtol=3e-12, atol=3e-12) + assert any(np.max(np.abs(x - y)) > 1e-4 for x, y in zip(h, l)) + + +def test_reused_context_invalidates_data_weights_and_intrinsic_basis(backend): + # One worker evaluates multiple intrinsic points. Reusing plans/storage is welcome; + # reusing any value that depends on data, PSD, or template is a correctness defect. + context = GPUPrecomputeContext(backend) + c1 = _case(seed=1) + c2 = _case(seed=2) + d1 = context.array(("H1", "data"), c1[2]) + w1 = context.array(("H1", "weights"), c1[3]) + d2 = context.array(("H1", "data"), c2[2]) + w2 = context.array(("H1", "weights"), c2[3]) + np.testing.assert_array_equal(to_host(d2), c2[2]) + np.testing.assert_array_equal(to_host(w2), c2[3]) + assert np.max(np.abs(to_host(d1) - to_host(d2))) > 1e-4 + assert np.max(np.abs(to_host(w1) - to_host(w2))) > 1e-4 + mutable = c1[2].copy() + old = context.array(("L1", "data"), mutable) + mutable[0] += 100 + 20j + new = context.array(("L1", "data"), mutable) + assert to_host(new)[0] == mutable[0] + assert to_host(old)[0] != to_host(new)[0] + # An exact repeat is a hit; a changed grid shape replaces the role's old allocation. + before = context.stats() + again = context.array(("L1", "data"), mutable) + after_hit = context.stats() + assert after_hit["cache_hits"] == before["cache_hits"] + 1 + assert after_hit["uploads"] == before["uploads"] + np.testing.assert_array_equal(to_host(again), mutable) + shorter = mutable[:-2].copy() + context.array(("L1", "data"), shorter) + after_grid_change = context.stats() + assert after_grid_change["uploads"] == before["uploads"] + 1 + # Three roles remain: H1 data/weights and the replaced L1 data entry only once. + assert after_grid_change["retained_arrays"] == 3 + + +def test_two_intrinsics_in_sequence_match_independent_oracle(backend): + c1, c2 = _case(seed=120), _case(seed=121) + _call(backend, c1) + got = tuple(map(to_host, _unpack(_call(backend, c2)))) + expected = (q_oracle(c2[0], c2[2], c2[3], 0.125, 5, 13),) + \ + uv_oracle(c2[0], c2[1], c2[3], 0.125) + for actual, oracle in zip(got, expected): + np.testing.assert_allclose(actual, oracle, rtol=3e-12, atol=3e-12) + + +def test_shape_and_grid_contracts_fail_fast(backend): + case = _case() + bad = list(case) + bad[1] = bad[1][..., :-1] + with pytest.raises((AssertionError, ValueError)): + _call(backend, tuple(bad)) + with pytest.raises((AssertionError, ValueError)): + _call(backend, case, delta_t=0.12345) + with pytest.raises((AssertionError, ValueError)): + _call(backend, case, n_window=33) + bad = list(case) + bad[3] = bad[3].copy() + bad[3][2] = -1.0 + with pytest.raises((AssertionError, ValueError)): + _call(backend, tuple(bad)) + + +def test_return_device_contract(backend): + q, u, v = _unpack(_call(backend, _case(), return_device=True)) + assert type(q).__module__.split(".")[0] == backend.__name__.split(".")[0] + assert type(u).__module__.split(".")[0] == backend.__name__.split(".")[0] + assert type(v).__module__.split(".")[0] == backend.__name__.split(".")[0] diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_basis_reuse.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_basis_reuse.py new file mode 100644 index 000000000..783a67dc1 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_basis_reuse.py @@ -0,0 +1,34 @@ +"""Verify sidereal-index reuse against the independent per-index recipe.""" +import numpy as np +import pytest + +from RIFT.likelihood import gpu_precompute as gp +from conftest import to_host + + +@pytest.mark.parametrize("batch", [1, 2, 4]) +def test_reverse_fft_reused_preserving_order_and_duplicates(backend, monkeypatch, batch): + rng = np.random.default_rng(9421) + base = backend.asarray(rng.normal(size=(3, 32)) + 1j*rng.normal(size=(3, 32))) + response = backend.asarray(rng.normal(size=(2, 32)) + 1j*rng.normal(size=(2, 32))) + indices = [(1, 0, 2), (0, 1, -1), (1, 0, -2), (1, 0, 2), (0, 1, 1)] + df, dt, epoch, sidereal = .125, .25, -.37, .012 + reverse = gp._lal_reverse + frequency = gp.lal_frequency_axis(32, df, xp=backend) + times = epoch + backend.arange(32)*dt + expected = backend.stack([ + gp._lal_forward(reverse(base*response[b][None, :]* + gp._derivative_weight(frequency, p, backend)[None, :], backend)* + backend.exp(2j*np.pi*n*sidereal*times)[None, :], backend) + for b, p, n in indices]) + calls = [] + + def counted_reverse(values, xp): + calls.append(values.shape) + return reverse(values, xp) + + monkeypatch.setattr(gp, "_lal_reverse", counted_reverse) + got = gp.build_compound_basis(base, response, indices, df, dt, epoch, + f_sidereal=sidereal, backend=backend, fft_batch=batch) + np.testing.assert_allclose(to_host(got), to_host(expected), rtol=3e-12, atol=3e-12) + assert calls == [(3, 32), (3, 32)] # Two unique (b,p), not five sidereal rows. diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_classic_compound_waveform_forwarding.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_classic_compound_waveform_forwarding.py new file mode 100644 index 000000000..7416580d4 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_classic_compound_waveform_forwarding.py @@ -0,0 +1,85 @@ +"""Guard waveform-option parity at the classic compound-driver boundary.""" + +import ast +from collections import Counter +from pathlib import Path +from types import SimpleNamespace + + +DRIVER = Path(__file__).resolve().parents[2] / "bin" / "integrate_likelihood_extrinsic_batchmode" + + +def _dotted_name(node): + if isinstance(node, ast.Name): + return node.id + if isinstance(node, ast.Attribute): + prefix = _dotted_name(node.value) + return prefix + "." + node.attr if prefix else node.attr + return "" + + +def _analyze_event_tree(): + tree = ast.parse(DRIVER.read_text(), filename=str(DRIVER)) + return next(node for node in tree.body + if isinstance(node, ast.FunctionDef) and node.name == "analyze_event") + + +def test_classic_compound_reuses_baseline_waveform_generation_kwargs(): + function = _analyze_event_tree() + assignments = [ + node for node in ast.walk(function) + if isinstance(node, ast.Assign) + and any(isinstance(target, ast.Name) + and target.id == "waveform_generation_kwargs" + for target in node.targets) + ] + assert len(assignments) == 1 + + option_fields = { + "use_gwsignal": "use_gwsignal", + "use_gwsignal_approx": "approximant", + "use_external_EOB": "use_external_EOB", + "nr_lookup": "nr_lookup", + "nr_lookup_valid_groups": "nr_lookup_group", + "perturbative_extraction": "nr_perturbative_extraction", + "perturbative_extraction_full": "nr_perturbative_extraction_full", + "use_provided_strain": "nr_use_provided_strain", + "hybrid_use": "nr_hybrid_use", + "hybrid_method": "nr_hybrid_method", + "ROM_group": "rom_group", + "ROM_param": "rom_param", + "ROM_use_basis": "rom_use_basis", + "ROM_limit_basis_size": "rom_limit_basis_size_to", + "no_memory": "no_memory", + "force_22_mode": "force_hyperbolic_22", + } + sentinels = {field: object() for field in option_fields.values()} + opts = SimpleNamespace(**sentinels) + nr_group, nr_param, nested = object(), object(), object() + expression = ast.Expression(assignments[0].value) + ast.fix_missing_locations(expression) + forwarded = eval( + compile(expression, str(DRIVER), "eval"), + dict(opts=opts, NR_template_group=nr_group, + NR_template_param=nr_param, extra_waveform_kwargs=nested), + ) + + assert forwarded["NR_group"] is nr_group + assert forwarded["NR_param"] is nr_param + assert forwarded["extra_waveform_kwargs"] is nested + for keyword, field in option_fields.items(): + assert forwarded[keyword] is sentinels[field] + + expanded_calls = [] + for call in (node for node in ast.walk(function) if isinstance(node, ast.Call)): + if any(keyword.arg is None + and isinstance(keyword.value, ast.Name) + and keyword.value.id == "waveform_generation_kwargs" + for keyword in call.keywords): + expanded_calls.append(_dotted_name(call.func)) + + assert Counter(expanded_calls) == Counter({ + "factored_likelihood.PrecomputeLikelihoodTerms": 1, + "PrecomputeLikelihoodTermsRotatingFreqResponseGPU": 1, + "factored_likelihood_rotating_freqresponse.PrecomputeLikelihoodTermsRotatingFreqResponse": 1, + }) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_cross_driver_parity.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_cross_driver_parity.py new file mode 100644 index 000000000..0041daee7 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_cross_driver_parity.py @@ -0,0 +1,70 @@ +"""Pin classic/JAX compound-likelihood algebra at production response orders.""" +from types import SimpleNamespace + +import jax +import jax.numpy as jnp +import lal +import numpy as np + +from RIFT.likelihood import factored_likelihood_rotating_freqresponse as compound +from RIFT.likelihood import slowrot_freqresponse +from RIFT.likelihood.jax_ile.banded import build_rotating_freqresponse_data +from RIFT.likelihood.jax_ile.core import fused_log_likelihood + +jax.config.update("jax_enable_x64", True) + + +def test_p1_q1_cubic_fixed_point_matches_classic_likelihood(): + """Same Q/U/V, geometry and parameters must agree before any sampler runs.""" + rng = np.random.default_rng(12291) + qmax, pmax = 1, 1 + modes = [(2, 2), (2, -2)] + a_list = compound.compound_index_set(qmax, pmax) + A, K, N = len(a_list), len(modes), 256 + delta_t = 1.0/1024 + tref = 1000000000.0 + epoch = tref - 128*delta_t + det = "H1" + + q = rng.normal(size=(A, K, N)) + 1j*rng.normal(size=(A, K, N)) + u = rng.normal(size=(A, A, K, K)) + 1j*rng.normal(size=(A, A,K,K)) + v = rng.normal(size=(A, A, K, K)) + 1j*rng.normal(size=(A,A,K,K)) + meta = dict( + feature="rotation_freqresponse", post_phase_required=True, + event_time_geo=tref, modes=modes, a_list=a_list, + Qmax=qmax, p_max=pmax, f_sidereal=compound.flwr.F_SIDEREAL, + L_arm=4000.0) + lookup = {det: np.asarray(modes, dtype=int)} + rho = {det: {a: q[i] for i, a in enumerate(a_list)}} + U, V, epochs = {det: u}, {det: v}, {det: epoch} + geometry = {det: slowrot_freqresponse.detector_geometry(det, L_arm=4000.0)} + tvals = np.arange(-4, 5)*delta_t + data = build_rotating_freqresponse_data( + meta, lookup, rho, U, V, epochs, delta_t, tvals, geometry) + + ra = np.asarray([0.7, 5.8]) + dec = np.asarray([-0.2, 0.85]) + psi = np.asarray([0.3, 2.7]) + incl = np.asarray([0.6, 2.5]) + phiref = np.asarray([0.4, 5.0]) + dist_mpc = np.asarray([100.0, 900.0]) + params = SimpleNamespace( + phi=ra, theta=dec, psi=psi, incl=incl, phiref=phiref, + dist=dist_mpc*1.0e6*lal.PC_SI, tref=tref, deltaT=delta_t) + + classic_t = compound.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, params, meta, lookup, rho, U, V, epochs, Lmax=2, + array_output=True, time_interp="cubic") + jax_t = np.asarray(fused_log_likelihood( + data, *(jnp.asarray(x) for x in + (ra, dec, psi, incl, phiref, dist_mpc)), + interp="cubic", return_lnLt=True)) + np.testing.assert_allclose(jax_t, classic_t, rtol=2e-13, atol=2e-13) + + classic_marg = compound.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, params, meta, lookup, rho, U, V, epochs, Lmax=2, + array_output=False, time_interp="cubic") + jax_marg = np.asarray(fused_log_likelihood( + data, *(jnp.asarray(x) for x in + (ra, dec, psi, incl, phiref, dist_mpc)), interp="cubic")) + np.testing.assert_allclose(jax_marg, classic_marg, rtol=2e-13, atol=2e-13) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_device_handoff_adversarial.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_device_handoff_adversarial.py new file mode 100644 index 000000000..60475079e --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_device_handoff_adversarial.py @@ -0,0 +1,304 @@ +"""Adversarial residency, ownership, and dispatch tests for GPU-to-JAX handoff.""" + +import gc +import os +from types import SimpleNamespace + +import numpy as np +import pytest + + +def _require_cupy_jax_gpu(): + try: + import cupy as cp + import jax + import jax.numpy as jnp + if cp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("CuPy reports no CUDA device") + if not any(device.platform == "gpu" for device in jax.devices()): + raise RuntimeError("JAX reports no GPU device") + if not bool(jax.config.x64_enabled): + raise RuntimeError("JAX x64 is disabled") + cp.cuda.Stream.null.synchronize() + return cp, jax, jnp + except Exception as exc: + if os.environ.get("RIFT_REQUIRE_GPU_PRECOMPUTE") == "1": + pytest.fail("real CuPy/JAX GPU handoff required: %r" % (exc,)) + pytest.skip("no usable CuPy/JAX GPU handoff: %r" % (exc,)) + + +def _synthetic_bank(seed=20260912): + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + + rng = np.random.default_rng(seed) + modes = [(2, 2), (2, -2)] + a_list = fr.compound_index_set(0, 0) + a_count, mode_count, n_time = len(a_list), len(modes), 128 + q = (rng.normal(size=(a_count, mode_count, n_time)) + + 1j * rng.normal(size=(a_count, mode_count, n_time))) + # A positive Hermitian U and zero V keep the synthetic likelihood finite + # without borrowing the implementation's contraction to make the fixture. + u = np.zeros((a_count, a_count, mode_count, mode_count), complex) + for a in range(a_count): + u[a, a] = np.eye(mode_count) * (0.2 + 0.01 * a) + v = np.zeros_like(u) + dt = 1.0 / 1024.0 + # Center the stored Q buffer on the geocentric event. H1's Earth-center + # delay and the short integration window then remain well inside support. + epoch = {"H1": 1_000_000_000.0 - (n_time // 2) * dt} + meta = dict( + feature="rotation_freqresponse", gpu_precompute=True, + device_resident=True, post_phase_required=True, + event_time_geo=1_000_000_000.0, modes=modes, a_list=a_list, + Qmax=0, p_max=0, f_sidereal=1.160576e-5, + ) + tvals = (np.arange(9) - 4) * dt + return q, u, v, epoch, dt, meta, tvals + + +def _conventional_data(q, u, v, epoch, dt, meta, tvals, geometry): + from RIFT.likelihood.jax_ile.banded import build_rotating_freqresponse_data + + a_list = meta["a_list"] + lookup = {"H1": np.asarray(meta["modes"], dtype=int)} + rho = {"H1": {a: q[i] for i, a in enumerate(a_list)}} + u_dict = {"H1": {(a, ap): u[i, j] + for i, a in enumerate(a_list) + for j, ap in enumerate(a_list)}} + v_dict = {"H1": {(a, ap): v[i, j] + for i, a in enumerate(a_list) + for j, ap in enumerate(a_list)}} + return build_rotating_freqresponse_data( + meta, lookup, rho, u_dict, v_dict, epoch, dt, tvals, geometry) + + +def _likelihood_arguments(jnp): + return [jnp.asarray(value) for value in ( + [1.2, 1.25], [0.3, 0.26], [0.2, 0.7], + [0.7, 1.05], [0.1, 0.55], [100.0, 160.0], + )] + + +def test_cupy_handoff_has_no_bulk_host_copy_and_survives_source_deletion(monkeypatch): + cp, jax, jnp = _require_cupy_jax_gpu() + from RIFT.likelihood import slowrot_freqresponse as sfr + from RIFT.likelihood.gpu_jax_handoff import ( + build_jax_rotating_freqresponse_data_from_device, + ) + from RIFT.likelihood.jax_ile.core import fused_log_likelihood + + q, u, v, epoch, dt, meta, tvals = _synthetic_bank() + geometry = {"H1": sfr.detector_geometry("H1", L_arm=4000.0)} + conventional = _conventional_data(q, u, v, epoch, dt, meta, tvals, geometry) + args = _likelihood_arguments(jnp) + expected = fused_log_likelihood(conventional, *args, interp="nearest") + zero_q = _conventional_data(np.zeros_like(q), u, v, epoch, dt, meta, + tvals, geometry) + zero_q_likelihood = fused_log_likelihood(zero_q, *args, interp="nearest") + jax.block_until_ready(expected) + + q_device, u_device, v_device = cp.asarray(q), cp.asarray(u), cp.asarray(v) + packed = dict( + q={"H1": q_device}, U={"H1": u_device}, V={"H1": v_device}, + epoch=epoch, delta_t=dt, modes=meta["modes"], a_list=meta["a_list"], + ) + + # Any cp.asnumpy call here is a bulk-copy regression. Small geometry and + # index tables originate on the host and do not need this escape hatch. + def forbidden_asnumpy(*unused_args, **unused_kwargs): + raise AssertionError("GPU handoff copied a CuPy array to host") + + monkeypatch.setattr(cp, "asnumpy", forbidden_asnumpy) + direct = build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, geometry, require_gpu=True) + for key in ("Q_bank", "U_bank", "V_bank"): + value = direct.detectors["H1"][key] + assert all(device.platform == "gpu" for device in value.devices()) + assert direct.gpu_handoff["contract_Q_U_V_host_copies"] == 0 + + # DLPack ownership must outlive every producer-side reference. Releasing + # the CuPy pool and churning same-sized blocks makes a borrowed-buffer bug + # deterministic enough to catch without allocating a long waveform bank. + shapes = [q_device.shape, u_device.shape, v_device.shape] + del packed, q_device, u_device, v_device + gc.collect() + cp.get_default_memory_pool().free_all_blocks() + churn = [cp.full(shape, 17.0 + i, dtype=cp.complex128) + for i, shape in enumerate(shapes)] + cp.cuda.Stream.null.synchronize() + del churn + gc.collect() + + got = fused_log_likelihood(direct, *args, interp="nearest") + jax.block_until_ready(got) + np.testing.assert_allclose( + np.asarray(direct.detectors["H1"]["Q_bank"]), + np.transpose(q, (0, 2, 1)), rtol=0.0, atol=0.0) + np.testing.assert_allclose(np.asarray(got), np.asarray(expected), + rtol=3e-12, atol=3e-12) + assert np.max(np.abs(np.asarray(got) - np.asarray(zero_q_likelihood))) > 1e-8 + + +def test_handoff_rejects_numpy_banks_before_jax_evaluation(): + from RIFT.likelihood import slowrot_freqresponse as sfr + from RIFT.likelihood.gpu_jax_handoff import ( + build_jax_rotating_freqresponse_data_from_device, + ) + + q, u, v, epoch, dt, meta, tvals = _synthetic_bank() + geometry = {"H1": sfr.detector_geometry("H1", L_arm=4000.0)} + packed = dict( + q={"H1": q}, U={"H1": u}, V={"H1": v}, epoch=epoch, + delta_t=dt, modes=meta["modes"], a_list=meta["a_list"], + ) + with pytest.raises(TypeError, match="device-resident"): + build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, geometry, require_gpu=False) + + +@pytest.mark.parametrize("interp", ["nearest", "cubic"]) +def test_highlevel_device_precompute_to_classic_likelihood_has_no_bulk_d2h(monkeypatch, interp): + cp, unused_jax, unused_jnp = _require_cupy_jax_gpu() + from test_highlevel_integration import _synthetic_problem + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + from RIFT.likelihood.gpu_precompute import ( + GPUPrecomputeContext, PrecomputeLikelihoodTermsRotatingFreqResponseGPU, + pack_device_precompute, + ) + + lal, fl, event, p, data, psd, modes, modes_c = _synthetic_problem() + monkeypatch.setattr(fl, "internal_hlm_generator", + lambda *args, **kwargs: (modes, modes_c)) + common = dict( + event_time_geo=event, t_window=0.25, P=p, data_dict=data, + psd_dict=psd, Lmax=2, fMax=24.0, Qmax=1, p_max=1, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, + verbose=False, quiet=True, skip_interpolation=True, + ) + monkeypatch.delenv("RIFT_GPU_PRECOMPUTE", raising=False) + cpu = fr.PrecomputeLikelihoodTermsRotatingFreqResponse(**common) + cpu_packed = fr.pack_rotating_freqresponse_arrays( + cpu[4], cpu[3], cpu[1], cpu[2]) + + pvec = SimpleNamespace( + phi=np.array([1.17, 1.21]), theta=np.array([-0.31, -0.28]), + incl=np.array([0.7, 1.0]), phiref=np.array([0.2, 1.1]), + psi=np.array([0.4, 0.9]), + dist=np.array([110.0, 170.0]) * 1e6 * lal.PC_SI, + tref=lal.LIGOTimeGPS(event), deltaT=p.deltaT, + ) + tvals = np.array([-p.deltaT, 0.0, p.deltaT]) + expected = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, cpu[4], *cpu_packed, Lmax=2, array_output=True, + time_interp=interp, xpy=np) + expected_marginal = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, cpu[4], *cpu_packed, Lmax=2, array_output=False, + time_interp=interp, xpy=np) + + real_asnumpy = cp.asnumpy + transfers = [] + + def scalar_only_asnumpy(value, *args, **kwargs): + ndim = int(getattr(value, "ndim", -1)) + transfers.append((ndim, int(getattr(value, "nbytes", -1)))) + if ndim != 0: + raise AssertionError("bulk GPU-to-host transfer during device precompute/packing") + return real_asnumpy(value, *args, **kwargs) + + monkeypatch.setattr(cp, "asnumpy", scalar_only_asnumpy) + packed, meta = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, backend=cp, context=GPUPrecomputeContext(cp), + return_device=True, fft_batch=7, q_row_batch=9, + frequency_chunk=37) + device_packed = pack_device_precompute(packed, meta) + lookup, rho, u_bank, v_bank, unused_epoch = device_packed + assert u_bank["H1"] is packed["U"]["H1"] + assert v_bank["H1"] is packed["V"]["H1"] + for index, a in enumerate(meta["a_list"]): + assert cp.shares_memory(rho["H1"][a], packed["q"]["H1"][index]) + got_device = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, meta, *device_packed, Lmax=2, array_output=True, + time_interp=interp, xpy=cp) + got_marginal = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, meta, *device_packed, Lmax=2, array_output=False, + time_interp=interp, xpy=cp) + cp.cuda.Stream.null.synchronize() + assert all(ndim == 0 for ndim, unused_nbytes in transfers) + np.testing.assert_allclose(real_asnumpy(got_device), expected, + rtol=3e-10, atol=2e-8) + np.testing.assert_allclose(real_asnumpy(got_marginal), expected_marginal, + rtol=3e-10, atol=2e-8) + + +def test_wrapper_gpu_dispatch_bypasses_legacy_pack_and_preserves_options(monkeypatch): + from RIFT.likelihood import gpu_jax_handoff + from RIFT.likelihood import gpu_precompute + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + from RIFT.likelihood.jax_ile import banded, wrapper + + sentinel_data = object() + a_list = [(0, 0, 0)] + modes = [(2, 2)] + packed = dict( + q={"H1": np.zeros((1, 1, 3), complex)}, + U={"H1": np.zeros((1, 1, 1, 1), complex)}, + V={"H1": np.zeros((1, 1, 1, 1), complex)}, + epoch={"H1": 999.8}, delta_t=0.125, modes=modes, a_list=a_list, + ) + meta = dict(feature="rotation_freqresponse", device_resident=True, + modes=modes, a_list=a_list) + calls = {} + + def fake_gpu(*args, **kwargs): + calls["gpu"] = (args, kwargs) + return packed, meta + + def fake_handoff(*args, **kwargs): + calls["handoff"] = (args, kwargs) + return sentinel_data + + monkeypatch.setenv("RIFT_GPU_PRECOMPUTE", "1") + monkeypatch.setattr(gpu_precompute, + "PrecomputeLikelihoodTermsRotatingFreqResponseGPU", fake_gpu) + monkeypatch.setattr(gpu_jax_handoff, + "build_jax_rotating_freqresponse_data_from_device", fake_handoff) + monkeypatch.setattr(fr, "pack_rotating_freqresponse_arrays", + lambda *a, **k: pytest.fail("legacy pack was called")) + monkeypatch.setattr(banded, "build_rotating_freqresponse_data", + lambda *a, **k: pytest.fail("legacy JAX builder was called")) + monkeypatch.setattr(wrapper.factored_likelihood, "marginalization_time_grid", + lambda half, dt, xpy=None: np.array([-dt, 0.0, dt])) + + p = SimpleNamespace(deltaT=0.125) + result, extras = wrapper.build_rotating_freqresponse_data_from_precompute( + p, {"H1": object()}, {"H1": object()}, 1000.0, 0.125, + 2, 256.0, t_window=0.25, Qmax=3, L_arm={"H1": 4000.0}, + p_max=2, analyticPSD_Q=False, inv_spec_trunc_Q=True, T_spec=0.75, + verbose=True, custom_waveform_option="kept", + ) + assert result is sentinel_data + gpu_args, gpu_kwargs = calls["gpu"] + assert gpu_args[:2] == (1000.0, 0.25) + assert gpu_kwargs["return_device"] is True + assert gpu_kwargs["Qmax"] == 3 and gpu_kwargs["p_max"] == 2 + assert gpu_kwargs["inv_spec_trunc_Q"] is True and gpu_kwargs["T_spec"] == 0.75 + assert gpu_kwargs["custom_waveform_option"] == "kept" + handoff_args, handoff_kwargs = calls["handoff"] + assert handoff_args[0] is packed and handoff_args[1] is meta + assert extras["meta"] is meta + assert extras["U_by_aa"] is packed["U"] + assert extras["V_by_aa"] is packed["V"] + + +def test_gpu_order_control_fails_before_allocating_reference_bank(monkeypatch): + from RIFT.likelihood import gpu_precompute + from RIFT.likelihood.jax_ile import wrapper + monkeypatch.setenv("RIFT_GPU_PRECOMPUTE", "1") + monkeypatch.setattr(gpu_precompute, + "PrecomputeLikelihoodTermsRotatingFreqResponseGPU", + lambda *a, **k: pytest.fail("unexpected reference precompute")) + with pytest.raises(NotImplementedError, match="response-order selection"): + wrapper.build_rotating_freqresponse_data_from_precompute( + SimpleNamespace(deltaT=0.125), {"H1": object()}, {"H1": object()}, + 1000.0, 0.125, 2, 256.0, + order_control={"choose_p": True, "p_reference": 3}) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_dispatch_contract.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_dispatch_contract.py new file mode 100644 index 000000000..b3a11a0f5 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_dispatch_contract.py @@ -0,0 +1,55 @@ +"""Opt-in dispatch and native-provider rejection before expensive work.""" +from types import SimpleNamespace + +import numpy as np +import pytest + + +def test_opt_in_dispatch_preserves_arguments(monkeypatch): + from RIFT.likelihood import gpu_precompute as gpu + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + seen = {} + + def replacement(*args, **kwargs): + seen.update(args=args, kwargs=kwargs) + return 'replacement sentinel' + + monkeypatch.setenv('RIFT_GPU_PRECOMPUTE','1') + monkeypatch.setenv('RIFT_GPU_WAVEFORM','lal') + monkeypatch.setattr(gpu,'PrecomputeLikelihoodTermsRotatingFreqResponseGPU',replacement) + result = fr.PrecomputeLikelihoodTermsRotatingFreqResponse( + 100.,0.2,None,{}, {},2,512.,Qmax=1,p_max=1,skip_interpolation=True) + assert result == 'replacement sentinel' + assert seen['args'] == (100.,0.2,None,{}, {},2,512.) + assert seen['kwargs']['Qmax'] == 1 and seen['kwargs']['p_max'] == 1 + assert seen['kwargs']['skip_interpolation'] is True + + +def test_unvalidated_waveform_env_fails_closed(monkeypatch): + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + monkeypatch.setenv('RIFT_GPU_PRECOMPUTE','1') + monkeypatch.setenv('RIFT_GPU_WAVEFORM','ripple') + with pytest.raises(ValueError,match='not yet validated'): + fr.PrecomputeLikelihoodTermsRotatingFreqResponse(100.,0.2,None,{}, {},2,512.) + + +@pytest.mark.parametrize('mismatch', ['epoch','delta_t','delta_f']) +def test_native_provider_grid_mismatch_rejected(mismatch): + from RIFT.likelihood.gpu_precompute import PrecomputeLikelihoodTermsRotatingFreqResponseGPU + n,df,dt = 32,0.125,0.25 + common = dict(modes={(2,2):np.ones(n,dtype=complex)},delta_f=df, + delta_t=dt,epoch=1e9,conditioned=True) + main = SimpleNamespace(**common) + conjugate = SimpleNamespace(**common) + if mismatch == 'epoch': + conjugate.epoch += 1. # default np.isclose used to accept this at GPS epochs + elif mismatch == 'delta_t': + conjugate.delta_t *= 1.000001 + else: + conjugate.delta_f *= 1.000001 + P = SimpleNamespace(deltaT=dt,dist=1.) + data = {'H1':SimpleNamespace(deltaF=df)} + with pytest.raises(ValueError,match='different grids or modes'): + PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + 1e9,0.1,P,data,{'H1':None},2,2.,backend=np, + waveform_provider=lambda *a,**k:(main,conjugate)) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_highlevel_integration.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_highlevel_integration.py new file mode 100644 index 000000000..2edf1f7ca --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_highlevel_integration.py @@ -0,0 +1,158 @@ +"""High-level CPU/GPU wrapper and downstream ILE likelihood parity. + +The waveform generator is replaced with a deterministic synthetic LAL mode bank. Detector +geometry, finite-arm response weights, PSD handling, epochs, packing, and the maintained +rotation/frequency-response likelihood are real production code. +""" + +from types import SimpleNamespace + +import numpy as np +import pytest + + +def _fd_series(lal, name, values, delta_f, epoch): + out = lal.CreateCOMPLEX16FrequencySeries( + name, lal.LIGOTimeGPS(float(epoch)), 0.0, float(delta_f), + lal.DimensionlessUnit, len(values)) + out.data.data[:] = values + return out + + +def _synthetic_problem(): + lal = pytest.importorskip("lal") + from RIFT.likelihood import factored_likelihood as fl + + n, df = 256, 0.25 + dt = 1.0 / (n * df) + event = 1_420_000_000.125 + rng = np.random.default_rng(1509) + # Full non-Hermitian RIFT spectra, including unequal positive/negative halves. + h = rng.normal(size=n) + 1j * rng.normal(size=n) + hc = rng.normal(size=n) + 1j * rng.normal(size=n) + d = 0.8 * h + rng.normal(size=n) + 1j * rng.normal(size=n) + h *= np.linspace(0.7, 1.4, n) + hc *= np.linspace(1.3, 0.6, n) + data = _fd_series(lal, "H1 synthetic data", d, df, event - 2.0) + mode = (2, 2) + modes = {mode: _fd_series(lal, "h22", h, df, -2.0)} + modes_c = {mode: _fd_series(lal, "hc22", hc, df, -2.0)} + psd = lal.CreateREAL8FrequencySeries( + "H1 synthetic PSD", lal.LIGOTimeGPS(0), 0.0, df, + lal.DimensionlessUnit, n // 2 + 1) + psd.data.data[:] = 1.0 + 0.01 * np.arange(n // 2 + 1) + p = SimpleNamespace( + dist=100.0 * 1e6 * lal.PC_SI, deltaF=df, deltaT=dt, + fmin=2.0, phi=1.17, theta=-0.31, + ) + return lal, fl, event, p, {"H1": data}, {"H1": psd}, modes, modes_c + + +@pytest.mark.parametrize("backend_name", ["numpy", "cupy"]) +def test_highlevel_precompute_pack_epoch_and_downstream_likelihood(monkeypatch, request, backend_name): + lal, fl, event, p, data, psd, modes, modes_c = _synthetic_problem() + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + from RIFT.likelihood.gpu_precompute import ( + GPUPrecomputeContext, PrecomputeLikelihoodTermsRotatingFreqResponseGPU) + + if backend_name == "numpy": + xp = np + else: + try: + import cupy as xp + if xp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("no CUDA device") + except Exception as exc: + if request.config.getoption("--require-gpu") or pytestconfig_requires_gpu(): + pytest.fail("real CUDA device required: %r" % (exc,)) + pytest.skip("no usable CUDA device: %r" % (exc,)) + + monkeypatch.setattr(fl, "internal_hlm_generator", + lambda *args, **kwargs: (modes, modes_c)) + common = dict( + event_time_geo=event, t_window=0.25, P=p, data_dict=data, + psd_dict=psd, Lmax=2, fMax=24.0, Qmax=1, p_max=1, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, + verbose=False, quiet=True, skip_interpolation=True, + ) + # Force the maintained CPU implementation even if the submission environment set the + # feature variable globally. + monkeypatch.delenv("RIFT_GPU_PRECOMPUTE", raising=False) + cpu = fr.PrecomputeLikelihoodTermsRotatingFreqResponse(**common) + timings = [] + got = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, backend=xp, context=GPUPrecomputeContext(xp), return_device=False, + fft_batch=7, q_row_batch=9, frequency_chunk=37, + timing_callback=lambda stage, elapsed, details: timings.append((stage, elapsed))) + assert timings[0][0] == "initialization" + assert [stage for stage, _ in timings[:4]] == [ + "initialization", "waveform_generation", "waveform_pack_upload", "waveform"] + assert all(np.isfinite(elapsed) and elapsed >= 0 for _, elapsed in timings) + + cpu_i, cpu_u, cpu_v, cpu_q, cpu_meta = cpu + got_i, got_u, got_v, got_q, got_meta = got + assert cpu_meta["a_list"] == got_meta["a_list"] + assert cpu_meta["modes"] == got_meta["modes"] + assert got_meta["post_phase_required"] is True + for a in cpu_meta["a_list"]: + for mode in cpu_meta["modes"]: + cts, gts = cpu_q["H1"][a][mode], got_q["H1"][a][mode] + assert float(cts.epoch) == pytest.approx(float(gts.epoch), abs=1e-12) + assert cts.deltaT == pytest.approx(gts.deltaT, rel=0, abs=1e-15) + np.testing.assert_allclose(gts.data.data, cts.data.data, + rtol=2e-10, atol=2e-10) + for ap in cpu_meta["a_list"]: + for pair in cpu_u["H1"][(a, ap)]: + np.testing.assert_allclose(got_u["H1"][(a, ap)][pair], + cpu_u["H1"][(a, ap)][pair], + rtol=2e-10, atol=2e-10) + np.testing.assert_allclose(got_v["H1"][(a, ap)][pair], + cpu_v["H1"][(a, ap)][pair], + rtol=2e-10, atol=2e-10) + + cpu_packed = fr.pack_rotating_freqresponse_arrays(cpu_meta, cpu_q, cpu_u, cpu_v) + got_packed = fr.pack_rotating_freqresponse_arrays(got_meta, got_q, got_u, got_v) + # Evaluate three times around the detector arrival. This exercises epoch placement, + # Q slicing, response coefficients, post-phases, U and V through maintained ILE code. + pvec = SimpleNamespace( + phi=np.array([1.17, 1.21]), theta=np.array([-0.31, -0.28]), + incl=np.array([0.7, 1.0]), phiref=np.array([0.2, 1.1]), + psi=np.array([0.4, 0.9]), + dist=np.array([110.0, 170.0]) * 1e6 * lal.PC_SI, + tref=lal.LIGOTimeGPS(event), deltaT=p.deltaT, + ) + tvals = np.array([-p.deltaT, 0.0, p.deltaT]) + ln_cpu = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, cpu_meta, *cpu_packed, Lmax=2, array_output=True, + time_interp="nearest", xpy=np) + ln_got = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pvec, got_meta, *got_packed, Lmax=2, array_output=True, + time_interp="nearest", xpy=np) + np.testing.assert_allclose(ln_got, ln_cpu, rtol=2e-10, atol=2e-8) + + +def test_context_replaces_old_response_orders_and_cutoffs(monkeypatch): + from RIFT.likelihood.gpu_precompute import ( + GPUPrecomputeContext, PrecomputeLikelihoodTermsRotatingFreqResponseGPU) + _, fl, event, p, data, psd, modes, modes_c = _synthetic_problem() + monkeypatch.setattr(fl, "internal_hlm_generator", lambda *a, **k: (modes, modes_c)) + context = GPUPrecomputeContext(np) + for order, cutoff, arm in [(0, 24., 4000.), (1, 20., 3000.), (0, 22., 3500.)]: + common = dict(event_time_geo=event, t_window=.25, P=p, + data_dict=data, psd_dict=psd, Lmax=2, fMax=cutoff, + Qmax=order, p_max=0, L_arm=arm, backend=np, + return_device=True, verbose=False, quiet=True) + reused, _ = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, context=context) + fresh, _ = PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, context=GPUPrecomputeContext(np)) + assert context.stats()["retained_arrays"] == 3 + for name in ("q", "U", "V"): + np.testing.assert_allclose(reused[name]["H1"], fresh[name]["H1"]) + + +def pytestconfig_requires_gpu(): + # pytest's fixture object is deliberately not threaded through the scientific helper; + # the environment is the stable Condor/container gate used by README.md. + import os + return os.environ.get("RIFT_REQUIRE_GPU_PRECOMPUTE") == "1" diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_lal_fft_oracle.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_lal_fft_oracle.py new file mode 100644 index 000000000..a57d360ef --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_lal_fft_oracle.py @@ -0,0 +1,23 @@ +import numpy as np +import pytest + +from oracle import q_oracle + + +def test_numpy_oracle_exactly_matches_lal_reverse_fft_roll_cut(): + lal = pytest.importorskip("lal") + rng = np.random.default_rng(6103) + n, df, n_shift, n_window = 32, 0.125, -5, 17 + h = rng.normal(size=n) + 1j * rng.normal(size=n) + d = rng.normal(size=n) + 1j * rng.normal(size=n) + w = rng.uniform(size=n) + integrand = 2 * np.conj(h) * d * w + hf = lal.CreateCOMPLEX16FrequencySeries( + "integrand", lal.LIGOTimeGPS(100.25), 0, df, lal.DimensionlessUnit, n) + hf.data.data[:] = integrand + ht = lal.CreateCOMPLEX16TimeSeries( + "q", lal.LIGOTimeGPS(0), 0, 1.0 / (n * df), lal.DimensionlessUnit, n) + lal.COMPLEX16FreqTimeFFT(ht, hf, lal.CreateReverseCOMPLEX16FFTPlan(n, 0)) + lal_q = np.roll(np.asarray(ht.data.data).copy(), -n_shift)[:n_window] + got = q_oracle(h[None, None, :], d, w, df, n_shift, n_window)[0, 0] + np.testing.assert_allclose(got, lal_q, rtol=3e-15, atol=3e-15) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_smoke_config_contract.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_smoke_config_contract.py new file mode 100644 index 000000000..4b4b11171 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_smoke_config_contract.py @@ -0,0 +1,39 @@ +"""Keep paired smoke tests independent of differing executable defaults.""" +import ast +import importlib.util +from pathlib import Path + +import pytest + + +def test_paired_smokes_pin_the_same_reference_frequency(): + for name in ("run_short_av_ile.py", "run_short_jax_av_ile.py"): + path = Path(__file__).with_name(name) + tree = ast.parse(path.read_text(), filename=str(path)) + values = [ + ast.literal_eval(call.args[2]) + for call in ast.walk(tree) + if isinstance(call, ast.Call) + and isinstance(call.func, ast.Name) + and call.func.id == "set_option" + and len(call.args) == 3 + and isinstance(call.args[1], ast.Str) + and call.args[1].s == "--reference-freq" + ] + assert values == [100.0], name + + +def test_fairdraw_guard_rejects_samples_outside_the_requested_box(tmpdir): + script = Path(__file__).with_name("run_short_jax_av_ile.py") + spec = importlib.util.spec_from_file_location("short_jax_smoke_contract", script) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + boxes = dict(right_ascension=(1.1, 1.3), declination=(0.2, 0.4), + inclination=(0.2, 0.6), psi=(0.3, 0.7), distance_mpc=(300., 500.)) + path = Path(str(tmpdir)) / "synthetic_fairdraw.dat" + header = "# right_ascension declination distance inclination psi phi_orb loglikelihood\n" + path.write_text(header + "1.2 0.3 400 0.4 0.5 0 10\n") + module.validate_fairdraw_bounds(path, boxes) + path.write_text(header + "1.4 0.3 400 0.4 0.5 0 10\n") + with pytest.raises(RuntimeError, match="right_ascension outside"): + module.validate_fairdraw_bounds(path, boxes) diff --git a/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_streamed_v_weighting.py b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_streamed_v_weighting.py new file mode 100644 index 000000000..a26c05103 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/gpu_precompute/test_streamed_v_weighting.py @@ -0,0 +1,143 @@ +"""Regression tests for memory-bounded weighting in streamed V contraction.""" + +import numpy as np +import pytest + +from RIFT.likelihood import gpu_precompute as gpu + +from conftest import to_host + + +def _problem(seed=6061): + rng = np.random.default_rng(seed) + nfreq, nmodes, nresponse = 38, 2, 3 + delta_f = 0.125 + delta_t = 1.0 / (nfreq * delta_f) + base = (rng.normal(size=(nmodes, nfreq)) + + 1j * rng.normal(size=(nmodes, nfreq))) + base_conj = (rng.normal(size=(nmodes, nfreq)) + + 1j * rng.normal(size=(nmodes, nfreq))) + response = (rng.normal(size=(nresponse, nfreq)) + + 1j * rng.normal(size=(nresponse, nfreq))) + # Deliberately nonsymmetric compound rows and a partial final a-block. + a_list = [(0, 0, -2), (2, 1, 1), (1, 0, 3), + (0, 2, -1), (2, 1, 2)] + weights = rng.uniform(0.01, 2.0, size=nfreq) + weights[[0, 3, nfreq // 2 + 1, nfreq - 1]] = 0.0 + epoch = -13.37 + f_sidereal = 0.019 + return (base, base_conj, response, a_list, weights, + delta_f, delta_t, epoch, f_sidereal) + + +@pytest.mark.parametrize("a_block,frequency_chunk", [(1, 1), (2, 7), (4, 64)]) +def test_streamed_weighted_left_matches_dense_complex_oracle( + backend, a_block, frequency_chunk): + (base, base_conj, response, a_list, weights, + delta_f, delta_t, epoch, f_sidereal) = _problem() + primary = gpu.build_compound_basis( + backend.asarray(base), backend.asarray(response), a_list, + delta_f, delta_t, epoch, f_sidereal=f_sidereal, + backend=backend, fft_batch=2) + conjugate = gpu.build_compound_basis( + backend.asarray(base_conj), backend.asarray(response), a_list, + delta_f, delta_t, epoch, f_sidereal=f_sidereal, + backend=backend, fft_batch=3) + + a_count, mode_count, nfreq = primary.shape + primary_flat = to_host(primary).reshape(a_count * mode_count, nfreq) + conjugate_flat = to_host(conjugate).reshape(a_count * mode_count, nfreq) + expected = ((np.conj(conjugate_flat) * weights[None, :]) + @ primary_flat.T) * (2.0 * delta_f) + expected = expected.reshape( + a_count, mode_count, a_count, mode_count).transpose(0, 2, 1, 3) + + got = gpu.streamed_v_matrix( + backend.asarray(base_conj), backend.asarray(response), a_list, primary, + backend.asarray(weights), delta_f, delta_t, epoch, + f_sidereal=f_sidereal, backend=backend, a_block=a_block, + fft_batch=min(2, a_block), frequency_chunk=frequency_chunk, + return_device=True) + np.testing.assert_allclose( + to_host(got), expected, rtol=3e-11, atol=3e-10) + + +class _TrackingArray(np.ndarray): + """NumPy view that records every broadcast multiply by a row weight.""" + + def __new__(cls, value, tracker): + obj = np.asarray(value).view(cls) + obj.tracker = tracker + return obj + + def __array_finalize__(self, source): + self.tracker = getattr(source, "tracker", None) + + def __mul__(self, other): + other_shape = np.shape(other) + if len(other_shape) == 2 and other_shape[0] == 1: + self.tracker.append((tuple(self.shape), tuple(other_shape))) + return _TrackingArray( + np.asarray(self) * np.asarray(other), self.tracker) + + def __rmul__(self, other): + return self.__mul__(other) + + +class _TrackingNumpy: + """Small backend shim used only to audit temporary working-set shapes.""" + + __name__ = "tracking_numpy" + complex128 = np.complex128 + + def __init__(self): + self.weighted_multiplies = [] + + def asarray(self, value, dtype=None): + if isinstance(value, _TrackingArray) and dtype is None: + return value + return _TrackingArray( + np.asarray(value, dtype=dtype), self.weighted_multiplies) + + def zeros(self, shape, dtype=None): + return _TrackingArray( + np.zeros(shape, dtype=dtype), self.weighted_multiplies) + + def conj(self, value): + return _TrackingArray( + np.conj(np.asarray(value)), self.weighted_multiplies) + + +def test_each_v_block_weights_only_its_small_left_working_set(monkeypatch): + """Never allocate a weighted (A*M, frequency_chunk) primary temporary.""" + rng = np.random.default_rng(9902) + a_count, mode_count, nfreq, a_block = 5, 2, 37, 2 + a_list = [(i, 0, 0) for i in range(a_count)] + primary = (rng.normal(size=(a_count, mode_count, nfreq)) + + 1j * rng.normal(size=(a_count, mode_count, nfreq))) + conjugate = (rng.normal(size=(a_count, mode_count, nfreq)) + + 1j * rng.normal(size=(a_count, mode_count, nfreq))) + backend = _TrackingNumpy() + + positions = {a: i for i, a in enumerate(a_list)} + + def fake_build(unused_base, unused_response, block, *unused_args, **unused_kwargs): + rows = [positions[tuple(a)] for a in block] + return backend.asarray(conjugate[rows]) + + monkeypatch.setattr(gpu, "build_compound_basis", fake_build) + got = gpu.streamed_v_matrix( + np.zeros((mode_count, nfreq), dtype=complex), + np.zeros((a_count, nfreq), dtype=complex), a_list, primary, + np.linspace(0.0, 1.0, nfreq), 0.25, 1.0 / (nfreq * 0.25), 0.0, + backend=backend, a_block=a_block, fft_batch=1, frequency_chunk=11, + return_device=True) + + assert got.shape == (a_count, a_count, mode_count, mode_count) + weighted_shapes = backend.weighted_multiplies + assert weighted_shapes + assert max(shape[0][0] for shape in weighted_shapes) == a_block * mode_count + assert all(shape[0][0] < a_count * mode_count for shape in weighted_shapes) + # Across all blocks/chunks, every streamed-left element is weighted once. + assert sum(left[0] * left[1] for left, unused_weight in weighted_shapes) == \ + a_count * mode_count * nfreq diff --git a/MonteCarloMarginalizeCode/Code/test/integrators/test_portfolio_balance_heuristic.py b/MonteCarloMarginalizeCode/Code/test/integrators/test_portfolio_balance_heuristic.py index 0384196be..6bfe3b2c0 100644 --- a/MonteCarloMarginalizeCode/Code/test/integrators/test_portfolio_balance_heuristic.py +++ b/MonteCarloMarginalizeCode/Code/test/integrators/test_portfolio_balance_heuristic.py @@ -34,7 +34,8 @@ This test builds AV(decoy) + GMM(broad) on a correlated-Gaussian target where a cold AV converges, and checks: * OLD estimator -> badly biased low, - * NEW estimator -> unbiased (matches true integral within a few percent), + * NEW estimator -> statistically consistent with the true integral at its + measured Monte Carlo uncertainty, * and a no-regression control: a NORMAL portfolio (cold AV + GMM, both sane) stays unbiased under the NEW estimator. @@ -177,6 +178,7 @@ def run(target, n_chunk, nmax, neff, use_mixture, decoy=None, seed=1234, lnI = float(B._asnumpy(lnI)) ln_wt = B.log_weights_from_rvs(port._rvs) return dict(lnI=lnI, bias=lnI - float(target.true_lnZ), + sigma_over_I=float(np.exp(0.5 * float(B._asnumpy(logvar)) - lnI)), n_eval=int(getattr(port, "ntotal", 0)) or nmax, n_ess=B.n_ess_kish(ln_wt), final_weights=np.array(port.portfolio_weights)) @@ -233,10 +235,16 @@ def main(): if not (old["bias"] < -0.7): print(" FAIL: old stratified estimator not badly biased low ({:+.3f}); " "decoy not exercised".format(old["bias"])); ok = False - # 2. the NEW estimator must be unbiased within a few percent (a few % in - # the integral is ~0.03-0.20 in ln); allow a modest gate - if abs(new["bias"]) > 0.20: - print(" FAIL: new q_mix estimator biased ({:+.3f} > 0.20)".format(new["bias"])); ok = False + # 2. The covering GMM receives only about 1% of draws. The resulting + # decoy run can have single-digit n_ess, where a fixed 0.20-log-unit + # threshold rejects ordinary Monte Carlo fluctuations. Compare the + # known integral to the run's own uncertainty in linear Z units. + # The no-decoy control below retains its tighter absolute gate. + new_zscore = abs(1.0 - np.exp(-new["bias"])) / new["sigma_over_I"] + if not np.isfinite(new_zscore) or new_zscore > 3.0: + print(" FAIL: new q_mix estimator differs from truth by {:.2f} sigma" + " (bias {:+.3f}, sigma/I {:.3f})".format( + new_zscore, new["bias"], new["sigma_over_I"])); ok = False # 3. the new estimator must be dramatically better than the old if not (abs(new["bias"]) < abs(old["bias"]) - 0.5): print(" FAIL: q_mix did not fix the decoy bias"); ok = False @@ -246,9 +254,10 @@ def main(): "({:+.3f})".format(ctl["bias"])); ok = False if not ok: raise SystemExit(1) - print("\n PASS: q_mix balance heuristic keeps the portfolio unbiased with a " - "decoy member (old {:+.3f} -> new {:+.3f}); control unbiased " - "({:+.3f}).".format(old["bias"], new["bias"], ctl["bias"])) + print("\n PASS: q_mix balance heuristic is consistent with the true " + "integral at {:.2f} sigma with a decoy member (old {:+.3f} -> " + "new {:+.3f}); control bias {:+.3f}.".format( + new_zscore, old["bias"], new["bias"], ctl["bias"])) if __name__ == "__main__": diff --git a/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_jax.py b/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_jax.py new file mode 100644 index 000000000..5dc879cfd --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_jax.py @@ -0,0 +1,456 @@ +"""Static-cost, JIT/AD-compatible multipeak wiring and contract.""" + +import numpy as np +import pytest + +jax = pytest.importorskip("jax") +import jax.numpy as jnp + +from RIFT.likelihood.jax_ile import anglemarg as AM +from RIFT.likelihood.jax_ile import direct_marginalization_policy as DP +from RIFT.likelihood.jax_ile.wrapper import JAXDistPhiPsiMargLikelihood +from test_angle_marg_exact import make_synth, RA, DEC, INCL, INTERP + + +def test_multipeak_jax_is_explicit_only_and_records_bounded_contract(): + assert "multipeak-jax" in AM.ANGLE_MARG_CHOICES + for amp in (1.0, 500.0, 5.0e6): + assert AM.choose_angle_marg_scheme(amp)[0] != "multipeak-jax" + + data = make_synth(scale=2.0) + like = JAXDistPhiPsiMargLikelihood( + data, 30.0, 3000.0, nphi=32, npsi=8, interp=INTERP, + angle_marg="multipeak-jax") + assert "amp_sizing" not in like.angle_marg_info + assert "sample_grid" not in like.angle_marg_info + assert like.angle_marg_info["config"] == DP.BoundedMultipeakConfig()._asdict() + assert like.angle_marg_info["bounded_cost"] is True + assert like.angle_marg_info["dense_reserve"] is False + assert like.angle_marg_info["fixed_plan_autodiff_only"] is True + assert like.angle_marg_info["derivative_warrant_certified"] is False + + +def test_multipeak_jax_wrapper_uses_jit_and_ad(monkeypatch): + def differentiable_stub(data, ra, dec, incl, *args, **kwargs): + return ra + 2.0 * dec + 3.0 * incl + + monkeypatch.setattr( + DP, "fused_log_likelihood_four_axis_bounded", differentiable_stub) + data = make_synth(scale=2.0) + like = JAXDistPhiPsiMargLikelihood( + data, 30.0, 3000.0, nphi=32, npsi=8, interp=INTERP, + angle_marg="multipeak-jax") + value = np.asarray(like._batched( + jnp.asarray(RA), jnp.asarray(DEC), jnp.asarray(INCL))) + assert value.shape == np.shape(RA) + val, grad = like._value_and_grad( + jnp.asarray([RA[0], DEC[0], INCL[0]])) + assert np.isfinite(np.asarray(val)) + np.testing.assert_allclose(np.asarray(grad), [1.0, 2.0, 3.0]) + + +def test_multipeak_jax_refuses_time_series_and_amp_grid_outputs(monkeypatch): + monkeypatch.setattr( + DP, "fused_log_likelihood_four_axis_bounded", + lambda data, ra, dec, incl, *args, **kwargs: ra * 0.0) + data = make_synth(scale=2.0) + like = JAXDistPhiPsiMargLikelihood( + data, 30.0, 3000.0, nphi=32, npsi=8, interp=INTERP, + angle_marg="multipeak-jax") + args = (data, jnp.asarray(RA), jnp.asarray(DEC), jnp.asarray(INCL)) + with pytest.raises(ValueError, match="no lnL"): + like._fused(*args, return_lnLt=True) + with pytest.raises(ValueError, match="static cost envelope"): + like._fused(*args, return_amp=True) + + +def test_bounded_config_refuses_nonstatic_or_invalid_envelopes(): + with pytest.raises(ValueError, match="time_guard"): + DP.validate_bounded_multipeak_config( + DP.BoundedMultipeakConfig(time_guard=1)) + with pytest.raises(ValueError, match="enriched_oversample"): + DP.validate_bounded_multipeak_config( + DP.BoundedMultipeakConfig(base_oversample=2, + enriched_oversample=2)) + with pytest.raises(ValueError, match="mode caps"): + DP.validate_bounded_multipeak_config( + DP.BoundedMultipeakConfig(base_max_starts=1, + enriched_max_modes=3)) + + +@pytest.mark.parametrize("invalid", [None, "norm", "table"]) +def test_device_function_is_jittable_differentiable_and_forwards_caps( + monkeypatch, invalid): + seen = [] + + def fake_tables(data, ra, dec, incl, interp, guard): + del data, interp + ntime = 2 * int(guard) + 3 + signal = ra + 2.0 * dec + 3.0 * incl + ca = jnp.broadcast_to(signal[None, None, :, None], + (1, 1, signal.size, ntime)).astype(jnp.complex128) + cb = jnp.ones_like(ca) + if invalid == "norm": + cb = cb.at[..., -1].set(2.0) + if invalid == "table": + ca = ca.at[..., 0].set(jnp.nan) + return ca, cb, {"m_max": 0} + + def fake_rank(table, norm, x_min, x_max, **kwargs): + del table, norm, x_min, x_max + seen.append((kwargs["max_starts"], kwargs["max_time_nodes"], + kwargs["angular_oversample"])) + return jnp.asarray(0.0) + + def fake_plan(table, norm, base, extra, x_min, x_max, **kwargs): + del norm, base, extra, x_min, x_max + token = {"token": jnp.real(jnp.sum(table))} + one = jnp.asarray(1, dtype=jnp.int32) + ok = jnp.asarray(True) + plan = dict(n_selected_modes=one, n_optimizer_starts=one, + n_lattice_evaluations=one, n_candidates_before_cap=one, + start_capacity_ok=ok, time_capacity_ok=ok) + shared = {"n_optimizer_starts_executed": one} + assert kwargs["max_modes"] == 1 + assert kwargs["enriched_max_modes"] == 2 + return token, token, plan, plan, shared + + def fake_integral(table, norm, base_plan, enriched_plan, + x_min, x_max, **kwargs): + del norm, base_plan, enriched_plan, x_min, x_max, kwargs + # Return a finite candidate even for the bad A table, isolating the + # outer tables_finite guard from both norm invariance and inner gates. + value = jnp.real(jnp.sum(jnp.nan_to_num(table))) + return value, jnp.asarray(True), {"accepted_local": jnp.asarray(True)} + + monkeypatch.setattr(DP._anglemarg, "angle_coefficient_tables", fake_tables) + monkeypatch.setattr(DP._anglemarg, "_runtime_amp_failsafe", + lambda *args, **kwargs: jnp.asarray(0.0)) + monkeypatch.setattr(DP._aap, "rank_joint_starts_from_uvq_device", fake_rank) + monkeypatch.setattr(DP._aap, "make_all_axis_mode_plan_pair_device", fake_plan) + monkeypatch.setattr(DP._aap, "empirical_enrichment_marginalize", fake_integral) + + cfg = DP.BoundedMultipeakConfig( + time_guard=2, base_max_starts=2, max_time_nodes=3, + base_oversample=1, enriched_oversample=2, + max_modes=1, enriched_max_modes=2, refine_iterations=1, + base_order=2, base_check_order=3, + enriched_order=3, enriched_check_order=4) + xg = jnp.asarray([0.5, 1.0]) + lwg = jnp.asarray([-np.log(2.0), -np.log(2.0)]) + + def scalar(theta): + value, ledger = DP.fused_log_likelihood_four_axis_bounded( + object(), theta[0:1], theta[1:2], theta[2:3], xg, lwg, + interp=None, amp_sizing=10.0, config=cfg, + local_log_normalization=0.0, x_bounds=(0.5, 1.0), + return_ledger=True) + return value[0], ledger + + (value, ledger), grad = jax.jit(jax.value_and_grad( + scalar, has_aux=True))( + jnp.asarray([0.1, 0.2, 0.3])) + if invalid is None: + assert np.isfinite(np.asarray(value)) + np.testing.assert_allclose(np.asarray(grad), 7.0 * np.asarray([1., 2., 3.])) + else: + assert np.isnan(np.asarray(value)) + assert not bool(ledger["usable"][0]) + assert not bool(ledger["norm_time_invariant" if invalid == "norm" else "tables_finite"][0]) + assert bool(np.asarray(ledger["bounded_cost"][0])) + assert not bool(np.asarray(ledger["derivative_warrant_certified"][0])) + assert (2, 3, 1) in seen and (2, 3, 2) in seen + + +@pytest.mark.parametrize("changes", [ + {"base_max_starts": 0}, {"max_time_nodes": 1}, + {"base_oversample": 0}, {"max_modes": 0}, + {"enriched_max_modes": 7}, {"enriched_max_modes": 257}, + {"base_order": 1}, {"base_check_order": 11}, + {"enriched_order": 12}, {"enriched_check_order": 13}, + {"local_radius": 0}, {"refine_iterations": 0}, + {"convergence_tol_nats": 0}, {"time_guard_tol_nats": 0}, + {"total_value_error_budget_nats": 0}, + {"time_outside_tol_nats": 0}, {"time_outside_tol_nats": -np.inf}, + {"norm_invariance_rtol": -1}, {"norm_invariance_rtol": np.nan}, + {"batch_rows": -1}, {"batch_rows": 1.5}, +]) +def test_invalid_static_envelope_is_refused(changes): + with pytest.raises(ValueError): + DP.validate_bounded_multipeak_config(DP.BoundedMultipeakConfig(**changes)) + + +def test_bounded_config_type_and_shared_defaults(): + with pytest.raises(TypeError, match="BoundedMultipeakConfig"): + DP.validate_bounded_multipeak_config(DP.PolicyConfig()) + cfg = DP.BoundedMultipeakConfig() + separate = {"time_guard", "base_max_starts", "max_time_nodes", + "max_modes", "enriched_max_modes", + "batch_rows", "base_order", "base_check_order", + "enriched_order", "enriched_check_order", + "convergence_tol_nats", "total_value_error_budget_nats"} + for field in cfg._fields: + if field not in separate: + assert getattr(cfg, field) == getattr(DP.PolicyConfig(), field) + + +@pytest.mark.parametrize("bounds", [(0., 1.), (-1., 1.), (1., 1.), (2., 1.)]) +def test_invalid_distance_support_refused_before_tables(bounds): + with pytest.raises(ValueError, match="x_bounds"): + DP.fused_log_likelihood_four_axis_bounded( + None, None, None, None, None, None, interp=None, amp_sizing=1., + local_log_normalization=0., x_bounds=bounds) + + +@pytest.mark.parametrize("kwargs,match", [ + ({"time_quadrature": "bandlimited"}, + "time_quadrature='bandlimited' is not valid for distance/phase/polarization marginalization"), + ({"dist_grid": "loguniform"}, "uniform-in-distance"), +]) +def test_wrapper_refuses_incompatible_measures(kwargs, match): + with pytest.raises(ValueError, match=match): + JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, + angle_marg="multipeak-jax", **kwargs) + + +def test_real_kernel_default_envelope_accepts_and_matches_reference(monkeypatch): + from test_direct_marginalization_policy import ( + _guarded_problem, _fake_data, _grid, _fine_reference, _N) + cfg = DP.BoundedMultipeakConfig() + table, norm, constants = _guarded_problem(_N, cfg.time_guard) + data = _fake_data(_N) + xg, lw = _grid() + monkeypatch.setattr(DP._core, "_DISTMARG_GH_N", 0) + + def tables(data, ra, dec, incl, interp, guard): + assert guard == cfg.time_guard + ca = jnp.asarray(table)[:, :, None, :] * ra[None, None, :, None] + cb = jnp.broadcast_to(jnp.asarray(norm)[:, :, None, None], + norm.shape + (ra.size, ca.shape[-1])) + return ca, cb, {"m_max": 2} + + # Only inject analytic coefficient tables: discovery, refinement, all + # quadratures and acceptance guards below are the actual production code. + monkeypatch.setattr(AM, "angle_coefficient_tables", tables) + lln, _ = DP.policy_log_normalization(data, xg, lw) + bounds = (float(xg.min()), float(xg.max())) + + def evaluate(scales, config=cfg): + return DP.fused_log_likelihood_four_axis_bounded( + data, scales, jnp.zeros_like(scales), jnp.zeros_like(scales), xg, lw, + interp=INTERP, amp_sizing=40., config=config, + local_log_normalization=lln, x_bounds=bounds, return_ledger=True) + + value, ledger = jax.jit(evaluate)(jnp.array([1., 1.01])) + assert np.all(ledger["usable"]), { + k: np.asarray(v) for k, v in ledger.items() if k.startswith("decline")} + for i, scale in enumerate((1., 1.01)): + fine = _fine_reference(constants, norm, data, xg, lw, 40., scale=scale) + assert abs(float(value[i]) - fine) < 1.e-3 + # The old orders decline the same problem, even with the corrected caps. + old_orders = cfg._replace(base_order=7, base_check_order=9, + enriched_order=9, enriched_check_order=11, + convergence_tol_nats=1.e-3, + total_value_error_budget_nats=1.e-2) + declined, why = jax.jit(lambda x: evaluate(x, old_orders))(jnp.array([1.])) + assert np.isnan(declined[0]) and bool(why["decline_quadrature"][0]) + # AD is exercised through the real accepted fixed-plan quadrature. + scalar = lambda x: evaluate(x[None])[0][0] + grad = jax.jit(jax.grad(scalar))(jnp.array(1.)) + h = 1.e-5 + fd = (float(jax.jit(scalar)(1. + h)) - float(jax.jit(scalar)(1. - h))) / (2*h) + np.testing.assert_allclose(grad, fd, rtol=2.e-5, atol=1.e-5) + + +def test_cli_envelope_reaches_wrapper_and_provenance(monkeypatch): + from test_direct_marginalization_policy_cli import _load_driver + mod = _load_driver() + parser = mod.build_parser() + opts, _ = parser.parse_args([]) + assert mod.bounded_multipeak_config_from_options(opts) == DP.BoundedMultipeakConfig() + opts, _ = parser.parse_args([ + "--multipeak-jax-time-guard", "64", + "--multipeak-jax-base-max-starts", "256", + "--multipeak-jax-max-time-nodes", "512", + "--multipeak-jax-base-order", "12"]) + cfg = mod.bounded_multipeak_config_from_options(opts) + seen = [] + def kernel(data, ra, dec, incl, *args, config, **kwargs): + seen.append(config) + values = ra * 0. + return (values, {"selected_value": values}) if kwargs.get("return_ledger") else values + monkeypatch.setattr(DP, "fused_log_likelihood_four_axis_bounded", kernel) + like = JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, + angle_marg="multipeak-jax", bounded_multipeak_config=cfg) + like.log_likelihood(RA, DEC, INCL).block_until_ready() + assert seen == [cfg] + assert cfg.time_guard == 64 and cfg.base_max_starts == 256 + assert cfg.max_time_nodes == 512 and cfg.base_order == 12 + assert like.angle_marg_info["config"] == cfg._asdict() + assert "dense_reserve=false" in mod.angle_grid_suspect_note("multipeak-jax") + + +@pytest.mark.parametrize("bad", [np.nan, np.inf, -np.inf]) +def test_host_decline_is_latched_and_cannot_publish_survivors(monkeypatch, bad): + from test_direct_marginalization_policy_cli import _load_driver + mod = _load_driver() + def kernel(data, ra, *args, **kwargs): + values = jnp.where(ra > 1., bad, 0.) + return (values, {"selected_value": values}) if kwargs.get("return_ledger") else values + monkeypatch.setattr(DP, "fused_log_likelihood_four_axis_bounded", kernel) + like = JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, angle_marg="multipeak-jax", + bounded_multipeak_decline_action="refuse") + with pytest.raises(RuntimeError, match="refusing"): + like.log_likelihood(jnp.array([0., 2.]), jnp.zeros(2), jnp.zeros(2)) + # Even a caller swallowing the exception and keeping finite rows cannot + # evade the final publication guard. + with pytest.raises(RuntimeError, match="refusing to publish"): + mod.require_bounded_multipeak_rows(like, np.array([0., 1.])) + fresh = type("Like", (), {"angle_marg_scheme": "multipeak-jax"})() + with pytest.raises(RuntimeError, match="refusing to publish"): + mod.require_bounded_multipeak_rows(fresh, np.array([0., bad])) + mod.require_bounded_multipeak_rows(fresh, np.array([0., 1.])) + + +@pytest.mark.parametrize("diagnostic", [-100., 10., np.nan]) +def test_drop_preserves_proposal_count_and_records_material_declines(monkeypatch, diagnostic): + from RIFT.likelihood.jax_ile.samplers import evidence_from_logweights + def kernel(data, ra, *args, **kwargs): + bad = ra > 1. + values = jnp.where(bad, jnp.nan, 0.) + ledger = {"selected_value": jnp.where(bad, diagnostic, 0.), + "decline_quadrature": bad} + return (values, ledger) if kwargs.get("return_ledger") else values + monkeypatch.setattr(DP, "fused_log_likelihood_four_axis_bounded", kernel) + like = JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, angle_marg="multipeak-jax") + values = like.log_likelihood(jnp.array([0., 2.]), jnp.zeros(2), jnp.zeros(2)) + np.testing.assert_array_equal(values, [0., -1.e30]) + logZ, _, _ = evidence_from_logweights(values) + assert logZ == pytest.approx(-np.log(2.)) # not logZ=0 from discarding the denominator + audit = like.bounded_multipeak_audit + assert audit["evaluated"] == 2 and audit["declined"] == 1 + assert audit["reasons"] == {"decline_quadrature": 1} + assert audit["max_accepted"] == 0. + if np.isfinite(diagnostic): + assert audit["max_declined_diagnostic"] == diagnostic + else: + assert audit["diagnostic_unknown"] == 1 + # Jitted sampler path uses the same target, without host callbacks. + assert float(jax.jit(like._scalar)(jnp.array([2., 0., 0.]))) == -1.e30 + + +@pytest.mark.parametrize("action", ["drop", "refuse"]) +def test_driver_mixed_cloud_publication_and_cli_wiring(monkeypatch, tmp_path, action): + import types + from test_direct_marginalization_policy_cli import _load_driver + from RIFT.likelihood.jax_ile import samplers, wrapper + mod = _load_driver() + parser = mod.build_parser() + opts, _ = parser.parse_args([ + "--mode", "flowmc-phipsimarg", "--distance-marginalization", + "--angle-marg-scheme", "multipeak-jax", + "--multipeak-jax-decline-action", action, + "--multipeak-jax-time-guard", "64", + "--multipeak-jax-max-time-nodes", "512", + "--output-file", str(tmp_path / "event"), "--save-samples"]) + seen = [] + def constructor(data, *args, **kwargs): + cfg = kwargs["bounded_multipeak_config"] + assert cfg.time_guard == 64 and cfg.max_time_nodes == 512 + assert kwargs["bounded_multipeak_decline_action"] == action + return types.SimpleNamespace( + angle_marg_scheme="multipeak-jax", direct_marginalization_policy="off", + policy_info=None, + angle_marg_info={"config": cfg._asdict()}, time_quadrature="simpson", + bounded_multipeak_decline_action=action, bounded_multipeak_audit={}) + monkeypatch.setattr(wrapper, "JAXDistPhiPsiMargLikelihood", constructor) + data = types.SimpleNamespace(lms=np.array([[2, 2]]), q_time_pregrid_factor=1) + monkeypatch.setattr(mod, "build_data_from_precompute", + lambda *a, **k: (data, {"guess_snr": 10.})) + result = dict(theta=np.zeros((3, 3)), lnL=np.array([0., -1.e30, 1.]), + logZ=.1, sigma_over_Z=.1, neff=2., post_weight=np.ones(3)/3) + monkeypatch.setattr(samplers, "flowmc_sample_phimarg", lambda *a, **k: result) + def samples(opts, idx, theta, lnL, with_distance, **kwargs): + seen.append(("samples", np.asarray(lnL), kwargs)) + def dat(*args, **kwargs): + seen.append(("dat", kwargs)) + monkeypatch.setattr(mod, "write_samples", samples) + monkeypatch.setattr(mod, "write_dat", dat) + P = types.SimpleNamespace(copy=lambda: None) + if action == "refuse": + with pytest.raises(RuntimeError, match="refusing to publish"): + mod.analyze_one(opts, P, {}, {}, False, 0., np.random.default_rng(1), 0, 1) + assert not seen + else: + mod.analyze_one(opts, P, {}, {}, False, 0., np.random.default_rng(1), 0, 1) + assert len(seen) == 2 + np.testing.assert_array_equal(seen[0][1], [0., 1.]) + assert len(seen[0][2]["logw"]) == 2 + note = seen[0][2]["angle_note"] + assert "output_rows_dropped=1" in note and "omitted_mass=unbounded" in note + assert "audit_scope=log_likelihood-batches" in note + assert "audit_excludes=scalar-MAP-Fisher-MALA" in note + assert "refuse_latch_scope=log_likelihood-batches" in note + assert "scalar_refuse_decline=nan-without-raise-or-latch" in note + assert seen[1][1]["angle_note"] == note + + +@pytest.mark.parametrize("args,expected", [ + (["--multipeak-jax-base-order", "15"], "requires --angle-marg-scheme"), + (["--angle-marg-scheme", "multipeak-jax", "--multipeak-jax-base-order", "1"], + "base_order"), + (["--mode", "prior-mc", "--angle-marg-scheme", "multipeak-jax"], + "requires --mode flowmc-phipsimarg"), +]) +def test_cli_refuses_inert_or_invalid_envelope(args, expected): + from test_direct_marginalization_policy_cli import _run + rc, out = _run("--mode", "flowmc-phipsimarg", *args) + assert rc != 0 and expected in out, out[-2000:] + + +def test_drop_refuses_an_entirely_declined_output_cloud(): + from types import SimpleNamespace + from test_direct_marginalization_policy_cli import _load_driver + mod = _load_driver() + like = SimpleNamespace(angle_marg_scheme="multipeak-jax", + bounded_multipeak_decline_action="drop") + with pytest.raises(RuntimeError, match="no accepted output rows"): + mod.require_bounded_multipeak_rows(like, np.array([-1.e30, np.nan])) + + +@pytest.mark.parametrize("kwargs,match", [ + ({"bounded_multipeak_decline_action": "invalid"}, + "bounded_multipeak_decline_action must be drop or refuse"), + ({"multipeak_guard": 32, + "bounded_multipeak_config": DP.BoundedMultipeakConfig(time_guard=16)}, + "multipeak_guard conflicts with bounded_multipeak_config"), +]) +def test_wrapper_rejects_invalid_python_configuration(kwargs, match): + with pytest.raises(ValueError, match=match): + JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, + angle_marg="multipeak-jax", **kwargs) + + +@pytest.mark.parametrize("action", ["drop", "refuse"]) +def test_scalar_declines_are_explicitly_outside_host_audit(monkeypatch, action): + def kernel(data, ra, *args, **kwargs): + return jnp.full_like(ra, jnp.nan) + monkeypatch.setattr(DP, "fused_log_likelihood_four_axis_bounded", kernel) + like = JAXDistPhiPsiMargLikelihood( + make_synth(scale=2.), 30., 3000., interp=INTERP, + angle_marg="multipeak-jax", bounded_multipeak_decline_action=action) + theta = jnp.zeros(3) + value, grad = like.value_and_grad(theta) + like.fisher(theta) + scalar = jax.jit(like._scalar)(theta) + if action == "refuse": + assert np.isnan(value) and np.isnan(scalar) + else: + assert value == -1.e30 and scalar == -1.e30 + assert like.bounded_multipeak_audit["evaluated"] == 0 + assert not getattr(like, "bounded_multipeak_declined", False) diff --git a/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_wiring.py b/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_wiring.py index c1d8c1b6e..7e8356c6a 100644 --- a/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_wiring.py +++ b/MonteCarloMarginalizeCode/Code/test/jax/test_angle_marg_multipeak_wiring.py @@ -56,6 +56,37 @@ def test_multipeak_returns_one_finite_value_per_sample(): assert np.all(np.isfinite(v)), v +def test_multipeak_runs_through_the_BATCHED_seam_the_sampler_uses(): + """THE TEST THAT WAS MISSING, and the reason the first landing broke every + Table 3 cell in ~30 s. + + The other tests here call `_fused` directly, which is EAGER. The sampler + reaches the likelihood through `_batched`, and for every other scheme that is + `jax.jit(_batched)`. multipeak is a host-side numpy/scipy planner: under a + trace its `np.asarray` on the coefficient tables raises + TracerArrayConversionError. Exercising `_fused` proves nothing about the + path production uses. + """ + data = make_synth(scale=2.0) + like = JAXDistPhiPsiMargLikelihood(data, 30.0, 3000.0, nphi=32, npsi=8, + interp=INTERP, angle_marg="multipeak") + v = np.asarray(like._batched(jnp.asarray(RA), jnp.asarray(DEC), + jnp.asarray(INCL))) + assert v.shape == np.shape(RA), v.shape + assert np.all(np.isfinite(v)), v + + +def test_multipeak_refuses_gradients_rather_than_inventing_one(): + """A numpy planner has no AD. A silent zero or wrong gradient would reach + --fisher-precondition, which swallows exceptions and falls back to raw + coordinates with the flag still recorded as supplied.""" + data = make_synth(scale=2.0) + like = JAXDistPhiPsiMargLikelihood(data, 30.0, 3000.0, nphi=32, npsi=8, + interp=INTERP, angle_marg="multipeak") + with pytest.raises(ValueError, match="not differentiable"): + like._value_and_grad(jnp.asarray([RA[0], DEC[0], INCL[0]])) + + def test_multipeak_records_its_provenance(): """This pipeline has a history of silently-inert flags: the scheme actually used must be visible in the record, not inferred from the request.""" diff --git a/MonteCarloMarginalizeCode/Code/test/jax/test_jax_av.py b/MonteCarloMarginalizeCode/Code/test/jax/test_jax_av.py index 58b8be1a4..bd9fdf1bd 100644 --- a/MonteCarloMarginalizeCode/Code/test/jax/test_jax_av.py +++ b/MonteCarloMarginalizeCode/Code/test/jax/test_jax_av.py @@ -28,6 +28,45 @@ def _driver_module(): return module +def test_waveform_precompute_kwargs_match_production_ile_controls(): + driver = _driver_module() + parser = driver.build_parser() + + opts, _ = parser.parse_args([ + "--approximant", "IMRPhenomXPHM", + "--internal-waveform-fd-L-frame", + "--internal-waveform-fd-no-condition", + ]) + got = driver._waveform_precompute_kwargs(opts) + assert got == { + "use_gwsignal": False, + "use_gwsignal_approx": None, + "ignore_threshold": None, + "no_memory": False, + "extra_waveform_kwargs": { + "fd_alignment_postevent_time": 2, + "e_freq": 1, + "fd_L_frame": True, + "no_condition": True, + }, + } + + defaults, _ = parser.parse_args(["--approximant", "IMRPhenomD"]) + assert driver._waveform_precompute_kwargs(defaults)[ + "extra_waveform_kwargs"] == { + "fd_alignment_postevent_time": 2, "e_freq": 1} + + +def test_internal_sample_rate_controls_precompute_cadence(): + driver = _driver_module() + parser = driver.build_parser() + defaults, _ = parser.parse_args(["--srate", "1024"]) + internal, _ = parser.parse_args( + ["--srate", "1024", "--srate-internal", "4096"]) + assert driver._analysis_delta_t(defaults) == 1.0 / 1024.0 + assert driver._analysis_delta_t(internal) == 1.0 / 4096.0 + + class _ToySkyLikelihood: ANGULAR_PARAM_ORDER = ("ra", "dec", "incl") @@ -77,6 +116,30 @@ def test_physical_coordinate_priors_are_normalized(): rtol=2e-6, atol=2e-6) +def test_pseudo_cosmo_distance_prior_density_and_draw_are_consistent(): + from RIFT.likelihood import priors_utils + + lo, hi = 1.0, 10000.0 + _, _, density = samplers._av_prior_spec( + "distMpc", lo, hi, distance_prior="pseudo_cosmo") + grid = np.linspace(lo, hi, 50001) + pdf = density(grid) + np.testing.assert_allclose(_trapezoid(pdf, grid), 1.0, rtol=2e-7) + norm = priors_utils.dist_prior_pseudo_cosmo_eval_norm(lo, hi) + np.testing.assert_allclose( + pdf, priors_utils.dist_prior_pseudo_cosmo(grid, nm=norm), rtol=1e-14) + + draw = samplers._av_prior_draw( + ("distMpc",), 30000, np.random.default_rng(240426), lo, hi, + distance_prior="pseudo_cosmo")[:, 0] + cdf = np.concatenate(([0.0], np.cumsum( + 0.5 * (pdf[1:] + pdf[:-1]) * np.diff(grid)))) + cdf /= cdf[-1] + expected_median = np.interp(0.5, cdf, grid) + assert abs(np.median(draw) - expected_median) < 0.015 * expected_median + assert np.all((draw >= lo) & (draw <= hi)) + + def test_sampling_window_does_not_renormalize_physical_prior(): bounds = {"ra": (1.17, 1.23), "dec": (0.27, 0.33)} ra_lo, ra_hi, ra_pdf = samplers._av_prior_spec( @@ -265,6 +328,33 @@ def test_pure_av_runs_inside_a_narrow_sky_sampling_window(): (result["theta"][:, 1] <= 0.5)) +def test_pure_av_never_evaluates_or_retains_points_outside_sampling_window(): + """A live bin at the upper edge must not extend beyond the declared box.""" + class OutwardRisingLikelihood: + ANGULAR_PARAM_ORDER = ("ra",) + + def __init__(self): + self.evaluated = [] + + def log_likelihood(self, ra): + values = np.asarray(ra, dtype=float) + self.evaluated.append(values.copy()) + # Force the retained live volume against the upper boundary, where + # fractional bin counts used to let the final bin overshoot. + return 2000.0 * values + + like = OutwardRisingLikelihood() + bounds = {"ra": (1.1, 1.3)} + result = samplers.adaptive_volume_sample( + like, 1.0, 100.0, sampler_method="AV", sample_bounds=bounds, + nmax=4000, neff=1000000, n_chunk=400, eval_chunk=128, seed=1409) + + evaluated = np.concatenate(like.evaluated) + assert np.all((evaluated >= 1.1) & (evaluated <= 1.3)) + assert np.all((result["theta"][:, 0] >= 1.1) & + (result["theta"][:, 0] <= 1.3)) + + def test_caller_supplied_oracle_cloud_bootstraps_portfolio(): rng = np.random.default_rng(31) centre = np.array([2.0, 0.2, 1.0]) @@ -317,6 +407,20 @@ def test_driver_exposes_sampler_as_an_orthogonal_backend(monkeypatch): assert opts.n_eff == 321 +def test_driver_accepts_pseudo_cosmo_only_for_av_backend(monkeypatch): + monkeypatch.delenv("JAX_ILE_DISTMARG_GH", raising=False) + driver = _driver_module() + parser = driver.build_parser() + opts, _ = parser.parse_args([ + "--sampler-method", "AV", "--d-prior", "pseudo_cosmo"]) + driver.check_critical_and_report(opts, parser) + + unsupported, _ = parser.parse_args([ + "--sampler-method", "AV", "--d-prior", "cosmo_sourceframe"]) + with pytest.raises(SystemExit): + driver.check_critical_and_report(unsupported, parser) + + def test_driver_validates_and_activates_av_sky_limits(monkeypatch): monkeypatch.delenv("JAX_ILE_DISTMARG_GH", raising=False) driver = _driver_module() diff --git a/MonteCarloMarginalizeCode/Code/test/jax/test_jax_banded_data_term.py b/MonteCarloMarginalizeCode/Code/test/jax/test_jax_banded_data_term.py new file mode 100644 index 000000000..7dcd8c2a9 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/jax/test_jax_banded_data_term.py @@ -0,0 +1,184 @@ +"""Regression tests for the compact banded ```` contraction. + +The production compound bank has many response basis elements. Expanding the +two basis/mode reductions as Python loops made the staged JAX program grow with +``A * K`` and dominated cold-start compilation. These tests pin both the +numerical contract and the requirement that changing ``A`` does not change the +top-level JAX program size. +""" + +import numpy as np +import pytest + +import jax +jax.config.update("jax_enable_x64", True) +import jax.numpy as jnp + +from RIFT.likelihood.jax_ile import core as JC + + +def _problem(A=5, K=3, S=2): + rng = np.random.default_rng(1905 + A + K) + nfull, npts, M = max(32, 8 + 5 * (S - 1) + 6 + 4), 6, 4 + + def complex_normal(shape): + return rng.normal(size=shape) + 1j * rng.normal(size=shape) + + q = jnp.asarray(complex_normal((A, nfull, K)), dtype=jnp.complex128) + conj_y = jnp.asarray(complex_normal((S, K)), dtype=jnp.complex128) + coeff = jnp.asarray(complex_normal((A, S)), dtype=jnp.complex128) + # Include distinct sub-sample positions for the two samples while staying + # away from the buffer edge for every interpolation stencil. + pos = jnp.asarray([[8.2 + 5 * i + j for j in range(npts)] + for i in range(S)], dtype=jnp.float64) + u = pos - jnp.floor(pos[:, :1]) - jnp.arange(npts)[None, :] + pp_t1 = jnp.asarray(np.arange(A) % M, dtype=jnp.int32) + pe = jnp.asarray(np.exp(1j * rng.normal(size=(M, S))), dtype=jnp.complex128) + pt = jnp.asarray(np.exp(1j * rng.normal(size=(M, npts))), dtype=jnp.complex128) + return q, conj_y, coeff, pos, u, pp_t1, pe, pt + + +def _python_oracle(q, conj_y, coeff, gather, pos, u, pp_t1=None, pe=None, pt=None): + """Literal pre-optimization A-by-K contraction, retained only as an oracle.""" + A, _, K = q.shape + out = jnp.zeros(pos.shape, dtype=jnp.complex128) + for a in range(A): + inner = jnp.zeros(pos.shape, dtype=jnp.complex128) + for k in range(K): + inner = inner + conj_y[:, k, None] * gather(q[a, :, k], pos, u) + if pp_t1 is None: + out = out + jnp.conj(coeff[a])[:, None] * inner + else: + im = int(pp_t1[a]) + out = out + ((jnp.conj(coeff[a]) * pe[im])[:, None] + * (pt[im][None, :] * inner)) + return out + + +@pytest.mark.parametrize("interp", ["nearest", "linear", "cubic", "sinc"]) +@pytest.mark.parametrize("post_phase", [False, True]) +def test_compact_contraction_matches_literal_oracle(interp, post_phase): + q, conj_y, coeff, pos, u, pp_t1, pe, pt = _problem() + gather = JC._GATHERERS[interp] + kwargs = dict(pp_t1=pp_t1, pe=pe, pt=pt) if post_phase else {} + expected = _python_oracle(q, conj_y, coeff, gather, pos, u, **kwargs) + got = JC._contract_banded_data_term( + q, conj_y, coeff, gather, pos, u, **kwargs) + got_jit = jax.jit( + lambda qq, cc: JC._contract_banded_data_term( + qq, conj_y, cc, gather, pos, u, **kwargs))(q, coeff) + np.testing.assert_allclose(np.asarray(got), np.asarray(expected), + rtol=2e-14, atol=2e-14) + np.testing.assert_allclose(np.asarray(got_jit), np.asarray(expected), + rtol=2e-14, atol=2e-14) + + +@pytest.mark.parametrize("interp", ["linear", "cubic"]) +def test_compact_contraction_retains_extrinsic_reverse_mode_ad(interp): + q, conj_y, coeff, pos, u, pp_t1, pe, pt = _problem() + gather = JC._GATHERERS[interp] + coeff_phase = jnp.linspace(-0.7, 0.9, coeff.shape[0])[:, None] + + def scalar(contract, pos_shift, angle): + shifted_pos = pos + pos_shift + # The wired separable-u path differentiates its common fractional p0 + # offset. We remain away from integer knots, so u + shift is identical + # to recomputing that offset in this neighborhood. + shifted_u = u + pos_shift + shifted_coeff = coeff * jnp.exp(1j * angle * coeff_phase) + value = contract( + q, conj_y, shifted_coeff, gather, shifted_pos, shifted_u, + pp_t1=pp_t1, pe=pe, pt=pt) + return jnp.real(jnp.sum(value * jnp.conj(value))) + + compact = lambda shift, angle: scalar( + JC._contract_banded_data_term, shift, angle) + literal = lambda shift, angle: scalar( + _python_oracle, shift, angle) + point = (jnp.asarray(0.031), jnp.asarray(-0.23)) + compact_value = compact(*point) + literal_value = literal(*point) + compact_grad = jax.grad(compact, argnums=(0, 1))(*point) + literal_grad = jax.grad(literal, argnums=(0, 1))(*point) + assert np.isfinite(float(compact_value)) + np.testing.assert_allclose(np.asarray(compact_value), np.asarray(literal_value), + rtol=2e-13, atol=2e-13) + np.testing.assert_allclose(np.asarray(compact_grad), np.asarray(literal_grad), + rtol=3e-12, atol=3e-12) + assert np.all(np.abs(np.asarray(compact_grad)) > 1e-8) + + +def test_jax_program_does_not_grow_with_number_of_bands(): + """A=40 must remain a loop bound, not become forty copies of the body.""" + + def trace(A): + q, conj_y, coeff, pos, u, pp_t1, pe, pt = _problem(A=A, K=2) + return jax.make_jaxpr( + lambda qq: JC._contract_banded_data_term( + qq, conj_y, coeff, JC._gather_cubic, pos, u, + pp_t1=pp_t1, pe=pe, pt=pt))(q).jaxpr + + small = trace(2) + production_sized = trace(40) + assert len(small.eqns) == len(production_sized.eqns) + assert any(eqn.primitive.name in ("scan", "while") + for eqn in production_sized.eqns) + + +def test_partial_post_phase_contract_is_rejected(): + q, conj_y, coeff, pos, u, pp_t1, _, _ = _problem() + with pytest.raises(ValueError, match="must be supplied together"): + JC._contract_banded_data_term( + q, conj_y, coeff, JC._gather_cubic, pos, u, pp_t1=pp_t1) + + +def test_chunked_rows_and_samples_preserve_padded_tail(monkeypatch): + # Both A*K and S are nonmultiples of their tile sizes. Repeated tail + # indices must not leak into the returned samples or their derivatives. + q, conj_y, coeff, pos, u, pp_t1, pe, pt = _problem(S=5) + gather = JC._gather_cubic + expected = _python_oracle( + q, conj_y, coeff, gather, pos, u, pp_t1, pe, pt) + monkeypatch.setattr(JC, "_banded_chunk_shape", + lambda *args: (2, 3, 0)) + + def contracted(offset): + return JC._contract_banded_data_term( + q, conj_y, coeff, gather, pos + offset, u + offset, + pp_t1=pp_t1, pe=pe, pt=pt) + + got = jax.jit(contracted)(0.0) + np.testing.assert_allclose(np.asarray(got), np.asarray(expected), + rtol=2e-14, atol=2e-14) + literal_grad = jax.grad(lambda x: jnp.real(jnp.sum( + jnp.abs(_python_oracle(q, conj_y, coeff, gather, pos + x, u + x, + pp_t1, pe, pt)) ** 2)))(0.031) + tiled_grad = jax.grad(lambda x: jnp.real(jnp.sum( + jnp.abs(contracted(x)) ** 2)))(0.031) + np.testing.assert_allclose(np.asarray(tiled_grad), np.asarray(literal_grad), + rtol=3e-12, atol=3e-12) + + +def test_chunk_shape_respects_forward_scratch_budget(): + budget = 2 * 1024**2 + rows, samples, estimated = JC._banded_chunk_shape( + A=100, K=20, S=2000, npts=128, nfull=4096, taps=16, + budget=budget) + assert 1 <= rows < 2000 + assert 1 <= samples < 2000 + assert estimated <= budget + with pytest.raises(ValueError, match="exceeds the scratch budget"): + JC._banded_chunk_shape(100, 20, 2000, 128, 4096, 16, budget=1024) + + +def test_empty_sample_batch_preserves_empty_result(): + q, conj_y, coeff, pos, u, pp_t1, pe, pt = _problem() + args = (q, conj_y[:0], coeff[:, :0], JC._gather_cubic, + pos[:0], u[:0]) + kwargs = dict(pp_t1=pp_t1, pe=pe[:, :0], pt=pt) + got = JC._contract_banded_data_term(*args, **kwargs) + got_jit = jax.jit(lambda: JC._contract_banded_data_term( + *args, **kwargs))() + assert got.shape == (0, pos.shape[1]) + assert got_jit.shape == got.shape + np.testing.assert_array_equal(np.asarray(got_jit), np.asarray(got)) diff --git a/MonteCarloMarginalizeCode/Code/test/jax/test_smc_evidence.py b/MonteCarloMarginalizeCode/Code/test/jax/test_smc_evidence.py new file mode 100644 index 000000000..777bcaa78 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/jax/test_smc_evidence.py @@ -0,0 +1,43 @@ +"""The SMC ladder normalizes over every prior walker, including zero likelihood.""" + +import numpy as np + + +def test_smc_evidence_counts_nonfinite_walkers_in_prior_mean(monkeypatch): + from RIFT.likelihood.jax_ile import samplers + + class ThreeAngleLike: + ANGULAR_PARAM_ORDER = ("ra", "dec", "incl") + + cloud = np.array([[0.1, 0.0, 0.5], [0.2, 0.0, 0.5], + [0.3, 0.0, 0.5], [0.4, 0.0, 0.5]]) + ln_likelihood = np.array([0.0, -1.0, -np.inf, -np.inf]) + monkeypatch.setattr(samplers, "sample_prior_3", lambda n, rng: cloud.copy()) + monkeypatch.setattr(samplers, "eval_lnL_3", + lambda like, theta, desc: ln_likelihood.copy()) + result = samplers.smc_puffball_sample( + ThreeAngleLike(), 1.0, 1000.0, n_walkers=4, n_move=0, + max_stages=1, max_dbeta=1.0, ess_frac=0.4, is_evidence=False, seed=1) + assert result["inv_T"] == 1.0 + assert np.isclose(result["logZ_laplace"], + np.log((1.0 + np.exp(-1.0)) / 4.0)) + + +def test_smc_all_finite_reference_is_unchanged(monkeypatch): + from RIFT.likelihood.jax_ile import samplers + + class ThreeAngleLike: + ANGULAR_PARAM_ORDER = ("ra", "dec", "incl") + + cloud = np.array([[0.1, 0.0, 0.5], [0.2, 0.0, 0.5], + [0.3, 0.0, 0.5], [0.4, 0.0, 0.5]]) + ln_likelihood = np.array([0.0, -1.0, -2.0, -3.0]) + monkeypatch.setattr(samplers, "sample_prior_3", lambda n, rng: cloud.copy()) + monkeypatch.setattr(samplers, "eval_lnL_3", + lambda like, theta, desc: ln_likelihood.copy()) + result = samplers.smc_puffball_sample( + ThreeAngleLike(), 1.0, 1000.0, n_walkers=4, n_move=0, + max_stages=1, max_dbeta=1.0, ess_frac=0.2, is_evidence=False, seed=1) + assert result["inv_T"] == 1.0 + assert np.isclose(result["logZ_laplace"], + np.log(np.mean(np.exp(ln_likelihood)))) diff --git a/MonteCarloMarginalizeCode/Code/test/test_gmm_truncated_score.py b/MonteCarloMarginalizeCode/Code/test/test_gmm_truncated_score.py new file mode 100644 index 000000000..8f23e02b0 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/test_gmm_truncated_score.py @@ -0,0 +1,67 @@ +"""The GMM score must describe its component-conditioned bounded draws.""" + +import numpy as np +from scipy.stats import multivariate_normal, norm + +from RIFT.integrators import gaussian_mixture_model as GMM + + +def test_component_truncated_score_matches_draws(): + model = GMM.gmm(2, np.array([[0.0, 1.0]])) + model.d = 1 + model.means = [np.array([0.0]), np.array([0.9])] + model.covariances = [np.array([[0.2**2]]), np.array([[0.5**2]])] + model.weights = np.array([0.5, 0.5]) + + component_mass = np.array([ + norm.cdf((1.0 - mu) / sd) - norm.cdf((-1.0 - mu) / sd) + for mu, sd in ((0.0, 0.2), (0.9, 0.5)) + ]) + assert component_mass[0] > 0.99 and component_mass[1] < 0.6 + + x = np.array([0.25, 0.5, 0.9, 0.97]) + y = 2.0 * x - 1.0 + expected = 2.0 * sum( + 0.5 * norm.pdf(y, loc=mu, scale=sd) / mass + for (mu, sd), mass in zip(((0.0, 0.2), (0.9, 0.5)), component_mass) + ) + np.testing.assert_allclose(model.score(x[:, None]), expected, rtol=1e-12) + + # This bin is supplied mainly by the second, more strongly truncated + # component. Verify that sample() actually follows the component-wise + # conditional law that score() now reports. + rng_state = np.random.get_state() + try: + np.random.seed(19) + draws = np.asarray(model.sample(10000)).reshape(-1) + finally: + np.random.set_state(rng_state) + empirical = np.mean((draws >= 0.85) & (draws <= 0.95)) + expected_bin = 0.5 * sum( + (norm.cdf((2.0 * 0.95 - 1.0 - mu) / sd) + - norm.cdf((2.0 * 0.85 - 1.0 - mu) / sd)) / mass + for (mu, sd), mass in zip(((0.0, 0.2), (0.9, 0.5)), component_mass) + ) + assert abs(empirical - expected_bin) < 0.015 + + +def test_multivariate_component_truncation_score(): + model = GMM.gmm(2, np.array([[0.0, 1.0], [0.0, 1.0]])) + model.d = 2 + model.means = [np.array([0.0, 0.0]), np.array([0.8, 0.8])] + model.covariances = [np.eye(2) * 0.25**2, np.eye(2) * 0.6**2] + model.weights = np.array([0.5, 0.5]) + + mass = [ + (norm.cdf((1.0 - mu) / sd) + - norm.cdf((-1.0 - mu) / sd))**2 + for mu, sd in ((0.0, 0.25), (0.8, 0.6)) + ] + assert mass[0] > 0.99 and mass[1] < 0.5 + x = np.array([[0.5, 0.5], [0.9, 0.9]]) + y = 2.0 * x - 1.0 + expected = 4.0 * sum( + 0.5 * multivariate_normal.pdf(y, mean=[mu, mu], cov=np.eye(2) * sd**2) / c + for (mu, sd), c in zip(((0.0, 0.25), (0.8, 0.6)), mass) + ) + np.testing.assert_allclose(model.score(x), expected, rtol=1e-7) diff --git a/MonteCarloMarginalizeCode/Code/test/test_gpu_jax_handoff.py b/MonteCarloMarginalizeCode/Code/test/test_gpu_jax_handoff.py new file mode 100644 index 000000000..74d25228c --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/test_gpu_jax_handoff.py @@ -0,0 +1,237 @@ +"""Structural and downstream parity tests for the direct GPU-to-JAX adapter.""" +import os +from types import SimpleNamespace + +os.environ.setdefault("JAX_ENABLE_X64", "1") + +import jax +import jax.numpy as jnp +import lal +import numpy as np +import pytest + +from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr +from RIFT.likelihood import slowrot_freqresponse as sfr +from RIFT.likelihood.gpu_jax_handoff import ( + build_jax_rotating_freqresponse_data_from_device, +) +from RIFT.likelihood.gpu_precompute import pack_device_precompute +from RIFT.likelihood.jax_ile.banded import build_rotating_freqresponse_data +from RIFT.likelihood.jax_ile.core import fused_log_likelihood + +jax.config.update("jax_enable_x64", True) + + +def _fixture(seed=71): + rng = np.random.default_rng(seed) + modes = [(2, 2), (2, -2)] + a_list = fr.compound_index_set(0, 0) + A, K, N = len(a_list), len(modes), 128 + q = rng.normal(size=(A, K, N)) + 1j*rng.normal(size=(A, K, N)) + u = rng.normal(size=(A, A, K, K)) + 1j*rng.normal(size=(A, A, K, K)) + v = rng.normal(size=(A, A, K, K)) + 1j*rng.normal(size=(A, A, K, K)) + # The two builders need only identical arrays for this routing test; the + # likelihood need not represent a physical positive norm. + packed = dict(q={"H1": jnp.asarray(q)}, U={"H1": jnp.asarray(u)}, + V={"H1": jnp.asarray(v)}, epoch={"H1": 1000.-64/1024.}, + delta_t=1/1024., modes=modes, a_list=a_list) + meta = dict(feature="rotation_freqresponse", gpu_precompute=True, + device_resident=True, post_phase_required=True, + event_time_geo=1000., modes=modes, a_list=a_list, + Qmax=0, p_max=0, f_sidereal=1.160576e-5) + geom = {"H1": sfr.detector_geometry("H1", L_arm=4000.)} + tvals = (np.arange(9)-4)/1024. + return packed, meta, geom, tvals, q, u, v + + +def test_device_builder_matches_existing_banded_builder_downstream(): + packed, meta, geom, tvals, q, u, v = _fixture() + direct = build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, geom, require_gpu=False) + + a_list = meta["a_list"] + lookup = {"H1": np.asarray(meta["modes"], dtype=int)} + rho = {"H1": {a: q[i] for i, a in enumerate(a_list)}} + U = {"H1": {(a, ap): u[i, j] for i, a in enumerate(a_list) + for j, ap in enumerate(a_list)}} + V = {"H1": {(a, ap): v[i, j] for i, a in enumerate(a_list) + for j, ap in enumerate(a_list)}} + conventional = build_rotating_freqresponse_data( + meta, lookup, rho, U, V, packed["epoch"], packed["delta_t"], + tvals, geom) + + args = [jnp.asarray(x) for x in ( + [1.2, 1.21], [0.3, 0.31], [0.2, 0.4], + [0.7, 0.8], [0.1, 0.5], [100., 120.])] + got = fused_log_likelihood(direct, *args, interp="nearest") + expected = fused_log_likelihood(conventional, *args, interp="nearest") + np.testing.assert_allclose(np.asarray(got), np.asarray(expected), + rtol=2e-13, atol=2e-13) + assert direct.detectors["H1"]["Q_bank"].shape == (len(a_list), 128, 2) + q_devices = direct.detectors["H1"]["Q_bank"].devices() + for key in ("location", "response", "x_arm", "y_arm"): + assert direct.detectors["H1"][key].devices() == q_devices + assert direct.gpu_handoff["contract_Q_U_V_host_copies"] == 0 + zero_packed = dict(packed) + zero_packed["q"] = {"H1": jnp.zeros_like(packed["q"]["H1"])} + zero_data = build_jax_rotating_freqresponse_data_from_device( + zero_packed, meta, tvals, geom, require_gpu=False) + zero_lnL = fused_log_likelihood(zero_data, *args, interp="nearest") + assert np.max(np.abs(np.asarray(got-zero_lnL))) > 1e-6, \ + "fixture did not exercise the sampled Q data term" + + +def test_two_detector_unequal_arm_handoff_matches_host_and_classic_cubic(): + """Keep detector banks and geometry distinct through both adapter routes.""" + rng = np.random.default_rng(88291) + detectors = ("H1", "L1") + arm_lengths = {"H1": 40000.0, "L1": 20000.0} + qmax = pmax = 1 + modes = [(2, 2), (2, -2)] + a_list = fr.compound_index_set(qmax, pmax) + A, K, N = len(a_list), len(modes), 256 + delta_t = 1.0 / 1024.0 + tref = 1000000000.0 + epochs = {"H1": tref - 128 * delta_t, + "L1": tref - 127 * delta_t} + tvals = np.arange(-4, 5) * delta_t + + # Independent detector arrays make an accidental H1/L1 alias or overwrite + # observable. The Q support is centred on the arrival-time stencil. + q = {}; u = {}; v = {} + for det in detectors: + q[det] = (rng.normal(size=(A, K, N)) + + 1j * rng.normal(size=(A, K, N))) + u[det] = (rng.normal(size=(A, A, K, K)) + + 1j * rng.normal(size=(A, A, K, K))) + v[det] = (rng.normal(size=(A, A, K, K)) + + 1j * rng.normal(size=(A, A, K, K))) + + meta = dict(feature="rotation_freqresponse", gpu_precompute=True, + device_resident=True, post_phase_required=True, + event_time_geo=tref, modes=modes, a_list=a_list, + Qmax=qmax, p_max=pmax, f_sidereal=fr.flwr.F_SIDEREAL, + L=dict(arm_lengths), L_arm=dict(arm_lengths)) + geometry = {det: sfr.detector_geometry(det, L_arm=arm_lengths[det]) + for det in detectors} + packed = dict(q={det: jnp.asarray(q[det]) for det in detectors}, + U={det: jnp.asarray(u[det]) for det in detectors}, + V={det: jnp.asarray(v[det]) for det in detectors}, + epoch=dict(epochs), delta_t=delta_t, + modes=modes, a_list=a_list) + + ra = np.asarray([0.7, 5.8]) + dec = np.asarray([-0.2, 0.85]) + psi = np.asarray([0.3, 2.7]) + incl = np.asarray([0.6, 2.5]) + phiref = np.asarray([0.4, 5.0]) + dist_mpc = np.asarray([100.0, 900.0]) + jax_args = tuple(jnp.asarray(x) for x in + (ra, dec, psi, incl, phiref, dist_mpc)) + params = SimpleNamespace( + phi=ra, theta=dec, psi=psi, incl=incl, phiref=phiref, + dist=dist_mpc * 1.0e6 * lal.PC_SI, tref=tref, deltaT=delta_t) + + def compare(selected): + packed_part = dict( + packed, + q={det: packed["q"][det] for det in selected}, + U={det: packed["U"][det] for det in selected}, + V={det: packed["V"][det] for det in selected}, + epoch={det: epochs[det] for det in selected}) + geom_part = {det: geometry[det] for det in selected} + lookup = {det: np.asarray(modes, dtype=int) for det in selected} + rho = {det: {a: q[det][i] for i, a in enumerate(a_list)} + for det in selected} + U = {det: u[det] for det in selected} + V = {det: v[det] for det in selected} + epoch = {det: epochs[det] for det in selected} + + direct = build_jax_rotating_freqresponse_data_from_device( + packed_part, meta, tvals, geom_part, require_gpu=False) + host = build_rotating_freqresponse_data( + meta, lookup, rho, U, V, epoch, delta_t, tvals, geom_part) + + direct_t = np.asarray(fused_log_likelihood( + direct, *jax_args, interp="cubic", return_lnLt=True)) + host_t = np.asarray(fused_log_likelihood( + host, *jax_args, interp="cubic", return_lnLt=True)) + classic_t = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, params, meta, lookup, rho, U, V, epoch, Lmax=2, + array_output=True, time_interp="cubic") + np.testing.assert_allclose(direct_t, host_t, rtol=2e-13, atol=2e-13) + np.testing.assert_allclose(direct_t, classic_t, rtol=2e-13, atol=2e-13) + + direct_marg = np.asarray(fused_log_likelihood( + direct, *jax_args, interp="cubic")) + host_marg = np.asarray(fused_log_likelihood( + host, *jax_args, interp="cubic")) + classic_marg = fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, params, meta, lookup, rho, U, V, epoch, Lmax=2, + array_output=False, time_interp="cubic") + np.testing.assert_allclose(direct_marg, host_marg, + rtol=2e-13, atol=2e-13) + np.testing.assert_allclose(direct_marg, classic_marg, + rtol=2e-13, atol=2e-13) + for det in selected: + assert direct.detectors[det]["L_arm"] == arm_lengths[det] + return direct_t, direct_marg + + h1_t, _ = compare(("H1",)) + l1_t, _ = compare(("L1",)) + network_t, _ = compare(detectors) + np.testing.assert_allclose(network_t, h1_t + l1_t, + rtol=2e-13, atol=2e-13) + assert not np.allclose(h1_t, l1_t), "detector fixtures are not distinguishable" + + zero_packed = dict( + packed, q={det: jnp.zeros_like(packed["q"][det]) for det in detectors}) + zero_data = build_jax_rotating_freqresponse_data_from_device( + zero_packed, meta, tvals, geometry, require_gpu=False) + zero_t = np.asarray(fused_log_likelihood( + zero_data, *(x[:1] for x in jax_args), interp="cubic", return_lnLt=True)) + assert np.max(np.abs(network_t[:1] - zero_t)) > 1e-6, \ + "fixture did not exercise the centred two-detector Q data term" + + +def test_handoff_fails_closed_for_host_arrays_cpu_jax_and_pregrid(): + packed, meta, geom, tvals, q, u, v = _fixture() + if all(d.platform != "gpu" for d in jax.devices()): + with pytest.raises(RuntimeError, match="non-GPU"): + build_jax_rotating_freqresponse_data_from_device(packed, meta, tvals, geom) + host = dict(packed) + host["q"] = {"H1": q} + host["U"] = {"H1": u} + host["V"] = {"H1": v} + with pytest.raises(TypeError, match="device-resident"): + build_jax_rotating_freqresponse_data_from_device( + host, meta, tvals, geom, require_gpu=False) + for bad in (1.5, 2): + with pytest.raises((ValueError, NotImplementedError)): + build_jax_rotating_freqresponse_data_from_device( + packed, meta, tvals, geom, q_time_pregrid_factor=bad, + require_gpu=False) + + +def test_classic_ile_device_pack_is_view_only_and_dense(): + packed, meta, _geom, _tvals, _q, _u, _v = _fixture() + lookup, rho, U, V, epoch = pack_device_precompute( + packed, meta, require_gpu=False) + assert lookup["H1"].tolist() == [[2, 2], [2, -2]] + for i, a in enumerate(meta["a_list"]): + np.testing.assert_array_equal(np.asarray(rho["H1"][a]), + np.asarray(packed["q"]["H1"][i])) + assert U["H1"] is packed["U"]["H1"] + assert V["H1"] is packed["V"]["H1"] + assert epoch == packed["epoch"] and epoch is not packed["epoch"] + + bad = dict(packed) + bad["U"] = {"H1": packed["U"]["H1"][1:]} + with pytest.raises(ValueError, match="U/V"): + pack_device_precompute(bad, meta, require_gpu=False) + bad_dt = dict(packed, delta_t=0.0) + with pytest.raises(ValueError, match="delta_t"): + pack_device_precompute(bad_dt, meta, require_gpu=False) + if all(d.platform != "gpu" for d in jax.devices()): + with pytest.raises(RuntimeError, match="not on a GPU"): + pack_device_precompute(packed, meta) diff --git a/MonteCarloMarginalizeCode/Code/test/test_jax_fairdraw_postprocess.py b/MonteCarloMarginalizeCode/Code/test/test_jax_fairdraw_postprocess.py new file mode 100644 index 000000000..bcfa77d32 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/test_jax_fairdraw_postprocess.py @@ -0,0 +1,122 @@ +"""Strict conversion of JAX-ILE tabular fair draws.""" + +import subprocess +import sys +import json +from pathlib import Path + +import numpy as np + + +CODE = Path(__file__).resolve().parents[1] +CONVERTER = CODE / "bin" / "util_ConvertJAXILEFairdraws.py" + + +def _write_pair(directory, event, draws=2, spins=None): + stem = "EXTR_out-{}.xml_0".format(event) + spins = [0, 0, 0, 0, 0, 0] if spins is None else spins + intrinsic = np.array([[event, 30.0, 20.0] + list(spins) + [ + 12.0, 0.1, 1000, 25]]) + np.savetxt(directory / (stem + "_.dat"), intrinsic) + values = np.column_stack([ + np.linspace(1.0, 1.1, draws), np.linspace(0.2, 0.3, draws), + np.linspace(300, 320, draws), np.linspace(0.4, 0.5, draws), + np.linspace(0.6, 0.7, draws), np.linspace(0.8, 0.9, draws), + np.linspace(10, 11, draws), + ]) + np.savetxt(directory / (stem + "_samples.dat"), values, + header="right_ascension declination distance inclination psi phi_orb loglikelihood") + + +def test_converter_joins_intrinsic_and_extrinsic_rows(tmpdir): + tmp_path = Path(str(tmpdir)) + _write_pair(tmp_path, 0) + _write_pair(tmp_path, 1) + output = tmp_path / "posterior.dat" + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--draws-per-intrinsic", "2", "--expected-intrinsic", "2", + "--output", str(output), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode == 0, result.stderr + assert " time " not in (" " + output.read_text().splitlines()[0] + " ") + table = np.loadtxt(output) + assert table.shape == (4, 25) + assert np.allclose(table[:, 0], 30.0) + assert np.allclose(table[:, 1], 20.0) + assert np.allclose(table[:, 20], 25.0) + provenance = json.loads((tmp_path / "posterior.provenance.json").read_text()) + assert provenance["intrinsic_points"] == 2 + assert provenance["posterior_rows"] == 4 + assert provenance["equal_weight_columns"] == {"p": 1.0, "ps": 1.0} + assert provenance["omitted_unavailable_coordinates"] == [ + "time", "redshift", "source_frame_masses"] + + +def test_converter_refuses_wrong_draw_count(tmpdir): + tmp_path = Path(str(tmpdir)) + _write_pair(tmp_path, 0, draws=1) + output = tmp_path / "bad.dat" + provenance = tmp_path / "bad.provenance.json" + output.write_text("stale\n") + provenance.write_text("stale\n") + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--draws-per-intrinsic", "2", "--output", str(output), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode != 0 + assert "expected 2 rows" in result.stderr + assert not output.exists() + assert not provenance.exists() + + +def test_converter_refuses_intrinsic_record_without_samples(tmpdir): + tmp_path = Path(str(tmpdir)) + _write_pair(tmp_path, 0) + _write_pair(tmp_path, 1) + (tmp_path / "EXTR_out-1.xml_0_samples.dat").unlink() + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--draws-per-intrinsic", "2", "--output", str(tmp_path / "bad.dat"), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode != 0 + assert "records_without_samples" in result.stderr + + +def test_converter_refuses_missing_both_and_stale_excess_pairs(tmpdir): + tmp_path = Path(str(tmpdir)) + _write_pair(tmp_path, 0) + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--expected-intrinsic", "2", "--output", str(tmp_path / "missing.dat"), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode != 0 + assert "missing=[1]" in result.stderr + + _write_pair(tmp_path, 1) + _write_pair(tmp_path, 2) + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--expected-intrinsic", "2", "--output", str(tmp_path / "excess.dat"), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode != 0 + assert "excess=[2]" in result.stderr + + +def test_converter_computes_spin_summaries(tmpdir): + tmp_path = Path(str(tmpdir)) + spins = [0.3, 0.4, 0.2, 0.0, 0.6, -0.1] + _write_pair(tmp_path, 0, spins=spins) + output = tmp_path / "spinning.dat" + result = subprocess.run([ + sys.executable, str(CONVERTER), "--directory", str(tmp_path), + "--draws-per-intrinsic", "2", "--expected-intrinsic", "1", + "--output", str(output), + ], universal_newlines=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + assert result.returncode == 0, result.stderr + table = np.loadtxt(output) + assert np.allclose(table[:, 23], (30 * 0.2 + 20 * -0.1) / 50) + q = 20.0 / 30.0 + expected_chip = max((2 + 1.5 * q) * 30**2 * 0.5, + (2 + 1.5 / q) * 20**2 * 0.6) / ((2 + 1.5 * q) * 30**2) + assert np.allclose(table[:, 24], expected_chip) diff --git a/MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py b/MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py index e264eabca..917ade0e2 100644 --- a/MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py +++ b/MonteCarloMarginalizeCode/Code/test/test_jax_ile_selectable.py @@ -168,6 +168,10 @@ def test_use_jax_ile_threads_into_every_ile_stage_sub(tmp_path): line = _executable_line(rundir / sub) assert line.endswith(JAX_EXE), (sub, line) assert "integrate_likelihood_extrinsic_batchmode" not in line + converter = (rundir / "allinone_convert.sh").read_text() + assert "util_ConvertJAXILEFairdraws.py" in converter + assert "--expected-intrinsic" in converter + assert "util_JoinExtrXML.py" not in converter def test_explicit_ile_exe_path_is_used_verbatim(tmp_path): diff --git a/MonteCarloMarginalizeCode/Code/test/test_jax_template_finalization.py b/MonteCarloMarginalizeCode/Code/test/test_jax_template_finalization.py new file mode 100644 index 000000000..51b31670a --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/test_jax_template_finalization.py @@ -0,0 +1,89 @@ +"""Regression tests for the JAX driver's intrinsic-template finalization. + +The executable is intentionally not imported: it parses command-line options at +module import. Extracting the real function by AST exercises its implementation +without constructing a fake copy of the logic under test. +""" +import ast +from pathlib import Path +from types import SimpleNamespace +import sys + +import lal +import lalsimulation as lalsim +import numpy as np + +from RIFT import lalsimutils + + +DRIVER = (Path(__file__).resolve().parents[1] / + "bin" / "integrate_likelihood_extrinsic_jax") + + +def _load_templates_function(): + source = DRIVER.read_text() + tree = ast.parse(source, filename=str(DRIVER)) + fn = next(node for node in tree.body + if isinstance(node, ast.FunctionDef) and node.name == "load_templates") + module = ast.Module(body=[fn], type_ignores=[]) + ast.fix_missing_locations(module) + namespace = { + "np": np, "sys": sys, "lalsimutils": lalsimutils, "lalsim": lalsim, + "MSUN": lal.MSUN_SI, "PC": lal.PC_SI, + } + exec(compile(module, str(DRIVER), "exec"), namespace) + return namespace["load_templates"] + + +def _opts(**updates): + values = dict( + reference_freq=100.0, fmin_template=20.0, approximant="TaylorF2", + sim_xml=None, sim_grid=None, random_event=False, event=0, + n_events_to_analyze=1, mass1=None, mass2=None, + ) + values.update(updates) + return SimpleNamespace(**values) + + +def _assert_finalized(P, m1, m2, s1z, lambda1, epoch, delta_f, delta_t): + assert P.phiref == 0.0 + assert P.psi == 0.0 + assert P.incl == 0.0 + np.testing.assert_allclose([P.m1, P.m2], [m1, m2], rtol=2e-7) + np.testing.assert_allclose(P.s1z, s1z, rtol=0, atol=1e-12) + np.testing.assert_allclose(P.lambda1, lambda1, rtol=2e-7) + assert float(P.tref) == float(epoch) + assert P.deltaF == delta_f and P.deltaT == delta_t + assert P.dist == 1000.0e6 * lal.PC_SI + + +def test_xml_template_zeroes_extrinsics_and_preserves_intrinsics(tmp_path): + original = lalsimutils.ChooseWaveformParams( + m1=1.45 * lal.MSUN_SI, m2=1.22 * lal.MSUN_SI, + s1z=0.031, s2z=-0.014, lambda1=527.0, lambda2=811.0, + phiref=0.73, psi=1.17, incl=2.02) + base = tmp_path / "nonzero-extrinsics" + lalsimutils.ChooseWaveformParams_array_to_xml([original], str(base)) + xml = str(base) + ".xml.gz" + # Compare to serialized values: XML spin fields have finite precision. + serialized = lalsimutils.xml_to_ChooseWaveformParams_array(xml)[0] + assert serialized.phiref != 0 and serialized.psi != 0 and serialized.incl != 0 + + load_templates = _load_templates_function() + epoch, delta_f, delta_t = 1000000000.25, 0.25, 1.0/4096 + got = load_templates(_opts(sim_xml=xml), epoch, delta_f, delta_t)[0] + _assert_finalized(got, serialized.m1, serialized.m2, serialized.s1z, + serialized.lambda1, epoch, delta_f, delta_t) + + +def test_grid_template_zeroes_extrinsics_and_preserves_intrinsics(tmp_path): + grid = tmp_path / "nonzero-extrinsics-grid.dat" + grid.write_text( + "m1 m2 s1z s2z lambda1 lambda2 phiref psi incl\n" + "1.47 1.19 0.027 -0.011 493 902 0.61 1.03 2.21\n") + + load_templates = _load_templates_function() + epoch, delta_f, delta_t = 1000000001.5, 0.125, 1.0/8192 + got = load_templates(_opts(sim_grid=str(grid)), epoch, delta_f, delta_t)[0] + _assert_finalized(got, 1.47*lal.MSUN_SI, 1.19*lal.MSUN_SI, + 0.027, 493.0, epoch, delta_f, delta_t) diff --git a/MonteCarloMarginalizeCode/Code/test/test_response_order.py b/MonteCarloMarginalizeCode/Code/test/test_response_order.py new file mode 100644 index 000000000..116d4b095 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/test_response_order.py @@ -0,0 +1,156 @@ +import unittest +from unittest import mock +import importlib.util +from pathlib import Path +import sys +import types + +import numpy as np + +_MODULE = (Path(__file__).resolve().parents[1] / "RIFT" / "likelihood" + / "response_order.py") +_rift = types.ModuleType("RIFT") +_likelihood = types.ModuleType("RIFT.likelihood") +_likelihood.__path__ = [] +_fl = types.ModuleType("RIFT.likelihood.factored_likelihood") +_fl.ComputeYlmsArrayVector = lambda modes, incl, phase: np.ones( + (len(modes), len(incl)), dtype=complex) +sys.modules.setdefault("RIFT", _rift) +sys.modules.setdefault("RIFT.likelihood", _likelihood) +sys.modules.setdefault("RIFT.likelihood.factored_likelihood", _fl) +_SPEC = importlib.util.spec_from_file_location( + "RIFT.likelihood.response_order", _MODULE) +response_order = importlib.util.module_from_spec(_SPEC) +sys.modules[_SPEC.name] = response_order +_SPEC.loader.exec_module(response_order) + + +class ResponseOrderTest(unittest.TestCase): + def products(self): + meta = dict(Qmax=2, p_list=[0, 1, 2, 3], modes=[(2, 2)], + event_time_geo=1000000000.0, L_arm=None) + diagonal = [1.0, 1.0e-2, 1.0e-4, 1.0e-8] + U = {'H1': {}} + V = {'H1': {}} + for p in meta['p_list']: + for pp in meta['p_list']: + U['H1'][(p, pp)] = np.array( + [[diagonal[p] if p == pp else 0.0]], dtype=complex) + V['H1'][(p, pp)] = np.zeros((1, 1), dtype=complex) + primary = {'H1': {p: 'p%d' % p for p in meta['p_list']}} + return primary, U, V, primary, meta + + def test_snr_tightens_selected_order(self): + products = self.products() + coeff = lambda meta, det, ra, dec, psi: ( + np.ones((len(ra), 4), dtype=complex), + np.ones((len(ra), 4), dtype=complex)) + with mock.patch.object(response_order, '_coefficients', coeff): + low = response_order.estimate_response_orders( + products[4], products[1], products[2], 10.0, + lnL_tolerance=0.1, n_samples=8, selected_q=0, + vary_p=False, vary_q=True) + high = response_order.estimate_response_orders( + products[4], products[1], products[2], 100.0, + lnL_tolerance=0.1, n_samples=8, selected_q=0, + vary_p=False, vary_q=True) + self.assertEqual(low['chosen']['Qmax'], 0) + self.assertEqual(high['chosen']['Qmax'], 1) + self.assertTrue(low['selected_passes']) + self.assertFalse(high['selected_passes']) + + def test_truncate_products(self): + out = response_order.truncate_precompute_products( + self.products(), p_max=0, q_max=0) + self.assertEqual(out[4]['p_list'], [0, 1]) + self.assertEqual(out[4]['Qmax'], 0) + self.assertEqual(set(out[0]['H1']), {0, 1}) + self.assertEqual(set(out[1]['H1']), {(0, 0), (0, 1), (1, 0), (1, 1)}) + + def test_rotation_truncation_removes_reference_only_harmonics(self): + indices = [(p, n) for p in (0, 1) for n in range(-3, 4)] + meta = dict(p_max=1, harmonics=tuple(range(-3, 4)), a_list=indices, + modes=[(2, 2)], event_time_geo=1000000000.0, + f_sidereal=1.0, post_phase_required=True) + primary = {'H1': {a: a for a in indices}} + cross = {'H1': {(a, ap): {(2, 2): 0j} for a in indices for ap in indices}} + out = response_order.truncate_precompute_products( + (primary, cross, cross, primary, meta), p_max=0, q_max=0) + self.assertEqual(out[4]['a_list'], [(0, n) for n in range(-2, 3)]) + self.assertEqual(out[4]['harmonics'], tuple(range(-2, 3))) + + def test_pack_uv_from_raw_does_not_touch_data_bank(self): + products = self.products() + raw_u = {'H1': {(p, pp): {((2, 2), (2, 2)): values[0, 0]} + for (p, pp), values in products[1]['H1'].items()}} + raw_v = {'H1': {(p, pp): {((2, 2), (2, 2)): values[0, 0]} + for (p, pp), values in products[2]['H1'].items()}} + U, V = response_order.pack_uv_from_raw(products[4], raw_u, raw_v) + self.assertEqual(U['H1'][(1, 1)].shape, (1, 1)) + self.assertEqual(U['H1'][(1, 1)][0, 0], 1.0e-2) + self.assertEqual(V['H1'][(2, 2)][0, 0], 0.0) + + def test_angular_design_spans_independent_coordinates(self): + design = np.asarray(response_order._angular_design(256)) + unit = np.vstack((design[0] / (2 * np.pi), + (np.sin(design[1]) + 1) / 2, + (np.cos(design[2]) + 1) / 2, + design[3] / np.pi, + design[4] / (2 * np.pi))) + corr = np.corrcoef(unit) + self.assertLess(np.max(np.abs(corr - np.eye(5))), 0.15) + + def test_reference_bank_guard_refuses_large_compound_bank(self): + size = response_order.reference_bank_size( + 'combined', p_max=4, q_max=10, lmax=4, n_detectors=5) + self.assertEqual(size['basis'], 1090) + self.assertGreater(size['uv_gib'], 10.0) + with self.assertRaises(ValueError): + response_order.guard_reference_bank( + 'combined', 4, 10, 4, 5, max_bank_gib=4.0) + + def test_one_shell_is_not_a_resolved_reference(self): + products = response_order.truncate_precompute_products( + self.products(), p_max=0, q_max=1) + coeff = lambda meta, det, ra, dec, psi: ( + np.ones((len(ra), len(meta['p_list'])), dtype=complex), + np.ones((len(ra), len(meta['p_list'])), dtype=complex)) + with mock.patch.object(response_order, '_coefficients', coeff): + report = response_order.estimate_response_orders( + products[4], products[1], products[2], 10.0, + lnL_tolerance=0.1, n_samples=8, selected_q=0, + vary_p=False, vary_q=True) + self.assertFalse(report['reference_resolved']) + + def test_combined_one_axis_resolution_ignores_unvaried_axis(self): + indices = [(b, 0, 0) for b in range(4)] + meta = dict(feature='rotation_freqresponse', p_max=0, Qmax=2, + a_list=indices, modes=[(2, 2)], + event_time_geo=1000000000.0, f_sidereal=1.0) + diagonal = [1.0, 1.0e-2, 1.0e-4, 1.0e-8] + U = {'H1': {}} + V = {'H1': {}} + for i, a in enumerate(indices): + for j, ap in enumerate(indices): + U['H1'][(a, ap)] = np.array( + [[diagonal[i] if i == j else 0.0]], dtype=complex) + V['H1'][(a, ap)] = np.zeros((1, 1), dtype=complex) + coeff = lambda meta, det, ra, dec, psi: ( + np.ones((len(ra), len(indices)), dtype=complex), + np.ones((len(ra), len(indices)), dtype=complex)) + with mock.patch.object(response_order, '_coefficients', coeff): + report = response_order.estimate_response_orders( + meta, U, V, 10.0, n_samples=8, selected_p=0, + selected_q=0, vary_p=False, vary_q=True) + self.assertTrue(report['reference_resolved']) + self.assertEqual([row['axis'] for row in report['reference_shells']], ['q']) + for row in report['rows']: + expected = (np.sqrt(row['finite_reference_mu']) + + np.sqrt(report['reference_tail_mu'])) ** 2 + self.assertAlmostEqual(row['max_mu'], expected) + self.assertGreater(report['rows'][0]['max_mu'], + report['rows'][0]['finite_reference_mu']) + + +if __name__ == '__main__': + unittest.main() diff --git a/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_legacy_compat.py b/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_legacy_compat.py new file mode 100644 index 000000000..002ea23f3 --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_legacy_compat.py @@ -0,0 +1,402 @@ +"""Regression coverage for legacy CPU waveforms in GPU precompute. + +The GPU path must default to RIFT's existing ``internal_hlm_generator`` for +every approximant and waveform option, then upload both ordinary and +conjugated conditioned mode banks. Native waveform providers are opt-in. +""" + +import importlib.util +import os +import unittest +from unittest import mock + +import numpy as np + + +def _imports(): + import lal + import lalsimulation as lalsim + from RIFT import lalsimutils as lsu + from RIFT.likelihood import factored_likelihood as fl + from RIFT.likelihood import factored_likelihood_rotating_freqresponse as fr + from RIFT.likelihood import gpu_precompute as gpu + return lal, lalsim, lsu, fl, fr, gpu + + +def _fd_series(lal, name, values, df, epoch): + out = lal.CreateCOMPLEX16FrequencySeries( + name, lal.LIGOTimeGPS(epoch), 0.0, df, lal.DimensionlessUnit, len(values) + ) + out.data.data[:] = values + return out + + +def _short_problem(approximant, precessing=False): + lal, lalsim, lsu, _, _, _ = _imports() + event, dt, df = 1.0e9, 1.0 / 512.0, 0.25 + p = lsu.ChooseWaveformParams( + m1=30 * lal.MSUN_SI, m2=25 * lal.MSUN_SI, + fmin=40.0, fref=60.0, deltaT=dt, deltaF=df, + approx=approximant, radec=True, phi=1.2, theta=0.3, + incl=0.7, psi=0.5, phiref=0.4, tref=event, + dist=200e6 * lal.PC_SI, detector="H1", + ) + if precessing: + p.s1x, p.s1y, p.s1z = 0.2, -0.04, 0.1 + p.s2x, p.s2y, p.s2z = 0.03, 0.1, -0.15 + data = {"H1": lsu.non_herm_hoff(p.manual_copy())} + n = data["H1"].data.length + psd = lal.CreateREAL8FrequencySeries( + "H1 PSD", lal.LIGOTimeGPS(0), 0.0, df, lal.SecondUnit, n // 2 + 1 + ) + frequencies = np.arange(n // 2 + 1) * df + psd.data.data[:] = [ + lalsim.SimNoisePSDaLIGOZeroDetHighPower(max(10.0, f)) + for f in frequencies + ] + return event, p, data, {"H1": psd} + + +def _packed(fr, result): + return fr.pack_rotating_freqresponse_arrays( + result[4], result[3], result[1], result[2] + ) + + +def _assert_precompute_close(testcase, fr, cpu, candidate, rtol): + cp, gp = _packed(fr, cpu), _packed(fr, candidate) + testcase.assertEqual(cpu[4]["modes"], candidate[4]["modes"]) + testcase.assertEqual(cpu[4]["a_list"], candidate[4]["a_list"]) + for det in cp[1]: + for a in cp[1][det]: + np.testing.assert_allclose(gp[1][det][a], cp[1][det][a], + rtol=rtol, atol=1e-9) + np.testing.assert_allclose(gp[2][det], cp[2][det], rtol=rtol, atol=1e-8) + np.testing.assert_allclose(gp[3][det], cp[3][det], rtol=rtol, atol=1e-8) + testcase.assertAlmostEqual(gp[4][det], cp[4][det], places=12) + + +class TestLegacyWaveformContract(unittest.TestCase): + @classmethod + def setUpClass(cls): + try: + cls.lal, cls.lalsim, cls.lsu, cls.fl, cls.fr, cls.gpu = _imports() + except ImportError as exc: + raise unittest.SkipTest("RIFT waveform dependencies unavailable: %s" % exc) + + def test_default_forwards_all_options_and_uploads_both_banks(self): + lal, fl, gpu = self.lal, self.fl, self.gpu + n, df, dt, epoch = 64, 1.0, 1.0 / 64.0, -0.75 + rng = np.random.default_rng(891) + ordinary_values = rng.normal(size=n) + 1j * rng.normal(size=n) + conjugate_values = 3 + rng.normal(size=n) + 1j * rng.normal(size=n) + ordinary = {(2, 2): _fd_series(lal, "ordinary", ordinary_values, df, epoch)} + conjugate = {(2, 2): _fd_series(lal, "conjugate", conjugate_values, df, epoch)} + + from types import SimpleNamespace + p = SimpleNamespace( + dist=100e6 * lal.PC_SI, deltaF=df, deltaT=dt, fmin=4.0, + phi=1.1, theta=0.2, + ) + data_values = rng.normal(size=n) + 1j * rng.normal(size=n) + data = {"H1": _fd_series(lal, "data", data_values, df, 100.0)} + psd = lal.CreateREAL8FrequencySeries( + "PSD", lal.LIGOTimeGPS(0), 0, df, lal.SecondUnit, n // 2 + 1 + ) + psd.data.data[:] = 1.0 + forwarded = dict( + extra_waveform_kwargs={"fd_standoff_factor": 0.91, "token": "nested"}, + use_gwsignal=True, use_gwsignal_approx="SEOBNRv5PHM", + use_external_EOB=True, nr_lookup=True, + NR_group="nr-group", NR_param="nr-param", + ROM_group="rom-group", ROM_param="rom-param", + force_22_mode=True, perturbative_extraction=True, + ) + calls, uploads = [], [] + real_device_asarray = gpu._device_asarray + + def generator(*args, **kwargs): + calls.append((args, kwargs)) + return ordinary, conjugate + + def record_upload(value, xp, dtype=None): + uploads.append(np.asarray(value).copy()) + return real_device_asarray(value, xp, dtype=dtype) + + with mock.patch.object(fl, "internal_hlm_generator", side_effect=generator), \ + mock.patch.object(gpu, "_device_asarray", side_effect=record_upload): + gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + 100.75, 0.125, p, data, {"H1": psd}, 2, 24.0, + Qmax=0, p_max=0, backend=np, skip_interpolation=True, + quiet=True, verbose=False, **forwarded + ) + + self.assertEqual(len(calls), 1) + args, kwargs = calls[0] + self.assertIs(args[0], p) + self.assertEqual(args[1], 2) + self.assertFalse(kwargs["verbose"]) + self.assertTrue(kwargs["quiet"]) + for key, value in forwarded.items(): + self.assertEqual(kwargs[key], value) + # The first two conversions are the distinct ordinary and conjugate + # waveform arrays. This guards against regenerating, aliasing, or + # synthesizing conjugates in the H2D default path. + np.testing.assert_array_equal(uploads[0], ordinary_values) + np.testing.assert_array_equal(uploads[1], conjugate_values) + + def test_rejects_legacy_modes_with_different_frequency_grids(self): + values = np.ones(64, dtype=complex) + modes = { + (2, 2): _fd_series(self.lal, "h22", values, 1.0, -0.5), + (3, 3): _fd_series(self.lal, "h33", values, 0.5, -0.5), + } + with self.assertRaisesRegex(ValueError, "do not share a frequency grid"): + self.gpu._series_arrays(modes, np) + + def test_rejects_legacy_modes_with_different_epochs(self): + values = np.ones(64, dtype=complex) + modes = { + (2, 2): _fd_series(self.lal, "h22", values, 1.0, -0.5), + (3, 3): _fd_series(self.lal, "h33", values, 1.0, -0.49), + } + with self.assertRaisesRegex(ValueError, "do not share a common epoch"): + self.gpu._series_arrays(modes, np) + + def test_aligns_modes_by_labels_not_dictionary_order(self): + values = np.arange(64, dtype=complex) + modes = { + (2, -2): _fd_series(self.lal,"hm",3*values,1.,-0.5), + (2, 2): _fd_series(self.lal,"hp",values,1.,-0.5), + } + keys, arrays, *_ = self.gpu._series_arrays( + modes,np,mode_order=[(2,2),(2,-2)]) + self.assertEqual(keys,[(2,2),(2,-2)]) + np.testing.assert_array_equal(arrays[0],values) + np.testing.assert_array_equal(arrays[1],3*values) + + +class TestSecondPhysicalModel(unittest.TestCase): + @classmethod + def setUpClass(cls): + try: + cls.lal, cls.lalsim, cls.lsu, cls.fl, cls.fr, cls.gpu = _imports() + except ImportError as exc: + raise unittest.SkipTest("RIFT waveform dependencies unavailable: %s" % exc) + + def _run(self, xp): + event, p, data, psd = _short_problem(self.lalsim.TaylorF2) + common = dict( + event_time_geo=event, t_window=0.05, P=p, data_dict=data, + psd_dict=psd, Lmax=2, fMax=200.0, Qmax=0, p_max=0, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, + verbose=False, quiet=True, skip_interpolation=True, + ) + real_generator = self.fl.internal_hlm_generator + + def reversed_conjugate_order(*args, **kwargs): + ordinary, conjugate = real_generator(*args, **kwargs) + return ordinary, dict(reversed(list(conjugate.items()))) + + with mock.patch.object(self.fl,"internal_hlm_generator",side_effect=reversed_conjugate_order): + cpu = self.fr.PrecomputeLikelihoodTermsRotatingFreqResponse(**common) + candidate = self.gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, backend=xp, context=self.gpu.GPUPrecomputeContext(xp) + ) + _assert_precompute_close(self, self.fr, cpu, candidate, + 3e-10 if xp is np else 3e-9) + + def test_taylorf2_numpy_precompute_matches_cpu(self): + self._run(np) + + def test_taylorf2_cupy_precompute_matches_cpu(self): + try: + import cupy as cp + if cp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("no CUDA device") + except Exception as exc: + self.skipTest("no usable CuPy GPU: %s" % exc) + self._run(cp) + + +class TestGenericMultimodeModel(unittest.TestCase): + @classmethod + def setUpClass(cls): + try: + cls.lal, cls.lalsim, cls.lsu, cls.fl, cls.fr, cls.gpu = _imports() + cls.approximant = cls.lalsim.IMRPhenomXPHM + except (ImportError, AttributeError) as exc: + raise unittest.SkipTest("IMRPhenomXPHM unavailable: %s" % exc) + + def _run(self, xp): + event, p, data, psd = _short_problem(self.approximant, precessing=True) + # Inspect the real legacy generator result itself: the compatibility + # claim requires generic labels and a common grid, not a 22-only proxy. + modes, conjugate = self.fl.internal_hlm_generator( + p, 4, verbose=False, quiet=True + ) + labels = list(modes) + expected_labels = { + (ell, m) for ell in range(2, 5) for m in range(-ell, ell + 1) + } + self.assertEqual(set(labels), expected_labels) + self.assertEqual(len(labels), 21) + self.assertEqual(set(labels), set(conjugate)) + first = modes[labels[0]] + for bank in (modes, conjugate): + for series in bank.values(): + self.assertEqual(series.data.length, first.data.length) + self.assertAlmostEqual(series.deltaF, first.deltaF, places=14) + self.assertAlmostEqual(float(series.epoch), float(first.epoch), places=12) + + common = dict( + event_time_geo=event, t_window=0.05, P=p, data_dict=data, + psd_dict=psd, Lmax=4, fMax=200.0, Qmax=0, p_max=0, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, + verbose=False, quiet=True, skip_interpolation=True, + ) + cpu = self.fr.PrecomputeLikelihoodTermsRotatingFreqResponse(**common) + candidate = self.gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, backend=xp, context=self.gpu.GPUPrecomputeContext(xp) + ) + _assert_precompute_close(self, self.fr, cpu, candidate, + 5e-10 if xp is np else 5e-9) + + cpu_packed, gpu_packed = _packed(self.fr, cpu), _packed(self.fr, candidate) + pv = p.manual_copy() + # Row zero is near the injected point; row one is deliberately offset + # in sky, orientation, phase, and distance. + pv.phi = np.array([p.phi + 1e-4, p.phi + 0.35]) + pv.theta = np.array([p.theta - 1e-4, p.theta - 0.22]) + pv.incl = np.array([p.incl + 1e-4, 1.15]) + pv.psi = np.array([p.psi + 1e-4, p.psi + 0.31]) + pv.phiref = np.array([p.phiref + 1e-4, p.phiref + 0.47]) + pv.dist = np.array([p.dist * 1.001, p.dist * 1.7]) + pv.tref = event + pv.deltaT = p.deltaT + tvals = np.array([-p.deltaT, 0.0, p.deltaT]) + ln_cpu = self.fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pv, cpu[4], *cpu_packed, Lmax=4, array_output=True, + time_interp="nearest", xpy=np, + ) + ln_gpu = self.fr.DiscreteFactoredLogLikelihoodRotatingFreqResponseNoLoop( + tvals, pv, candidate[4], *gpu_packed, Lmax=4, array_output=True, + time_interp="nearest", xpy=np, + ) + self.assertTrue(np.all(np.isfinite(ln_cpu))) + self.assertGreater(float(np.max(np.abs(np.asarray(ln_cpu)[0] - + np.asarray(ln_cpu)[1]))), 0.0) + np.testing.assert_allclose(ln_gpu, ln_cpu, + rtol=5e-10 if xp is np else 5e-9, + atol=2e-8) + + def test_xphm_numpy_precompute_and_likelihood_match_cpu(self): + self._run(np) + + def test_xphm_cupy_precompute_and_likelihood_match_cpu(self): + try: + import cupy as cp + if cp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("no CUDA device") + except Exception as exc: + self.skipTest("no usable CuPy GPU: %s" % exc) + self._run(cp) + + +class TestGWSignalSEOBNRv5PHM(unittest.TestCase): + """Short real-model guard for the legacy GWSignal-to-device route.""" + + @classmethod + def setUpClass(cls): + try: + cls.lal, cls.lalsim, cls.lsu, cls.fl, cls.fr, cls.gpu = _imports() + except ImportError as exc: + raise unittest.SkipTest("RIFT waveform dependencies unavailable: %s" % exc) + if importlib.util.find_spec("pyseobnr") is None: + raise unittest.SkipTest("SEOBNRv5PHM needs the optional pyseobnr backend") + # Once the advertised backend exists, a broken GWSignal import is a + # compatibility failure rather than an optional-dependency skip. + if not cls.fl.has_GWS: + raise RuntimeError("pyseobnr is installed but RIFT could not import GWSignal") + + def _run(self, xp): + event, p, data, psd = _short_problem( + self.lalsim.IMRPhenomXPHM, precessing=True + ) + waveform_kwargs = dict( + use_gwsignal=True, + use_gwsignal_approx="SEOBNRv5PHM", + # This deliberately short 512 Hz test is a transport/parity gate, + # not a high-mode accuracy study. Disable the model's per-mode + # ringdown/Nyquist veto exactly as the maintained diagnostic does. + extra_waveform_kwargs={"lmax_nyquist": 1}, + ) + common = dict( + event_time_geo=event, t_window=0.05, P=p, data_dict=data, + psd_dict=psd, Lmax=4, fMax=200.0, Qmax=0, p_max=0, + analyticPSD_Q=False, inv_spec_trunc_Q=False, T_spec=0.0, + verbose=False, quiet=True, skip_interpolation=True, + **waveform_kwargs + ) + + routed_calls = [] + real_gwsignal = self.fl.rgws.std_and_conj_hlmoff + + def record_gwsignal_route(*args, **kwargs): + banks = real_gwsignal(*args, **kwargs) + routed_calls.append((args, kwargs, banks)) + return banks + + # P.approx is intentionally an ordinary available LAL approximant: + # SEOBNRv5PHM is selected only by the explicit GWSignal string. Thus + # observing both calls here proves neither precompute silently fell + # back to the default LAL route. + with mock.patch.object( + self.fl.rgws, "std_and_conj_hlmoff", + side_effect=record_gwsignal_route): + cpu = self.fr.PrecomputeLikelihoodTermsRotatingFreqResponse(**common) + candidate = self.gpu.PrecomputeLikelihoodTermsRotatingFreqResponseGPU( + **common, backend=xp, context=self.gpu.GPUPrecomputeContext(xp) + ) + + self.assertEqual(len(routed_calls), 2) + for args, kwargs, (ordinary, conjugate) in routed_calls: + self.assertIsInstance(args[0], self.lsu.ChooseWaveformParams) + self.assertEqual(args[1], 4) + self.assertEqual(kwargs["approx_string"], "SEOBNRv5PHM") + self.assertEqual(kwargs["lmax_nyquist"], 1) + labels = set(ordinary) + self.assertEqual(labels, set(conjugate)) + self.assertGreater(len(labels), 2) + self.assertTrue(any(ell > 2 for ell, _ in labels)) + first = ordinary[next(iter(labels))] + self.assertEqual(first.data.length, data["H1"].data.length) + self.assertAlmostEqual(first.deltaF, data["H1"].deltaF, places=14) + for bank in (ordinary, conjugate): + for series in bank.values(): + self.assertEqual(series.data.length, first.data.length) + self.assertAlmostEqual(series.deltaF, first.deltaF, places=14) + self.assertAlmostEqual(float(series.epoch), + float(first.epoch), places=12) + + self.assertEqual(set(cpu[4]["modes"]), set(candidate[4]["modes"])) + _assert_precompute_close(self, self.fr, cpu, candidate, + 5e-10 if xp is np else 5e-9) + + def test_seobnrv5phm_gwsignal_numpy_precompute_matches_cpu(self): + self._run(np) + + def test_seobnrv5phm_gwsignal_cupy_precompute_matches_cpu(self): + try: + import cupy as cp + if cp.cuda.runtime.getDeviceCount() < 1: + raise RuntimeError("no CUDA device") + except Exception as exc: + if os.environ.get("RIFT_REQUIRE_GPU_PRECOMPUTE") == "1": + self.fail("mandatory CuPy device gate unavailable: %s" % exc) + self.skipTest("no usable CuPy GPU: %s" % exc) + self._run(cp) + + +if __name__ == "__main__": + unittest.main() diff --git a/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_waveform.py b/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_waveform.py new file mode 100644 index 000000000..d7336edcd --- /dev/null +++ b/MonteCarloMarginalizeCode/Code/test/waveforms/test_gpu_waveform.py @@ -0,0 +1,232 @@ +"""Fail-closed and Fourier-contract tests for the optional Ripple adapter.""" + +import unittest +from unittest import mock + +import numpy as np + +from RIFT.likelihood import gpu_waveform as gw + + +class TestFourierConvention(unittest.TestCase): + def test_continuous_roundtrip(self): + rng = np.random.default_rng(20260912) + ht = rng.normal(size=32) + 1j * rng.normal(size=32) + hf = gw._continuous_forward(ht, 1.0 / 16.0, np) + recovered = gw._continuous_inverse(hf, 1.0 / 16.0, np) + np.testing.assert_allclose(recovered, ht, rtol=2e-14, atol=2e-14) + + def test_forward_matches_lal_rift_packing(self): + try: + import lal + except ImportError: + self.skipTest("LAL is optional in lightweight unit environments") + rng = np.random.default_rng(17) + n, dt = 32, 1.0 / 16.0 + values = rng.normal(size=n) + 1j * rng.normal(size=n) + ts = lal.CreateCOMPLEX16TimeSeries( + "test", lal.LIGOTimeGPS(0), 0, dt, lal.DimensionlessUnit, n + ) + ts.data.data[:] = values + fs = lal.CreateCOMPLEX16FrequencySeries( + "test", ts.epoch, 0, 1.0 / (n * dt), lal.DimensionlessUnit, n + ) + lal.COMPLEX16TimeFreqFFT(fs, ts, lal.CreateForwardCOMPLEX16FFTPlan(n, 0)) + np.testing.assert_allclose( + gw._continuous_forward(values, dt, np), fs.data.data, + rtol=2e-14, atol=2e-14, + ) + + def test_tdfromfd_shift_and_real_ifft_matches_lal(self): + try: + import lal + except ImportError: + self.skipTest("LAL is optional in lightweight unit environments") + rng = np.random.default_rng(20260913) + n, dt = 64, 1.0 / 128.0 + df = 1.0 / (n * dt) + # A real inverse transform requires real DC and Nyquist coefficients. + values = rng.normal(size=n // 2 + 1) + 1j * rng.normal(size=n // 2 + 1) + values[[0, -1]] = values[[0, -1]].real + epoch, extra_time = -7.25, 10.5 * dt + got, got_epoch, shift_samples = gw._lal_tdfromfd_shift_and_irfft( + values, df, dt, epoch, extra_time, np + ) + # C round is half away from zero, unlike Python's ties-to-even round. + self.assertEqual(shift_samples, 11) + self.assertEqual(got_epoch, epoch + 11 * dt) + fs = lal.CreateCOMPLEX16FrequencySeries( + "shifted", lal.LIGOTimeGPS(got_epoch), 0.0, df, + lal.DimensionlessUnit, len(values), + ) + k = np.arange(len(values)) + fs.data.data[:] = values * np.exp(2j * np.pi * k * df * 11 * dt) + ts = lal.CreateREAL8TimeSeries( + "inverse", fs.epoch, 0.0, dt, lal.DimensionlessUnit, n + ) + lal.REAL8FreqTimeFFT(ts, fs, lal.CreateReverseREAL8FFTPlan(n, 0)) + np.testing.assert_allclose(got, ts.data.data, rtol=2e-14, atol=2e-14) + + def test_tdfromfd_numpy_jax_agree(self): + try: + import jax.numpy as jnp + except ImportError: + self.skipTest("JAX is optional in lightweight unit environments") + n, dt = 32, 1.0 / 64.0 + df = 1.0 / (n * dt) + values = np.linspace(0.0, 1.0, n // 2 + 1).astype(np.complex128) + expected = gw._lal_tdfromfd_shift_and_irfft( + values, df, dt, -1.0, 3.2 * dt, np + ) + got = gw._lal_tdfromfd_shift_and_irfft( + jnp.asarray(values), df, dt, -1.0, 3.2 * dt, jnp + ) + np.testing.assert_allclose(np.asarray(got[0]), expected[0], + rtol=2e-13, atol=2e-13) + self.assertEqual(got[1:], expected[1:]) + + +class TestRIFTPostprocessing(unittest.TestCase): + def test_grow_appends_right_and_preserves_epoch(self): + modes = {(2, 2): np.arange(12, dtype=np.complex128)} + got, epoch, ntaper = gw._rift_postprocess_td_modes( + modes, -3.25, 1.0 / 16.0, 20, 2.0, np + ) + self.assertEqual(epoch, -3.25) + self.assertEqual(ntaper, 8) + expected = np.pad(modes[(2, 2)], (0, 8)) + j = np.arange(ntaper) + expected[:ntaper] *= 0.5 - 0.5 * np.cos(np.pi * j / ntaper) + np.testing.assert_allclose(got[(2, 2)], expected) + + def test_shrink_discards_left_and_advances_epoch(self): + source = np.arange(24, dtype=np.float64).astype(np.complex128) + got, epoch, ntaper = gw._rift_postprocess_td_modes( + {(2, 2): source, (2, -2): source.conj()}, + -2.0, 1.0 / 16.0, 16, 2.0, np, + ) + self.assertEqual(epoch, -1.5) + self.assertEqual(ntaper, 8) + window = np.ones(16) + window[:ntaper] = 0.5 - 0.5 * np.cos(np.pi * np.arange(ntaper) / ntaper) + np.testing.assert_allclose(got[(2, 2)], source[-16:] * window) + + def test_rejects_noncommon_grid(self): + with self.assertRaisesRegex(gw.WaveformCompatibilityError, "share a grid"): + gw._rift_postprocess_td_modes( + {(2, 2): np.zeros(8), (2, -2): np.zeros(10)}, + 0.0, 1.0 / 16.0, 16, 2.0, np, + ) + + def test_matches_actual_lalsimutils_post_lal_stage(self): + try: + import lal + import lalsimulation as lalsim + from RIFT import lalsimutils as lsu + except ImportError: + self.skipTest("LAL and RIFT waveform dependencies are optional") + dt, df, fmin = 1.0 / 512.0, 0.25, 40.0 + p = lsu.ChooseWaveformParams( + m1=30 * lal.MSUN_SI, m2=25 * lal.MSUN_SI, + s1z=0.1, s2z=-0.2, fmin=fmin, fref=60.0, + deltaT=dt, deltaF=df, approx=lalsim.IMRPhenomD, + dist=200e6 * lal.PC_SI, phiref=0.4, psi=0.0, + ) + raw_struct = lsu.hlmoft_FromFD_dict(p.manual_copy(), Lmax=2) + raw = lsu.SphHarmTimeSeries_to_dict(raw_struct, 2) + expected = lsu.hlmoft(p.manual_copy(), Lmax=2, silent=True) + source = {label: np.asarray(series.data.data).copy() + for label, series in raw.items()} + got, epoch, ntaper = gw._rift_postprocess_td_modes( + source, float(raw[(2, 2)].epoch), dt, + int(1.0 / (dt * df)), fmin, np, + ) + self.assertEqual(ntaper, max( + int(0.01 * min(len(source[(2, 2)]), len(got[(2, 2)]))), + int(1.0 / (fmin * dt)), + )) + for label in expected: + self.assertAlmostEqual(float(expected[label].epoch), epoch, places=12) + target = np.asarray(expected[label].data.data) + scale = np.max(np.abs(target)) + if scale == 0: + np.testing.assert_array_equal(got[label], target) + continue + # Physical strain is ~1e-21: a unit-scale absolute tolerance would + # silently accept a missing or sign-flipped waveform. Normalize + # both sides, and pin that the comparison rejects those defects. + np.testing.assert_allclose(got[label] / scale, target / scale, + rtol=2e-14, atol=2e-14) + for broken in (np.zeros_like(target), -target): + with self.assertRaises(AssertionError): + np.testing.assert_allclose(broken / scale, target / scale, + rtol=2e-14, atol=2e-14) + self.assertGreater(np.max(np.abs(expected[(2, 2)].data.data)), 0) + + +class _Params: + deltaT = 1.0 / 16.0 + deltaF = 0.5 + fmin = 2.0 + fmax = 8.0 + fref = 2.0 + phiref = 0.3 + m1 = 30.0 * 1.9884099021470416e30 + m2 = 25.0 * 1.9884099021470416e30 + s1x = s1y = s2x = s2y = 0.0 + s1z = 0.1 + s2z = -0.2 + dist = 200.0 * 3.085677581491367e22 + + +class TestProviderGuards(unittest.TestCase): + def setUp(self): + try: + import jax.numpy as jnp + import lalsimulation as lalsim + except ImportError: + self.skipTest("JAX and LAL are required for provider contract tests") + self.jnp = jnp + self.P = _Params() + self.P.approx = lalsim.IMRPhenomD + + def test_rift_conditioning_fails_closed(self): + with self.assertRaisesRegex(gw.WaveformCompatibilityError, "not certified"): + gw.generate_imrphenomd_fd(self.P, backend=self.jnp) + + def test_rejects_wrong_approximant_before_ripple_import(self): + import lalsimulation as lalsim + self.P.approx = lalsim.TaylorF2 + with self.assertRaisesRegex(gw.WaveformCompatibilityError, "only supports"): + gw.generate_imrphenomd_fd( + self.P, backend=self.jnp, conditioning="direct_fd" + ) + + def test_direct_fd_is_unconditioned_and_has_exact_symmetries(self): + jnp = self.jnp + + class FakeRipple: + @staticmethod + def gen_IMRPhenomD(f, params, fref): + return (1.0 + 0.25j) * f ** (-7.0 / 6.0) + + with mock.patch.object(gw, "_load_ripple", return_value=FakeRipple): + bank = gw.generate_imrphenomd_fd( + self.P, backend=jnp, conditioning="direct_fd" + ) + self.assertFalse(bank.conditioned) + self.assertEqual(bank.epoch, 0.0) + n = bank.modes[(2, 2)].shape[0] + reflection = (-np.arange(n)) % n + h22 = np.asarray(bank.modes[(2, 2)]) + h2m2 = np.asarray(bank.modes[(2, -2)]) + np.testing.assert_allclose(h2m2, np.conj(h22[reflection])) + for lm, mode in bank.modes.items(): + np.testing.assert_allclose( + np.asarray(bank.conjugate_modes[lm]), + np.conj(np.asarray(mode)[reflection]), + ) + + +if __name__ == "__main__": + unittest.main()