Coverage for tests/test_sourceconfig.py: 100%
65 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:58 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:58 +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 numpy as np
23import pytest
25import lsst.gauss2d.fit as g2f
26from lsst.multiprofit.componentconfig import (
27 GaussianComponentConfig,
28 ParameterConfig,
29 SersicComponentConfig,
30 SersicIndexParameterConfig,
31)
32from lsst.multiprofit.sourceconfig import ComponentGroupConfig, SourceConfig
33from lsst.multiprofit.utils import get_params_uniq
36@pytest.fixture(scope="module")
37def centroid_limits():
38 """Return trivial limits for centroids."""
39 limits = g2f.LimitsD(min=-np.inf, max=np.inf)
40 return limits
43@pytest.fixture(scope="module")
44def centroid(centroid_limits):
45 """Return centroid parameters with trivial limits."""
46 cenx = g2f.CentroidXParameterD(0, limits=centroid_limits, fixed=True)
47 ceny = g2f.CentroidYParameterD(0, limits=centroid_limits, fixed=True)
48 centroid = g2f.CentroidParameters(cenx, ceny)
49 return centroid
52@pytest.fixture(scope="module")
53def channels():
54 """Return dict of generic RGB channels."""
55 return {band: g2f.Channel.get(band) for band in ("R", "G", "B")}
58def test_ComponentGroupConfig(centroid):
59 """Test that ComponentGroupConfig initializes correctly."""
60 config = ComponentGroupConfig(
61 components_gauss={"x": GaussianComponentConfig()},
62 components_sersic={"y": SersicComponentConfig()},
63 )
64 config.validate()
65 with pytest.raises(ValueError):
66 config = ComponentGroupConfig(
67 components_gauss={"x": GaussianComponentConfig()},
68 components_sersic={"x": SersicComponentConfig()},
69 )
70 config.validate()
73def test_SourceConfig_base():
74 """Test that SourceConfig initializes correctly."""
75 config = SourceConfig(
76 component_groups={"": ComponentGroupConfig(components_gauss={"": GaussianComponentConfig()})}
77 )
78 config.validate()
80 with pytest.raises(ValueError):
81 config = SourceConfig()
82 config.validate()
84 with pytest.raises(ValueError):
85 config = SourceConfig(component_groups={})
86 config.validate()
89def test_SourceConfig_fractional(centroid):
90 """Test that SourceConfig works with fractional ComponentGroup."""
91 rho, size_x, size_y = -0.3, 1.4, 1.6
92 drho, dsize_x, dsize_y = 0.5, 1.6, 1.3
94 n_components = 2
95 config = SourceConfig(
96 component_groups={
97 "src": ComponentGroupConfig(
98 components_gauss={
99 str(idx): GaussianComponentConfig(
100 rho=ParameterConfig(value_initial=rho + idx * drho),
101 size_x=ParameterConfig(value_initial=size_x + idx * dsize_x),
102 size_y=ParameterConfig(value_initial=size_y + idx * dsize_y),
103 )
104 for idx in range(n_components)
105 },
106 is_fractional=True,
107 )
108 },
109 )
110 config.validate()
111 channel = g2f.Channel.NONE
112 psf_model, priors = config.make_psf_model(
113 [
114 [
115 {channel: 1.0},
116 {channel: 0.5},
117 ]
118 ],
119 )
120 assert len(priors) == 0
121 assert len(psf_model.components) == n_components
124def test_SourceConfig_linear(centroid, channels):
125 """Test that SourceConfig works with a regular (linear) ComponentGroup."""
126 rho, size_x, size_y, sersicn, flux = 0.4, 1.5, 1.9, 0.5, 4.7
127 drho, dsize_x, dsize_y, dsersicn, dflux = -0.9, 2.5, 5.4, 2.8, 13.9
129 names = ("PS", "Sersic")
130 config = SourceConfig(
131 component_groups={
132 "src": ComponentGroupConfig(
133 components_sersic={
134 name: SersicComponentConfig(
135 rho=ParameterConfig(value_initial=rho + idx * drho),
136 size_x=ParameterConfig(value_initial=size_x + idx * dsize_x),
137 size_y=ParameterConfig(value_initial=size_y + idx * dsize_y),
138 sersic_index=SersicIndexParameterConfig(
139 value_initial=sersicn + idx * dsersicn,
140 fixed=idx == 0,
141 prior_mean=None,
142 ),
143 )
144 for idx, name in enumerate(names)
145 }
146 ),
147 }
148 )
149 fluxes = [
150 {
151 channel: flux + idx_channel * dflux * idx_comp
152 for idx_channel, channel in enumerate(channels.values())
153 }
154 for idx_comp in range(len(config.component_groups["src"].components_sersic))
155 ]
156 source, priors = config.make_source([fluxes])
157 assert len(priors) == 0
158 for idx, component in enumerate(source.components):
159 params = get_params_uniq(component)
160 values_init = {
161 g2f.RhoParameterD: rho + idx * drho,
162 g2f.ReffXParameterD: size_x + idx * dsize_x,
163 g2f.ReffYParameterD: size_y + idx * dsize_y,
164 g2f.SersicIndexParameterD: sersicn + idx * dsersicn,
165 }
166 for name_group, component_group in config.component_groups.items():
167 fluxes_comp = fluxes[idx]
168 name_comp = names[idx]
169 config_comp = component_group.components_sersic[name_comp]
170 fluxes_label = {
171 config.format_label(
172 component_group.format_label(
173 label=config_comp.format_label(
174 label=config.get_integral_label_default(), name_channel=channel.name
175 ),
176 name_component=name_comp,
177 ),
178 name_group=name_group,
179 ): fluxes_comp[channel]
180 for channel in channels.values()
181 }
182 for param in params:
183 if isinstance(param, g2f.IntegralParameterD):
184 assert fluxes_label[param.label] == param.value
185 elif value_init := values_init.get(param.__class__):
186 assert param.value == value_init