fix: instantiation of beam centre model
CI / lint (push) Skipped
CI / test (3.11) (push) Skipped
CI / test (3.12) (push) Skipped
CI / test (3.13) (push) Skipped
CI / lint (pull_request) Successful in 23s
CI / test (3.11) (pull_request) Failing after 25s
CI / test (3.12) (pull_request) Failing after 22s
CI / test (3.13) (pull_request) Failing after 22s

This commit is contained in:
David Perl
2026-08-21 14:45:05 +02:00
parent 0c106e213b
commit 7be19b0a50
2 changed files with 12 additions and 13 deletions
+8 -9
View File
@@ -23,22 +23,21 @@ class BeamCentreMeasurements(BaseModel):
class BeamCentre(BaseModel):
measurements: BeamCentreMeasurements
measured: BeamCentreMeasurements
model: BeamCenterFromDetectorStage
@model_validator(mode="before")
@classmethod
def _generate_model(cls, data: Any) -> Any:
if not isinstance(data, dict) or "model" in data or "measurements" not in data:
if not isinstance(data, dict) or "model" in data or "measured" not in data:
return data
measurements = BeamCentreMeasurements.model_validate(data["measurements"])
measured = BeamCentreMeasurements.model_validate(data["measured"])
return {
**data,
"measurements": measurements,
"measured": measured,
"model": fit_beam_centre_model(
det_z_mm=measurements.det_z_mm,
det_y_mm=measurements.det_y_mm,
beam_x_px=measurements.beam_x_px,
beam_y_px=measurements.beam_y_px,
det_z_mm=measured.det_z_mm,
det_y_mm=measured.det_y_mm,
beam_x_px=measured.beam_x_px,
beam_y_px=measured.beam_y_px,
),
}
+4 -4
View File
@@ -258,7 +258,7 @@ def test_beam_centre_fits_model_from_measurements(stage_grid):
beam_x_px, beam_y_px = synth(det_z_mm, det_y_mm)
beam_centre = BeamCentre.model_validate(
{
"measurements": {
"measured": {
"det_z_mm": det_z_mm.tolist(),
"det_y_mm": det_y_mm.tolist(),
"beam_x_px": beam_x_px.tolist(),
@@ -276,7 +276,7 @@ def test_beam_centre_keeps_an_explicitly_supplied_model(synthetic_fit, stage_gri
beam_x_px, beam_y_px = synth(det_z_mm, det_y_mm)
beam_centre = BeamCentre.model_validate(
{
"measurements": {
"measured": {
"det_z_mm": det_z_mm.tolist(),
"det_y_mm": det_y_mm.tolist(),
"beam_x_px": (beam_x_px + 1000.0).tolist(),
@@ -293,7 +293,7 @@ def test_beam_centre_revalidation_is_idempotent(stage_grid):
beam_x_px, beam_y_px = synth(det_z_mm, det_y_mm)
beam_centre = BeamCentre.model_validate(
{
"measurements": {
"measured": {
"det_z_mm": det_z_mm.tolist(),
"det_y_mm": det_y_mm.tolist(),
"beam_x_px": beam_x_px.tolist(),
@@ -309,7 +309,7 @@ def test_beam_centre_propagates_fit_errors():
with pytest.raises(ValueError, match="same length"):
BeamCentre.model_validate(
{
"measurements": {
"measured": {
"det_z_mm": [170.0, 250.0, 320.0],
"det_y_mm": [62.0, 62.0, 62.0],
"beam_x_px": [2090.0, 2095.0],