From 7685a51c4b8eea858e5ca354679527b8d5a84914 Mon Sep 17 00:00:00 2001 From: Clare Saunders Date: Mon, 13 Jul 2026 14:00:43 -0400 Subject: [PATCH] Add healpix based class and fix table joins --- .../tools/actions/keyedData/calcDistances.py | 13 +- .../tools/atools/astrometricRepeatability.py | 2 + .../tasks/associatedSourcesTractAnalysis.py | 218 +++++++++++++++--- 3 files changed, 195 insertions(+), 38 deletions(-) diff --git a/python/lsst/analysis/tools/actions/keyedData/calcDistances.py b/python/lsst/analysis/tools/actions/keyedData/calcDistances.py index db46089c2..6918635b7 100644 --- a/python/lsst/analysis/tools/actions/keyedData/calcDistances.py +++ b/python/lsst/analysis/tools/actions/keyedData/calcDistances.py @@ -104,6 +104,7 @@ def __call__(self, data: KeyedData, **kwargs) -> KeyedData: "AMx": np.nan, "ADx": np.nan, "AFx": np.nan, + "nPairs": 0, } if len(data[self.groupKey]) == 0: @@ -111,16 +112,7 @@ def __call__(self, data: KeyedData, **kwargs) -> KeyedData: rng = np.random.RandomState(seed=self.randomSeed) - def _compressArray(arrayIn): - h, rev = esutil.stat.histogram(arrayIn, rev=True) - arrayOut = np.zeros(len(arrayIn), dtype=np.int32) - (good,) = np.where(h > 0) - for counter, ind in enumerate(good): - arrayOut[rev[rev[ind] : rev[ind + 1]]] = counter - return arrayOut - - groupId = _compressArray(data[self.groupKey]) - + _, groupId = np.unique(data[self.groupKey], return_inverse=True) nObj = groupId.max() + 1 # Compute the meanRa/meanDec. @@ -268,6 +260,7 @@ def _compressArray(arrayIn): distanceParams["AMx"] = AMx.value distanceParams["ADx"] = ADx.value distanceParams["AFx"] = AFx.value + distanceParams["nPairs"] = len(rmsDistances) return distanceParams diff --git a/python/lsst/analysis/tools/atools/astrometricRepeatability.py b/python/lsst/analysis/tools/atools/astrometricRepeatability.py index 0f017b42c..d444edf8e 100644 --- a/python/lsst/analysis/tools/atools/astrometricRepeatability.py +++ b/python/lsst/analysis/tools/atools/astrometricRepeatability.py @@ -246,6 +246,7 @@ def setDefaults(self): "AMx": "mas", "AFx": "percent", "ADx": "mas", + "nPairs": "count", } self.produce.plot = HistPlot() @@ -274,6 +275,7 @@ def finalize(self): "AMx": f"{{band}}_AM{self.xValue}", "AFx": f"{{band}}_AF{self.xValue}", "ADx": f"{{band}}_AD{self.xValue}", + "nPairs": "{band}_nPairs", } diff --git a/python/lsst/analysis/tools/tasks/associatedSourcesTractAnalysis.py b/python/lsst/analysis/tools/tasks/associatedSourcesTractAnalysis.py index ccd4ca814..d74de3e2d 100644 --- a/python/lsst/analysis/tools/tasks/associatedSourcesTractAnalysis.py +++ b/python/lsst/analysis/tools/tasks/associatedSourcesTractAnalysis.py @@ -20,13 +20,17 @@ # along with this program. If not, see . from __future__ import annotations -__all__ = ("AssociatedSourcesTractAnalysisConfig", "AssociatedSourcesTractAnalysisTask") +__all__ = ( + "AssociatedSourcesTractAnalysisConfig", + "AssociatedSourcesTractAnalysisTask", + "AssociatedSourcesHealpix3AnalysisTask", +) import astropy.time import astropy.units as u import numpy as np -from astropy.table import Table, hstack -from scipy.spatial import KDTree +from astropy.coordinates import SkyCoord +from astropy.table import Table, hstack, vstack import lsst.pex.config as pexConfig from lsst.daf.butler import DatasetProvenance @@ -34,6 +38,7 @@ from lsst.pipe.base import NoWorkFound from lsst.pipe.base import connectionTypes as ct from lsst.skymap import BaseSkyMap +from lsst.sphgeom import HealpixPixelization from ..interfaces import AnalysisBaseConfig, AnalysisBaseConnections, AnalysisPipelineTask @@ -130,6 +135,18 @@ class AssociatedSourcesTractAnalysisConfig( }, doc="Column names for position and motion parameters in the astrometric correction catalogs.", ) + maxVisitCount = pexConfig.Field( + dtype=int, + default=None, + doc="Maximum number of visits to use in calculating the metrics.", + optional=True, + ) + maxObjects = pexConfig.Field( + dtype=int, + default=None, + doc="Maximum number of associated objects to use in calculating the metrics.", + optional=True, + ) class AssociatedSourcesTractAnalysisTask(AnalysisPipelineTask): @@ -147,8 +164,6 @@ def getBoxWcs(skymap, tract): def callback(self, inputs, dataId): """Callback function to be used with reconstructor.""" return self.prepareAssociatedSources( - inputs["skyMap"], - dataId["tract"], inputs["sourceCatalogs"], inputs["associatedSources"], inputs["associatedSourceIds"], @@ -158,8 +173,6 @@ def callback(self, inputs, dataId): def prepareAssociatedSources( self, - skymap, - tract, sourceCatalogs, associatedSources, associatedSourceIds, @@ -167,6 +180,7 @@ def prepareAssociatedSources( visitTable=None, ): """Concatenate source catalogs and join on associated source IDs.""" + rng = np.random.default_rng() # Strip any provenance from tables before merging to prevent # warnings from conflicts being issued by astropy.utils.merge. @@ -178,12 +192,26 @@ def prepareAssociatedSources( index = associatedSources["obj_index"] associatedSources["isolated_star_id"] = associatedSourceIds["isolated_star_id"][index] + if self.config.maxObjects: + objectChoice = rng.permutation(associatedSourceIds["isolated_star_id"])[: self.config.maxObjects] + objectChoice.sort() + sub1 = np.clip( + np.searchsorted(objectChoice, associatedSources["isolated_star_id"]), + 0, + len(objectChoice) - 1, + ) + matched = objectChoice[sub1] == associatedSources["isolated_star_id"] + associatedSources = associatedSources[matched] + trimmedSourceCatalogs = [] fullCatLen = 0 # It would be preferable to use astropy's built in functions - # but they are too slow so we have this wonderful masterpiece - # Which is still not fast but two thirds of the time is the butler get - reshapedAssocSources = associatedSources["sourceId"].reshape(len(associatedSources), 1) + # but they are too slow so instead we use a numpy searchsorted + # maneuver. + sortedAssocSources = associatedSources["sourceId"].copy() + assocSourcesSort = associatedSources["sourceId"].argsort() + sortedAssocSources.sort() + nAssocSources = len(sortedAssocSources) colsNeeded = list(self.collectInputNames()) # Only get the columns needed for the source catalogues. # The isolated_star_id and the obj_index are added later @@ -196,19 +224,28 @@ def prepareAssociatedSources( if "obj_index" in colsNeeded: colsNeeded.remove("obj_index") colsNeeded += ["sourceId", "coord_ra", "coord_dec"] + + if self.config.maxVisitCount: + sourceCatalogs = rng.permutation(sourceCatalogs)[: self.config.maxVisitCount] + for sourceCatalogRef in sourceCatalogs: sourceCatalog = sourceCatalogRef.get(parameters={"columns": set(colsNeeded)}) DatasetProvenance.strip_provenance_from_flat_dict(sourceCatalog.meta) - reshapedSourceCat = sourceCatalog["sourceId"].reshape(len(sourceCatalog), 1) - tree = KDTree(reshapedSourceCat) - _, inds = tree.query(reshapedAssocSources, distance_upper_bound=0.1) - ids = inds < len(sourceCatalog) + sub = np.clip( + np.searchsorted(sortedAssocSources, sourceCatalog["sourceId"]), 0, nAssocSources - 1 + ) + sourceCatalogInds = sortedAssocSources[sub] == sourceCatalog["sourceId"] + assocCatalogInds = sub[sourceCatalogInds] # Keep only the sources in groups that are fully contained within # the tract by matching to the associated sources table - trimmedSourceCatalogs.append(hstack([associatedSources[ids], sourceCatalog[inds[ids]]])) - fullCatLen += np.sum(ids) + trimmedSourceCatalogs.append( + hstack( + [associatedSources[assocSourcesSort][assocCatalogInds], sourceCatalog[sourceCatalogInds]] + ) + ) + fullCatLen += np.sum(sourceCatalogInds) columns = trimmedSourceCatalogs[0].columns dtypes = trimmedSourceCatalogs[0].dtype @@ -219,7 +256,7 @@ def prepareAssociatedSources( fullCat[n : n + len(trimmedSourceCatalog)] = trimmedSourceCatalog n += len(trimmedSourceCatalog) - if astrometricCorrectionCatalog is not None: + if (astrometricCorrectionCatalog is not None) and (len(fullCat) != 0): self.applyAstrometricCorrections(fullCat, astrometricCorrectionCatalog, visitTable) # Keep only finite ras and decs @@ -254,16 +291,18 @@ def applyAstrometricCorrections(self, dataJoined, astrometricCorrectionCatalog, astrometricCorrectionCatalog["pmDec"] *= u.mas / u.yr astrometricCorrectionCatalog["parallax"] *= u.mas - # Again using astropy join would have been great but this is four - # times faster - lenAstroCorrCat = len(astrometricCorrectionCatalog) - tree = KDTree(astrometricCorrectionCatalog["isolated_star_id"].reshape(lenAstroCorrCat, 1)) - _, inds = tree.query( - dataJoined["isolated_star_id"].reshape(len(dataJoined), 1), distance_upper_bound=0.5 + # Join the dataJoined catalog with the astrometricCorrectionCatalog. + sourceIds = dataJoined["isolated_star_id"].copy() + sortOrder = np.argsort(sourceIds) + sub = np.clip( + np.searchsorted(sourceIds, astrometricCorrectionCatalog["isolated_star_id"], sorter=sortOrder), + 0, + len(sourceIds) - 1, ) - ids = inds < lenAstroCorrCat + correctionsCatalogInds = sourceIds[sortOrder[sub]] == astrometricCorrectionCatalog["isolated_star_id"] + sourceIdsInds = sortOrder[sub[correctionsCatalogInds]] - dataWithPM = hstack([dataJoined[ids], astrometricCorrectionCatalog[inds[ids]]]) + dataWithPM = hstack([dataJoined[sourceIdsInds], astrometricCorrectionCatalog[correctionsCatalogInds]]) mjds = visitTable.loc[dataWithPM["visit"]]["expMidptMJD"] times = astropy.time.Time(mjds, format="mjd", scale="tai") @@ -271,9 +310,8 @@ def applyAstrometricCorrections(self, dataJoined, astrometricCorrectionCatalog, medianMJD = astropy.time.Time(np.median(mjds), format="mjd", scale="tai") raCorrection, decCorrection = calculate_apparent_motion(dataWithPM, medianMJD) - - dataJoined["coord_ra"][ids] = dataWithPM["coord_ra"] - raCorrection.value - dataJoined["coord_dec"][ids] = dataWithPM["coord_dec"] - decCorrection.value + dataJoined["coord_ra"][sourceIdsInds] = dataWithPM["coord_ra"] - raCorrection.value + dataJoined["coord_dec"][sourceIdsInds] = dataWithPM["coord_dec"] - decCorrection.value def runQuantum(self, butlerQC, inputRefs, outputRefs): inputs = butlerQC.get(inputRefs) @@ -309,3 +347,127 @@ def runQuantum(self, butlerQC, inputRefs, outputRefs): kwargs = {"data": data, "plotInfo": plotInfo, "skymap": inputs["skyMap"], "camera": inputs["camera"]} outputs = self.run(**kwargs) self.putByBand(butlerQC, outputs, outputRefs) + + +class AssociatedSourcesHealpix3AnalysisConnections( + AssociatedSourcesTractAnalysisConnections, + dimensions=("healpix3", "instrument"), +): + associatedSources = ct.Input( + doc="Table of associated sources", + name="{associatedSourcesInputName}", + storageClass="ArrowAstropy", + deferLoad=True, + dimensions=("instrument", "skymap", "tract"), + multiple=True, + ) + + associatedSourceIds = ct.Input( + doc="Table containing unique ids for the associated sources", + name="{associatedSourceIdsInputName}", + storageClass="ArrowAstropy", + deferLoad=True, + dimensions=("instrument", "skymap", "tract"), + multiple=True, + ) + astrometricCorrectionCatalog = ct.Input( + doc="Catalog with proper motion and parallax information.", + name="isolated_star_stellar_motions", + storageClass="ArrowAstropy", + deferLoad=True, + dimensions=("instrument", "skymap", "tract"), + multiple=True, + ) + + +class AssociatedSourcesHealpix3AnalysisConfig( + AssociatedSourcesTractAnalysisConfig, pipelineConnections=AssociatedSourcesHealpix3AnalysisConnections +): + pass + + +class AssociatedSourcesHealpix3AnalysisTask(AssociatedSourcesTractAnalysisTask): + ConfigClass = AssociatedSourcesHealpix3AnalysisConfig + _DefaultName = "associatedSourcesHealpix3Analysis" + + def getHealpixOverlap(self, sources, sourceIds, pixelId, astrometricCorrections=None): + + pixelization = HealpixPixelization(3) + pixelRegion = pixelization.pixel(pixelId) + + sourceCoords = SkyCoord(sourceIds["ra"] * u.degree, sourceIds["dec"] * u.degree).cartesian.xyz + + inPixel = pixelRegion.contains(*sourceCoords.value) + if not inPixel.any(): + return sources[:0] + pixelIds = sourceIds[inPixel]["isolated_star_id"] + + sub1 = np.clip(np.searchsorted(pixelIds, sources["isolated_star_id"]), 0, len(pixelIds) - 1) + matched = pixelIds[sub1] == sources["isolated_star_id"] + + return sources[matched] + + def runQuantum(self, butlerQC, inputRefs, outputRefs): + inputs = butlerQC.get(inputRefs) + + # Load specified columns from source catalogs + names = self.collectInputNames() + names |= {"sourceId", "coord_ra", "coord_dec"} + for item in ["obj_index", "isolated_star_id"]: + if item in names: + names.remove(item) + + dataId = butlerQC.quantum.dataId + plotInfo = self.parsePlotInfo( + {"associatedSources": inputs["associatedSources"][0]}, dataId, connectionName="associatedSources" + ) + + # Loop over tract inputs, keeping only objects that in this healpix, + # then stack in one big table. + pixelId = dataId["healpix3"] + associatedSourceRefs = { + assocRef.dataId["tract"]: assocRef for assocRef in inputs["associatedSources"] + } + associatedSourceIdRefs = { + assocRef.dataId["tract"]: assocRef for assocRef in inputs["associatedSourceIds"] + } + astrometricCorrectionRefs = { + assocRef.dataId["tract"]: assocRef for assocRef in inputs["astrometricCorrectionCatalog"] + } + data = [] + for tract in associatedSourceRefs: + tractAssociatedSources = self.loadData(associatedSourceRefs[tract], ["obj_index", "sourceId"]) + tractAssociatedSourceIds = self.loadData( + associatedSourceIdRefs[tract], ["isolated_star_id", "ra", "dec"] + ) + if self.config.applyAstrometricCorrections: + astromCorrections = astrometricCorrectionRefs[tract].get( + parameters={"columns": self.config.astrometricCorrectionParameters.values()} + ) + else: + astromCorrections = None + tractInput = { + "associatedSources": tractAssociatedSources, + "associatedSourceIds": tractAssociatedSourceIds, + "astrometricCorrectionCatalog": astromCorrections, + "sourceCatalogs": inputs["sourceCatalogs"], + "visitTable": inputs["visitTable"], + } + tractData = self.callback(tractInput, dataId) + if len(tractData) == 0: + continue + trimmedData = self.getHealpixOverlap(tractData, tractAssociatedSourceIds, pixelId) + trimmedData["tract"] = tract + data.append(trimmedData) + data = vstack(data) + + if len(data["associatedSources"]) == 0: + raise NoWorkFound(f"No associated sources in healpix {dataId.healpix3.id}") + + kwargs = { + "data": data, + "plotInfo": plotInfo, + "camera": inputs["camera"], + } + outputs = self.run(**kwargs) + self.putByBand(butlerQC, outputs, outputRefs)