Coverage for tests/test_diaPipe.py: 99%

491 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-25 23:05 +0000

1# This file is part of ap_association. 

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 contextlib 

23import tempfile 

24import unittest 

25from unittest.mock import patch, MagicMock, DEFAULT 

26import warnings 

27 

28import numpy as np 

29import pandas as pd 

30import astropy.table as tb 

31import astropy.units as u 

32 

33import lsst.afw.table as afwTable 

34import lsst.dax.apdb as daxApdb 

35from lsst.meas.base import DetectorVisitIdGeneratorConfig, IdGenerator 

36import lsst.pex.config as pexConfig 

37import lsst.pipe.base as pipeBase 

38import lsst.utils.tests 

39from lsst.pipe.base.testUtils import assertValidOutput 

40 

41from lsst.ap.association import DiaPipelineTask 

42from lsst.ap.association.utils import getRegion 

43from lsst.pipe.tasks.schemaUtils import convertDataFrameToSdmSchema 

44from utils_tests import makeExposure, makeDiaObjects, makeDiaSources, makeDiaForcedSources, \ 

45 makeSolarSystemSources 

46 

47 

48def _makeMockDataFrame(): 

49 """Create a new mock of a DataFrame. 

50 

51 Returns 

52 ------- 

53 mock : `unittest.mock.Mock` 

54 A mock guaranteed to accept all operations used by `pandas.DataFrame`. 

55 """ 

56 with warnings.catch_warnings(): 

57 # spec triggers deprecation warnings on DataFrame, but will 

58 # automatically adapt to any removals. 

59 warnings.simplefilter("ignore", category=DeprecationWarning) 

60 return MagicMock(spec=pd.DataFrame()) 

61 

62 

63def _makeMockTable(): 

64 """Create a new mock of a Table. 

65 

66 Returns 

67 ------- 

68 mock : `unittest.mock.Mock` 

69 A mock guaranteed to accept all operations used by `astropy.table.Table`. 

70 """ 

71 with warnings.catch_warnings(): 

72 # spec triggers deprecation warnings on DataFrame, but will 

73 # automatically adapt to any removals. 

74 warnings.simplefilter("ignore", category=DeprecationWarning) 

75 return MagicMock(spec=tb.Table()) 

76 

77 

78class TestDiaPipelineTask(unittest.TestCase): 

79 

80 @classmethod 

81 def _makeDefaultConfig(cls, config_file, **kwargs): 

82 config = DiaPipelineTask.ConfigClass() 

83 config.apdb_config_url = config_file 

84 config.update(**kwargs) 

85 return config 

86 

87 def setUp(self): 

88 # Create an instance of random generator with fixed seed. 

89 rng = np.random.default_rng(1234) 

90 self.rng = rng 

91 

92 # schemas are persisted in both Gen 2 and Gen 3 butler as prototypical catalogs 

93 srcSchema = afwTable.SourceTable.makeMinimalSchema() 

94 srcSchema.addField("base_PixelFlags_flag", type="Flag") 

95 srcSchema.addField("base_PixelFlags_flag_offimage", type="Flag") 

96 self.srcSchema = afwTable.SourceCatalog(srcSchema) 

97 self.exposure = makeExposure(False, False) 

98 self.diffim = makeExposure(False, False) 

99 self.template = makeExposure(False, False) 

100 self.diaObjects = makeDiaObjects(20, self.exposure, rng) 

101 self.diaSources = makeDiaSources( 

102 100, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, rng) 

103 self.diaForcedSources = makeDiaForcedSources( 

104 200, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, rng) 

105 self.ssSources = makeSolarSystemSources( 

106 20, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, rng) 

107 

108 sqlite_file = tempfile.NamedTemporaryFile() 

109 self.addCleanup(sqlite_file.close) 

110 self.config_file = tempfile.NamedTemporaryFile() 

111 self.addCleanup(self.config_file.close) 

112 apdb_config = daxApdb.ApdbSql.init_database(db_url=f"sqlite:///{sqlite_file.name}") 

113 apdb_config.save(self.config_file.name) 

114 

115 def testRun(self): 

116 """Test running while creating and packaging alerts. 

117 """ 

118 self._testRun(doPackageAlerts=True, doSolarSystemAssociation=True, doReloadDiaObjects=False) 

119 

120 def testRunWithSolarSystemAssociation(self): 

121 """Test running while creating and packaging alerts. 

122 """ 

123 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=True, doReloadDiaObjects=False) 

124 

125 def testRunWithAlerts(self): 

126 """Test running while creating and packaging alerts. 

127 """ 

128 self._testRun(doPackageAlerts=True, doSolarSystemAssociation=False, doReloadDiaObjects=False) 

129 

130 def testRunWithoutAlertsOrSolarSystem(self): 

131 """Test running without creating and packaging alerts. 

132 """ 

133 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=False, doReloadDiaObjects=False) 

134 

135 def testRunWithReload(self): 

136 """Test running with reloading DiaObjects. 

137 """ 

138 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=False, doReloadDiaObjects=True) 

139 

140 def testRunWithReloadAndSolarSystem(self): 

141 """Test running with solar system association and reloading DiaObjects. 

142 """ 

143 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=True, doReloadDiaObjects=True) 

144 

145 def testRunWithReloadAndAlerts(self): 

146 """Test running with reloading DiaObjects while creating and packaging alerts. 

147 """ 

148 self._testRun(doPackageAlerts=True, doSolarSystemAssociation=False, doReloadDiaObjects=True) 

149 

150 def testRunDisableDeprecatedDoRunForcedMeasurement(self): 

151 """Test running with forced sources disabled. 

152 """ 

153 self._testRun(doPackageAlerts=True, doSolarSystemAssociation=False, doReloadDiaObjects=True, 

154 doRunForcedMeasurement=False, subtasksToMock=["diaCalculation", ] 

155 ) 

156 

157 def testRunWithReloadAllApdbCatalogs(self): 

158 """Test running with reloading the DiaSource history. 

159 """ 

160 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=False, doReloadDiaObjects=True, 

161 doReloadAllApdbCatalogs=True) 

162 

163 def testRunWithReloadAllApdbCatalogsAndAlerts(self): 

164 """Test reloading the DiaSource history while packaging alerts. 

165 """ 

166 self._testRun(doPackageAlerts=True, doSolarSystemAssociation=False, doReloadDiaObjects=True, 

167 doReloadAllApdbCatalogs=True) 

168 

169 def testRunWithReloadAllApdbCatalogsOnly(self): 

170 """Test that reloading everything implies the DiaObject reload. 

171 """ 

172 self._testRun(doPackageAlerts=False, doSolarSystemAssociation=False, doReloadDiaObjects=False, 

173 doReloadAllApdbCatalogs=True) 

174 

175 def _testRun(self, doPackageAlerts=False, doSolarSystemAssociation=False, 

176 doReloadDiaObjects=False, doReloadAllApdbCatalogs=False, subtasksToMock=None, **kwargs): 

177 """Test the normal workflow of each ap_pipe step. 

178 """ 

179 config = self._makeDefaultConfig( 

180 config_file=self.config_file.name, 

181 doPackageAlerts=doPackageAlerts, 

182 doSolarSystemAssociation=doSolarSystemAssociation, 

183 doReloadDiaObjects=doReloadDiaObjects, 

184 doReloadAllApdbCatalogs=doReloadAllApdbCatalogs, 

185 **kwargs 

186 ) 

187 task = DiaPipelineTask(config=config) 

188 # Set DataFrame index testing to always return False. Mocks return 

189 # true for this check otherwise. 

190 task.testDataFrameIndex = lambda x: False 

191 diaSrc = _makeMockDataFrame() 

192 ssObjects = _makeMockTable() 

193 

194 # Each of these subtasks should be called once during diaPipe 

195 # execution. We use mocks here to check they are being executed 

196 # appropriately. 

197 if subtasksToMock is None: 

198 subtasksToMock = ["diaCalculation", "diaForcedSource", ] 

199 if doPackageAlerts: 

200 subtasksToMock.append("alertPackager") 

201 else: 

202 self.assertFalse(hasattr(task, "alertPackager")) 

203 

204 if not doSolarSystemAssociation: 

205 self.assertFalse(hasattr(task, "solarSystemAssociator")) 

206 

207 def concatMock(_data, **_kwargs): 

208 return _makeMockDataFrame() 

209 

210 # Mock out the run() methods of these Tasks to ensure they 

211 # return data in the correct form. 

212 def solarSystemAssociator_run(unAssocDiaSources, solarSystemObjectTable, visitInfo, 

213 bbox, wcs): 

214 return lsst.pipe.base.Struct(nTotalSsObjects=42, 

215 nAssociatedSsObjects=30, 

216 ssoAssocDiaSources=_makeMockTable(), 

217 unAssocDiaSources=_makeMockTable(), 

218 associatedSsSources=_makeMockTable(), 

219 unassociatedSsObjects=_makeMockTable()) 

220 

221 def associator_run(table, diaObjects, schema=None): 

222 return lsst.pipe.base.Struct(nUpdatedDiaObjects=2, nUnassociatedDiaObjects=3, 

223 matchedDiaSources=_makeMockDataFrame(), 

224 unAssocDiaSources=_makeMockDataFrame()) 

225 

226 def loadObjects_run(region, preloadedDiaObjects): 

227 task.metadata['loadRefreshedDiaObjectsStartUtc'] = 1.234 

228 task.metadata['loadRefreshedDiaObjectsEndUtc'] = 5.678 

229 return self.diaObjects 

230 

231 def loadSources_run(region, diaObjects, visitTime): 

232 task.metadata['loadRefreshedDiaSourcesStartUtc'] = 1.234 

233 task.metadata['loadRefreshedDiaSourcesEndUtc'] = 5.678 

234 return self.diaSources 

235 

236 def loadForcedSources_run(region, diaObjects, visitTime): 

237 task.metadata['loadRefreshedDiaForcedSourcesStartUtc'] = 1.234 

238 task.metadata['loadRefreshedDiaForcedSourcesEndUtc'] = 5.678 

239 return self.diaForcedSources 

240 

241 def updateObjectTableMock(diaObjects, diaSources): 

242 pass 

243 

244 def _selectGoodDiaObjects(diaObjectCat, mergedDiaSourceHistory): 

245 return diaObjectCat.copy(deep=True) 

246 

247 # apdb isn't a subtask, but still needs to be mocked out for correct 

248 # execution in the test environment. 

249 with patch.multiple(task, **{task: DEFAULT for task in subtasksToMock + ["apdb"]}), \ 

250 patch('lsst.ap.association.diaPipe.pd.concat', side_effect=concatMock), \ 

251 patch('lsst.ap.association.diaPipe.DiaPipelineTask.updateObjectTable', 

252 side_effect=updateObjectTableMock), \ 

253 patch('lsst.ap.association.diaPipe.DiaPipelineTask._selectGoodDiaObjects', 

254 side_effect=_selectGoodDiaObjects), \ 

255 patch('lsst.ap.association.diaPipe.DiaPipelineTask.loadRefreshedDiaObjects', 

256 side_effect=loadObjects_run) as loadObjectsRun, \ 

257 patch('lsst.ap.association.diaPipe.DiaPipelineTask.loadRefreshedDiaSources', 

258 side_effect=loadSources_run) as loadSourcesRun, \ 

259 patch('lsst.ap.association.diaPipe.DiaPipelineTask.loadRefreshedDiaForcedSources', 

260 side_effect=loadForcedSources_run) as loadForcedSourcesRun, \ 

261 patch('lsst.ap.association.association.AssociationTask.run', 

262 side_effect=associator_run) as mainRun, \ 

263 patch('lsst.pipe.tasks.ssoAssociation.SolarSystemAssociationTask.run', 

264 side_effect=solarSystemAssociator_run) as ssRun: 

265 

266 result = task.run(diaSrc, 

267 None, 

268 self.diffim, 

269 self.exposure, 

270 self.template, 

271 # `run` sees None when the preloaded 

272 # connections have been removed. 

273 preloadedDiaObjects=None if doReloadAllApdbCatalogs else self.diaObjects, 

274 preloadedDiaSources=None if doReloadAllApdbCatalogs else self.diaSources, 

275 preloadedDiaForcedSources=( 

276 None if doReloadAllApdbCatalogs else self.diaForcedSources), 

277 band="g", 

278 idGenerator=IdGenerator(), 

279 solarSystemObjectTable=ssObjects) 

280 for subtaskName in subtasksToMock: 

281 getattr(task, subtaskName).run.assert_called_once() 

282 assertValidOutput(task, result) 

283 # Exact type and contents of apdbMarker are undefined. 

284 self.assertIsInstance(result.apdbMarker, pexConfig.Config) 

285 meta = task.getFullMetadata() 

286 # Check that the expected metadata has been set. 

287 self.assertEqual(meta["diaPipe.numUpdatedDiaObjects"], 2) 

288 self.assertEqual(meta["diaPipe.numUnassociatedDiaObjects"], 3) 

289 # and that associators ran once or not at all. 

290 mainRun.assert_called_once() 

291 if doSolarSystemAssociation: 

292 ssRun.assert_called_once() 

293 else: 

294 ssRun.assert_not_called() 

295 # doReloadAllApdbCatalogs implies the DiaObject reload. 

296 if doReloadDiaObjects or doReloadAllApdbCatalogs: 

297 loadObjectsRun.assert_called_once() 

298 else: 

299 loadObjectsRun.assert_not_called() 

300 # These key names are shared with LoadDiaCatalogsTask so that both 

301 # tasks publish to the same Sasquatch topics. 

302 if doReloadAllApdbCatalogs: 

303 loadSourcesRun.assert_called_once() 

304 loadForcedSourcesRun.assert_called_once() 

305 self.assertGreater(meta["diaPipe.loadDiaSourcesDuration"], 0) 

306 self.assertGreater(meta["diaPipe.loadDiaForcedSourcesDuration"], 0) 

307 else: 

308 loadSourcesRun.assert_not_called() 

309 loadForcedSourcesRun.assert_not_called() 

310 self.assertEqual(meta["diaPipe.loadDiaSourcesDuration"], -1) 

311 self.assertEqual(meta["diaPipe.loadDiaForcedSourcesDuration"], -1) 

312 

313 def test_reloadAllApdbCatalogsRemovesPreloadedConnections(self): 

314 """Test that reloading drops the preloaded catalog connections. 

315 """ 

316 preloaded = {"preloadedDiaObjects", "preloadedDiaSources", "preloadedDiaForcedSources"} 

317 

318 config = self._makeDefaultConfig(config_file=self.config_file.name, 

319 doReloadAllApdbCatalogs=False) 

320 connections = config.connections.ConnectionsClass(config=config) 

321 self.assertTrue(preloaded <= set(connections.inputs)) 

322 

323 config = self._makeDefaultConfig(config_file=self.config_file.name, 

324 doReloadAllApdbCatalogs=True) 

325 connections = config.connections.ConnectionsClass(config=config) 

326 self.assertFalse(preloaded & set(connections.inputs)) 

327 

328 def test_runQuantumOmitsRemovedPreloadedInputs(self): 

329 """Test that runQuantum passes only the preloaded inputs it has. 

330 

331 `run` defaults the preloaded catalogs to `None`, so runQuantum passes 

332 just the connections that survived, and must leave the loaded catalogs 

333 alone when they are present. 

334 """ 

335 preloaded = ["preloadedDiaObjects", "preloadedDiaSources", "preloadedDiaForcedSources"] 

336 loaded = {"preloadedDiaObjects": self.diaObjects, 

337 "preloadedDiaSources": self.diaSources, 

338 "preloadedDiaForcedSources": self.diaForcedSources} 

339 

340 def runQuantumKwargs(doReloadAllApdbCatalogs): 

341 config = self._makeDefaultConfig(config_file=self.config_file.name, 

342 doReloadAllApdbCatalogs=doReloadAllApdbCatalogs) 

343 task = DiaPipelineTask(config=config) 

344 # The quantum only carries the connections that survived. 

345 inputs = {"diaSourceTable": _makeMockDataFrame(), 

346 "diffIm": self.diffim, 

347 "exposure": self.exposure, 

348 "template": self.template, 

349 "solarSystemObjectTable": _makeMockTable()} 

350 if not doReloadAllApdbCatalogs: 

351 inputs.update(loaded) 

352 butlerQC = MagicMock() 

353 butlerQC.get.return_value = inputs 

354 with patch.object(DetectorVisitIdGeneratorConfig, "apply", return_value=IdGenerator()), \ 

355 patch.object(DiaPipelineTask, "run") as mockRun: 

356 task.runQuantum(butlerQC, MagicMock(), MagicMock()) 

357 return mockRun.call_args.kwargs 

358 

359 kwargs = runQuantumKwargs(True) 

360 for name in preloaded: 

361 self.assertNotIn(name, kwargs, msg=f"{name} should be left to the default") 

362 

363 kwargs = runQuantumKwargs(False) 

364 for name in preloaded: 

365 self.assertIs(kwargs[name], loaded[name], msg=f"{name} should be the loaded catalog") 

366 

367 def test_runRequiresBandAndIdGenerator(self): 

368 """Test that `run` rejects the `None` defaults on required arguments. 

369 """ 

370 config = self._makeDefaultConfig(config_file=self.config_file.name) 

371 task = DiaPipelineTask(config=config) 

372 kwargs = {"diaSourceTable": _makeMockDataFrame(), 

373 "legacySolarSystemTable": None, 

374 "diffIm": self.diffim, 

375 "exposure": self.exposure, 

376 "template": self.template, 

377 "band": "g", 

378 "idGenerator": IdGenerator()} 

379 

380 with self.assertRaisesRegex(ValueError, "band"): 

381 task.run(**(kwargs | {"band": None})) 

382 with self.assertRaisesRegex(ValueError, "idGenerator"): 

383 task.run(**(kwargs | {"idGenerator": None})) 

384 

385 def test_loadRefreshedDiaObjectsNoPreloaded(self): 

386 """Test reloading DiaObjects when no preloaded catalog is available. 

387 """ 

388 config = self._makeDefaultConfig(config_file=self.config_file.name, 

389 doReloadAllApdbCatalogs=True) 

390 task = DiaPipelineTask(config=config) 

391 task.apdb.store(self.exposure.visitInfo.date.toAstropy() - 30*u.day, 

392 self.diaObjects, 

393 self.diaSources, 

394 self.diaForcedSources) 

395 

396 diaObjects = task.loadRefreshedDiaObjects(getRegion(self.exposure)) 

397 

398 self.assertEqual(len(diaObjects), len(self.diaObjects)) 

399 self.assertEqual(diaObjects.index.name, "diaObjectId") 

400 

401 def test_loadRefreshedDiaSources(self): 

402 """Test that the DiaSource history is reloaded from the APDB. 

403 """ 

404 config = self._makeDefaultConfig(config_file=self.config_file.name, 

405 doReloadAllApdbCatalogs=True) 

406 task = DiaPipelineTask(config=config) 

407 visitTime = self.exposure.visitInfo.date.toAstropy() 

408 # Store the history as though it were observed a month earlier. 

409 task.apdb.store(visitTime - 30*u.day, 

410 self.diaObjects, 

411 self.diaSources, 

412 self.diaForcedSources) 

413 

414 region = getRegion(self.exposure) 

415 diaSources = task.loadRefreshedDiaSources(region, self.diaObjects, visitTime) 

416 diaForcedSources = task.loadRefreshedDiaForcedSources(region, self.diaObjects, visitTime) 

417 

418 self.assertEqual(len(diaSources), len(self.diaSources)) 

419 self.assertEqual(len(diaForcedSources), len(self.diaForcedSources)) 

420 self.assertEqual(list(diaSources.index.names), ["diaObjectId", "band", "diaSourceId"]) 

421 self.assertEqual(list(diaForcedSources.index.names), ["diaObjectId", "diaForcedSourceId"]) 

422 self.assertFalse(diaSources.index.has_duplicates) 

423 self.assertFalse(diaForcedSources.index.has_duplicates) 

424 

425 def test_loadRefreshedDiaSourcesNoDiaObjects(self): 

426 """Test reloading the history when no DiaObjects are in range. 

427 """ 

428 config = self._makeDefaultConfig(config_file=self.config_file.name, 

429 doReloadAllApdbCatalogs=True) 

430 task = DiaPipelineTask(config=config) 

431 emptyDiaObjects = self.diaObjects.iloc[:0] 

432 region = getRegion(self.exposure) 

433 visitTime = self.exposure.visitInfo.date.toAstropy() 

434 

435 diaSources = task.loadRefreshedDiaSources(region, emptyDiaObjects, visitTime) 

436 diaForcedSources = task.loadRefreshedDiaForcedSources(region, emptyDiaObjects, visitTime) 

437 

438 self.assertTrue(diaSources.empty) 

439 self.assertTrue(diaForcedSources.empty) 

440 

441 def test_tooManyDiaObjectsError(self): 

442 maxNewDiaObjects = 100 

443 

444 nDiaSources = maxNewDiaObjects + 1 

445 diaSources = makeDiaSources(nDiaSources, np.zeros(nDiaSources), self.exposure, self.rng) 

446 

447 def runAndTestWithContextManager(threshold): 

448 config = self._makeDefaultConfig(config_file=self.config_file.name, 

449 doSolarSystemAssociation=False, 

450 filterUnAssociatedSources=False, 

451 maxNewDiaObjects=threshold, 

452 ) 

453 task = DiaPipelineTask(config=config) 

454 contextManager = self.assertRaises(pipeBase.AlgorithmError) if nDiaSources > threshold > 0 \ 

455 else contextlib.nullcontext() 

456 with contextManager: 

457 task.associateDiaSources( 

458 diaSources, 

459 None, 

460 None, 

461 self.diaObjects, 

462 ) 

463 # Test cases at, above, and below the threshold as well as at 0. 

464 runAndTestWithContextManager(0) 

465 runAndTestWithContextManager(maxNewDiaObjects - 1) 

466 runAndTestWithContextManager(maxNewDiaObjects) 

467 runAndTestWithContextManager(maxNewDiaObjects + 1) 

468 

469 def test_createDiaObjects(self): 

470 """Test that creating new DiaObjects works as expected. 

471 """ 

472 nSources = 5 

473 config = self._makeDefaultConfig(config_file=self.config_file.name, doPackageAlerts=False) 

474 task = DiaPipelineTask(config=config) 

475 diaSources = pd.DataFrame(data=[ 

476 {"ra": 0.04*idx, "dec": 0.04*idx, 

477 "diaSourceId": idx + 1 + nSources, "diaObjectId": 0, 

478 "ssObjectId": 0} 

479 for idx in range(nSources)]) 

480 

481 result = task.createNewDiaObjects(convertDataFrameToSdmSchema(task.schema, diaSources, "DiaSource", 

482 skipIndex=True)) 

483 self.assertEqual(nSources, len(result.newDiaObjects)) 

484 self.assertTrue(np.all(np.equal( 

485 result.diaSources["diaObjectId"].to_numpy(), 

486 result.diaSources["diaSourceId"].to_numpy()))) 

487 self.assertTrue(np.all(np.equal( 

488 result.newDiaObjects["diaObjectId"].to_numpy(), 

489 result.diaSources["diaSourceId"].to_numpy()))) 

490 

491 def test_purgeDiaObjects(self): 

492 """Remove diaOjects that are outside an image's bounding box. 

493 """ 

494 

495 config = self._makeDefaultConfig(config_file=self.config_file.name, doPackageAlerts=False) 

496 task = DiaPipelineTask(config=config) 

497 exposure = makeExposure(False, False) 

498 nObj0 = 20 

499 

500 # Create diaObjects 

501 diaObjects = makeDiaObjects(nObj0, exposure, self.rng) 

502 # Shrink the bounding box so that some of the diaObjects will be outside 

503 bbox = exposure.getBBox() 

504 size = np.minimum(bbox.getHeight(), bbox.getWidth()) 

505 bbox.grow(-size//4) 

506 exposureCut = exposure[bbox] 

507 sizeCut = np.minimum(bbox.getHeight(), bbox.getWidth()) 

508 buffer = 10 

509 bbox.grow(buffer) 

510 

511 def check_diaObjects(bbox, wcs, diaObjects): 

512 raVals = diaObjects.ra.to_numpy() 

513 decVals = diaObjects.dec.to_numpy() 

514 xVals, yVals = wcs.skyToPixelArray(raVals, decVals, degrees=True) 

515 selector = bbox.contains(xVals, yVals) 

516 return selector 

517 

518 selector0 = check_diaObjects(bbox, exposureCut.getWcs(), diaObjects) 

519 nIn0 = np.count_nonzero(selector0) 

520 nOut0 = np.count_nonzero(~selector0) 

521 self.assertEqual(nObj0, nIn0 + nOut0) 

522 

523 # Add an ID that is not in the diaObject table. It should not get removed. 

524 diaObjectIds0 = diaObjects["diaObjectId"].copy(deep=True) 

525 diaObjectIds0[max(diaObjectIds0.index) + 1] = 999 

526 diaObjects1, objIds = task.purgeDiaObjects(exposureCut.getBBox(), exposureCut.getWcs(), diaObjects, 

527 diaObjectIds=diaObjectIds0, buffer=buffer) 

528 diaObjectIds1 = diaObjects1["diaObjectId"] 

529 # Verify that the bounding box was not changed 

530 sizeCheck = np.minimum(exposureCut.getBBox().getHeight(), exposureCut.getBBox().getWidth()) 

531 self.assertEqual(sizeCut, sizeCheck) 

532 selector1 = check_diaObjects(bbox, exposureCut.getWcs(), diaObjects1) 

533 nIn1 = np.count_nonzero(selector1) 

534 nOut1 = np.count_nonzero(~selector1) 

535 nObj1 = len(diaObjects1) 

536 self.assertEqual(nObj1, nIn0) 

537 # Verify that not all diaObjects were removed 

538 self.assertGreater(nObj1, 0) 

539 # Check that some diaObjects were removed 

540 self.assertLess(nObj1, nObj0) 

541 # Verify that no objects outside the bounding box remain 

542 self.assertEqual(nOut1, 0) 

543 # Verify that no objects inside the bounding box were removed 

544 self.assertEqual(nIn1, nIn0) 

545 # The length of the updated object IDs should equal the number of objects 

546 # plus one, since we added an extra ID. 

547 self.assertEqual(nObj1 + 1, len(objIds)) 

548 # All of the object IDs extracted from the catalog should be in the pruned object IDs 

549 self.assertTrue(set(objIds).issuperset(diaObjectIds1)) 

550 # The pruned object IDs should contain entries that are not in the catalog 

551 self.assertFalse(set(diaObjectIds1).issuperset(objIds)) 

552 # Some IDs should have been removed 

553 self.assertLess(len(objIds), len(diaObjectIds0)) 

554 

555 def test_filterDiaObjects(self): 

556 """Unassociated diaSources that are filtered should have good reliability and SNR. 

557 Glint trail sources should also be filtered out. 

558 """ 

559 

560 config = self._makeDefaultConfig(config_file=self.config_file.name, 

561 doPackageAlerts=False, 

562 filterUnAssociatedSources=True) 

563 

564 configBadFilter = self._makeDefaultConfig(config_file=self.config_file.name, 

565 doPackageAlerts=False, 

566 filterUnAssociatedSources=True, 

567 newObjectFluxField="notAFlux") 

568 configBadFlag = self._makeDefaultConfig(config_file=self.config_file.name, 

569 doPackageAlerts=False, 

570 filterUnAssociatedSources=True, 

571 newObjectBadFlags=("junkSource", "notUsed")) 

572 with self.assertRaises(pipeBase.InvalidQuantumError): 

573 DiaPipelineTask(config=configBadFilter) 

574 with self.assertRaises(pipeBase.InvalidQuantumError): 

575 DiaPipelineTask(config=configBadFlag) 

576 task = DiaPipelineTask(config=config) 

577 nUnassociatedDiaSources = 234 

578 

579 # Create diaSources 

580 diaSources = makeDiaSources(nUnassociatedDiaSources, 

581 np.zeros(nUnassociatedDiaSources), 

582 self.exposure, 

583 self.rng, 

584 flagList=task.config.newObjectBadFlags) 

585 reliability = self.rng.random(nUnassociatedDiaSources) 

586 flux = (self.rng.random(nUnassociatedDiaSources)**2)*100 

587 fluxErr = np.sqrt(flux) 

588 glint_trail = np.zeros(nUnassociatedDiaSources, dtype=bool) 

589 glint_trail[12:16] = True # add 4 glint trail sources 

590 diaSources["reliability"] = reliability 

591 diaSources["reliabilityVersion"] = '99.42' 

592 diaSources[config.newObjectFluxField] = flux 

593 diaSources[config.newObjectFluxField + "Err"] = fluxErr 

594 diaSources["glint_trail"] = glint_trail 

595 badFlagName = task.config.newObjectBadFlags[0] 

596 badFlags = np.zeros(nUnassociatedDiaSources, dtype=bool) 

597 nBadFlags = 20 

598 badFlags[0:nBadFlags] = True 

599 diaSources[badFlagName] = badFlags 

600 

601 def runAndCheckFilter(diaSources, snrThreshold=None, lowReliabilitySnrThreshold=None, 

602 reliabilityThreshold=None, lowSnrReliabilityThreshold=None, 

603 badFlags=None, 

604 ): 

605 

606 filterResults = task.filterSources( 

607 diaSources.copy(deep=True), 

608 snrThreshold=snrThreshold, 

609 lowReliabilitySnrThreshold=lowReliabilitySnrThreshold, 

610 reliabilityThreshold=reliabilityThreshold, 

611 lowSnrReliabilityThreshold=lowSnrReliabilityThreshold, 

612 badFlags=badFlags, 

613 ) 

614 self.assertEqual(len(filterResults.goodSources) + len(filterResults.badSources), 

615 nUnassociatedDiaSources) 

616 goodFlux = filterResults.goodSources[config.newObjectFluxField] 

617 goodFluxErr = filterResults.goodSources[config.newObjectFluxField + "Err"] 

618 goodSnr = np.array(goodFlux/goodFluxErr) 

619 self.assertTrue(np.all(goodSnr > snrThreshold)) 

620 goodReliability = np.array(filterResults.goodSources["reliability"]) 

621 self.assertTrue(np.all(goodReliability > reliabilityThreshold)) 

622 goodLowSnrFlag = goodSnr < lowReliabilitySnrThreshold 

623 lowSnrReliability = goodReliability[goodLowSnrFlag] 

624 self.assertTrue(np.all(lowSnrReliability > lowSnrReliabilityThreshold)) 

625 glintTrailSources = np.array(filterResults.goodSources["glint_trail"]) 

626 self.assertTrue(not any(glintTrailSources)) 

627 

628 # No sources should be removed if the thresholds are turned off 

629 runAndCheckFilter(diaSources, 

630 snrThreshold=0, lowReliabilitySnrThreshold=0, 

631 reliabilityThreshold=0, lowSnrReliabilityThreshold=0) 

632 runAndCheckFilter(diaSources, 

633 snrThreshold=0, lowReliabilitySnrThreshold=0, 

634 reliabilityThreshold=0, lowSnrReliabilityThreshold=0, 

635 badFlags=[badFlagName]) 

636 runAndCheckFilter(diaSources, 

637 snrThreshold=2, lowReliabilitySnrThreshold=8, 

638 reliabilityThreshold=0, lowSnrReliabilityThreshold=0) 

639 runAndCheckFilter(diaSources, 

640 snrThreshold=2, lowReliabilitySnrThreshold=8, 

641 reliabilityThreshold=0, lowSnrReliabilityThreshold=0.5) 

642 runAndCheckFilter(diaSources, 

643 snrThreshold=0, lowReliabilitySnrThreshold=0, 

644 reliabilityThreshold=0.1, lowSnrReliabilityThreshold=0.5) 

645 runAndCheckFilter(diaSources, 

646 snrThreshold=2, lowReliabilitySnrThreshold=8, 

647 reliabilityThreshold=0.1, lowSnrReliabilityThreshold=0.5) 

648 runAndCheckFilter(diaSources, 

649 snrThreshold=2, lowReliabilitySnrThreshold=8, 

650 reliabilityThreshold=0.1, lowSnrReliabilityThreshold=0.5, 

651 badFlags=[badFlagName]) 

652 

653 def testRunWithForcedMeasurement(self): 

654 """Test running association with forced photometry.""" 

655 

656 reliabilityThreshold = 0.5 

657 trailLengthThreshold = 1.0 

658 config = self._makeDefaultConfig(config_file=self.config_file.name, 

659 doPackageAlerts=False, 

660 forcedReliabilityThreshold=reliabilityThreshold, 

661 forcedTrailLengthThreshold=trailLengthThreshold) 

662 task = DiaPipelineTask(config=config) 

663 nDeepObjects = 20 

664 nShallowObjects = 20 

665 nGoodDiaSourcesDeep = 100 

666 nGoodDiaSourcesShallow = nShallowObjects 

667 nGoodDiaSources = nGoodDiaSourcesDeep + nGoodDiaSourcesShallow 

668 nBadDiaSources = 200 

669 

670 # Create diaObjects 

671 diaObjectsDeep = makeDiaObjects(nDeepObjects, self.exposure, self.rng, startId=1) 

672 diaObjectsShallow = makeDiaObjects(nShallowObjects, self.exposure, self.rng, startId=1 + nDeepObjects) 

673 diaObjects = task.mergeCatalogs(diaObjectsDeep, diaObjectsShallow, tableName="DiaObject") 

674 

675 diaSourcesGoodDeep = makeDiaSources(nGoodDiaSourcesDeep, diaObjectsDeep["diaObjectId"].to_numpy(), 

676 self.exposure, self.rng, startId=1, 

677 flagList=config.forcedBadFlags) 

678 diaSourcesGoodShallow = makeDiaSources(nGoodDiaSourcesShallow, 

679 diaObjectsShallow["diaObjectId"].to_numpy(), 

680 self.exposure, self.rng, startId=1 + nGoodDiaSourcesDeep, 

681 flagList=config.forcedBadFlags) 

682 

683 diaSourcesGoodDeep = convertDataFrameToSdmSchema(task.schema, diaSourcesGoodDeep, 

684 tableName="DiaSource", skipIndex=True) 

685 diaSourcesGood = task.mergeCatalogs(diaSourcesGoodDeep, diaSourcesGoodShallow, tableName="DiaSource") 

686 

687 diaSourcesBad = makeDiaSources(nBadDiaSources, diaObjects["diaObjectId"].to_numpy(), self.exposure, 

688 self.rng, randomizeObjects=True, startId=1 + nGoodDiaSources, 

689 flagList=config.forcedBadFlags) 

690 diaSourcesGood['reliability'] = (self.rng.random(nGoodDiaSources)*(1 - reliabilityThreshold) 

691 + reliabilityThreshold) 

692 diaSourcesGood['reliabilityVersion'] = '99.42' 

693 diaSourcesGood['trailLength'] = self.rng.random(nGoodDiaSources)*trailLengthThreshold 

694 diaSourcesBad = convertDataFrameToSdmSchema(task.schema, diaSourcesBad, 

695 tableName="DiaSource", skipIndex=True) 

696 

697 # Set some "bad" diaSources to have bad reliability, some good 

698 diaSourcesBad['reliability'] = self.rng.random(nBadDiaSources) 

699 diaSourcesBad['reliabilityVersion'] = '99.42' 

700 # Set some "bad" diaSources to have too long trail lengths, some acceptible 

701 diaSourcesBad['trailLength'] = self.rng.random(nBadDiaSources)*trailLengthThreshold + 0.5 

702 missingBadFlags = ((diaSourcesBad['reliability'] > reliabilityThreshold) 

703 & (diaSourcesBad['trailLength'] < trailLengthThreshold) 

704 ) 

705 for badFlag in config.forcedBadFlags: 

706 # Set a fraction of the bad diaSources to have each flag, assigned randomly 

707 diaSourcesBad[badFlag] = self.rng.random(nBadDiaSources) > 0.8 

708 missingBadFlags &= ~diaSourcesBad[badFlag] 

709 # Catch any "bad" diaSources that have not yet been flagged 

710 if np.any(missingBadFlags): 710 ↛ 712line 710 didn't jump to line 712 because the condition on line 710 was always true

711 diaSourcesBad[badFlag] |= missingBadFlags 

712 diaObjectsForcedGoodOnly = task._selectGoodDiaObjects(diaObjects, diaSourcesGood) 

713 diaObjectsForcedBadOnly = task._selectGoodDiaObjects(diaObjects, diaSourcesBad) 

714 # None of the diaObjects should be selected if only given the bad diaSources 

715 self.assertTrue(diaObjectsForcedBadOnly.empty) 

716 

717 diaSources = task.mergeCatalogs(diaSourcesGood, diaSourcesBad, tableName="DiaSource") 

718 diaObjectsForced = task._selectGoodDiaObjects(diaObjects, diaSources) 

719 # The number of diaObjects selected should be the same regardless of 

720 # whether the bad diaSources are included 

721 self.assertEqual(len(diaObjectsForced), len(diaObjectsForcedGoodOnly)) 

722 # All of the deep diaObjects should be selected, and none of the shallow 

723 self.assertEqual(len(diaObjectsForced), nDeepObjects) 

724 fSrc = task.runForcedMeasurement(diaObjectsForced, diaObjectsForced, self.exposure, self.exposure, 

725 IdGenerator()) 

726 self.assertEqual(set(fSrc['diaObjectId']), set(diaObjectsDeep['diaObjectId'])) 

727 

728 def test_selectGoodObjects(self): 

729 """Test the diaObject selection funtion used for forced photometry. 

730 """ 

731 reliabilityThreshold = 0.5 

732 trailLengthThreshold = 1.0 

733 config = self._makeDefaultConfig(config_file=self.config_file.name, 

734 doPackageAlerts=False, 

735 forcedReliabilityThreshold=reliabilityThreshold, 

736 forcedTrailLengthThreshold=trailLengthThreshold) 

737 task = DiaPipelineTask(config=config) 

738 nObjects = 20 

739 nDiaSources = 100 

740 

741 # Create diaObjects 

742 diaObjects = makeDiaObjects(nObjects, self.exposure, self.rng, startId=1) 

743 

744 diaSources = makeDiaSources(nDiaSources, diaObjects["diaObjectId"].to_numpy(), 

745 self.exposure, self.rng, startId=1, 

746 flagList=config.forcedBadFlags) 

747 # Since flagList is not specified, the columns will be missing, and added and filled with NaNs 

748 # when put through convertDataFrameToSdmSchema. 

749 diaSourcesMissingColumns = makeDiaSources(nDiaSources, diaObjects["diaObjectId"].to_numpy(), 

750 self.exposure, self.rng, startId=1 + nDiaSources) 

751 # Should raise an error if run before adding the required columns 

752 with self.assertRaises(RuntimeError): 

753 task._selectGoodDiaObjects(diaObjects, diaSources) 

754 

755 diaSources['reliability'] = (self.rng.random(nDiaSources)*(1 - reliabilityThreshold) 

756 + reliabilityThreshold) 

757 diaSources['reliabilityVersion'] = '99.42' 

758 diaSources['trailLength'] = self.rng.random(nDiaSources)*trailLengthThreshold 

759 

760 diaSourcesMissingColumns['reliability'] = (self.rng.random(nDiaSources)*(1 - reliabilityThreshold) 

761 + reliabilityThreshold) 

762 diaSourcesMissingColumns['trailLength'] = self.rng.random(nDiaSources)*trailLengthThreshold 

763 

764 diaSourcesNotMatched = makeDiaSources(nDiaSources, diaObjects["diaObjectId"].to_numpy() + 999, 

765 self.exposure, self.rng, startId=1) 

766 

767 diaSourcesNotMatched['reliability'] = (self.rng.random(nDiaSources)*(1 - reliabilityThreshold) 

768 + reliabilityThreshold) 

769 diaSourcesNotMatched['reliabilityVersion'] = '99.42' 

770 diaSourcesNotMatched['trailLength'] = self.rng.random(nDiaSources)*trailLengthThreshold 

771 

772 diaSources = convertDataFrameToSdmSchema(task.schema, diaSources, 

773 tableName="DiaSource", skipIndex=True) 

774 diaSourcesMissingColumns = convertDataFrameToSdmSchema(task.schema, diaSourcesMissingColumns, 

775 tableName="DiaSource", skipIndex=True) 

776 diaSourcesNotMatched = convertDataFrameToSdmSchema(task.schema, diaSourcesNotMatched, 

777 tableName="DiaSource", skipIndex=True) 

778 

779 # Run the method with missing flags 

780 diaObjectsSelected = task._selectGoodDiaObjects(diaObjects, diaSources) 

781 diaObjectsBadSelection = task._selectGoodDiaObjects(diaObjects, diaSourcesNotMatched) 

782 diaObjectsNaNSelection = task._selectGoodDiaObjects(diaObjects, diaSourcesMissingColumns) 

783 # The matched catalog of good diaSources should not drop any diaObjects 

784 self.assertTrue(diaObjects.equals(diaObjectsSelected)) 

785 # The mis-matched catalog of good diaSources should drop every diaObject 

786 self.assertTrue(diaObjectsBadSelection.empty) 

787 # All diaObjects should be dropped for the diaSource catalog that has all NaN values for the flags 

788 self.assertTrue(diaObjectsNaNSelection.empty) 

789 

790 def test_selectGoodObjectsWithWrongFlags(self): 

791 """Test the diaObject selection funtion used for forced photometry. 

792 """ 

793 reliabilityThreshold = 0.5 

794 trailLengthThreshold = 1.0 

795 config = self._makeDefaultConfig(config_file=self.config_file.name, 

796 doPackageAlerts=False, 

797 forcedReliabilityThreshold=reliabilityThreshold, 

798 forcedTrailLengthThreshold=trailLengthThreshold, 

799 forcedBadFlags=['foo', 'bar']) 

800 task = DiaPipelineTask(config=config) 

801 nObjects = 20 

802 nDiaSources = 100 

803 

804 # Create diaObjects 

805 diaObjects = makeDiaObjects(nObjects, self.exposure, self.rng, startId=1) 

806 

807 diaSources = makeDiaSources(nDiaSources, diaObjects["diaObjectId"].to_numpy(), 

808 self.exposure, self.rng, startId=1) 

809 

810 diaSources['reliability'] = (self.rng.random(nDiaSources)*(1 - reliabilityThreshold) 

811 + reliabilityThreshold) 

812 diaSources['reliabilityVersion'] = '99.42' 

813 diaSources['trailLength'] = self.rng.random(nDiaSources)*trailLengthThreshold 

814 

815 diaSourcesNotMatched = makeDiaSources(nDiaSources, diaObjects["diaObjectId"].to_numpy() + 999, 

816 self.exposure, self.rng, startId=1) 

817 

818 diaSourcesNotMatched['reliability'] = (self.rng.random(nDiaSources)*(1 - reliabilityThreshold) 

819 + reliabilityThreshold) 

820 diaSourcesNotMatched['reliabilityVersion'] = '99.42' 

821 diaSourcesNotMatched['trailLength'] = self.rng.random(nDiaSources)*trailLengthThreshold 

822 

823 diaSources = convertDataFrameToSdmSchema(task.schema, diaSources, 

824 tableName="DiaSource", skipIndex=True) 

825 diaSourcesNotMatched = convertDataFrameToSdmSchema(task.schema, diaSourcesNotMatched, 

826 tableName="DiaSource", skipIndex=True) 

827 

828 # Verify that the flags are in fact missing 

829 for flag in config.forcedBadFlags: 

830 self.assertNotIn(flag, diaSources.columns) 

831 # Run the method with missing flags 

832 diaObjectsSelected = task._selectGoodDiaObjects(diaObjects, diaSources) 

833 diaObjectsBadSelection = task._selectGoodDiaObjects(diaObjects, diaSourcesNotMatched) 

834 # The matched catalog of good diaSources should not drop any diaObjects 

835 self.assertTrue(diaObjects.equals(diaObjectsSelected)) 

836 # The mis-matched catalog of good diaSources should drop every diaObject 

837 self.assertTrue(diaObjectsBadSelection.empty) 

838 

839 def test_mergeEmptyCatalog(self): 

840 """Test that a catalog is unchanged if it is merged with an empty 

841 catalog. 

842 """ 

843 diaSourcesBase = self.diaSources 

844 

845 config = self._makeDefaultConfig(config_file=self.config_file.name, doPackageAlerts=False) 

846 task = DiaPipelineTask(config=config) 

847 # Include some but not all columns that should be in diaSourcesBase, and some that are mis-matched 

848 diaSourcesEmpty = pd.DataFrame(columns=["ra", "dec", "foo"]) 

849 diaSourcesTest = task.mergeCatalogs(diaSourcesBase, diaSourcesEmpty, tableName="DiaSource") 

850 self.assertTrue(diaSourcesBase.equals(diaSourcesTest)) 

851 

852 def test_mergeCatalogs(self): 

853 """Test that a merged catalog is concatenated correctly. 

854 """ 

855 config = self._makeDefaultConfig(config_file=self.config_file.name, doPackageAlerts=False) 

856 task = DiaPipelineTask(config=config) 

857 

858 diaSourcesBase = convertDataFrameToSdmSchema(task.schema, self.diaSources, "DiaSource", 

859 skipIndex=True) 

860 nBase = len(diaSourcesBase) 

861 nNew = int(nBase/2) 

862 

863 diaSourcesNew = makeDiaSources(nNew, self.diaObjects["diaObjectId"].to_numpy(), self.exposure, 

864 self.rng) 

865 diaSourcesNew = convertDataFrameToSdmSchema(task.schema, diaSourcesNew, "DiaSource", skipIndex=True) 

866 diaSourcesTest = task.mergeCatalogs(diaSourcesBase, diaSourcesNew, tableName="DiaSource") 

867 self.assertEqual(len(diaSourcesTest), nBase + nNew) 

868 diaSourcesExtract1 = diaSourcesTest.iloc[:nBase] 

869 diaSourcesExtract2 = diaSourcesTest.iloc[nBase:] 

870 

871 pd.testing.assert_frame_equal(diaSourcesBase, diaSourcesExtract1) 

872 pd.testing.assert_frame_equal(diaSourcesNew, diaSourcesExtract2) 

873 

874 def test_updateObjectTable(self): 

875 """Test that the diaObject record is updated with the number of 

876 diaSources. 

877 """ 

878 config = self._makeDefaultConfig(config_file=self.config_file.name, doPackageAlerts=False) 

879 task = DiaPipelineTask(config=config) 

880 nObjects = 20 

881 nSrcPerObject = 10 

882 nExtraSources = 5 

883 nSources = nSrcPerObject*nObjects + nExtraSources 

884 expectedSourcesPerObject = nSrcPerObject*np.ones(nObjects) 

885 expectedSourcesPerObject[:nExtraSources] += 1 

886 diaObjects = makeDiaObjects(nObjects, self.exposure, self.rng) 

887 diaSources = makeDiaSources(nSources, diaObjects["diaObjectId"].to_numpy(), self.exposure, self.rng) 

888 updatedDiaObjects = task.updateObjectTable(diaObjects, diaSources) 

889 self.assertTrue(np.all(updatedDiaObjects.nDiaSources.values == expectedSourcesPerObject)) 

890 

891 def test_diaCalculationAllBands(self): 

892 """Test that the full DiaPipelineTask.run computes per-band 

893 psfFluxMean for all bands with source data and NaN where absent. 

894 """ 

895 config = self._makeDefaultConfig( 

896 config_file=self.config_file.name, 

897 doPackageAlerts=False, 

898 doSolarSystemAssociation=False, 

899 ) 

900 task = DiaPipelineTask(config=config) 

901 task.testDataFrameIndex = lambda x: False 

902 

903 nObjects = 3 

904 diaObjects = makeDiaObjects(nObjects, self.exposure, self.rng) 

905 diaObjectIds = diaObjects["diaObjectId"].to_numpy() 

906 

907 # Synthesize preloadedDiaSources in g, r, and i bands with known 

908 # psfFlux values. These represent the historical source catalog. 

909 bands_with_data = ["g", "r", "i"] 

910 rows = [] 

911 srcId = 1 

912 for band in bands_with_data: 

913 for i, objId in enumerate(diaObjectIds): 

914 nSrc = 3 

915 for j in range(nSrc): 

916 rows.append({ 

917 "diaSourceId": srcId, 

918 "diaObjectId": objId, 

919 "band": band, 

920 "ra": diaObjects.iloc[i]["ra"], 

921 "dec": diaObjects.iloc[i]["dec"], 

922 "midpointMjdTai": 58000.0 + srcId, 

923 "psfFlux": 1000.0 + srcId * 10, 

924 "psfFluxErr": 10.0, 

925 }) 

926 srcId += 1 

927 preloadedDiaSources = pd.DataFrame(rows) 

928 preloadedDiaSources = convertDataFrameToSdmSchema( 

929 task.schema, preloadedDiaSources, "DiaSource", skipIndex=True 

930 ) 

931 preloadedDiaSources.set_index( 

932 ["diaObjectId", "band", "diaSourceId"], inplace=True, drop=False 

933 ) 

934 

935 # The new diaSourceTable for this visit (band=g) needs to associate 

936 # with existing objects. Use the same positions as the diaObjects. 

937 newSrcRows = [] 

938 for i, objId in enumerate(diaObjectIds): 

939 newSrcRows.append({ 

940 "diaSourceId": srcId, 

941 "diaObjectId": 0, 

942 "ssObjectId": 0, 

943 "band": "g", 

944 "ra": diaObjects.iloc[i]["ra"], 

945 "dec": diaObjects.iloc[i]["dec"], 

946 "midpointMjdTai": 58100.0, 

947 "psfFlux": 1500.0, 

948 "psfFluxErr": 15.0, 

949 }) 

950 srcId += 1 

951 diaSourceTable = pd.DataFrame(newSrcRows) 

952 

953 preloadedDiaForcedSources = pd.DataFrame( 

954 columns=["diaObjectId", "diaForcedSourceId"] 

955 ).set_index(["diaObjectId", "diaForcedSourceId"], drop=False) 

956 

957 with patch.multiple(task, apdb=DEFAULT, diaForcedSource=DEFAULT): 

958 result = task.run( 

959 diaSourceTable, 

960 None, 

961 self.diffim, 

962 self.exposure, 

963 self.template, 

964 preloadedDiaObjects=diaObjects, 

965 preloadedDiaSources=preloadedDiaSources, 

966 preloadedDiaForcedSources=preloadedDiaForcedSources, 

967 band="g", 

968 idGenerator=IdGenerator(), 

969 ) 

970 

971 outputDiaObjects = result.diaObjects 

972 

973 # Bands with data should have finite psfFluxMean. 

974 for band in bands_with_data: 

975 col = f"{band}_psfFluxMean" 

976 self.assertIn(col, outputDiaObjects.columns) 

977 vals = outputDiaObjects[col].to_numpy() 

978 self.assertTrue( 

979 np.all(np.isfinite(vals)), 

980 f"{col} should be finite for all objects but got {vals}" 

981 ) 

982 

983 # Bands without data should have NaN psfFluxMean. 

984 bands_without_data = [b for b in config.validBands 

985 if b not in bands_with_data] 

986 for band in bands_without_data: 

987 col = f"{band}_psfFluxMean" 

988 if col in outputDiaObjects.columns: 988 ↛ 986line 988 didn't jump to line 986 because the condition on line 988 was always true

989 vals = outputDiaObjects[col].to_numpy() 

990 self.assertTrue( 

991 np.all(np.isnan(vals)), 

992 f"{col} should be NaN (no data in {band}) but got {vals}" 

993 ) 

994 

995 

996class MemoryTester(lsst.utils.tests.MemoryTestCase): 

997 pass 

998 

999 

1000def setup_module(module): 

1001 lsst.utils.tests.init() 

1002 

1003 

1004if __name__ == "__main__": 1004 ↛ 1005line 1004 didn't jump to line 1005 because the condition on line 1004 was never true

1005 lsst.utils.tests.init() 

1006 unittest.main()