Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 3 additions & 10 deletions python/lsst/analysis/tools/actions/keyedData/calcDistances.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,23 +104,15 @@ def __call__(self, data: KeyedData, **kwargs) -> KeyedData:
"AMx": np.nan,
"ADx": np.nan,
"AFx": np.nan,
"nPairs": 0,
}

if len(data[self.groupKey]) == 0:
return distanceParams

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.
Expand Down Expand Up @@ -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

Expand Down
2 changes: 2 additions & 0 deletions python/lsst/analysis/tools/atools/astrometricRepeatability.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,6 +246,7 @@ def setDefaults(self):
"AMx": "mas",
"AFx": "percent",
"ADx": "mas",
"nPairs": "count",
}

self.produce.plot = HistPlot()
Expand Down Expand Up @@ -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",
}


Expand Down
218 changes: 190 additions & 28 deletions python/lsst/analysis/tools/tasks/associatedSourcesTractAnalysis.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,20 +20,25 @@
# along with this program. If not, see <https://www.gnu.org/licenses/>.
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
from lsst.drp.tasks.gbdesAstrometricFit import calculate_apparent_motion
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

Expand Down Expand Up @@ -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):
Expand All @@ -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"],
Expand All @@ -158,15 +173,14 @@ def callback(self, inputs, dataId):

def prepareAssociatedSources(
self,
skymap,
tract,
sourceCatalogs,
associatedSources,
associatedSourceIds,
astrometricCorrectionCatalog=None,
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.
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -254,26 +291,27 @@ 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")
dataWithPM["MJD"] = times
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)
Expand Down Expand Up @@ -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)
Loading