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
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:
@@ -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,
|
||||
),
|
||||
}
|
||||
|
||||
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user