Coverage for tests/test_sourceconfig.py: 100%

65 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-19 09:17 +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/>. 

21 

22import numpy as np 

23import pytest 

24 

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 

34 

35 

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 

41 

42 

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 

50 

51 

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")} 

56 

57 

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() 

71 

72 

73def test_SourceConfig_base(): 

74 """Test that SourceConfig initializes correctly.""" 

75 config = SourceConfig( 

76 component_groups={"": ComponentGroupConfig(components_gauss={"": GaussianComponentConfig()})} 

77 ) 

78 config.validate() 

79 

80 with pytest.raises(ValueError): 

81 config = SourceConfig() 

82 config.validate() 

83 

84 with pytest.raises(ValueError): 

85 config = SourceConfig(component_groups={}) 

86 config.validate() 

87 

88 

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 

93 

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 

122 

123 

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 

128 

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