diff --git a/examples/cheetah_diag0_model.ipynb b/examples/cheetah_diag0_model.ipynb index ca8864a..7e3b5e7 100644 --- a/examples/cheetah_diag0_model.ipynb +++ b/examples/cheetah_diag0_model.ipynb @@ -1,15 +1,5 @@ { "cells": [ - { - "cell_type": "code", - "execution_count": null, - "id": "aeb55fa2", - "metadata": {}, - "outputs": [], - "source": [ - "import matplotlib.pyplot as plt" - ] - }, { "cell_type": "code", "execution_count": null, @@ -20,6 +10,7 @@ "# model = get_sc_diag0_cheetah_model()\n", "import torch\n", "from virtual_accelerator.models.sc_diag0 import get_sc_diag0_cheetah_model\n", + "import matplotlib.pyplot as plt\n", "\n", "model = get_sc_diag0_cheetah_model()" ] @@ -134,36 +125,35 @@ "vals = torch.tensor([-2.4370, -2, -1.6, -1.2, -0.8])\n", "vals = vals.unsqueeze(-1)\n", "print(vals.shape)\n", - "model.set({\"QUAD:HTR:120:BCTRL\": vals})" + "model.set({\"QUAD:DIAG0:190:BCTRL\": vals})" ] }, { "cell_type": "code", "execution_count": null, - "id": "355df5de", + "id": "7b9c97ba", "metadata": {}, "outputs": [], "source": [ - "vals = torch.tensor([-2.4370, -2, -1.6, -1.2, -0.8])\n", - "vals = vals.unsqueeze(-1)\n", - "print(vals.shape)\n", - "model.set({\"QUAD:HTR:120:BCTRL\": vals})" + "img = model.get([\"OTRS:DIAG0:420:Image:ArrayData\"])\n", + "print(img[\"OTRS:DIAG0:420:Image:ArrayData\"].shape)" ] }, { "cell_type": "code", "execution_count": null, - "id": "7b9c97ba", + "id": "0c13fd5c", "metadata": {}, "outputs": [], "source": [ - "model.get([\"OTRS:DIAG0:420:Image:ArrayData\"])" + "bpm = model.get([\"BPMS:DIAG0:190:XSCDTH\", \"BPMS:DIAG0:190:YSCDTH\"])\n", + "print(bpm)" ] }, { "cell_type": "code", "execution_count": null, - "id": "0c13fd5c", + "id": "eda3eb11", "metadata": {}, "outputs": [], "source": [] @@ -171,7 +161,7 @@ ], "metadata": { "kernelspec": { - "display_name": "Python 3", + "display_name": "linac-simulation", "language": "python", "name": "python3" }, diff --git a/virtual_accelerator/cheetah/diag0.py b/virtual_accelerator/cheetah/diag0.py new file mode 100644 index 0000000..a5c9573 --- /dev/null +++ b/virtual_accelerator/cheetah/diag0.py @@ -0,0 +1,108 @@ +import os +import torch +from cheetah.accelerator import Segment, Screen +from pathlib import Path + +from cheetah.accelerator.patch import Patch +from cheetah.accelerator.superimposed import SuperimposedElement +from cheetah.accelerator import Quadrupole + + +def get_diag0_beamline(): + # try to get LCLS_LATTICE -- returns none if not found + lcls_lattice_location = os.getenv("LCLS_LATTICE") + + # if LCLS_LATTICE not found, use the local json model + if lcls_lattice_location is None: + tracking_segment = Segment.from_lattice_json( + os.path.join(Path(__file__).parent, "sc_diag0.json") + ).subcell(start="bpmdg000") + else: + tracking_segment = Segment.from_lattice_json( + f"{lcls_lattice_location}/cheetah/sc_diag0.json" + ).subcell(start="bpmdg000") + + dyqdg001 = Patch( + name="dyqdg001", pitch=torch.tensor((0.0, 9.49758257820075558e-003)) + ) + dyqdg003 = Patch( + name="dyqdg003", pitch=torch.tensor((0.0, 5.88487966838956720e-004)) + ) + + elements = list(tracking_segment.elements) + + # create SuperimposedElements for qdg001 and qdg003 + quad = tracking_segment.qdg001[0] + # In bmad this (and other elements below) are referenced twice, for us we need to multiply the length of a single use element by 2 + quad.length = quad.length * 2 + super_qdg001 = SuperimposedElement( + name="qdg001", + base_element=quad, + superimposed_element=Segment([dyqdg001, tracking_segment.bpmdg001]), + ) + quad = tracking_segment.qdg003[0] + quad.length = quad.length * 2 + super_qdg003 = SuperimposedElement( + name="qdg003", + base_element=quad, + superimposed_element=Segment([dyqdg003, tracking_segment.bpmdg003]), + ) + for ele in [super_qdg001, super_qdg003]: + idx = [ele.name for ele in elements].index(ele.name) + elements[idx : idx + 3] = [ele] + + # create superimposed elements for other quads containing bpms + split_quads = [2, 4, 5, 8, 9, 11] + for idx in split_quads: + quad = getattr(tracking_segment, f"qdg{idx:0>3}")[0] + quad.length = quad.length * 2 + bpm = getattr(tracking_segment, f"bpmdg{idx:0>3}") + super_element = SuperimposedElement( + name=f"qdg{idx:0>3}", base_element=quad, superimposed_element=bpm + ) + idx = [ele.name for ele in elements].index(f"qdg{idx:0>3}") + elements[idx : idx + 3] = [super_element] + + # create the superimposed element for the transverse deflecting cavity + tdc_idx = [ele.name for ele in elements].index("tcxdg0") + tdc = tracking_segment.tcxdg0[0] + tdc.length = tdc.length * 2.0 + tdc.num_steps = 11 + vkick = tracking_segment.ycdgtcx + xkick = tracking_segment.xcdgtcx + tdc.voltage = torch.tensor(0.0) + + super_tdc = SuperimposedElement( + name="tcxdg0", base_element=tdc, superimposed_element=Segment([vkick, xkick]) + ) + elements[tdc_idx : tdc_idx + 4] = [super_tdc] + + tracking_segment = Segment(elements) + + # change the offset of these quads + tracking_segment.qdg001.base_element.misalignment = torch.tensor( + (0.0, -4.83865231890650768e-003) + ) + tracking_segment.qdg001.superimposed_element.misalignment = torch.tensor( + (0.0, -4.83865231890650768e-003) + ) + + tracking_segment.qdg003.base_element.misalignment = torch.tensor( + (0.0, -2.99813063290814820e-004) + ) + tracking_segment.qdg003.superimposed_element.misalignment = torch.tensor( + (0.0, -2.99813063290814820e-004) + ) + + # set screens to use kde + for ele in tracking_segment.elements: + if isinstance(ele, Screen): + ele.method = "charge_deposition" + elif isinstance(ele, Quadrupole): + ele.tracking_method = "second_order" + elif isinstance(ele, SuperimposedElement): + ele.base_element.tracking_method = "second_order" + elif hasattr(ele, "supported_tracking_methods"): + ele.tracking_method = "linear" + + return tracking_segment diff --git a/virtual_accelerator/cheetah/transformer.py b/virtual_accelerator/cheetah/transformer.py index 48806bd..6dc87d3 100644 --- a/virtual_accelerator/cheetah/transformer.py +++ b/virtual_accelerator/cheetah/transformer.py @@ -68,10 +68,18 @@ def get_cheetah_property(self, simulator, control_variable_name): f"No mapping found for control variable '{control_variable_name}'" ) - element = getattr(simulator.segment, element_name) - beam_energy_at_element = simulator.energies[element_name] - # due to getting beam energy this calc is very slow, maybe some list format should - # be passable for args. + try: + element = getattr(simulator.segment, element_name) + beam_energy_at_element = simulator.energies[element_name] + except AttributeError: + try: + flat_segment = simulator.segment.flattened() + element = getattr(flat_segment, element_name) + beam_energy_at_element = simulator.energies_flattened[element_name] + except AttributeError: + raise ValueError( + f"Element '{element_name}' not found in simulator.segment" + ) return access_cheetah_attribute(element, attribute, beam_energy_at_element) def set_cheetah_property(self, simulator, control_variable_name, value): @@ -95,9 +103,18 @@ def set_cheetah_property(self, simulator, control_variable_name, value): raise ValueError( f"No mapping found for control variable '{control_variable_name}'" ) - - element = getattr(simulator.segment, element_name) - beam_energy_at_element = simulator.energies[element_name] + try: + element = getattr(simulator.segment, element_name) + beam_energy_at_element = simulator.energies[element_name] + except AttributeError: + try: + flat_segment = simulator.segment.flattened() + element = getattr(flat_segment, element_name) + beam_energy_at_element = simulator.energies_flattened[element_name] + except AttributeError: + raise ValueError( + f"Element '{element_name}' not found in simulator.segment" + ) access_cheetah_attribute( element, attribute, beam_energy_at_element, set_value=value ) diff --git a/virtual_accelerator/cheetah/utils.py b/virtual_accelerator/cheetah/utils.py index 5f38ad3..5e370b3 100644 --- a/virtual_accelerator/cheetah/utils.py +++ b/virtual_accelerator/cheetah/utils.py @@ -2,6 +2,7 @@ import torch import os from pathlib import Path +from cheetah.accelerator.superimposed import SuperimposedElement class NoSetMethodError(Exception): @@ -92,10 +93,17 @@ def get_magnetic_rigidity(energy): "BCON": FieldAccessor(lambda e, energy: 1.0), "BDES": FieldAccessor(lambda e, energy: e.angle * get_magnetic_rigidity(energy)), } +# check set_cheetah_value works then update othe setattrs TRANSVERSE_DEFLECTING_CAVITY_MAPPING = { - "AREQ": "voltage", - "PREQ": "phase", + "AREQ": FieldAccessor( + lambda e, energy: e.voltage / 1e6, + lambda e, energy, v: set_cheetah_value(e, "voltage", v * 1e6), + ), + "PREQ": FieldAccessor( + lambda e, energy: e.phase * (360 / (2 * torch.pi)), + lambda e, energy, p: setattr(e, "phase", p * (2 * torch.pi) / 360), + ), "AFBENB": FieldAccessor(lambda e, energy: 0.0), "AFBST": FieldAccessor(lambda e, energy: 0.0), "AMPL_W0CH0": FieldAccessor(lambda e, energy: 0.0), @@ -107,21 +115,37 @@ def get_magnetic_rigidity(energy): } BPM_MAPPING = { - "X": FieldAccessor(lambda e, energy: e.reading[0]), - "Y": FieldAccessor(lambda e, energy: e.reading[1]), - "XSCDT1H": FieldAccessor(lambda e, energy: e.reading[0]), - "YSCDT1H": FieldAccessor(lambda e, energy: e.reading[1]), + "X": FieldAccessor( + lambda e, energy: e.reading[..., 0] * 1000 + ), # convert from m to mm + "Y": FieldAccessor( + lambda e, energy: e.reading[..., 1] * 1000 + ), # convert from m to mm + "XSCDT1H": FieldAccessor( + lambda e, energy: e.reading[..., 0] * 1000 + ), # convert from m to mm + "YSCDT1H": FieldAccessor( + lambda e, energy: e.reading[..., 1] * 1000 + ), # convert from m to mm + "XSCDTH": FieldAccessor( + lambda e, energy: e.reading[..., 0] * 1000 + ), # convert from m to mm + "YSCDTH": FieldAccessor( + lambda e, energy: e.reading[..., 1] * 1000 + ), # convert from m to mm "TMIT": FieldAccessor(lambda e, energy: 1.0), } # multiply image intensity by 16 bit number range (is similar to real machine?) SCREEN_MAPPING = { - "Image:ArrayData": FieldAccessor(lambda e, energy: e.reading.T * 65535), + "Image:ArrayData": FieldAccessor( + lambda e, energy: e.reading.transpose(-2, -1) * 65535 + ), "PNEUMATIC": "is_active", "Image:ArraySize1_RBV": FieldAccessor(lambda e, energy: e.resolution[0]), "Image:ArraySize0_RBV": FieldAccessor(lambda e, energy: e.resolution[1]), "RESOLUTION": FieldAccessor(lambda e, energy: e.pixel_size[0] * 1e6), - "IMAGE": FieldAccessor(lambda e, energy: e.reading.T * 65535), + "IMAGE": FieldAccessor(lambda e, energy: e.reading.transpose(-2, -1) * 65535), "N_OF_ROW": FieldAccessor(lambda e, energy: e.resolution[0]), "N_OF_COL": FieldAccessor(lambda e, energy: e.resolution[1]), } @@ -135,6 +159,7 @@ def get_magnetic_rigidity(energy): "BPM": BPM_MAPPING, "Screen": SCREEN_MAPPING, "TransverseDeflectingCavity": TRANSVERSE_DEFLECTING_CAVITY_MAPPING, + "Patch": {}, } LCLS_ELEMENTS = os.path.join( @@ -143,7 +168,7 @@ def get_magnetic_rigidity(energy): ) -def handle_quadrupole_composite(elements, pv_attribute, energy, set_value): +def handle_quadrupole_composite(element, pv_attribute, energy, set_value): """ Handle composite quadrupole devices split into multiple subelements. @@ -177,23 +202,23 @@ def handle_quadrupole_composite(elements, pv_attribute, energy, set_value): torch.Tensor or float or None: Attribute value if getting, otherwise None. """ - total_length = sum(e.length for e in elements) + length = element.base_element.length + # if length is vectorized need tests for broadcasting set_values/ length - if set_value is not None and pv_attribute in {"BCTRL", "BDES"}: - new_k1 = set_value / get_magnetic_rigidity(energy) / total_length - - for e in elements: - e.k1 = new_k1 - return + if set_value is not None and pv_attribute in {"BCTRL", "BACT", "BDES"}: + new_k1 = set_value / get_magnetic_rigidity(energy) / length + element.base_element.k1 = new_k1 if pv_attribute in {"BCTRL", "BACT", "BDES"}: - return elements[0].k1 * total_length * get_magnetic_rigidity(energy) + return element.base_element.k1 * length * get_magnetic_rigidity(energy) # fallback behavior - return default_composite_handler(elements, pv_attribute, energy, set_value) + return access_cheetah_attribute( + element.base_element, pv_attribute, energy, set_value + ) -def default_composite_handler(elements, pv_attribute, energy, set_value): +def default_composite_handler(element, pv_attribute, energy, set_value): """ Default handler for composite devices. Assumes all subelements share identical attributes. @@ -221,18 +246,18 @@ def default_composite_handler(elements, pv_attribute, energy, set_value): torch.Tensor or float or None: Attribute value if getting, otherwise None. """ - + e = element.base_element if set_value is not None: - for e in elements: - access_cheetah_attribute(e, pv_attribute, energy, set_value) + access_cheetah_attribute(e, pv_attribute, energy, set_value) return - return access_cheetah_attribute(elements[0], pv_attribute, energy) + return access_cheetah_attribute(e, pv_attribute, energy) COMPOSITE_HANDLERS = { "Quadrupole": handle_quadrupole_composite, "TransverseDeflectingCavity": default_composite_handler, + "Patch": default_composite_handler, } @@ -263,16 +288,21 @@ def access_cheetah_attribute(element, pv_attribute, energy, set_value=None): value: The corresponding Cheetah attribute value if `set_value` is None, otherwise sets the value and returns None. """ - # implementing fix for quads, will rethink for tcavs - # simplest case, each subelement has sub_length = length/len(sub_elements) - # handling composite elements - - if isinstance(element, list): - if len(element) == 0: + # need to think about writing/reading to and from nn.Parameters, + # if var 'TRAINABLE:PV:300' is set to nn.Parameter(torch.tensor(1.0)) + # if needs to make the corresponding element attribute a nn.Parameter as well, or if it can just set the data of the existing attribute + if isinstance(element, SuperimposedElement): + if len(element.superimposed_element.elements) == 0: raise ValueError("Cannot access attribute on empty element list") - element_type = type(element[0]).__name__ - if any(type(sub).__name__ != element_type for sub in element): - raise ValueError("All subelements in element list must have same type") + + element_type = type(element.base_element).__name__ + if any( + type(sub).__name__ not in MAPPINGS + for sub in element.superimposed_element.elements + ): + raise ValueError( + "All subelements in element list must have a supported element type" + ) handler = COMPOSITE_HANDLERS.get(element_type, default_composite_handler) return handler(element, pv_attribute, energy, set_value) @@ -315,6 +345,42 @@ def access_cheetah_attribute(element, pv_attribute, energy, set_value=None): ) from e +def set_cheetah_value(element, attr_name, value): + """ + Set a Cheetah element attribute safely. + + If the existing attribute is an nn.Parameter, update its value in-place + without replacing the Parameter object. + + If the existing attribute is a Tensor, update in-place when shape-compatible. + + Otherwise, fall back to setattr. + """ + existing = getattr(element, attr_name) + # value = nn.Parameter ( ) * tensor is not a param. + if isinstance(existing, torch.nn.Parameter): + print("is nn should update in place.") + value = torch.as_tensor( + value, + dtype=existing.dtype, + device=existing.device, + ) + + with torch.no_grad(): + existing.copy_(value) + + return + + if isinstance(existing, torch.Tensor): + value = torch.as_tensor( + value, + dtype=existing.dtype, + device=existing.device, + ) + + setattr(element, attr_name, value) + + def get_mad_control_mapping(fname: str | None = None): """ Create a mapping from madnames to control names and device types diff --git a/virtual_accelerator/cheetah/variables.py b/virtual_accelerator/cheetah/variables.py index 0201224..929385b 100644 --- a/virtual_accelerator/cheetah/variables.py +++ b/virtual_accelerator/cheetah/variables.py @@ -46,7 +46,7 @@ def get_variables_from_segment( See `get_variables_from_element_name` for details on the specification of `element_attr_mapping`. """ - from cheetah.accelerator import Screen + from cheetah.accelerator import Screen, SuperimposedElement all_variables = {} element_attr_mapping = element_attr_mapping or get_element_attr_mapping() @@ -60,9 +60,36 @@ def get_variables_from_segment( warnings.warn(f"Element {element.name} not found in device mapping") continue - element_variables = get_variables_from_element_name( - type(element).__name__, control_name, element_attr_mapping - ) + if isinstance(element, SuperimposedElement): + element_variables = get_variables_from_element_name( + type(element.base_element).__name__, + control_name, + element_attr_mapping, + ) + + # iterate through the superimposed elements and get variables for each + for sub_element in element.superimposed_element.elements: + if type(sub_element).__name__ in ["Drift", "Marker", "Cavity"]: + continue + elif sub_element.name.upper() in device_mapping: + sub_control_name = device_mapping[sub_element.name.upper()] + else: + warnings.warn( + f"Element {sub_element.name} not found in device mapping" + ) + continue + + sub_element_variables = get_variables_from_element_name( + type(sub_element).__name__, + sub_control_name, + element_attr_mapping, + ) + element_variables.update(sub_element_variables) + + else: + element_variables = get_variables_from_element_name( + type(element).__name__, control_name, element_attr_mapping + ) # if element type is a screen then modify the output variable if isinstance(element, Screen): diff --git a/virtual_accelerator/models/sc_diag0.py b/virtual_accelerator/models/sc_diag0.py index 9a5d989..0c5530d 100644 --- a/virtual_accelerator/models/sc_diag0.py +++ b/virtual_accelerator/models/sc_diag0.py @@ -1,5 +1,4 @@ import os - from virtual_accelerator.utils.optional_dependencies import import_optional from virtual_accelerator.utils.variables import ( get_epics_to_name_or_overlay_mapping, @@ -37,10 +36,10 @@ def get_sc_diag0_cheetah_model(): ) from lume_cheetah import LUMECheetahModel, CheetahSimulator - from cheetah.accelerator import Segment from cheetah.particles import ParticleBeam from virtual_accelerator.cheetah.transformer import SLACCheetahTransformer from virtual_accelerator.cheetah.variables import get_variables_from_segment + from virtual_accelerator.cheetah.diag0 import get_diag0_beamline import torch incoming_beam = ParticleBeam.from_twiss( @@ -53,14 +52,10 @@ def get_sc_diag0_cheetah_model(): energy=torch.tensor(90e6), ) incoming_beam.particle_charges = torch.tensor(1.0) - - # Get path to lattice files lcls_lattice = os.environ.get("LCLS_LATTICE") # Create lattice from file - segment = Segment.from_lattice_json( - os.path.join(lcls_lattice, "cheetah/sc_diag0.json") - ) + segment = get_diag0_beamline() # Ensure screen elements can support vectorization for element in segment.elements: diff --git a/virtual_accelerator/tests/conftest.py b/virtual_accelerator/tests/conftest.py new file mode 100644 index 0000000..7051553 --- /dev/null +++ b/virtual_accelerator/tests/conftest.py @@ -0,0 +1,50 @@ +import os +import sys + +import pytest + +sys.path.insert( + 0, + os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")), +) + +try: + from virtual_accelerator.models.cu_hxr import get_cu_hxr_cheetah_model +except ModuleNotFoundError: + get_cu_hxr_cheetah_model = None + +try: + from virtual_accelerator.models.sc_diag0 import get_sc_diag0_cheetah_model +except ModuleNotFoundError: + get_sc_diag0_cheetah_model = None + +CHEETAH_SUBMODULES = ["diag0", "transformer", "utils", "variables"] +CHEETAH_MODEL_FACTORIES = [] +if get_cu_hxr_cheetah_model is not None: + CHEETAH_MODEL_FACTORIES.append(("cu_hxr", get_cu_hxr_cheetah_model)) +if get_sc_diag0_cheetah_model is not None: + CHEETAH_MODEL_FACTORIES.append(("sc_diag0", get_sc_diag0_cheetah_model)) + + +def pytest_configure(config): + config.addinivalue_line( + "markers", + "for_every_cheetah_module: run the test for every virtual_accelerator.cheetah submodule", + ) + config.addinivalue_line( + "markers", + "for_every_cheetah_model: run the test for every Cheetah-based virtual accelerator model", + ) + + +def pytest_generate_tests(metafunc): + if metafunc.definition.get_closest_marker("for_every_cheetah_module"): + metafunc.parametrize("module_name", CHEETAH_SUBMODULES) + if metafunc.definition.get_closest_marker("for_every_cheetah_model"): + if not CHEETAH_MODEL_FACTORIES: + pytest.skip("No Cheetah model factories available in this environment") + metafunc.parametrize( + "model_name,model_factory", + CHEETAH_MODEL_FACTORIES, + ids=[name for name, _ in CHEETAH_MODEL_FACTORIES], + ) diff --git a/virtual_accelerator/tests/test_cheetah_imports.py b/virtual_accelerator/tests/test_cheetah_imports.py new file mode 100644 index 0000000..68f14c8 --- /dev/null +++ b/virtual_accelerator/tests/test_cheetah_imports.py @@ -0,0 +1,36 @@ +import importlib +import pkgutil +from pathlib import Path + +import pytest + + +CHEETAH_SUBMODULES = ["diag0", "transformer", "utils", "variables"] + + +def test_virtual_accelerator_package_import(): + package = importlib.import_module("virtual_accelerator") + assert package is not None + assert hasattr(package, "__file__") + assert Path(package.__file__).exists() + + +def test_virtual_accelerator_cheetah_package_import(): + cheetah_pkg = importlib.import_module("virtual_accelerator.cheetah") + assert cheetah_pkg is not None + assert hasattr(cheetah_pkg, "__path__") + assert any(Path(path).exists() for path in cheetah_pkg.__path__) + + +def test_cheetah_submodules_present(): + cheetah_pkg = importlib.import_module("virtual_accelerator.cheetah") + available = {module.name for module in pkgutil.iter_modules(cheetah_pkg.__path__)} + assert set(CHEETAH_SUBMODULES).issubset(available) + + +@pytest.mark.for_every_cheetah_module +def test_cheetah_submodule_import(module_name): + module = importlib.import_module(f"virtual_accelerator.cheetah.{module_name}") + assert module.__name__.endswith(f".{module_name}") + assert hasattr(module, "__file__") + assert Path(module.__file__).exists() diff --git a/virtual_accelerator/tests/test_cheetah_models.py b/virtual_accelerator/tests/test_cheetah_models.py new file mode 100644 index 0000000..c6b7eb3 --- /dev/null +++ b/virtual_accelerator/tests/test_cheetah_models.py @@ -0,0 +1,93 @@ +import pytest +from numbers import Number + +try: + import torch + from lume_cheetah import LUMECheetahModel +except ModuleNotFoundError as exc: + pytest.skip( + f"Skipping Cheetah model tests because dependency is missing: {exc.name}", + allow_module_level=True, + ) + + +@pytest.fixture +def cheetah_model(model_name, model_factory): + model = model_factory() + assert isinstance(model, LUMECheetahModel) + return model + + +class TestCheetahModelBasics: + @pytest.mark.for_every_cheetah_model + def test_has_variables(self, cheetah_model): + assert cheetah_model.control_variables + assert cheetah_model.observable_variables + assert cheetah_model.supported_variables + + for name in cheetah_model.control_variables: + assert name in cheetah_model.supported_variables + + for name in cheetah_model.observable_variables: + assert name in cheetah_model.supported_variables + + @pytest.mark.for_every_cheetah_model + def test_get_observable_variable(self, cheetah_model): + observable_name = next(iter(cheetah_model.observable_variables)) + output = cheetah_model.get([observable_name]) + + assert observable_name in output + assert output[observable_name] is not None + + value = output[observable_name] + if hasattr(value, "shape"): + assert value.shape is not None + + @pytest.mark.for_every_cheetah_model + def test_set_and_read_control_variable(self, cheetah_model): + control_name = next(iter(cheetah_model.control_variables)) + current_value = cheetah_model.get([control_name])[control_name] + + if isinstance(current_value, torch.Tensor): + current_value = float(current_value.item()) + elif isinstance(current_value, Number): + current_value = float(current_value) + else: + pytest.skip( + f"Skipping control variable {control_name} because its type is not numeric: {type(current_value)}" + ) + + new_value = ( + current_value + 1.0 if abs(current_value) < 1e6 else current_value * 0.9 + ) + cheetah_model.set({control_name: new_value}) + + updated_value = cheetah_model.get([control_name])[control_name] + if isinstance(updated_value, torch.Tensor): + updated_value = float(updated_value.item()) + + assert isinstance(updated_value, Number) + assert abs(updated_value - new_value) < 1e-6 + + @pytest.mark.for_every_cheetah_model + def test_reset_restores_initial_state(self, cheetah_model): + control_name = next(iter(cheetah_model.control_variables)) + original_value = cheetah_model.get([control_name])[control_name] + + if isinstance(original_value, torch.Tensor): + original_value = float(original_value.item()) + elif isinstance(original_value, Number): + original_value = float(original_value) + else: + pytest.skip( + f"Skipping reset test for control variable {control_name} because its type is not numeric: {type(original_value)}" + ) + + cheetah_model.set({control_name: original_value + 1.0}) + cheetah_model.reset() + + reset_value = cheetah_model.get([control_name])[control_name] + if isinstance(reset_value, torch.Tensor): + reset_value = float(reset_value.item()) + + assert abs(reset_value - original_value) < 1e-6 diff --git a/virtual_accelerator/utils/slac_variable_config.yaml b/virtual_accelerator/utils/slac_variable_config.yaml index e55c0ed..69aea2c 100644 --- a/virtual_accelerator/utils/slac_variable_config.yaml +++ b/virtual_accelerator/utils/slac_variable_config.yaml @@ -8,6 +8,16 @@ BPM: unit: mm read_only: true variable_class: ScalarVariable + XSCDTH: + unit: mm + read_only: true + variable_class: ScalarVariable + default_value: NULL + YSCDTH: + unit: mm + read_only: true + variable_class: ScalarVariable + Quadrupole: BCTRL: diff --git a/virtual_accelerator/utils/variables.py b/virtual_accelerator/utils/variables.py index c72623f..198c8be 100644 --- a/virtual_accelerator/utils/variables.py +++ b/virtual_accelerator/utils/variables.py @@ -62,7 +62,7 @@ def get_name_or_overlay_to_epics_mapping( df = df[df["Beampath"].str.contains(beampath, na=False)] # remove rows with `keyword` = `USEG` and `LCAV` - df = df[~df["Keyword"].str.contains("USEG|LCAV|TCAV", na=False)] + df = df[~df["Keyword"].str.contains("USEG|LCAV", na=False)] name_data = df[["Element", "Control System Name"]].dropna() return dict(zip(name_data["Element"], name_data["Control System Name"]))