Coverage for python/lsst/multiprofit/fitting/fit_source.py: 70%
518 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:54 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:54 +0000
1# This file is part of 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/>.
22import logging
23import time
24from abc import ABC, abstractmethod
25from collections.abc import Iterable, Mapping, Sequence
26from functools import cached_property
27from typing import Any, ClassVar, Self
29import astropy
30import astropy.units as u
31import numpy as np
32import pydantic
33from astropy.table import Table
35import lsst.gauss2d.fit as g2f
36import lsst.pex.config as pexConfig
37from lsst.utils.logging import PeriodicLogger
39from ..componentconfig import Fluxes, GaussianComponentConfig
40from ..errors import NoDataError, RaDecConversionNotImplementedError
41from ..modelconfig import ModelConfig
42from ..modeller import FitInputsDummy, Modeller
43from ..sourceconfig import ComponentGroupConfig, SourceConfig
44from ..utils import frozen_arbitrary_allowed_config, get_params_uniq
45from .fit_catalog import CatalogExposureABC, CatalogFitterConfig, ColumnInfo
47__all__ = [
48 "CatalogExposureSourcesABC",
49 "CatalogSourceFitterABC",
50 "CatalogSourceFitterConfig",
51 "CatalogSourceFitterConfigData",
52]
55class CatalogExposureSourcesABC(CatalogExposureABC):
56 """Interface for a CatalogExposure for source modelling."""
58 @property
59 def band(self) -> str:
60 """Return the name of the exposure's passband (e.g. 'r')."""
61 return self.channel.name
63 # Note: not named band because that's usually a string
64 @property
65 @abstractmethod
66 def channel(self) -> g2f.Channel:
67 """Return the exposure's associated channel object."""
69 @abstractmethod
70 def get_psf_model(self, params: Mapping[str, Any]) -> g2f.PsfModel | None:
71 """Get the PSF model for a given source row.
73 Parameters
74 ----------
75 params : Mapping[str, Any]
76 A mapping with parameter values for the best-fit PSF model at the
77 centroid of a single source.
79 Returns
80 -------
81 psf_model : `lsst.gauss2d.fit.PsfModel`
82 A PsfModel object initialized with the best-fit parameters, or None
83 if PSF rebuilding failed for an expected reason (i.e. the input PSF
84 fit table has a flag set).
85 """
87 @abstractmethod
88 def get_source_observation(self, source: Mapping[str, Any], **kwargs: Any) -> g2f.ObservationD | None:
89 """Get the Observation for a given source row.
91 Parameters
92 ----------
93 source : Mapping[str, Any]
94 A mapping with any values needed to retrieve an observation for a
95 single source.
96 **kwargs
97 Additional keyword arguments not used during fitting.
99 Returns
100 -------
101 observation : `lsst.gauss2d.fit.Observation`
102 An Observation object with suitable data for fitting parametric
103 models of the source, or None if the observation cannot be fit.
104 """
107class CatalogSourceFitterConfig(CatalogFitterConfig):
108 """Configuration for the MultiProFit profile fitter."""
110 centroid_pixel_offset = pexConfig.Field[float](
111 doc="Number to add to MultiProFit centroids (bottom-left corner is 0,0) to convert to catalog"
112 " coordinates (e.g. set to -0.5 if the bottom-left corner is -0.5, -0.5)",
113 default=0,
114 )
115 compute_radec_covariance = pexConfig.Field[bool](
116 doc="Whether to compute the RA/dec covariance. Ignore if convert_cen_xy_to_radec is False.",
117 default=False,
118 )
119 config_model = pexConfig.ConfigField[ModelConfig](doc="Source model configuration")
120 convert_cen_xy_to_radec = pexConfig.Field[bool](
121 doc="Convert pixel x/y centroid params to RA/dec",
122 default=True,
123 )
124 defer_radec_conversion = pexConfig.Field[bool](
125 doc="Whether to defer conversion of pixel x/y centroid params to RA/dec to compute_model_radec_err."
126 " Only effective if convert_cen_xy_to_radec and compute_errors is not NONE, and requires that the"
127 " overloaded compute_model_radec_err method sets RA/dec values itself.",
128 default=False,
129 )
130 fit_psmodel_final = pexConfig.Field[bool](
131 default=False,
132 doc="Fit a point source model after optimization",
133 )
134 prior_cen_x_stddev = pexConfig.Field[float](
135 default=0, doc="Prior std. dev. on x centroid (ignored if not >0)"
136 )
137 prior_cen_y_stddev = pexConfig.Field[float](
138 default=0, doc="Prior std. dev. on y centroid (ignored if not >0)"
139 )
140 unit_flux = pexConfig.Field[str](default=None, doc="Flux unit", optional=True)
142 def make_model_data(
143 self,
144 idx_row: int,
145 catexps: list[CatalogExposureSourcesABC],
146 ) -> tuple[g2f.DataD, list[g2f.PsfModel]]:
147 """Make data and psf_models for a catalog row.
149 Parameters
150 ----------
151 idx_row
152 The index of the row in each catalog.
153 catexps
154 Catalog-exposure pairs to initialize observations from.
156 Returns
157 -------
158 data
159 The resulting data object.
160 psf_models
161 A list of psf_models, one per catexp.
163 Notes
164 -----
165 Only observations with good data and valid PSF models will be
166 returned; bad data will be excluded from the return values.
167 """
168 observations = []
169 psf_models = []
171 for catexp in catexps:
172 catalog = catexp.get_catalog()
173 # This indicates that there's no corresponding exposure
174 # (the catexp interface expects a tabular type for catalog but
175 # no interface for an exposure has been defined, yet)
176 if len(catalog) == 0: 176 ↛ 177line 176 didn't jump to line 177 because the condition on line 176 was never true
177 continue
178 source = catalog[idx_row]
179 observation = catexp.get_source_observation(source)
180 # If the observation or PSF model is bad enough that it cannot be
181 # fit, do not add it to the data.
182 if observation is not None: 182 ↛ 171line 182 didn't jump to line 171 because the condition on line 182 was always true
183 psf_model = catexp.get_psf_model(source)
184 if psf_model is not None: 184 ↛ 171line 184 didn't jump to line 171 because the condition on line 184 was always true
185 observations.append(observation)
186 # PSF model parameters cannot be fit along with sources
187 for param in get_params_uniq(psf_model):
188 param.fixed = True
189 psf_models.append(psf_model)
191 data = g2f.DataD(observations)
192 return data, psf_models
194 def make_point_sources(
195 self,
196 channels: Iterable[g2f.Channel],
197 sources: list[g2f.Source],
198 ) -> tuple[list[g2f.Source], list[g2f.Prior]]:
199 """Make initialized point sources given channels.
201 Parameters
202 ----------
203 channels
204 The channels to initialize fluxes for.
205 sources
206 List of sources.
208 Returns
209 -------
210 sources
211 The list of initialized sources.
212 priors
213 The list of priors.
215 Notes
216 -----
217 The prior list is always empty, but is returned to keep this function
218 consistent with make_sources.
219 """
220 point_sources = []
221 fluxes = [[{channel: 1.0 for channel in channels}]]
223 for (name_src, config_src), source in zip(self.config_model.sources.items(), sources):
224 centroids = next(iter(config_src.component_groups.values())).centroids
225 config_src_psf = SourceConfig(
226 component_groups={
227 "": ComponentGroupConfig(
228 centroids=centroids,
229 components_gauss={"": GaussianComponentConfig()},
230 )
231 }
232 )
233 source, _ = config_src_psf.make_source(fluxes)
234 point_sources.append(source)
236 return point_sources, []
238 def make_sources(
239 self,
240 channels: Iterable[g2f.Channel],
241 source_fluxes: list[list[list[Fluxes]]] | None = None,
242 ) -> tuple[list[g2f.Source], list[g2f.Prior]]:
243 """Make initialized sources given channels using `self.config_model`.
245 Parameters
246 ----------
247 channels
248 The channels to initialize fluxes for.
249 source_fluxes
250 A list of fluxes by channel for each component group in each
251 source. The default is to initialize using
252 `ComponentGroupConfig.get_fluxes_default`.
254 Returns
255 -------
256 sources
257 The list of initialized sources.
258 priors
259 The list of priors.
260 """
261 n_sources = len(self.config_model.sources)
262 if source_fluxes is None: 262 ↛ 280line 262 didn't jump to line 280 because the condition on line 262 was always true
263 source_fluxes = [None] * n_sources
264 for idx, (config_source, component_group_fluxes) in enumerate(
265 zip(
266 self.config_model.sources.values(),
267 source_fluxes,
268 )
269 ):
270 component_group_fluxes = [
271 component_group.get_fluxes_default(
272 channels=channels,
273 component_configs=component_group.get_component_configs(),
274 is_fractional=component_group.is_fractional,
275 )
276 for component_group in config_source.component_groups.values()
277 ]
278 source_fluxes[idx] = component_group_fluxes
279 else:
280 if len(source_fluxes) != n_sources:
281 raise ValueError(f"{len(source_fluxes)=} != {len(self.config_model.sources)=}")
283 sources, priors = self.config_model.make_sources(
284 component_group_fluxes_srcs=source_fluxes,
285 )
287 has_prior_x = self.prior_cen_x_stddev > 0 and np.isfinite(self.prior_cen_x_stddev)
288 has_prior_y = self.prior_cen_y_stddev > 0 and np.isfinite(self.prior_cen_y_stddev)
289 if has_prior_x or has_prior_y: 289 ↛ 290line 289 didn't jump to line 290 because the condition on line 289 was never true
290 for source in sources:
291 for param in get_params_uniq(source, fixed=False):
292 if has_prior_x and isinstance(param, g2f.CentroidXParameterD):
293 priors.append(g2f.GaussianPrior(param.x_param_ptr, 0, self.prior_cen_x_stddev))
294 elif has_prior_y and isinstance(param, g2f.CentroidYParameterD):
295 priors.append(g2f.GaussianPrior(param.y_param_ptr, 0, self.prior_cen_y_stddev))
297 return sources, priors
299 def schema_configurable(self) -> list[ColumnInfo]:
300 columns = []
301 if self.config_fit.eval_residual: 301 ↛ 303line 301 didn't jump to line 303 because the condition on line 301 was always true
302 columns.append(ColumnInfo(key="n_eval_jac", dtype="i4"))
303 if self.fit_linear_final: 303 ↛ 305line 303 didn't jump to line 305 because the condition on line 303 was always true
304 columns.append(ColumnInfo(key="delta_lnL_fit_linear", dtype="f8"))
305 if self.fit_psmodel_final: 305 ↛ 306line 305 didn't jump to line 306 because the condition on line 305 was never true
306 columns.append(ColumnInfo(key="delta_lnL_fit_ps", dtype="f8"))
307 return columns
309 def schema(
310 self,
311 bands: list[str] | None = None,
312 ) -> list[ColumnInfo]:
313 if bands is None or not (len(bands) > 0): 313 ↛ 314line 313 didn't jump to line 314 because the condition on line 313 was never true
314 raise ValueError("CatalogSourceFitter must provide at least one band")
315 schema = super().schema(bands)
317 parameters = CatalogSourceFitterConfigData(
318 config=self,
319 channels=tuple(g2f.Channel.get(band) for band in bands),
320 ).parameters
321 unit_size = u.Unit("pix")
322 units = {
323 g2f.IntegralParameterD: self.unit_flux,
324 g2f.ReffXParameterD: unit_size,
325 g2f.ReffYParameterD: unit_size,
326 g2f.SizeXParameterD: unit_size,
327 g2f.SizeYParameterD: unit_size,
328 }
329 idx_start = len(schema)
330 schema.extend(
331 [
332 ColumnInfo(key=key, dtype="f8", unit=units.get(type(param)))
333 for key, param in parameters.items()
334 ]
335 )
336 # Keep track of covariance key by declination parameter indexs
337 # If we want to add RA/dec covariance, it'll need to come after decErr
338 keys_cov = {}
339 compute_errors = self.compute_errors != "NONE"
340 if self.convert_cen_xy_to_radec:
341 label_cen = self.get_key_cen()
342 cen_underscored = label_cen.startswith("_")
343 suffix_x, suffix_y, suffix_ra, suffix_dec = (
344 f"{label_cen}{suffix}"
345 for suffix in (
346 self.get_suffix_x(),
347 self.get_suffix_y(),
348 self.get_suffix_ra(),
349 self.get_suffix_dec(),
350 )
351 )
352 suffix_ra = f"{label_cen}{self.get_suffix_ra()}"
353 suffix_dec = f"{label_cen}{self.get_suffix_dec()}"
354 for key, param in parameters.items():
355 # TODO: Update if allowing x, y <-> dec, RA mappings
356 # ... or arbitrary rotations
357 is_y = isinstance(param, g2f.CentroidYParameterD)
358 suffix_radec, suffix_xy = (
359 (suffix_ra, suffix_x)
360 if isinstance(param, g2f.CentroidXParameterD)
361 else ((suffix_dec, suffix_y) if is_y else (None, None))
362 )
363 if suffix_radec is not None:
364 # Add whatever the corresponding prefix is, and also
365 # remove any leading underscore if there's no prefix
366 prefix, suffix = (
367 ("", suffix_radec[1:])
368 if (cen_underscored and (key == suffix_xy[1:]))
369 else (key.split(suffix_xy)[0], suffix_radec)
370 )
371 schema.append(ColumnInfo(key=f"{prefix}{suffix}", dtype="f8", unit=u.deg))
372 if compute_errors and is_y:
373 suffix_radec = f"{label_cen}{self.get_suffix_ra_dec_cov()}"
374 prefix, suffix = (
375 ("", suffix_radec[1:])
376 if (cen_underscored and (key == suffix_xy[1:]))
377 else (key.split(suffix_xy)[0], suffix_radec)
378 )
379 keys_cov[len(schema) - 1] = f"{prefix}{suffix}"
380 if compute_errors: 380 ↛ 389line 380 didn't jump to line 389 because the condition on line 380 was always true
381 suffix = self.suffix_error
382 idx_end = len(schema)
383 for idx in range(idx_start, idx_end):
384 column = schema[idx]
385 schema.append(ColumnInfo(key=f"{column.key}{suffix}", dtype=column.dtype, unit=column.unit))
386 if (key_cov := keys_cov.get(idx)) is not None:
387 schema.append(ColumnInfo(key=key_cov, dtype="f8", unit=u.deg**2))
389 schema.extend(self.schema_configurable())
390 return schema
393class CatalogSourceFitterConfigData(pydantic.BaseModel):
394 """Configuration data for a fitter that can initialize lsst.gauss2d.fit
395 models and images thereof.
397 This class relies on cached properties being computed once, mostly shortly
398 after initialization. Therefore, it and the config field must be frozen to
399 ensure that the model remains unchanged.
400 """
402 model_config: ClassVar[pydantic.ConfigDict] = frozen_arbitrary_allowed_config
404 channels: list[g2f.Channel] = pydantic.Field(title="The list of channels")
405 config: CatalogSourceFitterConfig = pydantic.Field(title="A CatalogSourceFitterConfig to be frozen")
407 @pydantic.model_validator(mode="after")
408 def validate_config(self) -> Self:
409 self.config.validate()
410 return self
412 @cached_property
413 def components(self) -> tuple[g2f.Component]:
414 sources = self.sources_priors[0]
415 components = []
416 for source in sources:
417 components.extend(source.components)
418 return components
420 @cached_property
421 def parameters(self) -> dict[str, g2f.ParameterD]:
422 config = self.config
423 config_model = config.config_model
424 idx_comp_first = 0
425 has_prefix_source = config_model.has_prefix_source()
426 n_channels = len(self.channels)
427 parameters = {}
429 label_cen = config.get_key_cen()
430 label_rho = config.get_key_rho()
431 label_sersic = config.get_key_sersicindex()
432 label_x, label_y = config.get_suffix_x(), config.get_suffix_y()
434 for name_source, config_source in config_model.sources.items():
435 prefix_source = f"{name_source}_" if has_prefix_source else ""
436 has_prefix_group = config_source.has_prefix_group()
438 for name_group, config_group in config_source.component_groups.items():
439 prefix_group = f"{prefix_source}{name_group}_" if has_prefix_group else prefix_source
440 multicen = len(config_group.centroids) > 1
441 configs_comp = config_group.get_component_configs().items()
443 is_multicomp = len(configs_comp) > 1
445 for idx_comp_group, (name_comp, config_comp) in enumerate(configs_comp):
446 component = self.components[idx_comp_first + idx_comp_group]
448 key_comp = name_comp if is_multicomp else ""
449 prefix_comp = f"{prefix_group}{key_comp}"
450 key_size = config.get_prefixed_label(
451 config.get_key_size(config_comp.get_size_label()),
452 prefix_comp,
453 )
454 key_rho = config.get_prefixed_label(label_rho, prefix_comp)
456 if multicen or (idx_comp_group == 0): 456 ↛ 463line 456 didn't jump to line 463 because the condition on line 456 was always true
457 prefix_cen = prefix_comp if multicen else prefix_group
458 # Avoid double-underscoring if there's nothing to
459 # prefix or an existing prefix
460 key_cen = config.get_prefixed_label(label_cen, prefix_cen)
461 parameters[f"{key_cen}{label_x}"] = component.centroid.x_param
462 parameters[f"{key_cen}{label_y}"] = component.centroid.y_param
463 if not config_comp.size_x.fixed: 463 ↛ 465line 463 didn't jump to line 465 because the condition on line 463 was always true
464 parameters[f"{key_size}{label_x}"] = component.ellipse.size_x_param
465 if not config_comp.size_y.fixed: 465 ↛ 467line 465 didn't jump to line 467 because the condition on line 465 was always true
466 parameters[f"{key_size}{label_y}"] = component.ellipse.size_y_param
467 if not config_comp.rho.fixed: 467 ↛ 469line 467 didn't jump to line 469 because the condition on line 467 was always true
468 parameters[key_rho] = component.ellipse.rho_param
469 if not config_comp.flux.fixed: 469 ↛ 478line 469 didn't jump to line 478 because the condition on line 469 was always true
470 # TODO: return this to component.integralmodel
471 # when binding for g2f.FractionalIntegralModel is fixed
472 params_flux = get_params_uniq(component, fixed=False, nonlinear=False)
473 if len(params_flux) != n_channels: 473 ↛ 474line 473 didn't jump to line 474 because the condition on line 473 was never true
474 raise ValueError(f"{params_flux=} len={len(params_flux)} != {n_channels=}")
475 for channel, param_flux in zip(self.channels, params_flux):
476 key_flux = config.get_key_flux(label=prefix_comp, band=channel.name)
477 parameters[key_flux] = param_flux
478 if hasattr(config_comp, "sersic_index") and not config_comp.sersic_index.fixed:
479 parameters[config.get_prefixed_label(label_sersic, prefix_comp)] = (
480 component.sersicindex_param
481 )
483 return parameters
485 @cached_property
486 def sources_priors(self) -> tuple[tuple[g2f.Source], tuple[g2f.Prior]]:
487 sources, priors = self.config.make_sources(channels=self.channels)
488 return tuple(sources), tuple(priors)
491class CatalogSourceFitterABC(ABC, pydantic.BaseModel):
492 """Fit a Gaussian mixture source model to an image with a PSF model.
494 Notes
495 -----
496 Any exceptions raised and not in errors_expected will be logged in a
497 generic unknown_flag failure column.
498 """
500 model_config: ClassVar[pydantic.ConfigDict] = frozen_arbitrary_allowed_config
502 errors_expected: dict[type[Exception], str] = pydantic.Field(
503 default_factory=dict,
504 title="A dictionary of Exceptions with the name of the flag column key to fill if raised.",
505 )
506 modeller: Modeller = pydantic.Field(
507 default_factory=Modeller,
508 title="A Modeller instance to use for fitting.",
509 )
511 def _get_columns_params_radec(
512 self,
513 params_radec: dict[str, tuple[g2f.CentroidXParameterD, g2f.CentroidYParameterD]],
514 compute_errors: bool,
515 config: CatalogSourceFitterConfig,
516 ) -> tuple[list[tuple[str, str, str, str]], list[tuple[str, str, str, str, str, str]]]:
517 """Get a list of the columns needed for conversion of x/y centroid
518 parameters into ra/dec.
520 Parameters
521 ----------
522 params_radec
523 Dict of tuple of x, y parameter objects by name.
524 compute_errors
525 Whether errors will be computed.
526 config
527 The configuration with column formatting parameters.
529 Returns
530 -------
531 columns_params_radec
532 Column names for RA, dec, x, and y.
533 columns_params_radec_err
534 Column names for RA_err, dec_err, x, y, x_err, y_err.
535 """
536 columns_params_radec = []
537 columns_params_radec_err = []
538 suffix_err = config.suffix_error
539 key_cen = config.get_key_cen()
540 suffix_x, suffix_y = config.get_suffix_x(), config.get_suffix_y()
541 suffix_ra, suffix_dec = config.get_suffix_ra(), config.get_suffix_dec()
543 for key_base, (param_cen_x, param_cen_y) in params_radec.items():
544 # This removes redundant underscores
545 key_base_cen = config.get_prefixed_label(key_cen, key_base)
547 if param_cen_y is None: 547 ↛ 548line 547 didn't jump to line 548 because the condition on line 547 was never true
548 raise RuntimeError(
549 f"Fitter failed to find corresponding cen_y param for {key_base=}; is it fixed?"
550 )
551 column_ra = f"{key_base_cen}{suffix_ra}"
552 column_dec = f"{key_base_cen}{suffix_dec}"
554 columns_params_radec.append(
555 (
556 column_ra,
557 column_dec,
558 f"{key_base_cen}{suffix_x}",
559 f"{key_base_cen}{suffix_y}",
560 )
561 )
562 if compute_errors: 562 ↛ 543line 562 didn't jump to line 543 because the condition on line 562 was always true
563 key_cov = (
564 None
565 if not config.compute_radec_covariance
566 else (f"{key_base_cen}{config.get_suffix_ra_dec_cov()}")
567 )
568 columns_params_radec_err.append(
569 (
570 f"{key_base_cen}{suffix_ra}{suffix_err}",
571 f"{key_base_cen}{suffix_dec}{suffix_err}",
572 f"{key_base_cen}{suffix_x}",
573 f"{key_base_cen}{suffix_y}",
574 f"{key_base_cen}{suffix_x}{suffix_err}",
575 f"{key_base_cen}{suffix_y}{suffix_err}",
576 key_cov,
577 column_ra,
578 column_dec,
579 )
580 )
581 return columns_params_radec, columns_params_radec_err
583 @staticmethod
584 def _get_logger() -> logging.Logger:
585 logger = logging.getLogger(__name__)
587 return logger
589 def _validate_errors_expected(self, config: CatalogSourceFitterConfig) -> None:
590 """Check that self.errors_expected is set correctly.
592 Parameters
593 ----------
594 config
595 The fitting configuration.
597 Raises
598 ------
599 ValueError
600 Raised if the configuration is invalid.
601 """
602 if len(self.errors_expected) != len(config.flag_errors): 602 ↛ 603line 602 didn't jump to line 603 because the condition on line 602 was never true
603 raise ValueError(f"{self.errors_expected=} keys not same len as {config.flag_errors=}")
604 errors_bad = {}
605 errors_recast = {}
606 for error_name, error_type in self.errors_expected.items():
607 if error_type in errors_recast: 607 ↛ 608line 607 didn't jump to line 608 because the condition on line 607 was never true
608 errors_bad[error_name] = error_type
609 else:
610 errors_recast[error_type] = error_name
611 if errors_bad: 611 ↛ 612line 611 didn't jump to line 612 because the condition on line 611 was never true
612 raise ValueError(f"{self.errors_expected=} keys contain duplicates from {config.flag_errors=}")
614 def compute_model_radec_err(
615 self,
616 source_multi: Mapping[str, Any],
617 results,
618 columns_params_radec_err,
619 idx: int,
620 set_radec: bool = False,
621 ) -> None:
622 """Compute right ascension and declination errors for a source.
624 This default implementation is naive, assuming only that
625 get_model_radec is implemented, and should be overridden.
627 Parameters
628 ----------
629 source_multi
630 A mapping with fields expected to be populated in the
631 corresponding multiband source catalog.
632 results
633 The output catalog to read/write from/to.
634 columns_params_radec_err
635 A list of tuples containing six keys for:
636 ra, dec: RA/Dec inputs.
637 ra_err, dec_err: RA/Dec error outputs.
638 cen_x, cen_y: Pixel x/y centroid inputs.
639 cen_x_err, cen_y_err: Pixel x/y centroid error inputs.
640 idx
641 The integer index of this source in the results catalog.
642 set_radec
643 Whether this method should set RA, dec values instead of reading
644 them (should be True if defer_radec_conversion is True).
645 """
646 for ( 646 ↛ exitline 646 didn't return from function 'compute_model_radec_err' because the loop on line 646 didn't complete
647 key_ra_err,
648 key_dec_err,
649 key_cen_x,
650 key_cen_y,
651 key_cen_x_err,
652 key_cen_y_err,
653 key_cen_ra_dec_cov,
654 key_ra,
655 key_dec,
656 ) in columns_params_radec_err:
657 cen_x, cen_y = results[key_cen_x][idx], results[key_cen_y][idx]
658 # TODO: improve this in DM-45682
659 # For one, it won't work right at limits:
660 # RA=359.99... or dec=+89.99...
661 # Could also consider dividing by sqrt(2)
662 # ...but that factor would multiply out later
663 ra_err, dec_err = self.get_model_radec(
664 source_multi,
665 cen_x + results[key_cen_x_err][idx],
666 cen_y + results[key_cen_y_err][idx],
667 )
668 ra, dec = results[key_ra][idx], results[key_dec][idx]
669 results[key_ra_err][idx], results[key_dec_err][idx] = abs(ra_err - ra), abs(dec_err - dec)
671 def copy_centroid_errors(
672 self,
673 columns_cenx_err_copy: tuple[str],
674 columns_ceny_err_copy: tuple[str],
675 results: Table,
676 catalog_multi: Sequence,
677 catexps: list[CatalogExposureSourcesABC],
678 config_data: CatalogSourceFitterConfigData,
679 ) -> None:
680 """Copy centroid errors from an input catalog.
682 This method exists to support fitting models with fixed centroids
683 derived from an input catalog. Implementers can simply copy an
684 existing column into the results catalog or use the data as needed;
685 however, there is no reasonable default implementation.
687 Parameters
688 ----------
689 columns_cenx_err_copy
690 X-axis result centroid columns to copy errors for.
691 columns_ceny_err_copy
692 Y-axis result centroid columns to copy errors for.
693 results
694 The table of fit results to copy errors into.
695 catalog_multi
696 The input multiband catalog.
697 catexps
698 The input data.
699 config_data
700 The fitter config and data.
702 Raises
703 ------
704 NotImplementedError
705 Raised if columns need to be copied but no implementation is
706 available.
707 """
708 if columns_cenx_err_copy or columns_ceny_err_copy:
709 raise NotImplementedError(
710 f"Fitter of {type(self)=} got {columns_cenx_err_copy=} and/or {columns_ceny_err_copy=}"
711 f" but has not overriden copy_centroid_errors"
712 )
714 def fit(
715 self,
716 catalog_multi: Sequence,
717 catexps: list[CatalogExposureSourcesABC],
718 config_data: CatalogSourceFitterConfigData | None = None,
719 logger: logging.Logger | None = None,
720 **kwargs: Any,
721 ) -> astropy.table.Table:
722 """Fit PSF-convolved source models with MultiProFit.
724 Each source has a single PSF-convolved model fit, given PSF model
725 parameters from a catalog, and a combination of initial source
726 model parameters and a deconvolved source image from the
727 CatalogExposureSources.
729 Parameters
730 ----------
731 catalog_multi
732 A multi-band source catalog to fit a model to.
733 catexps
734 A list of (source and psf) catalog-exposure pairs.
735 config_data
736 Configuration settings and data for fitting and output.
737 logger
738 The logger. Defaults to calling `_getlogger`.
739 **kwargs
740 Additional keyword arguments to pass to self.modeller.
742 Returns
743 -------
744 catalog : `astropy.Table`
745 A table with fit parameters for the PSF model at the location
746 of each source.
747 """
748 if config_data is None: 748 ↛ 749line 748 didn't jump to line 749 because the condition on line 748 was never true
749 config_data = CatalogSourceFitterConfigData(
750 config=CatalogSourceFitterConfig(),
751 channels=[catexp.channel for catexp in catexps],
752 )
753 if logger is None: 753 ↛ 756line 753 didn't jump to line 756 because the condition on line 753 was always true
754 logger = self._get_logger()
756 config = config_data.config
757 self._validate_errors_expected(config)
758 self.validate_fit_inputs(
759 catalog_multi=catalog_multi, catexps=catexps, config_data=config_data, logger=logger, **kwargs
760 )
762 model_sources, priors = config_data.sources_priors
764 # TODO: If free Observation params are ever supported, make null Data
765 # Because config_data knows nothing about the Observation(s)
766 params = config_data.parameters
767 values_init = {param: param.value for param in params.values() if param.free}
768 prefix = config.prefix_column
769 columns_param_fixed: dict[str, tuple[g2f.ParameterD, float]] = {}
770 columns_param_free: dict[str, tuple[g2f.ParameterD, float]] = {}
771 columns_param_flux: dict[str, g2f.IntegralParameterD] = {}
772 params_cen_x: dict[str, g2f.CentroidXParameterD] = {}
773 params_cen_y: dict[str, g2f.CentroidYParameterD] = {}
774 columns_err = []
776 errors_hessian: bool = config.compute_errors == "INV_HESSIAN"
777 errors_hessian_bestfit: bool = config.compute_errors == "INV_HESSIAN_BESTFIT"
778 compute_errors: bool = errors_hessian or errors_hessian_bestfit
780 columns_cenx_err_copy = []
781 columns_ceny_err_copy = []
783 suffix_err = config.suffix_error
784 key_cen = config.get_key_cen()
785 cen_underscored = key_cen.startswith("_")
786 suffix_cenx = f"{key_cen}{config.get_suffix_x()}"
787 suffix_ceny = f"{key_cen}{config.get_suffix_y()}"
789 # Add each param to appropriate and more specific pre-computed lists
790 for key, param in params.items():
791 key_full = f"{prefix}{key}"
792 is_cenx = isinstance(param, g2f.CentroidXParameterD)
793 is_ceny = isinstance(param, g2f.CentroidYParameterD)
795 # Add the corresponding error key to the appropriate list
796 if compute_errors: 796 ↛ 805line 796 didn't jump to line 805 because the condition on line 796 was always true
797 if param.free: 797 ↛ 799line 797 didn't jump to line 799 because the condition on line 797 was always true
798 columns_err.append(f"{key_full}{suffix_err}")
799 elif is_cenx:
800 columns_cenx_err_copy.append(f"{key_full}{suffix_err}")
801 elif is_ceny:
802 columns_ceny_err_copy.append(f"{key_full}{suffix_err}")
804 # Add this param to the appropriate dict
805 (columns_param_fixed if param.fixed else columns_param_free)[key_full] = (
806 param,
807 config_data.config.centroid_pixel_offset if (is_cenx or is_ceny) else 0,
808 )
809 if isinstance(param, g2f.IntegralParameterD):
810 columns_param_flux[key_full] = param
811 elif config.convert_cen_xy_to_radec:
812 # Infer the prefix if possible, after checking for a dropped
813 # leading underscore in case there's no prefix
814 if is_cenx:
815 prefix_cen, suffix_cen = (
816 ("", key_full)
817 if (cen_underscored and (key_full == suffix_cenx[1:]))
818 else key_full.split(suffix_cenx)
819 )
820 params_cen_x[prefix_cen] = param
821 elif is_ceny:
822 prefix_cen, suffix_cen = (
823 ("", key_full)
824 if (cen_underscored and (key_full == suffix_ceny[1:]))
825 else key_full.split(suffix_ceny)
826 )
827 params_cen_y[prefix_cen] = param
829 if config.convert_cen_xy_to_radec or config.fit_psmodel_final:
830 assert params_cen_x.keys() == params_cen_y.keys()
831 columns_params_radec, columns_params_radec_err = self._get_columns_params_radec(
832 {k: (x, params_cen_y[k]) for k, x in params_cen_x.items()},
833 compute_errors,
834 config=config,
835 )
837 fit_psmodel_final = False
838 if config.fit_psmodel_final: 838 ↛ 841line 838 didn't jump to line 841 because the condition on line 838 was never true
839 # This should never be True until DM-46497 is merged, but models
840 # in other/future derived classes might have multiple centroids
841 if (len(set(params_cen_x.values())) > 1) or (len(set(params_cen_y.values())) > 1):
842 raise ValueError(
843 f"Got {params_cen_x=} and {params_cen_y} with > 1 unique elements, so "
844 f"config.fit_psmodel_final may not be set to True"
845 )
846 fit_psmodel_final = True
848 key_cen_x_psmodel, key_cen_y_psmodel = columns_params_radec[0][2:4]
850 channels = config_data.channels
851 sources_psmodel, priors_psmodel = config.make_point_sources(channels, model_sources)
852 params_psmodel = sources_psmodel[0].parameters()
853 cenx_psmodel, ceny_psmodel = None, None
854 fluxes_psmodel = {}
855 idx_band = 0
856 for param in params_psmodel:
857 if isinstance(param, g2f.CentroidXParameterD):
858 if cenx_psmodel is not None:
859 raise RuntimeError("Point source model found multiple x centroids")
860 cenx_psmodel = param
861 elif isinstance(param, g2f.CentroidYParameterD):
862 if ceny_psmodel is not None:
863 raise RuntimeError("Point source model found multiple y centroids")
864 ceny_psmodel = param
865 elif isinstance(param, g2f.IntegralParameterD):
866 fluxes_psmodel[channels[idx_band]] = param
867 idx_band += 1
869 convert_cen_xy_to_radec_first = config.convert_cen_xy_to_radec and not (
870 config.compute_errors and config.defer_radec_conversion
871 )
873 # Setup the results table with correct column names
874 n_rows = len(catalog_multi)
875 channels = self.get_channels(catexps)
876 results, columns = config.make_catalog(n_rows, bands=list(channels.keys()))
878 # Copy centroid error columns into results ( if needed)
879 self.copy_centroid_errors(
880 columns_cenx_err_copy=columns_cenx_err_copy,
881 columns_ceny_err_copy=columns_ceny_err_copy,
882 results=results,
883 catalog_multi=catalog_multi,
884 catexps=catexps,
885 config_data=config_data,
886 )
888 # dummy size for first iteration
889 size, size_new = 0, 0
890 fitInputs = FitInputsDummy()
891 plot = False
893 # Configure default options for calls to compute_variances
894 # keys are for values of return_negative
895 kwargs_err_default = {
896 True: {
897 "options": g2f.HessianOptions(findiff_add=1e-3, findiff_frac=1e-3),
898 "use_diag_only": config.compute_errors_no_covar,
899 },
900 False: {"options": g2f.HessianOptions(findiff_add=1e-6, findiff_frac=1e-6)},
901 }
903 range_idx = range(n_rows)
905 # TODO: Do this check with dummy data
906 # It might not work with real data if the first row is bad
907 # data, psf_models = config.make_model_data(
908 # idx_row=range_idx[0], catexps=catexps)
909 # model = g2f.ModelD(data=data, psfmodels=psf_models,
910 # sources=model_sources, priors=priors)
911 # Remember to filter out fixed centroids from params
912 # assert list(params.values()) == get_params_uniq(model, fixed=False)
914 time_init_all = time.process_time()
915 logger_periodic = PeriodicLogger(logger)
916 n_skipfail = 0
918 for idx in range_idx:
919 time_init = time.process_time()
920 row = results[idx]
921 source_multi = catalog_multi[idx]
922 id_source = source_multi[config.column_id]
923 row[config.column_id] = id_source
924 time_final = time_init
926 try:
927 self.validate_source(idx_row=idx, catalog_multi=catalog_multi)
928 data, psf_models = config.make_model_data(idx_row=idx, catexps=catexps)
929 if data.size == 0: 929 ↛ 930line 929 didn't jump to line 930 because the condition on line 929 was never true
930 raise NoDataError("make_model_data returned empty data")
931 model = g2f.ModelD(data=data, psfmodels=psf_models, sources=model_sources, priors=priors)
932 self.initialize_model(
933 model,
934 source_multi,
935 catexps,
936 config_data=config_data,
937 values_init=values_init,
938 )
940 # Caches the jacobian residual if the data size is unchanged
941 # Note: this will need to change with priors
942 # (data should report its own size)
943 size_new = np.sum([datum.image.size for datum in data])
944 if size_new != size:
945 fitInputs = None
946 size = size_new
947 # Some algorithms might not even use fitInputs
948 elif fitInputs is not None: 948 ↛ 953line 948 didn't jump to line 953 because the condition on line 948 was always true
949 fitInputs = fitInputs if not fitInputs.validate_for_model(model) else None
951 # TODO: Check if flux param limits and transforms are set
952 # appropriately if config.fit_linear_init is False
953 if config.fit_linear_init: 953 ↛ 956line 953 didn't jump to line 956 because the condition on line 953 was always true
954 self.modeller.fit_model_linear(model=model, ratio_min=0.01)
956 for observation in data:
957 observation.image.data[~np.isfinite(observation.image.data)] = 0
959 result_full = self.modeller.fit_model(
960 model, fitinputs=fitInputs, config=config.config_fit, **kwargs
961 )
962 fitInputs = result_full.inputs
963 results[f"{prefix}n_iter"][idx] = result_full.n_eval_func
964 results[f"{prefix}time_eval"][idx] = result_full.time_eval
965 results[f"{prefix}time_fit"][idx] = result_full.time_run
966 if config.config_fit.eval_residual: 966 ↛ 969line 966 didn't jump to line 969 because the condition on line 966 was always true
967 results[f"{prefix}n_eval_jac"][idx] = result_full.n_eval_jac
969 params_free_missing = result_full.params_free_missing or tuple()
971 # Set all params to best fit values
972 # In case the optimizer doesn't
973 for (key, (param, offset)), value in zip(
974 columns_param_free.items(),
975 result_full.params_best,
976 ):
977 param.value_transformed = value
978 if param not in params_free_missing: 978 ↛ 973line 978 didn't jump to line 973 because the condition on line 978 was always true
979 results[key][idx] = param.value + offset
981 # Also add any offset to the fixed parameters
982 # (usually centroids, if any)
983 for key, (param, offset) in columns_param_fixed.items(): 983 ↛ 984line 983 didn't jump to line 984 because the loop on line 983 never started
984 results[key][idx] = param.value + offset
986 # Do a final linear fit
987 # If the nonlinear fit is good, the values won't change much
988 if config.fit_linear_final: 988 ↛ 1007line 988 didn't jump to line 1007 because the condition on line 988 was always true
989 loglike_init, loglike_new = self.modeller.fit_model_linear(
990 model=model, ratio_min=0.01, validate=True
991 )
992 loglike_final = max(loglike_init, loglike_new)
993 results[f"{prefix}delta_lnL_fit_linear"][idx] = np.sum(loglike_new) - np.sum(loglike_init)
995 if params_free_missing: 995 ↛ 996line 995 didn't jump to line 996 because the condition on line 995 was never true
996 columns_param_flux_fit = {
997 column: param
998 for column, param in columns_param_flux.items()
999 if param not in params_free_missing
1000 }
1001 else:
1002 columns_param_flux_fit = columns_param_flux
1004 for column, param in columns_param_flux_fit.items():
1005 results[column][idx] = param.value
1006 else:
1007 loglike_final = model.evaluate()
1009 if convert_cen_xy_to_radec_first:
1010 for key_ra, key_dec, key_cen_x, key_cen_y in columns_params_radec: 1010 ↛ 1016line 1010 didn't jump to line 1016 because the loop on line 1010 didn't complete
1011 # These will have been converted back if necessary
1012 cen_x, cen_y = results[key_cen_x][idx], results[key_cen_y][idx]
1013 radec = self.get_model_radec(source_multi, cen_x, cen_y)
1014 results[key_ra][idx], results[key_dec][idx] = radec
1016 if fit_psmodel_final: 1016 ↛ 1017line 1016 didn't jump to line 1017 because the condition on line 1016 was never true
1017 cen_x, cen_y = results[key_cen_x_psmodel][idx], results[key_cen_y_psmodel][idx]
1018 cenx_psmodel.value = cen_x
1019 ceny_psmodel.value = cen_y
1020 model_psf = g2f.ModelD(data=data, psfmodels=psf_models, sources=sources_psmodel)
1021 _ = self.modeller.fit_model_linear(model_psf)
1022 model_psf.setup_evaluators(evaluatormode=g2f.EvaluatorMode.loglike)
1023 loglike_psfmodel = model_psf.evaluate()
1024 # Reset fluxes for the next fit
1025 for param in fluxes_psmodel.values():
1026 param.value = 1.0
1027 results[f"{prefix}delta_lnL_fit_ps"][idx] = loglike_final[0] - loglike_psfmodel[0]
1029 if compute_errors: 1029 ↛ 1145line 1029 didn't jump to line 1145 because the condition on line 1029 was always true
1030 errors = []
1031 model_eval = model
1032 errors_iter = None
1033 for param in params_free_missing: 1033 ↛ 1034line 1033 didn't jump to line 1034 because the loop on line 1033 never started
1034 param.fixed = True
1036 if config.compute_errors_from_jacobian: 1036 ↛ 1050line 1036 didn't jump to line 1050 because the condition on line 1036 was always true
1037 try:
1038 errors_iter = np.sqrt(
1039 self.modeller.compute_variances(
1040 model_eval,
1041 transformed=False,
1042 use_diag_only=config.compute_errors_no_covar,
1043 )
1044 )
1045 errors.append((errors_iter, np.sum(~(errors_iter > 0))))
1046 except Exception:
1047 pass
1048 # If computing errors from the Jacobian didn't work, or if
1049 # it was disabled in the config, try the Hessian
1050 if errors_iter is None: 1050 ↛ 1051line 1050 didn't jump to line 1051 because the condition on line 1050 was never true
1051 img_data_old = []
1052 if errors_hessian_bestfit:
1053 # Model sans prior
1054 model_eval = g2f.ModelD(
1055 data=model.data, psfmodels=model.psfmodels, sources=model.sources
1056 )
1057 model_eval.setup_evaluators(evaluatormode=g2f.EvaluatorMode.image)
1058 model_eval.evaluate()
1059 # Compute the errors by setting the data to the
1060 # best-fit model (a quasi-parametric bootstrap
1061 # with one iteration)
1062 for obs, output in zip(model_eval.data, model_eval.outputs):
1063 img_data_old.append(obs.image.data.copy())
1064 img = obs.image.data
1065 img.flat = output.data.flat
1066 # To make this a real bootstrap, could do this
1067 # (but would need to iterate):
1068 # + rng.standard_normal(img.size)*(
1069 # obs.sigma_inv.data.flat)
1071 # Try without forcing all of the Hessian terms to be
1072 # negative first. At the optimum they should be, but
1073 # in practice the best-fit values are always at least
1074 # a little off and so the sign is equally likely to be
1075 # positive as negative.
1076 for return_negative in (False, True):
1077 kwargs_err = kwargs_err_default[return_negative]
1078 if errors and errors[-1][1] == 0:
1079 break
1080 try:
1081 errors_iter = np.sqrt(
1082 self.modeller.compute_variances(
1083 model_eval, transformed=False, **kwargs_err
1084 )
1085 )
1086 errors.append((errors_iter, np.sum(~(errors_iter > 0))))
1087 except Exception:
1088 try:
1089 errors_iter = np.sqrt(
1090 self.modeller.compute_variances(
1091 model_eval,
1092 transformed=False,
1093 use_svd=True,
1094 **kwargs_err,
1095 )
1096 )
1097 errors.append((errors_iter, np.sum(~(errors_iter > 0))))
1098 except Exception:
1099 pass
1100 # Return the data to its original noisy values
1101 # (it was replaced by the model earlier)
1102 if errors_hessian_bestfit:
1103 for obs, img_datum_old in zip(model.data, img_data_old):
1104 obs.image.data.flat = img_datum_old.flat
1105 # Save and optionally plot the errors
1106 if errors: 1106 ↛ 1145line 1106 didn't jump to line 1145 because the condition on line 1106 was always true
1107 idx_min = np.argmax([err[1] for err in errors])
1108 errors = errors[idx_min][0]
1109 if plot: 1109 ↛ 1110line 1109 didn't jump to line 1110 because the condition on line 1109 was never true
1110 errors_plot = np.clip(errors, 0, 1000)
1111 errors_plot[~np.isfinite(errors_plot)] = 0
1112 from ..plotting import ErrorValues, plot_loglike
1114 try:
1115 plot_loglike(model, errors={"err": ErrorValues(values=errors_plot)})
1116 except Exception:
1117 for param in params:
1118 param.fixed = False
1120 if params_free_missing: 1120 ↛ 1121line 1120 didn't jump to line 1121 because the condition on line 1120 was never true
1121 columns_err_fitted = [
1122 column
1123 for column, param in zip(columns_err, params.values())
1124 if param not in params_free_missing
1125 ]
1126 else:
1127 columns_err_fitted = columns_err
1129 for value, column_err in zip(errors, columns_err_fitted):
1130 results[column_err][idx] = value
1132 for param in params_free_missing: 1132 ↛ 1133line 1132 didn't jump to line 1133 because the loop on line 1132 never started
1133 param.fixed = False
1135 # Convert the x/y errors to ra/dec errors
1136 if config.convert_cen_xy_to_radec:
1137 self.compute_model_radec_err(
1138 source_multi,
1139 results,
1140 columns_params_radec_err,
1141 idx,
1142 set_radec=not convert_cen_xy_to_radec_first,
1143 )
1145 results[f"{prefix}chisq_reduced"][idx] = result_full.chisq_best / size
1146 time_final = time.process_time()
1147 results[f"{prefix}time_full"][idx] = time_final - time_init
1148 except Exception as e:
1149 n_skipfail += 1
1150 size = 0 if fitInputs is None else size_new
1151 column = self.errors_expected.get(e.__class__, "")
1152 if column: 1152 ↛ 1162line 1152 didn't jump to line 1162 because the condition on line 1152 was always true
1153 row[f"{prefix}{column}"] = True
1154 logger.debug(
1155 "id_source=%i (idx=%i/%i) fit failed with known exception: %s",
1156 id_source,
1157 idx,
1158 n_rows,
1159 e,
1160 )
1161 else:
1162 row[f"{prefix}unknown_flag"] = True
1163 logger.info(
1164 "id_source=%i (idx=%i/%i) fit failed with unexpected exception: %s",
1165 id_source,
1166 idx,
1167 n_rows,
1168 e,
1169 exc_info=1,
1170 )
1171 logger_periodic.log(
1172 "Fit idx=%i/%i sources (%i skipped/failed) in %.2f",
1173 idx,
1174 n_rows,
1175 n_skipfail,
1176 time_final - time_init_all,
1177 )
1179 n_unknown = np.sum(row[f"{prefix}unknown_flag"])
1180 if n_unknown > 0: 1180 ↛ 1181line 1180 didn't jump to line 1181 because the condition on line 1180 was never true
1181 logger.warning("%i/%i source fits failed with unexpected exceptions", n_unknown, n_rows)
1183 return results
1185 def get_channels(
1186 self,
1187 catexps: list[CatalogExposureSourcesABC],
1188 ) -> dict[str, g2f.Channel]:
1189 channels = {}
1190 for catexp in catexps:
1191 try:
1192 channel = catexp.channel
1193 except AttributeError:
1194 band = catexp.band
1195 if callable(band):
1196 band = band()
1197 channel = g2f.Channel.get(band)
1198 if channel not in channels: 1198 ↛ 1190line 1198 didn't jump to line 1190 because the condition on line 1198 was always true
1199 channels[channel.name] = channel
1200 return channels
1202 def get_model(
1203 self,
1204 idx_row: int,
1205 catalog_multi: Sequence,
1206 catexps: list[CatalogExposureSourcesABC],
1207 config_data: CatalogSourceFitterConfigData | None = None,
1208 results: astropy.table.Table | None = None,
1209 **kwargs: Any,
1210 ) -> g2f.ModelD:
1211 """Reconstruct the model for a single row of a fit catalog.
1213 Parameters
1214 ----------
1215 idx_row
1216 The index of the row in the catalog.
1217 catalog_multi
1218 The multi-band catalog originally used for initialization.
1219 catexps
1220 The catalog-exposure pairs to reconstruct the model for.
1221 config_data
1222 The configuration used to generate sources.
1223 Default-initialized if None.
1224 results
1225 The corresponding best-fit parameter catalog to initialize
1226 parameter values from. If None, the model params will be set by
1227 `self.initialize_model`, as they would be when calling `self.fit`.
1228 **kwargs
1229 Additional keyword arguments to pass to initialize_model. Not
1230 used during fitting.
1232 Returns
1233 -------
1234 model
1235 The reconstructed model.
1236 """
1237 channels = self.get_channels(catexps)
1238 if config_data is None: 1238 ↛ 1239line 1238 didn't jump to line 1239 because the condition on line 1238 was never true
1239 config_data = CatalogSourceFitterConfigData(
1240 config=CatalogSourceFitterConfig(),
1241 channels=list(channels.values()),
1242 )
1243 config = config_data.config
1245 if not idx_row >= 0: 1245 ↛ 1246line 1245 didn't jump to line 1246 because the condition on line 1245 was never true
1246 raise ValueError(f"{idx_row=} !>=0")
1247 if not len(catalog_multi) > idx_row: 1247 ↛ 1248line 1247 didn't jump to line 1248 because the condition on line 1247 was never true
1248 raise ValueError(f"{len(catalog_multi)=} !> {idx_row=}")
1249 if (results is not None) and not (len(results) > idx_row): 1249 ↛ 1250line 1249 didn't jump to line 1250 because the condition on line 1249 was never true
1250 raise ValueError(f"{len(results)=} !> {idx_row=}")
1252 model_sources, priors = config_data.sources_priors
1253 source_multi = catalog_multi[idx_row]
1255 data, psf_models = config.make_model_data(
1256 idx_row=idx_row,
1257 catexps=catexps,
1258 )
1259 model = g2f.ModelD(data=data, psfmodels=psf_models, sources=model_sources, priors=priors)
1260 self.initialize_model(model, source_multi, catexps, **kwargs)
1262 if results is not None: 1262 ↛ 1267line 1262 didn't jump to line 1267 because the condition on line 1262 was always true
1263 row = results[idx_row]
1264 for column, param in config_data.parameters.items():
1265 param.value = row[f"{config.prefix_column}{column}"]
1267 return model
1269 def get_model_radec(self, source: Mapping[str, Any], cen_x: float, cen_y: float) -> tuple[float, float]:
1270 """Return right ascension and declination values for a source.
1272 Implementing this method is necessary only when fitting data with
1273 accompanying WCS.
1275 Parameters
1276 ----------
1277 source
1278 A mapping with fields expected to be populated in the
1279 corresponding source catalog.
1280 cen_x
1281 The x-axis centroid in pixel coordinates.
1282 cen_y
1283 The y-axis centroid in pixel coordinates.
1285 Returns
1286 -------
1287 ra, dec
1288 The right ascension and declination.
1289 """
1290 raise RaDecConversionNotImplementedError("get_model_radec has no default implementation")
1292 @abstractmethod
1293 def initialize_model(
1294 self,
1295 model: g2f.ModelD,
1296 source: Mapping[str, Any],
1297 catexps: list[CatalogExposureSourcesABC],
1298 config_data: CatalogSourceFitterConfigData,
1299 values_init: Mapping[g2f.ParameterD, float] | None = None,
1300 **kwargs: Any,
1301 ) -> None:
1302 """Initialize a Model for a single source row.
1304 Parameters
1305 ----------
1306 model
1307 The model object to initialize.
1308 source
1309 A mapping with fields expected to be populated in the
1310 corresponding source catalog for initialization.
1311 catexps
1312 A list of (source and psf) catalog-exposure pairs.
1313 config_data
1314 Configuration settings and data for fitting and output.
1315 values_init
1316 Initial parameter values from the model configuration.
1317 **kwargs
1318 Additional keyword arguments that cannot be required for fitting.
1319 """
1321 @abstractmethod
1322 def validate_fit_inputs(
1323 self,
1324 catalog_multi: Sequence,
1325 catexps: list[CatalogExposureSourcesABC],
1326 config_data: CatalogSourceFitterConfigData = None,
1327 logger: logging.Logger = None,
1328 **kwargs: Any,
1329 ) -> None:
1330 """Validate inputs to self.fit.
1332 This method is called before any fitting is done. It may be used for
1333 any purpose, including checking that the inputs are a particular
1334 subclass of the base classes.
1336 Parameters
1337 ----------
1338 catalog_multi
1339 A multi-band source catalog to fit a model to.
1340 catexps
1341 A list of (source and psf) catalog-exposure pairs.
1342 config_data
1343 Configuration settings and data for fitting and output.
1344 logger
1345 The logger. Defaults to calling `_getlogger`.
1346 **kwargs
1347 Additional keyword arguments to pass to self.modeller.
1348 """
1349 pass
1351 def validate_source(
1352 self,
1353 idx_row: int,
1354 catalog_multi: Sequence,
1355 ) -> None:
1356 """Validate that the source is suitable to fit.
1358 Subclasses may override this method to raise a relevant exception
1359 if the source should be skipped.
1361 Parameters
1362 ----------
1363 idx_row
1364 The index of the row in the multiband catalog.
1365 catalog_multi
1366 The multiband input catalog.
1367 """