diff --git a/src/sp_validation/cosmo_val/real_space.py b/src/sp_validation/cosmo_val/real_space.py index 89dad2a9..fce6d55c 100644 --- a/src/sp_validation/cosmo_val/real_space.py +++ b/src/sp_validation/cosmo_val/real_space.py @@ -14,8 +14,30 @@ import treecorr from cs_util import plots as cs_plots +from sp_validation.statistics import jackknife_patch_centers + class RealSpaceMixin: + def _shear_catalog(self, ver, npatch): + """Calibrated shear catalogue of ``ver`` with seeded jackknife patches. + + Call inside ``self.results[ver].temporarily_read_data()``. + """ + positions = { + "ra": self.results[ver].dat_shear["RA"], + "dec": self.results[ver].dat_shear["Dec"], + "w": self._read_shear_cols(ver, "w_col"), + "ra_units": self.treecorr_config["ra_units"], + "dec_units": self.treecorr_config["dec_units"], + } + centers = None + if int(npatch) > 1: + centers = jackknife_patch_centers( + treecorr.Catalog(**positions), int(npatch) + ) + g1, g2 = self._calibrated_g(ver) + return treecorr.Catalog(**positions, g1=g1, g2=g2, patch_centers=centers) + def calculate_2pcf(self, ver, npatch=None, **treecorr_config): """ Calculate the two-point correlation function (2PCF) ξ± for a given catalog @@ -43,8 +65,10 @@ def calculate_2pcf(self, ver, npatch=None, **treecorr_config): Notes: - If the output file for the given configuration already exists, the calculation is skipped, and the results are loaded from the file. - - If a patch file for the given configuration does not exist, it is - created during the process. + - Jackknife patches come from a seeded k-means on a fixed-depth tree + (``statistics.jackknife_patch_centers``), so they are a pure + function of the catalogue's positions and weights, the same on + every run and thread count. - The ``.txt`` TreeCorr dump is the only raw byproduct written here. """ @@ -70,27 +94,7 @@ def calculate_2pcf(self, ver, npatch=None, **treecorr_config): else: # Load data and create a catalog with self.results[ver].temporarily_read_data(): - g1, g2 = self._calibrated_g(ver) - w = self._read_shear_cols(ver, "w_col") - - # Use patch file if it exists - patch_file = self._output_path(f"{ver}_patches_npatch={npatch}.dat") - - cat_gal = treecorr.Catalog( - ra=self.results[ver].dat_shear["RA"], - dec=self.results[ver].dat_shear["Dec"], - g1=g1, - g2=g2, - w=w, - ra_units=self.treecorr_config["ra_units"], - dec_units=self.treecorr_config["dec_units"], - npatch=npatch, - patch_centers=patch_file if os.path.exists(patch_file) else None, - ) - - # If no patch file exists, save the current patches - if not os.path.exists(patch_file): - cat_gal.write_patch_centers(patch_file) + cat_gal = self._shear_catalog(ver, npatch) # Process the catalog & write the correlation functions gg.process(cat_gal) @@ -346,23 +350,10 @@ def calculate_aperture_mass_dispersion( gg.read(out_fname) else: with self.results[ver].temporarily_read_data(): - g1, g2 = self._calibrated_g(ver) - cat_gal = treecorr.Catalog( - ra=self.results[ver].dat_shear["RA"], - dec=self.results[ver].dat_shear["Dec"], - g1=g1, - g2=g2, - w=self._read_shear_cols(ver, "w_col"), - ra_units=self.treecorr_config["ra_units"], - dec_units=self.treecorr_config["dec_units"], - npatch=npatch, - ) - + cat_gal = self._shear_catalog(ver, npatch) gg.process(cat_gal) gg.write(out_fname) del cat_gal - del g1 - del g2 mapsq, mapsq_im, mxsq, mxsq_im, varmapsq = gg.calculateMapSq( R=theta_map, diff --git a/src/sp_validation/rho_tau.py b/src/sp_validation/rho_tau.py index 62ae8bf8..62f40959 100644 --- a/src/sp_validation/rho_tau.py +++ b/src/sp_validation/rho_tau.py @@ -9,6 +9,7 @@ # SquareRootScale now lives in sp_validation.plots; re-exported here so that # `from sp_validation.rho_tau import SquareRootScale` keeps working. from sp_validation.plots import SquareRootScale # noqa: F401 +from sp_validation.statistics import jackknife_patch_centers def _extract_xip(correlations): @@ -369,10 +370,10 @@ def get_jackknife_cov( print(f"Computing the patch centers for patch {i + 1}/{ncov}") npatch = rho_stat_handler.catalogs._params["patch_number"] - field = rho_stat_handler.catalogs.catalogs_dict[ - f"psf_{version}{i}" - ].getNField(max_top=int.bit_length(npatch) - 1, coords="spherical") - patch, centers = field.run_kmeans(npatch) + centers = jackknife_patch_centers( + rho_stat_handler.catalogs.catalogs_dict[f"psf_{version}{i}"], + npatch, + ) # Update the patch centers of the catalogs for key, cat in rho_stat_handler.catalogs.catalogs_dict.items(): diff --git a/src/sp_validation/statistics.py b/src/sp_validation/statistics.py index dc71f680..23572ee5 100644 --- a/src/sp_validation/statistics.py +++ b/src/sp_validation/statistics.py @@ -3,8 +3,9 @@ :Name: statistics.py :Description: Cosmology-independent statistical helpers (jackknife resampling, - chi2/PTE, calibrated min-PTE across many null tests, - covariance<->correlation, OneCovariance reshaping). + jackknife patch centres, chi2/PTE, calibrated min-PTE across + many null tests, covariance<->correlation, OneCovariance + reshaping). """ from dataclasses import dataclass @@ -12,6 +13,42 @@ import numpy as np from scipy import stats +#: Depth of the ball-tree layers that seed the jackknife k-means. +PATCH_MIN_TOP = 6 + + +def jackknife_patch_centers(cat, npatch, seed=0): + """Seeded k-means jackknife patch centres for a TreeCorr catalogue. + + Seeding the k-means does not by itself fix the patches. TreeCorr starts the + k-means from the top layers of a ball tree whose depth, unless ``min_top`` + is given, grows with the OpenMP thread count (``Field._determine_top``), + and ``Catalog(npatch=..., rng=...)`` offers no way to set it. Pinning that + depth makes the centres a function of the catalogue's positions, weights + and ``seed`` alone, on every machine. + + Parameters + ---------- + cat : treecorr.Catalog + Catalogue with spherical (RA, Dec) positions. + npatch : int + Number of patches. + seed : int, optional + Seed for the k-means initialisation. + + Returns + ------- + numpy.ndarray + Patch centres, to pass as ``patch_centers`` to ``treecorr.Catalog``. + """ + field = cat.getNField( + min_top=PATCH_MIN_TOP, + max_top=int.bit_length(npatch) - 1, + coords="spherical", + ) + _, centers = field.run_kmeans(npatch, rng=np.random.default_rng(seed)) + return centers + def jackknif_weighted_average2( data, diff --git a/src/sp_validation/tests/test_cosmo_val.py b/src/sp_validation/tests/test_cosmo_val.py index c08c055e..a743161a 100644 --- a/src/sp_validation/tests/test_cosmo_val.py +++ b/src/sp_validation/tests/test_cosmo_val.py @@ -501,39 +501,60 @@ def test_a_patched_xi_dump_reads_back(self, tmp_path): getattr(read, column), getattr(measured, column), rtol=1e-4 ) - def test_calculate_2pcf_does_not_depend_on_thread_count(self, tmp_path): - """calculate_2pcf's ξ± is the same on 4 and on 48 TreeCorr threads. - - Production binning (default bin_slop/angle_slop), both runs on the - jackknife patches the first one writes, each from a fresh Catalog; they - must agree to far below the jackknife σ. + def test_calculate_2pcf_is_reproducible_across_machines( + self, tmp_path, monkeypatch + ): + """Two fresh output trees on 4- and 16-CPU machines measure the same ξ±. + + TreeCorr's default k-means tree depth follows the machine's OpenMP + thread count (3 levels on 4 CPUs, 4 on 16). With 100 patches (up to 6 + levels, and not a power of two) that changes the k-means start, so a + seed alone would split the catalogue differently. calculate_2pcf pins + the depth: patch labels, per-patch-pair counts, ξ± and its jackknife + variance all agree. """ import treecorr - - params, version = self._write_synthetic_catalogs( - tmp_path, n_gal=4000, coherent_shear=True - ) - cv = CosmologyValidation( - versions=[version], - npatch=8, - theta_min=15.0, - theta_max=70.0, - nbins=6, - **params, - ) - - xi = {} - for n_threads in (4, 48): - # calculate_2pcf reads back an existing text dump instead of measuring. - for dump in Path(params["output_dir"]).glob(f"{version}_xi_*.txt"): - dump.unlink() - gg = cv.calculate_2pcf(version, num_threads=n_threads) - assert treecorr.get_omp_threads() == n_threads # the count took effect - xi[n_threads] = np.concatenate([gg.xip, gg.xim]) - sigma = np.sqrt(np.concatenate([gg.varxip, gg.varxim])) - - shift = np.max(np.abs(xi[48] - xi[4]) / sigma) - assert shift < 1e-6, f"ξ± moves by {shift:.3g}σ between 4 and 48 threads" + import treecorr.field + + patches = [] + process = treecorr.GGCorrelation.process + + def recording_process(gg, cat, *args, **kwargs): + patches.append(np.array(cat.patch)) + return process(gg, cat, *args, **kwargs) + + monkeypatch.setattr(treecorr.GGCorrelation, "process", recording_process) + + xi, var, counts = {}, {}, {} + for tree, n_cpu in (("a", 4), ("b", 16)): + # The thread count TreeCorr's default tree depth reads. + monkeypatch.setattr(treecorr.field, "get_omp_threads", lambda n=n_cpu: n) + (tmp_path / tree).mkdir() + params, version = self._write_synthetic_catalogs( + tmp_path / tree, + n_gal=4000, + ra_range=(0.0, 60.0), + dec_range=(-10.0, 30.0), + coherent_shear=True, + ) + gg = CosmologyValidation( + versions=[version], + npatch=100, + theta_min=15.0, + theta_max=70.0, + nbins=6, + **params, + ).calculate_2pcf(version) + xi[tree] = np.concatenate([gg.xip, gg.xim]) + var[tree] = np.concatenate([gg.varxip, gg.varxim]) + counts[tree] = {k: r.npairs for k, r in gg.results.items()} + + np.testing.assert_array_equal(patches[0], patches[1]) + assert counts["a"].keys() == counts["b"].keys() + for k in counts["a"]: + np.testing.assert_array_equal(counts["a"][k], counts["b"][k]) + np.testing.assert_allclose(xi["a"], xi["b"], rtol=0, atol=1e-12) + np.testing.assert_allclose(var["a"], var["b"], rtol=1e-10) def test_treecorr_runs_on_the_cpus_the_process_holds(self, tmp_path): """By default TreeCorr takes the process's CPU affinity, not the node's count."""