Coverage for tests/test_plotting.py: 100%

55 statements  

« 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/>. 

21 

22import numpy as np 

23import pytest 

24 

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 

32 

33sigma_inv = 1e4 

34 

35 

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} 

40 

41 

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 

47 

48 dn_rows, dn_cols = 1, -2 

49 dx_min, dy_min = -2, 1 

50 

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) 

67 

68 

69@pytest.fixture(scope="module") 

70def psf_model(): 

71 """Return a trivial PSF model.""" 

72 return make_psf_model_null() 

73 

74 

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) 

79 

80 

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

85 

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 

117 

118 

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 

133 

134 

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