From 3d6b67b0237ae14a5864b5ce1d248967a695fac1 Mon Sep 17 00:00:00 2001 From: Ryan Roussel Date: Tue, 14 Jul 2026 11:58:15 -0500 Subject: [PATCH 1/4] add x/y variables for models --- virtual_accelerator/bmad/variables.py | 2 + virtual_accelerator/cheetah/actions.py | 28 +++++++ virtual_accelerator/cheetah/variables.py | 6 +- .../tests/_bmad_model_test_utils.py | 77 +++++++------------ virtual_accelerator/tests/test_cu_hxr.py | 12 ++- 5 files changed, 69 insertions(+), 56 deletions(-) diff --git a/virtual_accelerator/bmad/variables.py b/virtual_accelerator/bmad/variables.py index 127f8a1..9c8e176 100644 --- a/virtual_accelerator/bmad/variables.py +++ b/virtual_accelerator/bmad/variables.py @@ -412,6 +412,8 @@ def get_screen_variables( screen_spec=screen_spec, index=0, # need to reverse the order of the shape for the ArraySize0_RBV and ArraySize1_RBV variables since they are in row-major order ), + bmad_actions.BPMXVariable(name=f"{base_pv}:X", element_name=screen_name), + bmad_actions.BPMYVariable(name=f"{base_pv}:Y", element_name=screen_name), ] return variables diff --git a/virtual_accelerator/cheetah/actions.py b/virtual_accelerator/cheetah/actions.py index ac6808f..eb87d7b 100644 --- a/virtual_accelerator/cheetah/actions.py +++ b/virtual_accelerator/cheetah/actions.py @@ -5,6 +5,7 @@ classes. """ +from cheetah.accelerator import Screen from lume_cheetah.actions import ( CheetahReadOnlyEnumVariable, CheetahReadOnlyNDVariable, @@ -314,3 +315,30 @@ def _get(self, simulator): def _set(self, simulator, value): super()._set(simulator, 1.0 if bool(value) else 0.0) + +class ScreenCentroidVariable(TorchScalarVariable, _ReadOnlyActionMixin): + """Read-only scalar for beam centroid multiplied by 1e3 (to convert to mm, mrad).""" + + element_name: str + centroid_axis: str + unit: str + + def _get(self, simulator): + element = getattr(simulator.segment, self.element_name) + if not isinstance(element, Screen): + raise ValueError( + f"Element {self.element_name!r} is not a Screen and cannot provide X readback" + ) + return getattr(element.get_read_beam(), self.centroid_axis).mean().item() * 1e3 + +class ScreenXVariable(ScreenCentroidVariable): + """Read-only scalar for beam x centroid.""" + + centroid_axis: str = "x" + unit: str = "mm" + +class ScreenYVariable(ScreenCentroidVariable): + """Read-only scalar for beam y centroid.""" + + centroid_axis: str = "y" + unit: str = "mm" diff --git a/virtual_accelerator/cheetah/variables.py b/virtual_accelerator/cheetah/variables.py index 24dfb59..00b9e7f 100644 --- a/virtual_accelerator/cheetah/variables.py +++ b/virtual_accelerator/cheetah/variables.py @@ -37,15 +37,13 @@ "Image:ArraySize0_RBV": "ScreenImageArraySizeVariable", "RESOLUTION": "ScreenResolutionVariable", "IMAGE": "ScreenImageVariable", - "N_OF_ROW": "ScreenImageArraySizeVariable", - "N_OF_COL": "ScreenImageArraySizeVariable", + "X": "ScreenXVariable", + "Y": "ScreenYVariable", } SCREEN_ARRAY_SIZE_INDEX_BY_SUFFIX = { "Image:ArraySize1_RBV": 0, - "N_OF_ROW": 0, "Image:ArraySize0_RBV": 1, - "N_OF_COL": 1, } diff --git a/virtual_accelerator/tests/_bmad_model_test_utils.py b/virtual_accelerator/tests/_bmad_model_test_utils.py index 2c1c1f2..9102f46 100644 --- a/virtual_accelerator/tests/_bmad_model_test_utils.py +++ b/virtual_accelerator/tests/_bmad_model_test_utils.py @@ -11,6 +11,24 @@ TEST_BEAM_PATH = os.path.join(Path(__file__).parent, "../bmad", "test_beam") +DEFAULT_MAGNET_PV_ATTRS = ( + "BCTRL", + "BACT", + "BDES", + "BMIN", + "BMAX", + "STATCTRLSUB.T", + "CTRL", +) +DEFAULT_SCREEN_PV_ATTRS = ( + "Image:ArrayData", + "Image:ArraySize1_RBV", + "Image:ArraySize0_RBV", + "RESOLUTION", + "X", + "Y", +) +DEFAULT_BPM_PV_ATTRS = ("X", "Y", "TMIT") def _normalize_element_name(element_name: str) -> str: """Return an element name without any split-index suffix. @@ -264,15 +282,7 @@ def assert_magnet_pvs_match_tao_lattice( model, element_key: str, excluded_elements: Iterable[str] = (), - element_attrs: tuple[str, ...] = ( - "BCTRL", - "BACT", - "BDES", - "BMIN", - "BMAX", - "STATCTRLSUB.T", - "CTRL", - ), + element_attrs: tuple[str, ...] = DEFAULT_MAGNET_PV_ATTRS, ) -> None: """Assert magnet PV coverage for a Tao-backed model. @@ -302,15 +312,7 @@ def assert_magnet_pvs_match_cheetah_segment( model, element_key: str, excluded_elements: Iterable[str] = (), - element_attrs: tuple[str, ...] = ( - "BCTRL", - "BACT", - "BDES", - "BMIN", - "BMAX", - "STATCTRLSUB.T", - "CTRL", - ), + element_attrs: tuple[str, ...] = DEFAULT_MAGNET_PV_ATTRS, ) -> None: """Assert magnet PV coverage for a Cheetah-backed model. @@ -342,15 +344,7 @@ def assert_magnet_pvs_match_lattice_elements( element_names: Sequence[str], element_keys: Sequence[str], excluded_elements: Iterable[str] = (), - element_attrs: tuple[str, ...] = ( - "BCTRL", - "BACT", - "BDES", - "BMIN", - "BMAX", - "STATCTRLSUB.T", - "CTRL", - ), + element_attrs: tuple[str, ...] = DEFAULT_MAGNET_PV_ATTRS, ) -> None: """Assert magnet PV coverage from explicit lattice metadata sequences. @@ -400,12 +394,7 @@ def assert_magnet_pvs_match_lattice_elements( def assert_screen_image_pvs_in_supported_variables( model, screen_elements: tuple[str, ...] | list[str] | None = None, - screen_attrs: tuple[str, ...] = ( - "Image:ArrayData", - "Image:ArraySize1_RBV", - "Image:ArraySize0_RBV", - "RESOLUTION", - ), + screen_attrs: tuple[str, ...] = DEFAULT_SCREEN_PV_ATTRS, ) -> None: """ Verify image-related PVs for screen elements are present in supported variables. @@ -437,14 +426,9 @@ def assert_screen_image_pvs_in_supported_variables( ) -def assert_screen_image_pvs_match_tao_dump_locations( +def assert_screen_image_pvs_match_tao_lattice( model, - screen_attrs: tuple[str, ...] = ( - "Image:ArrayData", - "Image:ArraySize1_RBV", - "Image:ArraySize0_RBV", - "RESOLUTION", - ), + screen_attrs: tuple[str, ...] = DEFAULT_SCREEN_PV_ATTRS, ) -> None: """Assert screen image PV coverage using Tao ``dump_locations``. @@ -464,12 +448,7 @@ def assert_screen_image_pvs_match_tao_dump_locations( def assert_screen_image_pvs_match_cheetah_segment( model, - screen_attrs: tuple[str, ...] = ( - "Image:ArrayData", - "Image:ArraySize1_RBV", - "Image:ArraySize0_RBV", - "RESOLUTION", - ), + screen_attrs: tuple[str, ...] = DEFAULT_SCREEN_PV_ATTRS, ) -> None: """Assert screen image PV coverage for screen elements in a Cheetah segment. @@ -495,7 +474,7 @@ def assert_screen_image_pvs_match_cheetah_segment( def assert_bpm_pvs_match_tao_lattice( model, - bpm_attrs: tuple[str, ...] = ("X", "Y", "TMIT"), + bpm_attrs: tuple[str, ...] = DEFAULT_BPM_PV_ATTRS, ) -> None: """ Verify that mapped BPM elements expose expected BPM PVs. @@ -530,7 +509,7 @@ def assert_bpm_pvs_match_tao_lattice( def assert_bpm_pvs_match_cheetah_segment( model, - bpm_attrs: tuple[str, ...] = ("X", "Y", "TMIT"), + bpm_attrs: tuple[str, ...] = DEFAULT_BPM_PV_ATTRS, ) -> None: """Assert BPM PV coverage for BPM elements in a Cheetah segment. @@ -558,7 +537,7 @@ def assert_bpm_pvs_match_cheetah_segment( def assert_bpm_pvs_match_elements( model, bpm_elements: Iterable[str], - bpm_attrs: tuple[str, ...] = ("X", "Y", "TMIT"), + bpm_attrs: tuple[str, ...] = DEFAULT_BPM_PV_ATTRS, ) -> None: """Assert BPM PV coverage for an explicit BPM element name collection. diff --git a/virtual_accelerator/tests/test_cu_hxr.py b/virtual_accelerator/tests/test_cu_hxr.py index f66e488..78d40fd 100644 --- a/virtual_accelerator/tests/test_cu_hxr.py +++ b/virtual_accelerator/tests/test_cu_hxr.py @@ -25,7 +25,7 @@ assert_magnet_pvs_match_cheetah_segment, assert_magnet_pvs_match_tao_lattice, assert_roundtrip_pv_get_set, - assert_screen_image_pvs_in_supported_variables, + assert_screen_image_pvs_match_tao_lattice, ) CU_HXR_PROFMON_CONFIG_PATH = ( @@ -58,7 +58,7 @@ def test_initialization(self): model = get_cu_hxr_bmad_model( end_element="OTR4", track_beam=True, custom_beam_path=TEST_BEAM_PATH ) - assert_screen_image_pvs_in_supported_variables(model) + assert_screen_image_pvs_match_tao_lattice(model) # test getting all of the supported variables to ensure no errors with screen variable setup _ = model.get(list(model.supported_variables)) @@ -82,7 +82,7 @@ def test_cu_hxr_twiss(self): def test_sub_lattice(self): model = get_cu_hxr_bmad_model("QE04#1", "OTR2") - assert len(model.supported_variables) < 40 + assert len(model.supported_variables) < 50 # test getting partial lattice with beam tracking model = get_cu_hxr_bmad_model( @@ -150,6 +150,12 @@ def test_roundtrip_pv_get_set(self): ) assert_roundtrip_pv_get_set(model) + def test_screen_pvs_match_cheetah_segment(self): + model = get_cu_hxr_bmad_model( + custom_beam_path=TEST_BEAM_PATH, end_element="OTR4", track_beam=True + ) + assert_screen_image_pvs_match_tao_lattice(model) + class TestCUHXRCheetah: pytestmark = [ From 2dfdef8b69cb29651fb65556b3ed9e76573db07d Mon Sep 17 00:00:00 2001 From: Ryan Roussel Date: Tue, 14 Jul 2026 12:12:26 -0500 Subject: [PATCH 2/4] linting --- virtual_accelerator/cheetah/actions.py | 3 +++ virtual_accelerator/tests/_bmad_model_test_utils.py | 1 + 2 files changed, 4 insertions(+) diff --git a/virtual_accelerator/cheetah/actions.py b/virtual_accelerator/cheetah/actions.py index eb87d7b..dc857a0 100644 --- a/virtual_accelerator/cheetah/actions.py +++ b/virtual_accelerator/cheetah/actions.py @@ -316,6 +316,7 @@ def _get(self, simulator): def _set(self, simulator, value): super()._set(simulator, 1.0 if bool(value) else 0.0) + class ScreenCentroidVariable(TorchScalarVariable, _ReadOnlyActionMixin): """Read-only scalar for beam centroid multiplied by 1e3 (to convert to mm, mrad).""" @@ -331,12 +332,14 @@ def _get(self, simulator): ) return getattr(element.get_read_beam(), self.centroid_axis).mean().item() * 1e3 + class ScreenXVariable(ScreenCentroidVariable): """Read-only scalar for beam x centroid.""" centroid_axis: str = "x" unit: str = "mm" + class ScreenYVariable(ScreenCentroidVariable): """Read-only scalar for beam y centroid.""" diff --git a/virtual_accelerator/tests/_bmad_model_test_utils.py b/virtual_accelerator/tests/_bmad_model_test_utils.py index 485c93e..1c4439b 100644 --- a/virtual_accelerator/tests/_bmad_model_test_utils.py +++ b/virtual_accelerator/tests/_bmad_model_test_utils.py @@ -27,6 +27,7 @@ ) DEFAULT_BPM_PV_ATTRS = ("X", "Y", "TMIT") + def _normalize_element_name(element_name: str) -> str: """Return an element name without any split-index suffix. From 37dde667cb761003283e1b0f5f162b2892d22e27 Mon Sep 17 00:00:00 2001 From: Ryan Roussel Date: Tue, 14 Jul 2026 12:24:03 -0500 Subject: [PATCH 3/4] docstring update Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- virtual_accelerator/cheetah/actions.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/virtual_accelerator/cheetah/actions.py b/virtual_accelerator/cheetah/actions.py index dc857a0..0c3d051 100644 --- a/virtual_accelerator/cheetah/actions.py +++ b/virtual_accelerator/cheetah/actions.py @@ -328,7 +328,7 @@ def _get(self, simulator): element = getattr(simulator.segment, self.element_name) if not isinstance(element, Screen): raise ValueError( - f"Element {self.element_name!r} is not a Screen and cannot provide X readback" + f"Element {self.element_name!r} is not a Screen and cannot provide {self.centroid_axis.upper()} readback" ) return getattr(element.get_read_beam(), self.centroid_axis).mean().item() * 1e3 From b6f6260f6fe1534bbe2d0fd145b05b068b036602 Mon Sep 17 00:00:00 2001 From: Ryan Roussel Date: Tue, 14 Jul 2026 12:24:50 -0500 Subject: [PATCH 4/4] rename test Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- virtual_accelerator/tests/test_cu_hxr.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/virtual_accelerator/tests/test_cu_hxr.py b/virtual_accelerator/tests/test_cu_hxr.py index 78d40fd..5d30a7f 100644 --- a/virtual_accelerator/tests/test_cu_hxr.py +++ b/virtual_accelerator/tests/test_cu_hxr.py @@ -150,7 +150,7 @@ def test_roundtrip_pv_get_set(self): ) assert_roundtrip_pv_get_set(model) - def test_screen_pvs_match_cheetah_segment(self): + def test_screen_pvs_match_tao_lattice(self): model = get_cu_hxr_bmad_model( custom_beam_path=TEST_BEAM_PATH, end_element="OTR4", track_beam=True )