Coverage for python/lsst/meas/extensions/multiprofit/fit_coadd_multiband.py: 22%
518 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-25 23:22 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-25 23:22 +0000
1# This file is part of meas_extensions_multiprofit.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (https://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <https://www.gnu.org/licenses/>.
22__all__ = (
23 "BasicModelInitializer",
24 "CachedBasicModelInitializer",
25 "CatalogExposurePsfs",
26 "InitialInputData",
27 "MagnitudeDependentSizePriorConfig",
28 "MakeBasicInitializerAction",
29 "MakeCachedBasicInitializerAction",
30 "MakeInitializerActionBase",
31 "ModelInitializer",
32 "MultiProFitSourceConfig",
33 "MultiProFitSourceFitter",
34 "MultiProFitSourceTask",
35 "PsfComponentsActionBase",
36 "PsfFitSuccessActionBase",
37 "SourceTablePsfComponentsAction",
38 "SourceTablePsfFitSuccessAction",
39)
41import logging
42import math
43from abc import ABC, abstractmethod
44from collections.abc import Iterable, Mapping, Sequence
45from functools import cached_property
46from typing import Any, ClassVar
48import astropy.units as u
49import numpy as np
50import pydantic
51from astropy.table import Table
53import lsst.afw.geom
54import lsst.afw.table as afwTable
55import lsst.gauss2d as g2
56import lsst.gauss2d.fit as g2f
57import lsst.pex.config as pexConfig
58import lsst.pipe.base as pipeBase
59import lsst.pipe.tasks.fit_coadd_multiband as fitMB
60import lsst.utils.timer as utilsTimer
61from lsst.daf.butler.formatters.parquet import astropy_to_arrow
62from lsst.multiprofit.errors import NoDataError, PsfRebuildFitFlagError
63from lsst.multiprofit.fitting.fit_psf import CatalogPsfFitterConfig, CatalogPsfFitterConfigData
64from lsst.multiprofit.fitting.fit_source import (
65 CatalogExposureSourcesABC,
66 CatalogSourceFitterABC,
67 CatalogSourceFitterConfig,
68 CatalogSourceFitterConfigData,
69)
70from lsst.multiprofit.modeller import Model
71from lsst.multiprofit.utils import frozen_arbitrary_allowed_config, get_params_uniq, set_config_from_dict
72from lsst.pex.config.configurableActions import ConfigurableAction, ConfigurableActionField
74from .errors import IsParentError, NotPrimaryError
75from .input_config import InputConfig
76from .utils import get_spanned_image
78_LOG = logging.getLogger(__name__)
79TWO_SQRT_PI = 2 * math.sqrt(np.pi)
82class PsfFitSuccessActionBase(ConfigurableAction):
83 """Base action to return whether a source had a succesful PSF fit."""
85 def get_schema(self) -> list[str]:
86 """Return the list of columns required to call this action."""
87 raise NotImplementedError("This method must be overloaded in subclasses")
89 def __call__(self, source: Mapping[str, Any], *args: Any, **kwargs: Any) -> bool:
90 raise NotImplementedError("This method must be overloaded in subclasses")
93class PsfComponentsActionBase(ConfigurableAction):
94 """Base action to return a list of Gaussians from a source mapping.
96 This base class should be used as a sentinel when using a MultiProFit PSF
97 fit table, and only needs to be specialized for external PSF fitters.
98 """
100 def get_schema(self) -> list[str]:
101 """Return the list of columns required to call this action."""
102 raise NotImplementedError("This method must be overloaded in subclasses")
104 def __call__(self, source: Mapping[str, Any], *args: Any, **kwargs: Any) -> list[g2.Gaussian]:
105 raise NotImplementedError("This method must be overloaded in subclasses")
108class SourceTablePsfFitSuccessAction(PsfFitSuccessActionBase):
109 """Action to return PSF fit status from a SourceTable row."""
111 flag_format = pexConfig.Field[str](
112 doc="Format for the flag field; flag_prefix, flag_suffix and flag_sub are substituted",
113 default="{flag_prefix}{flag_suffix}{flag_sub}",
114 )
115 flag_prefix = pexConfig.Field[str](
116 doc="Prefix for the key for the summed flag field",
117 default="modelfit_DoubleShapeletPsfApprox",
118 )
119 flag_suffix = pexConfig.Field[str](
120 doc="Suffix for all flag fields",
121 default="_flag",
122 )
123 flags_sub = pexConfig.ListField[str](
124 doc="Suffixes for specific flag fields that must not be true",
125 default=["_invalidPointForPsf", "_invalidMoments", "_maxIterations"],
126 )
128 def _format(self, flag_sub: str) -> str:
129 return self.flag_format.format(
130 flag_prefix=self.flag_prefix,
131 flag_sub=flag_sub,
132 flag_suffix=self.flag_suffix,
133 )
135 def get_schema(self) -> Iterable[str]:
136 for flag_sub in self.flags_sub:
137 yield self._format(flag_sub=flag_sub)
139 def __call__(self, source: Mapping[str, Any], *args: Any, **kwargs: Any) -> bool:
140 good = True
141 for flag_sub in self.flags_sub:
142 good &= not source[self._format(flag_sub=flag_sub)]
143 return good
146class SourceTablePsfComponentsAction(PsfComponentsActionBase):
147 """Action to return PSF components from a SourceTable.
149 This is anticipated to be a deepCoadd_meas with PSF fit parameters from a
150 measurement plugin returning covariance matrix terms.
151 """
153 action_source = ConfigurableActionField[PsfFitSuccessActionBase](
154 doc="Action to return whether the PSF fit was successful for a single source row",
155 default=SourceTablePsfFitSuccessAction,
156 )
157 format = pexConfig.Field[str](
158 doc="Format for the field names, where {idx_comp} is the index of the component and {moment}"
159 "is the name of the moment (xx, xy or yy, integral)",
160 default="modelfit_DoubleShapeletPsfApprox_{idx_comp}_{moment}",
161 )
162 name_moment_xx = pexConfig.Field[str](doc="Name of the xx (2nd x-axis) moment", default="xx")
163 name_moment_xy = pexConfig.Field[str](doc="Name of the xy (covariance term) moment", default="xy")
164 name_moment_yy = pexConfig.Field[str](doc="Name of the yy (2nd y-axis) moment", default="yy")
165 name_moment_integral = pexConfig.Field[str](doc="Name of the integral (zeroth) moment", default="0")
166 n_components = pexConfig.Field[int](
167 doc="Number of Gaussian components",
168 default=2,
169 check=lambda x: x >= 2,
170 )
172 @staticmethod
173 def get_integral(moment_zero) -> float:
174 """Get the total integrated flux from a zeroth moment value.
176 The zeroth moment is simply the integrated flux divided by a
177 constant value of 2*sqrt(pi).
179 Parameters
180 ----------
181 moment_zero
182 The zeroth moment value.
184 Returns
185 -------
186 integral
187 The total integrated weight (flux).
188 """
189 return moment_zero * TWO_SQRT_PI
191 def get_schema(self) -> list[str]:
192 names_moments = (
193 self.name_moment_xx,
194 self.name_moment_yy,
195 self.name_moment_xy,
196 self.name_moment_integral,
197 )
198 columns = [
199 column
200 for idx_comp in range(self.n_components)
201 for column in (
202 self.format.format(name_moment=name_moment, idx_comp=idx_comp)
203 for name_moment in names_moments
204 )
205 ] + self.action_source.get_schema()
206 return columns
208 def __call__(self, source: Mapping[str, Any], *args: Any, **kwargs: Any) -> list[g2.Gaussian]:
209 if not self.action_source(source):
210 raise PsfRebuildFitFlagError(
211 f"PSF fit failed due to action based on schema: {self.action_source.get_schema()}"
212 )
213 gaussians = [None] * self.n_components
214 for idx_comp in range(self.n_components):
215 gaussian = g2.Gaussian(
216 ellipse=g2.Ellipse(
217 g2.Covariance(
218 sigma_x_sq=source[self.format.format(moment=self.name_moment_xx, idx_comp=idx_comp)],
219 sigma_y_sq=source[self.format.format(moment=self.name_moment_yy, idx_comp=idx_comp)],
220 cov_xy=source[self.format.format(moment=self.name_moment_xy, idx_comp=idx_comp)],
221 )
222 ),
223 integral=g2.GaussianIntegralValue(
224 value=self.get_integral(
225 source[self.format.format(moment=self.name_moment_integral, idx_comp=idx_comp)]
226 )
227 ),
228 )
229 gaussians[idx_comp] = gaussian
230 return gaussians
233class MagnitudeDependentSizePriorConfig(pexConfig.Config):
234 """Configuration for a magnitude-dependent size prior.
236 Defaults are for ugrizy total mag and log10(r_eff/arcsec).
237 """
239 intercept_mag = pexConfig.Field[float](
240 doc="The magnitude at which no adjustment is applied",
241 default=18.0,
242 )
243 slope_median_per_mag = pexConfig.Field[float](
244 doc="The slope in the median size, in dex per mag",
245 default=-0.15,
246 )
247 slope_stddev_per_mag = pexConfig.Field[float](
248 doc="The slope in the standard deviation of the size, in dex per mag",
249 default=0.0,
250 )
253class ModelInitializer(ABC, pydantic.BaseModel):
254 """An interface for a configurable model initializer based on priors
255 and optional external data.
256 """
258 model_config: ClassVar[pydantic.ConfigDict] = frozen_arbitrary_allowed_config
260 inputs: dict[str, Any] = pydantic.Field(
261 title="Additional external inputs used in initialization",
262 default_factory=dict,
263 )
264 priors_shape_mag: dict = pydantic.Field(
265 title="Magnitude-dependent shape prior configurations",
266 default_factory=dict,
267 )
269 @abstractmethod
270 def initialize_model(
271 self,
272 model: Model,
273 source: Mapping[str, Any],
274 catexps: list[CatalogExposureSourcesABC],
275 config_data: CatalogSourceFitterConfigData,
276 values_init: Mapping[g2f.ParameterD, float] | None = None,
277 **kwargs,
278 ):
279 """Initialize a MultiProFit model for a single object corresponding
280 to a row in a catalog.
282 Parameters
283 ----------
284 model
285 The model to initialize parameter values for.
286 source
287 A mapping with fields expected to be populated in the
288 corresponding source catalog for initialization.
289 catexps
290 Per-band catalog-exposure pairs.
291 config_data
292 Fitter configuration and data.
293 values_init
294 Default initial values for parameters.
295 **kwargs
296 Additional keyword arguments for any purpose.
297 """
298 raise NotImplementedError(f"{self.__name__} must implement initialize_model")
301class MakeInitializerActionBase(ConfigurableAction):
302 """An interface for an action that creates an initializer."""
304 def __call__(
305 self,
306 catalog_multi: Sequence,
307 catexps: list[fitMB.CatalogExposureInputs],
308 config_data: CatalogSourceFitterConfigData,
309 **kwargs,
310 ) -> ModelInitializer:
311 """Make a ModelInitializer object that can initialize model
312 parameter values for a given object in a catalog.
314 Parameters
315 ----------
316 catalog_multi
317 The multiband catalog with one row per object to fit.
318 catexps
319 Per-band catalog-exposure pairs.
320 config_data
321 Fitter configuration and data.
322 **kwargs
323 Additional arguments to pass to add to ModelInitializer.inputs.
325 Returns
326 -------
327 initializer
328 The configured ModelInitializer.
329 """
330 raise NotImplementedError(f"{self.__name__} must implement __call__")
333class BasicModelInitializerConfig(pexConfig.Config):
334 """Configuration for a BasicModelInitializer."""
336 psf_factor_shrink = pexConfig.Field[float](
337 doc="Multiplicative factor to shrink PSF sizes by for deconvolution",
338 default=0.9,
339 check=lambda x: 0.0 <= x < 1.0,
340 )
341 psf_factor_minimum = pexConfig.Field[float](
342 doc="Factor to multiply the PSF size by for a minimum initialize size",
343 default=0.5,
344 check=lambda x: x >= 0,
345 )
346 size_minimum = pexConfig.Field[float](
347 doc="Absolute minimum initial size in pixels",
348 default=0.5,
349 check=lambda x: x >= 0,
350 )
351 rho_abs_max = pexConfig.Field[float](
352 doc="Maximum absolute initial value of rho",
353 default=0.8,
354 check=lambda x: x >= 0,
355 )
358class BasicModelInitializer(ModelInitializer):
359 """A generic model initializer that should work on most kinds of models
360 with a single source.
361 """
363 config: BasicModelInitializerConfig = pydantic.Field(title="A BasicModelInitializerConfig to be frozen")
365 def _get_params_init(self, model_sources: tuple[g2f.Source]) -> tuple[g2f.ParameterD]:
366 """Return an ordered set of free parameters from a model's sources.
368 Parameters
369 ----------
370 model_sources
371 The sources in the model.
373 Returns
374 -------
375 params_init
376 The parameter objects for sources in the model.
378 Notes
379 -----
380 Only free and/or centroid parameters are returned (centroids are
381 always needed even if they are fixed).
382 """
383 # TODO: There ought to be a better way to not get the PSF centroids
384 # (those are part of model.data's fixed parameters)
385 params_init = (
386 tuple(
387 param
388 for param in get_params_uniq(model_sources[0])
389 if param.free
390 or (isinstance(param, g2f.CentroidXParameterD) or isinstance(param, g2f.CentroidYParameterD))
391 )
392 if (len(model_sources) == 1)
393 else tuple(
394 {
395 param: None
396 for source in model_sources
397 for param in get_params_uniq(source)
398 if param.free
399 or (
400 isinstance(param, g2f.CentroidXParameterD)
401 or isinstance(param, g2f.CentroidYParameterD)
402 )
403 }.keys()
404 )
405 )
406 return params_init
408 def _get_priors_type(
409 self,
410 priors: tuple[g2f.Prior],
411 ) -> tuple[tuple[g2f.GaussianPrior], tuple[g2f.ShapePrior]]:
412 """Return the list of priors of known type, by type.
414 Parameters
415 ----------
416 priors
417 A list of priors of any type, typically from a model.
419 Returns
420 -------
421 priors_gauss
422 A list of all of the Gaussian priors, in the order they occurred.
423 priors_shape
424 A list of all of the shape priors, in the order they occurred.
425 """
426 priors_gauss: list[g2f.GaussianPrior] = []
427 priors_shape: list[g2f.ShapePrior] = []
428 for prior in priors:
429 if isinstance(prior, g2f.GaussianPrior):
430 priors_gauss.append(prior)
431 elif isinstance(prior, g2f.ShapePrior):
432 priors_shape.append(prior)
433 return tuple(priors_gauss), tuple(priors_shape)
435 def get_centroid_and_shape(
436 self,
437 source: Mapping[str, Any],
438 catexps: list[CatalogExposureSourcesABC],
439 config_data: CatalogSourceFitterConfigData,
440 values_init: Mapping[g2f.ParameterD, float] | None = None,
441 ) -> tuple[tuple[float, float], tuple[float, float, float]]:
442 """Get the centroid and shape for a source.
444 Parameters
445 ----------
446 source
447 A mapping with fields expected to be populated in the
448 corresponding source catalog for initialization.
449 catexps
450 A list of (source and psf) catalog-exposure pairs.
451 config_data
452 Configuration settings and data for fitting and output.
453 values_init
454 Initial parameter values from the model configuration.
456 Returns
457 -------
458 centroid
459 The x- and y-axis centroid values.
460 sig_x, sig_y, rho
461 The x- and y-axis Gaussian sigma and rho values defining the
462 estimated elliptical shape of the source.
463 """
464 centroid = source["slot_Centroid_x"], source["slot_Centroid_y"]
465 # Attempt partial deconvolution of observed moments
466 psf_factor_shrink = self.config.psf_factor_shrink**2
467 psf_factor_minimum = self.config.psf_factor_minimum**2
468 rho_min, rho_max = -self.config.rho_abs_max, self.config.rho_abs_max
469 psf_xx = source["base_SdssShape_psf_xx"]
470 psf_yy = source["base_SdssShape_psf_yy"]
471 sig_x, sig_y = (
472 math.sqrt(
473 np.nanmax(
474 (
475 source[f"slot_Shape_{suffix}"] - moment_sq * psf_factor_shrink,
476 moment_sq * psf_factor_minimum,
477 self.config.size_minimum,
478 )
479 )
480 )
481 for suffix, moment_sq in (("xx", psf_xx), ("yy", psf_yy))
482 )
483 psf_xy = source["base_SdssShape_psf_xy"]
484 sig_xy = sig_x * sig_y
485 if not (sig_xy > 0):
486 rho = 0
487 else:
488 rho = np.clip((source["slot_Shape_xy"] - psf_xy * psf_factor_shrink) / sig_xy, rho_min, rho_max)
489 shape = sig_x, sig_y, rho
490 return centroid, shape
492 def get_params_init(self, model: Model) -> tuple[g2f.ParameterD]:
493 """Return the free and/or centroid parameters for a model.
495 Parameters
496 ----------
497 model
498 The model to return parameters for.
500 Returns
501 -------
502 parameters
503 The ordered list of parameters for the model.
504 """
505 return self._get_params_init(model_sources=model.sources)
507 def get_priors_type(self, model: Model) -> tuple[tuple[g2f.GaussianPrior], tuple[g2f.ShapePrior]]:
508 """Return the list of priors of known type, by type.
510 Parameters
511 ----------
512 model
513 The model to return priors for.
515 Returns
516 -------
517 priors_gauss
518 A list of all of the Gaussian priors, in the order they occurred.
519 priors_shape
520 A list of all of the shape priors, in the order they occurred.
521 """
522 return self._get_priors_type(model.priors)
524 def initialize_model(
525 self,
526 model: Model,
527 source: Mapping[str, Any],
528 catexps: list[CatalogExposureSourcesABC],
529 config_data: CatalogSourceFitterConfigData,
530 values_init: Mapping[g2f.ParameterD, float] | None = None,
531 **kwargs,
532 ):
533 if values_init is None:
534 values_init = {}
535 set_flux_limits = kwargs.pop("set_flux_limits", True)
536 flux_init_min = kwargs.pop("value_init_min", 1e-10)
537 flux_limit_min = kwargs.pop("flux_limit_min", 1e-12)
538 if kwargs:
539 raise ValueError(f"Unexpected {kwargs=}")
540 centroid_pixel_offset = config_data.config.centroid_pixel_offset
541 (cen_x, cen_y), (sig_x, sig_y, rho) = self.get_centroid_and_shape(
542 source,
543 catexps,
544 config_data,
545 values_init=values_init,
546 )
547 # If we couldn't get a shape at all, make it small and roundish
548 if not np.isfinite(rho):
549 # Note rho=0 (circular) is generally disfavoured by shape priors
550 # However, setting it to a non-zero value seems to make scipy
551 # fail to move off initial conditions, as do sizes below 2 pixels
552 sig_x, sig_y, rho = 2.0, 2.0, 0.0
554 # Make restrictive centroid limits (intersection, not union)
555 x_min, y_min, x_max, y_max = -np.inf, -np.inf, np.inf, np.inf
557 fluxes_init = {}
558 fluxes_limits = {}
560 # This is the maximum number of potential observations
561 # They might not all have made it into the data
562 n_catexps = len(catexps)
563 n_components = len(model.sources[0].components)
565 # If not true, some bands must have no data to fit
566 if len(catexps) != len(model.data):
567 catexps_obs = []
568 for catexp in catexps:
569 fluxes_init[catexp.channel] = flux_init_min
570 fluxes_limits[catexp.channel] = (0, np.inf)
571 # No associated catalog means we can't fit (and should be
572 # because there's no exposure for this band in this patch)
573 if len(catexp.get_catalog()) > 0:
574 catexps_obs.append(catexp)
575 else:
576 catexps_obs = catexps
578 for idx_obs, observation in enumerate(model.data):
579 coordsys = observation.image.coordsys
580 catexp = catexps_obs[idx_obs]
581 band = catexp.band
583 x_min = max(x_min, coordsys.x_min)
584 y_min = max(y_min, coordsys.y_min)
585 x_max = min(x_max, coordsys.x_min + float(observation.image.n_cols))
586 y_max = min(y_max, coordsys.y_min + float(observation.image.n_rows))
588 flux_total = np.nansum(observation.image.data[observation.mask_inv.data])
590 column_ref = f"merge_measurement_{band}"
591 if column_ref in source.schema.getNames() and source[column_ref]:
592 row = source
593 else:
594 row = catexp.catalog.find(source["id"])
596 if not row["base_SdssShape_flag"]:
597 flux_init = row["base_SdssShape_instFlux"]
598 else:
599 flux_init = row["slot_GaussianFlux_instFlux"]
600 if not (flux_init > 0):
601 flux_init = row["slot_PsfFlux_instFlux"]
603 calib = catexp.exposure.photoCalib
604 flux_init = calib.instFluxToNanojansky(flux_init) if (flux_init > 0) else max(flux_total, 1.0)
605 if set_flux_limits:
606 flux_max = 10 * max((flux_init, flux_total))
607 flux_min = min(flux_limit_min, flux_max / 1000)
608 else:
609 flux_min, flux_max = 0, np.inf
610 if not (flux_init > flux_min):
611 flux_upper = flux_max if (flux_max < np.inf) else 10.0 * flux_min
612 flux_init = flux_min + 0.01 * (flux_upper - flux_min)
613 fluxes_init[observation.channel] = flux_init / n_components
614 fluxes_limits[observation.channel] = (flux_min, flux_max)
616 if not np.isfinite(cen_x):
617 cen_x = observation.image.n_cols / 2.0
618 else:
619 cen_x -= centroid_pixel_offset
620 if not np.isfinite(cen_y):
621 # TODO: Add bbox coords or remove
622 cen_y = observation.image.n_rows / 2.0
623 else:
624 cen_y -= centroid_pixel_offset
626 # An R_eff larger than the box size is problematic. This should also
627 # stop unreasonable size proposals; a log10 transform isn't enough.
628 # TODO: Try logit for r_eff?
629 size_major = g2.EllipseMajor(g2.Ellipse(sigma_x=sig_x, sigma_y=sig_y, rho=rho)).r_major
630 limits_size = max(5.0 * size_major, 2.0 * np.hypot(x_max - x_min, y_max - y_min))
631 limits_xy = (1e-5, limits_size)
632 params_limits_init = {
633 g2f.CentroidXParameterD: (cen_x, (x_min, x_max)),
634 g2f.CentroidYParameterD: (cen_y, (y_min, y_max)),
635 g2f.ReffXParameterD: (sig_x, limits_xy),
636 g2f.ReffYParameterD: (sig_y, limits_xy),
637 g2f.SigmaXParameterD: (sig_x, limits_xy),
638 g2f.SigmaYParameterD: (sig_y, limits_xy),
639 g2f.RhoParameterD: (rho, None),
640 # TODO: get guess from configs?
641 g2f.SersicMixComponentIndexParameterD: (1.0, None),
642 }
644 fluxes_init_tuple = tuple(fluxes_init.values())
645 fluxes_limits_tuple = tuple(fluxes_limits.values())
646 idx_obs = 0
647 for param in self.params_init:
648 if param.linear:
649 value_init = fluxes_init_tuple[idx_obs]
650 limits_new = fluxes_limits_tuple[idx_obs]
651 idx_obs += 1
652 if idx_obs == n_catexps:
653 idx_obs = 0
654 else:
655 type_param = type(param)
656 value_init, limits_new = params_limits_init.get(type_param, (values_init.get(param), None))
657 if limits_new:
658 param.limits = g2f.LimitsD(limits_new[0], limits_new[1])
659 if value_init is not None:
660 param.value = np.clip(value_init, param.limits.min, param.limits.max)
662 priors_shape_mag = self.priors_shape_mag
663 has_priors_mag = len(priors_shape_mag) > 0
664 if has_priors_mag:
665 mag_total = u.nJy.to(u.ABmag, np.nansum(fluxes_init_tuple))
667 # TODO: Add centroid prior
668 priors_gauss, priors_shape = self.get_priors_type(model)
669 for prior in priors_shape:
670 if has_priors_mag and ((prior_adjustments := priors_shape_mag.get(prior)) is not None):
671 mag_dep_prior, prior_shape_new = prior_adjustments
672 prior_size_new = prior_shape_new.prior_size
673 # the size-apparent mag relation probably flattens
674 # for very bright/faint objects - maybe not so
675 # sharply, but clipping a broad mag range ought to be fine
676 prior.prior_size.mean_parameter.value = prior_size_new.mean_parameter.value * 10 ** (
677 mag_dep_prior.slope_median_per_mag
678 * np.clip(
679 mag_total - mag_dep_prior.intercept_mag,
680 -12.5,
681 12.5,
682 )
683 )
684 # it's uncertain how the intrinsic scatter behaves
685 # educated guess is it doesn't change much, also
686 # one runs out of bright galaxies to measure it anyway
687 prior.prior_size.stddev_parameter.value = prior_size_new.stddev_parameter.value * 10 ** (
688 mag_dep_prior.slope_stddev_per_mag
689 * np.clip(
690 mag_total - mag_dep_prior.intercept_mag,
691 -12.5,
692 12.5,
693 )
694 )
695 else:
696 prior.prior_size.mean_parameter.value = size_major
699class CachedBasicModelInitializer(BasicModelInitializer):
700 """A basic initializer with a cached list of model sources and priors."""
702 priors: tuple[g2f.Prior, ...] = pydantic.Field(title="The gauss2d_fit model priors")
703 sources: tuple[g2f.Source, ...] = pydantic.Field(title="The gauss2d_fit model sources")
705 @cached_property
706 def params_init(self) -> tuple[g2f.ParameterD]:
707 """Return a cached reference to the result of _get_params_init."""
708 return self._get_params_init(model_sources=self.sources)
710 @cached_property
711 def priors_type(self) -> tuple[tuple[g2f.GaussianPrior], tuple[g2f.ShapePrior]]:
712 """Return a cached reference to the result of _get_priors_type."""
713 return self._get_priors_type(self.priors)
715 def get_params_init(self, model: Model) -> tuple[g2f.ParameterD]:
716 assert tuple(model.sources) == self.sources
717 return self.params_init
719 def get_priors_type(self, model: Model) -> tuple[tuple[g2f.GaussianPrior], tuple[g2f.ShapePrior]]:
720 assert tuple(model.priors) == self.priors
721 return self.priors_type
724class InitialInputData(pydantic.BaseModel):
725 """A configurable wrapper to retrieve formatted columns from a catalog.
727 This provides a common interface to typical MultiProFit table outputs.
728 """
730 model_config: ClassVar[pydantic.ConfigDict] = frozen_arbitrary_allowed_config
732 column_id: str | None = pydantic.Field(
733 title="Override for id column specified in config_input",
734 default=None,
735 )
736 config_input: InputConfig = pydantic.Field(title="Configuration for the data table")
737 data: Table = pydantic.Field(title="The data table")
738 name_model: str = pydantic.Field(title="The name of the model in columns")
739 prefix_column: str = pydantic.Field(title="The prefix for all fitted column names")
740 size_column: str = pydantic.Field(title="The name of the size column", default="reff")
742 def get_column_id(self):
743 """Return the name of the object ID column."""
744 return self.column_id or self.config_input.column_id
746 def get_column(self, name_column: str, data=None):
747 """Get the values from a column.
749 Parameters
750 ----------
751 name_column
752 The name of the column to retrieve.
753 data
754 The catalog to retrieve the column from. Default is self.data.
756 Returns
757 -------
758 values
759 The column values.
760 """
761 if data is None:
762 data = self.data
763 return data[f"{self.prefix_column}{name_column}"]
765 def model_post_init(self, __context: Any) -> None:
766 # Initialize a mapping of the row number for a given object ID value
767 # This is implemented in afw catalogs but not most other tabular types
768 id_index = {idnum: idx for idx, idnum in enumerate(self.data[self.get_column_id()])}
769 object.__setattr__(self, "id_index", id_index)
772class MakeBasicInitializerAction(MakeInitializerActionBase):
773 """An action to construct an initializer for a single-component,
774 single-source model.
775 """
777 config = pexConfig.ConfigField[BasicModelInitializerConfig](
778 doc="Configuration for the initializer to be constructed",
779 )
781 def _make_initializer(
782 self,
783 catalog_multi: Sequence,
784 catexps: list[fitMB.CatalogExposureInputs],
785 config_data: CatalogSourceFitterConfigData,
786 ) -> ModelInitializer:
787 return BasicModelInitializer(config=self.config)
789 def __call__(
790 self,
791 catalog_multi: Sequence,
792 catexps: list[fitMB.CatalogExposureInputs],
793 config_data: CatalogSourceFitterConfigData,
794 **kwargs,
795 ) -> ModelInitializer:
796 initializer = self._make_initializer(
797 catalog_multi=catalog_multi,
798 catexps=catexps,
799 config_data=config_data,
800 )
801 for name, (config_input, data) in kwargs.items():
802 if not isinstance(data, Table) and hasattr(data, "meta"):
803 _LOG.warning(
804 f"Ignoring extra input {name=} because it is of type {type(data)} and is either not an"
805 f" astropy.table.Table or missing a 'meta' attr"
806 )
807 config_data = data.meta["config"]
808 prefix_column = config_data["prefix_column"]
809 config_source = next(iter(config_data["config_model"]["sources"].values()))
810 config_group = next(iter(config_source["component_groups"].values()))
811 is_sersic = len(config_group["components_sersic"]) > 0
812 name_model = next(
813 iter(config_group["components_sersic"] if is_sersic else config_group["components_gaussian"])
814 )
815 initializer.inputs[name] = InitialInputData(
816 column_id=config_data.get("column_id"),
817 config_input=config_input,
818 data=data,
819 name_model=name_model,
820 prefix_column=prefix_column,
821 size_column="reff" if is_sersic else "sig",
822 )
823 return initializer
826class MakeCachedBasicInitializerAction(MakeBasicInitializerAction):
827 """A MakeBasicInitializerAction that caches references to the source
828 and prior objects of the model.
830 This is solely a performance optimization and should be favored over
831 MakeBasicInitializerAction unless the caching is shown to be slower.
832 """
834 def _make_initializer(
835 self,
836 catalog_multi: Sequence,
837 catexps: list[fitMB.CatalogExposureInputs],
838 config_data: CatalogSourceFitterConfigData,
839 ) -> ModelInitializer:
840 sources, priors = config_data.sources_priors
841 return CachedBasicModelInitializer(config=self.config, priors=priors, sources=sources)
844class MultiProFitSourceConfig(CatalogSourceFitterConfig, fitMB.CoaddMultibandFitSubConfig):
845 """Configuration for the MultiProFit profile fitter."""
847 action_initializer = ConfigurableActionField[MakeInitializerActionBase](
848 doc="The action to return an initializer",
849 default=MakeCachedBasicInitializerAction,
850 )
851 action_psf = ConfigurableActionField[PsfComponentsActionBase](
852 doc="The action to return PSF component values from catalogs, if implemented",
853 default=None,
854 )
855 columns_copy = pexConfig.DictField[str, str](
856 doc="Mapping of input/output column names to copy from the input"
857 "multiband catalog to the output fit catalog.",
858 default={},
859 dictCheck=lambda x: len(set(x.values())) == len(x.values()),
860 )
861 mask_names_zero = pexConfig.ListField[str](
862 doc="Mask bits to mask out",
863 default=["BAD", "EDGE", "SAT", "NO_DATA"],
864 )
865 psf_sigma_subtract = pexConfig.Field[float](
866 doc="PSF x/y sigma value to subtract in quadrature from best-fit values",
867 default=0.0,
868 check=lambda x: np.isfinite(x) and (x >= 0),
869 )
870 prefix_column = pexConfig.Field[str](default="mpf_", doc="Column name prefix")
871 size_priors = pexConfig.ConfigDictField[str, MagnitudeDependentSizePriorConfig](
872 doc="Per-component magnitude-dependent size prior configurations."
873 " Will be added to component with existing configs.",
874 default={},
875 )
877 def bands_read_only(self) -> set[str]:
878 # TODO: Re-implement determination of prior-only bands once
879 # data-driven priors are re-implemented (DM-4xxxx)
880 return set()
882 def requires_psf(self):
883 """Return whether the PSF action is not None."""
884 return type(self.action_psf) is PsfComponentsActionBase
886 def setDefaults(self):
887 super().setDefaults()
888 self.defer_radec_conversion = True
889 self.compute_radec_covariance = True
890 self.flag_errors = {
891 IsParentError.column_name(): "IsParentError",
892 NoDataError.column_name(): "NoDataError",
893 NotPrimaryError.column_name(): "NotPrimaryError",
894 PsfRebuildFitFlagError.column_name(): "PsfRebuildFitFlagError",
895 }
896 self.centroid_pixel_offset = -0.5
897 self.naming_scheme = "lsst"
898 self.prefix_column = ""
899 self.suffix_error = "Err"
902@pydantic.dataclasses.dataclass(frozen=True, kw_only=True, config=fitMB.CatalogExposureConfig)
903class CatalogExposurePsfs(fitMB.CatalogExposureInputs, CatalogExposureSourcesABC):
904 """Input data from lsst pipelines, parsed for MultiProFit."""
906 channel: g2f.Channel = pydantic.Field(title="Channel for the image's band")
907 config_fit: MultiProFitSourceConfig = pydantic.Field(title="Config for fitting options")
909 @cached_property
910 def _psf_flux_params(self) -> tuple[list[g2f.ParameterD], bool]:
911 psf_model = self.psf_model_data.psf_model
912 n_comps = len(psf_model.components)
913 params_flux = [None] * n_comps
914 is_frac = [False] * n_comps
915 for idx_comp, comp in enumerate(psf_model.components):
916 # TODO: Change to comp.integralmodel when DM-44344 is fixed
917 # integralmodels will still need to be handled differently
918 params_all = get_params_uniq(comp)
919 params_frac = [param for param in params_all if isinstance(param, g2f.ProperFractionParameterD)]
920 if params_frac:
921 is_last = idx_comp == (n_comps - 1)
922 if len(params_frac) != (idx_comp + 1 - is_last):
923 raise RuntimeError(
924 f"Got unexpected {params_frac=} for"
925 f" {self.psf_model_data.psf_model.components[idx_comp]=} ({idx_comp=});"
926 f" len should be idx_comp+1"
927 )
928 params_flux[idx_comp] = None if is_last else params_frac[idx_comp]
929 is_frac[idx_comp] = True
930 else:
931 params_integral = [param for param in params_all if isinstance(param, g2f.IntegralParameterD)]
932 if len(params_integral != 1):
933 raise RuntimeError(
934 f"Got unexpected {params_integral=} != 1 for"
935 f" {self.psf_model_data.psf_model.components[idx_comp]=} ({idx_comp=})"
936 )
937 params_flux[idx_comp] = params_integral[0]
938 is_frac_any = any(is_frac)
939 if is_frac_any and not all(is_frac):
940 # TODO: This should work by iterating through componentgroups
941 # But that's not trivial or supported now
942 raise RuntimeError("Got PSF model with a mix of fractional and linear models; cannot initialize")
944 return params_flux, is_frac_any
946 def get_psf_model(self, params: Mapping[str, Any]) -> g2f.PsfModel | None:
947 psf_model = self.psf_model_data.psf_model
948 # PsfComponentsActionBase is an abstract class, so check if the action
949 # is a subclass that needs to be called
950 if not self.config_fit.requires_psf():
951 try:
952 gaussians = self.config_fit.action_psf(params)
953 except PsfRebuildFitFlagError:
954 return None
955 n_comps = len(psf_model.components)
956 fluxes = [0.0] * n_comps
957 params_flux, is_frac = self._psf_flux_params
958 for idx_comp, (comp, gaussian) in enumerate(zip(psf_model.components, gaussians)):
959 ellipse_out = comp.ellipse
960 ellipse_in = gaussian.ellipse
961 ellipse_out.sigma_x = ellipse_in.sigma_x
962 ellipse_out.sigma_y = ellipse_in.sigma_y
963 ellipse_out.rho = ellipse_in.rho
964 fluxes[idx_comp] = gaussian.integral.value
965 # Apparently negative fluxes are possible. Not much can be done to
966 # fix that but set them to a tiny value (zero might work)
967 fluxes = np.clip(fluxes, 1e-3, np.inf)
968 flux_total = sum(fluxes)
969 if is_frac:
970 flux_remaining = 1.0
971 for flux, param_frac in zip(fluxes, params_flux[:-1]):
972 flux_component = flux / flux_total
973 param_frac.value = flux_component / flux_remaining
974 flux_remaining -= flux_component
975 else:
976 for flux, param_flux in zip(fluxes, params_flux):
977 param_flux.value = flux / flux_total
978 else:
979 # TODO: this should probably use .index or something
980 match = np.argwhere(
981 self.table_psf_fits[self.psf_model_data.config.column_id] == params[self.config_fit.column_id]
982 )[0][0]
983 psf_model = self.psf_model_data.psf_model
984 try:
985 self.psf_model_data.init_psf_model(self.table_psf_fits[match])
986 except PsfRebuildFitFlagError:
987 return None
989 sigma_subtract = self.config_fit.psf_sigma_subtract
990 if sigma_subtract > 0:
991 sigma_subtract_sq = sigma_subtract * sigma_subtract
992 # 1/10 of PSF sigma should suffice as a minimum size
993 sigma_min_sq = sigma_subtract_sq / 100.0
994 for param in self.psf_model_data.parameters.values():
995 if isinstance(
996 param,
997 g2f.SigmaXParameterD | g2f.SigmaYParameterD | g2f.ReffXParameterD | g2f.ReffYParameterD,
998 ):
999 param.value = math.sqrt(max(param.value**2 - sigma_subtract_sq, sigma_min_sq))
1000 return psf_model
1002 def get_source_observation(self, source, **kwargs) -> g2f.ObservationD | None:
1003 footprint = source.getFootprint()
1004 bbox = footprint.getBBox()
1005 if not (bbox.getArea() > 0):
1006 return None
1007 bitmask = 0
1008 mask = self.exposure.mask[bbox]
1009 spans = footprint.spans.asArray()
1010 for bitname in self.config_fit.mask_names_zero:
1011 bitval = mask.getPlaneBitMask(bitname)
1012 bitmask |= bitval
1013 mask = ((mask.array & bitmask) != 0) & (spans != 0)
1014 mask = ~mask
1016 is_deblended_child = source["parent"] != 0
1018 img, _, sigma_inv = get_spanned_image(
1019 exposure=self.exposure,
1020 footprint=footprint if is_deblended_child else None,
1021 bbox=bbox,
1022 spans=spans,
1023 get_sig_inv=True,
1024 )
1025 x_min_bbox, y_min_bbox = bbox.beginX, bbox.beginY
1026 # Crop to tighter box for deblended model if edges are unusable
1027 # ... this rarely ever seems to happen though
1028 if is_deblended_child:
1029 coords = np.argwhere(np.isfinite(img) & (sigma_inv > 0) & np.isfinite(sigma_inv))
1030 if len(coords) == 0:
1031 return None
1032 x_min, y_min = coords.min(axis=0)
1033 x_max, y_max = coords.max(axis=0)
1034 x_max += 1
1035 y_max += 1
1037 if (x_min > 0) or (y_min > 0) or (x_max < img.shape[0]) or (y_max < img.shape[1]):
1038 # Ensure the nominal centroid is still inside the box
1039 # ... although it's a bad sign if that row/column is all bad
1040 x_cen = source["slot_Centroid_x"] - x_min_bbox
1041 y_cen = source["slot_Centroid_y"] - y_min_bbox
1042 x_min = min(x_min, int(np.floor(x_cen)))
1043 x_max = max(x_max, int(np.ceil(x_cen)))
1044 y_min = min(y_min, int(np.floor(y_cen)))
1045 y_max = max(y_max, int(np.ceil(y_cen)))
1046 x_min_bbox += x_min
1047 y_min_bbox += y_min
1048 img = img[x_min:x_max, y_min:y_max]
1049 sigma_inv = sigma_inv[x_min:x_max, y_min:y_max]
1050 mask = mask[x_min:x_max, y_min:y_max]
1052 mask[~np.isfinite(img) | ~np.isfinite(sigma_inv)] = False
1053 sigma_inv[~mask] = 0
1055 coordsys = g2.CoordinateSystem(1.0, 1.0, x_min_bbox, y_min_bbox)
1057 obs = g2f.ObservationD(
1058 image=g2.ImageD(img, coordsys),
1059 sigma_inv=g2.ImageD(sigma_inv, coordsys),
1060 mask_inv=g2.ImageB(mask, coordsys),
1061 channel=self.channel,
1062 )
1063 return obs
1065 def __post_init__(self):
1066 # TODO: Can/should this be the derived type (MultiProFitPsfConfig)?
1067 config = CatalogPsfFitterConfig()
1068 config_dict = self.table_psf_fits.meta.get("config")
1069 if config_dict:
1070 set_config_from_dict(config, config_dict)
1071 else:
1072 # TODO: How should this be set?
1073 # If using external PSF fits, it needs to be configured normally
1074 pass
1075 config_data = CatalogPsfFitterConfigData(config=config)
1076 object.__setattr__(self, "psf_model_data", config_data)
1079class MultiProFitSourceFitter(CatalogSourceFitterABC):
1080 """A MultiProFit source fitter.
1082 Parameters
1083 ----------
1084 wcs
1085 A WCS solution that applies to all exposures.
1086 errors_expected
1087 A dictionary of exceptions that are expected to sometimes be raised
1088 during processing (e.g. for missing data) keyed by the name of the
1089 flag column used to record the failure.
1090 add_missing_errors
1091 Whether to add all of the standard MultiProFit errors with default
1092 column names to errors_expected, if not already present.
1093 **kwargs
1094 Keyword arguments to pass to the superclass constructor.
1095 """
1097 initializer: ModelInitializer = pydantic.Field(
1098 title="The model parameter initializer",
1099 default_factory=lambda: BasicModelInitializer(),
1100 )
1101 wcs: lsst.afw.geom.SkyWcs = pydantic.Field(
1102 title="The WCS object to use to convert pixel coordinates to RA/dec",
1103 )
1105 def __init__(
1106 self,
1107 wcs: lsst.afw.geom.SkyWcs,
1108 errors_expected: dict[str, Exception] | None = None,
1109 add_missing_errors: bool = True,
1110 **kwargs: Any,
1111 ):
1112 if errors_expected is None:
1113 errors_expected = {}
1114 if add_missing_errors:
1115 for error_catalog in (IsParentError, NoDataError, NotPrimaryError, PsfRebuildFitFlagError):
1116 if error_catalog not in errors_expected:
1117 errors_expected[error_catalog] = error_catalog.column_name()
1118 super().__init__(wcs=wcs, errors_expected=errors_expected, **kwargs)
1120 def copy_centroid_errors(
1121 self,
1122 columns_cenx_err_copy: tuple[str],
1123 columns_ceny_err_copy: tuple[str],
1124 results: Table,
1125 catalog_multi: Sequence,
1126 catexps: list[CatalogExposureSourcesABC],
1127 config_data: CatalogSourceFitterConfigData,
1128 ):
1129 for column in columns_cenx_err_copy:
1130 results[column] = catalog_multi["slot_Centroid_xErr"]
1131 for column in columns_ceny_err_copy:
1132 results[column] = catalog_multi["slot_Centroid_yErr"]
1134 def compute_model_radec_err(
1135 self,
1136 source_multi: Mapping[str, Any],
1137 results,
1138 columns_params_radec_err,
1139 idx: int,
1140 set_radec: bool = False,
1141 ) -> None:
1142 for (
1143 key_ra_err,
1144 key_dec_err,
1145 key_cen_x,
1146 key_cen_y,
1147 key_cen_x_err,
1148 key_cen_y_err,
1149 key_cen_ra_dec_cov,
1150 key_ra,
1151 key_dec,
1152 ) in columns_params_radec_err:
1153 (ra, dec), (ra_err, dec_err, ra_dec_cov) = afwTable.convertCentroid(
1154 self.wcs,
1155 results[key_cen_x][idx],
1156 results[key_cen_y][idx],
1157 results[key_cen_x_err][idx],
1158 results[key_cen_y_err][idx],
1159 0.0,
1160 )
1161 if set_radec:
1162 results[key_ra][idx], results[key_dec][idx] = ra, dec
1163 else:
1164 ra_in, dec_in = results[key_ra][idx], results[key_dec][idx]
1165 if not np.isclose((ra, dec), (ra_in, dec_in), rtol=1e-7, atol=1e-8):
1166 self._get_logger().warning(
1167 "idx=%i ra, dec = %f,%f differ significantly from convertCentroid ra, dec = %f, %f",
1168 idx,
1169 ra_in,
1170 dec_in,
1171 ra,
1172 dec,
1173 )
1174 results[key_ra_err][idx], results[key_dec_err][idx] = ra_err, dec_err
1175 if key_cen_ra_dec_cov is not None:
1176 results[key_cen_ra_dec_cov][idx] = ra_dec_cov
1178 def get_model_radec(self, source: Mapping[str, Any], cen_x: float, cen_y: float):
1179 # no extra conversions are needed here - cen_x, cen_y are in catalog
1180 # coordinates already
1181 ra, dec = self.wcs.pixelToSky(cen_x, cen_y)
1182 return ra.asDegrees(), dec.asDegrees()
1184 def initialize_model(
1185 self,
1186 model: g2f.ModelD,
1187 source: Mapping[str, Any],
1188 catexps: list[CatalogExposureSourcesABC],
1189 config_data: CatalogSourceFitterConfigData,
1190 values_init: Mapping[g2f.ParameterD, float] | None = None,
1191 **kwargs,
1192 ):
1193 self.initializer.initialize_model(
1194 model=model,
1195 source=source,
1196 catexps=catexps,
1197 config_data=config_data,
1198 values_init=values_init,
1199 **kwargs,
1200 )
1202 def make_CatalogExposurePsfs(
1203 self,
1204 catexp: fitMB.CatalogExposureInputs,
1205 config: MultiProFitSourceConfig,
1206 ) -> CatalogExposurePsfs:
1207 """Make a CatalogExposurePsfs from a list of inputs and a fit config.
1209 Parameters
1210 ----------
1211 catexp
1212 The input catalog-exposure pairs.
1213 config
1214 The MultiProFit source fitting config.
1216 Returns
1217 -------
1218 catexp_psf
1219 The resulting CatalogExposurePsfs.
1220 """
1221 catexp_psf = CatalogExposurePsfs(
1222 # dataclasses.asdict(catexp)_makes a recursive deep copy.
1223 # That must be avoided.
1224 **{key: getattr(catexp, key) for key in catexp.__dataclass_fields__.keys()},
1225 channel=g2f.Channel.get(catexp.band),
1226 config_fit=config,
1227 )
1228 return catexp_psf
1230 def validate_fit_inputs(
1231 self,
1232 catalog_multi: Sequence,
1233 catexps: list[CatalogExposurePsfs],
1234 config_data: CatalogSourceFitterConfigData = None,
1235 logger: logging.Logger = None,
1236 **kwargs: Any,
1237 ) -> None:
1238 errors = []
1239 for idx, catexp in enumerate(catexps):
1240 if not isinstance(catexp, CatalogExposurePsfs):
1241 errors.append(f"catexps[{idx=} {type(catexp)=} !isinstance(CatalogExposurePsfs)")
1242 # Pre-validate the model
1243 config_sources = config_data.config.config_model.sources
1244 model_sources, priors = config_data.sources_priors
1245 priors_shape = [prior for prior in priors if isinstance(prior, g2f.ShapePrior)]
1247 if len(config_sources.keys()) > 1:
1248 errors.append(f"model config has multiple sources: {list(config_sources.keys())}")
1249 elif len(priors_shape) > 0:
1250 idx_prior_found = 0
1251 name_source, config_source = next(iter(config_sources.items()))
1252 source = model_sources[0]
1253 config_groups = config_source.component_groups
1254 if len(config_groups.keys()) > 1:
1255 errors.append(f"model {name_source=} has multiple groups: {list(config_source.keys())}")
1256 else:
1257 name_group, config_group = next(iter(config_groups.items()))
1258 for idx_comp, (name_comp, config_comp) in enumerate(
1259 config_group.get_component_configs().items()
1260 ):
1261 ellipse = source.components[idx_comp].ellipse
1262 # component.ellipse returns a const ref and must be copied
1263 # The ellipse classes might need copy constructors
1264 ellipse_copy = type(ellipse)(
1265 # No kwargs here, since they are unfortunately not
1266 # standardized (e.g. Gaussian is sigma_x not size_x)
1267 # but the arg order is
1268 ellipse.size_x,
1269 ellipse.size_y,
1270 ellipse.rho,
1271 )
1272 prior_shape_new = config_comp.make_shape_prior(ellipse_copy)
1273 if prior_shape_new is not None:
1274 if idx_prior_found == len(priors_shape):
1275 errors.append(
1276 f"Could not validate prior for {name_source=} {name_group=} {name_comp=}"
1277 )
1278 break
1279 prior_shape_old = priors_shape[idx_prior_found]
1280 ll_new, ll_old = (
1281 prior.evaluate().loglike for prior in (prior_shape_new, prior_shape_old)
1282 )
1283 # The necessary tolerance for this check is uncertain
1284 if not np.isclose(ll_new, ll_old):
1285 logger.warning(
1286 f"shape prior for {name_comp=} got inconsistent {ll_new=} vs {ll_old}"
1287 )
1288 if (prior_shape_mod := config_data.config.size_priors.get(name_comp)) is not None:
1289 self.initializer.priors_shape_mag[prior_shape_old] = (
1290 prior_shape_mod,
1291 prior_shape_new,
1292 )
1294 if errors:
1295 raise RuntimeError("\n".join(errors))
1297 def validate_source(
1298 self,
1299 idx_row: int,
1300 catalog_multi: Sequence,
1301 ) -> None:
1302 source = catalog_multi[idx_row]
1303 if (not source["detect_isPrimary"]) or source["merge_peak_sky"]:
1304 raise NotPrimaryError(f"source {source['id']} has invalid flags for fit")
1307class MultiProFitSourceTask(fitMB.CoaddMultibandFitSubTask):
1308 """Run MultiProFit on Exposure/SourceCatalog pairs in multiple bands.
1310 This task uses MultiProFit to fit a single model to all sources in a coadd,
1311 using a previously-fit PSF model for each exposure. The task may also use
1312 prior measurements from single- or merged multiband catalogs for
1313 initialization.
1314 """
1316 ConfigClass: ClassVar = MultiProFitSourceConfig
1317 _DefaultName: ClassVar = "multiProFitSource"
1319 def make_default_fitter(
1320 self,
1321 catalog_multi: Sequence,
1322 catexps: list[fitMB.CatalogExposureInputs],
1323 config_data: CatalogSourceFitterConfigData,
1324 **kwargs,
1325 ) -> MultiProFitSourceFitter:
1326 """Make a default MultiProFitSourceFitter.
1328 Parameters
1329 ----------
1330 catalog_multi
1331 A multi-band, indexable source catalog.
1332 catexps
1333 Catalog-exposure-PSF model tuples to fit source models for.
1334 config_data
1335 Configuration and data for the initalizer.
1336 **kwargs
1337 Additional keyword arguments to pass to
1338 self.config.action_initializer.
1340 Returns
1341 -------
1342 fitter
1343 A MultiProFitSourceFitter using the first catexp's wcs.
1344 """
1345 initializer = self.config.action_initializer(
1346 catalog_multi=catalog_multi, catexps=catexps, config_data=config_data, **kwargs
1347 )
1348 # Look for the first WCS - they ought to be identical
1349 # If they are not, the patch coadd data model must have changed
1350 wcs = None
1351 for catexp in catexps:
1352 if catexp.exposure is not None:
1353 wcs = catexp.exposure.wcs
1354 break
1355 if wcs is None:
1356 raise RuntimeError(f"Could not find valid wcs in any of {catexps=}")
1357 fitter = MultiProFitSourceFitter(wcs=wcs, initializer=initializer)
1358 return fitter
1360 @utilsTimer.timeMethod
1361 def run(
1362 self,
1363 catalog_multi: Sequence,
1364 catexps: list[fitMB.CatalogExposureInputs],
1365 fitter: MultiProFitSourceFitter | None = None,
1366 **kwargs,
1367 ) -> pipeBase.Struct:
1368 """Run the MultiProFit source fit task on catalog-exposure pairs.
1370 Parameters
1371 ----------
1372 catalog_multi
1373 A multi-band, indexable source catalog.
1374 catexps
1375 Catalog-exposure-PSF model tuples to fit source models for.
1376 fitter
1377 The fitter instance to use. Default-initialized if not provided.
1378 **kwargs
1379 Additional keyword arguments to pass to self.fit.
1381 Returns
1382 -------
1383 catalog : `astropy.Table`
1384 A table with fit parameters for the PSF model at the location
1385 of each source.
1386 """
1387 n_catexps = len(catexps)
1388 if n_catexps == 0:
1389 raise ValueError("Must provide at least one catexp")
1390 catexps_conv: list[CatalogExposurePsfs] = [None] * n_catexps
1391 channels = [g2f.Channel.get(catexp.band) for catexp in catexps]
1392 config_data = CatalogSourceFitterConfigData(channels=channels, config=self.config)
1393 if fitter is None:
1394 inputs_init = kwargs.get("inputs_init")
1395 if inputs_init:
1396 del kwargs["inputs_init"]
1397 else:
1398 inputs_init = {}
1399 fitter = self.make_default_fitter(
1400 catalog_multi=catalog_multi, catexps=catexps, config_data=config_data, **inputs_init
1401 )
1402 for idx, catexp in enumerate(catexps):
1403 if not isinstance(catexp, CatalogExposurePsfs):
1404 catexp = fitter.make_CatalogExposurePsfs(catexp, config=self.config)
1405 catexps_conv[idx] = catexp
1406 catalog = fitter.fit(
1407 catalog_multi=catalog_multi, catexps=catexps_conv, config_data=config_data, **kwargs
1408 )
1409 for name_in, name_out in self.config.columns_copy.items():
1410 catalog[name_out] = catalog_multi[name_in]
1411 catalog[name_out].description = catalog_multi.schema.find(name_in).field.getDoc()
1412 return pipeBase.Struct(output=astropy_to_arrow(catalog))