Coverage for tests/test_plotting.py: 100%
55 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 02:06 -0700
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 02:06 -0700
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 numpy as np
23import pytest
25import lsst.gauss2d.fit as g2f
26from lsst.multiprofit.componentconfig import CentroidConfig, GaussianComponentConfig, ParameterConfig
27from lsst.multiprofit.model_utils import make_psf_model_null
28from lsst.multiprofit.modelconfig import ModelConfig
29from lsst.multiprofit.observationconfig import CoordinateSystemConfig, ObservationConfig
30from lsst.multiprofit.plotting import abs_mag_sol_lsst, bands_weights_lsst, plot_model_rgb
31from lsst.multiprofit.sourceconfig import ComponentGroupConfig, SourceConfig
33sigma_inv = 1e4
36@pytest.fixture(scope="module")
37def channels() -> dict[str, g2f.Channel]:
38 """Return dict of generic RGB channels."""
39 return {band: g2f.Channel.get(band) for band in bands_weights_lsst}
42@pytest.fixture(scope="module")
43def data(channels) -> g2f.DataD:
44 """Return initialized data in all bands."""
45 n_rows, n_cols = 16, 21
46 x_min, y_min = 0, 0
48 dn_rows, dn_cols = 1, -2
49 dx_min, dy_min = -2, 1
51 observations = []
52 for idx, band in enumerate(channels):
53 config = ObservationConfig(
54 band=band,
55 coordsys=CoordinateSystemConfig(
56 x_min=x_min + idx * dx_min,
57 y_min=y_min + idx * dy_min,
58 ),
59 n_rows=n_rows + idx * dn_rows,
60 n_cols=n_cols + idx * dn_cols,
61 )
62 observation = config.make_observation()
63 observation.sigma_inv.fill(sigma_inv)
64 observation.mask_inv.fill(1)
65 observations.append(observation)
66 return g2f.DataD(observations)
69@pytest.fixture(scope="module")
70def psf_model():
71 """Return a trivial PSF model."""
72 return make_psf_model_null()
75@pytest.fixture(scope="module")
76def psf_models(psf_model, channels) -> list[g2f.PsfModel]:
77 """Return the trivial PSF model for each band."""
78 return [psf_model] * len(channels)
81@pytest.fixture(scope="module")
82def model(channels, data, psf_models):
83 """Return a single-Gaussian model with a trivial PSF."""
84 fluxes_group = [{channels[band]: 10 ** (-0.4 * (mag - 8.9)) for band, mag in abs_mag_sol_lsst.items()}]
86 modelconfig = ModelConfig(
87 sources={
88 "src": SourceConfig(
89 component_groups={
90 "": ComponentGroupConfig(
91 centroids={
92 "default": CentroidConfig(
93 x=ParameterConfig(value_initial=6.0, fixed=True),
94 y=ParameterConfig(value_initial=11.0, fixed=True),
95 )
96 },
97 components_gauss={
98 "": GaussianComponentConfig(
99 rho=ParameterConfig(value_initial=0.1),
100 size_x=ParameterConfig(value_initial=3.8),
101 size_y=ParameterConfig(value_initial=5.1),
102 )
103 },
104 )
105 }
106 ),
107 },
108 )
109 model = modelconfig.make_model([[fluxes_group]], data=data, psf_models=psf_models)
110 model.setup_evaluators(g2f.EvaluatorMode.image)
111 model.evaluate()
112 rng = np.random.default_rng(1)
113 for output, obs in zip(model.outputs, model.data):
114 img = obs.image.data
115 img.flat = output.data.flat + rng.standard_normal(img.size) / sigma_inv
116 return model
119def test_plot_model_rgb(model):
120 """Test that RGB model plotting works."""
121 fig, ax, fig_gs, ax_gs, *_ = plot_model_rgb(
122 model,
123 minimum=0,
124 stretch=0.15,
125 Q=4,
126 weights=bands_weights_lsst,
127 plot_chi_hist=True,
128 )
129 assert fig is not None
130 assert ax is not None
131 assert fig_gs is not None
132 assert ax_gs is not None
135def test_plot_model_rgb_auto(model):
136 """Test that RGB model plotting with automatic stretching works."""
137 fig, ax, *_ = plot_model_rgb(
138 model,
139 Q=6,
140 weights=bands_weights_lsst,
141 rgb_min_auto=True,
142 rgb_stretch_auto=True,
143 plot_singleband=False,
144 plot_chi_hist=False,
145 )
146 assert fig is not None
147 assert ax is not None