diff --git a/CHANGELOG.md b/CHANGELOG.md index ac3824bdd..3a72fdec9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,8 @@ ### 🐛 Bug fixes +- Fix Bmad lattice conversion issues with overlay and group definitions, wildcard element references, scientific notation in expressions, unary signs, and mapping of `type` and `alias` fields to element `metadata`. (see #663) (@roussel-ryan, @jank324) + ### 🐆 Other - Update openPMD dependency to renamed package (see #684) (@jank324) diff --git a/cheetah/converters/bmad.py b/cheetah/converters/bmad.py index 36d53bec0..6c4107fa8 100644 --- a/cheetah/converters/bmad.py +++ b/cheetah/converters/bmad.py @@ -40,6 +40,11 @@ def convert_element( "dtype": dtype or torch.get_default_dtype(), } bmad_parsed = context[name] + metadata = ( + {k: bmad_parsed[k] for k in ["alias", "type"] if k in bmad_parsed} + if isinstance(bmad_parsed, dict) + else {} + ) shared_properties = ["element_type", "alias", "type"] @@ -55,7 +60,9 @@ def convert_element( elif isinstance(bmad_parsed, dict) and "element_type" in bmad_parsed: if bmad_parsed["element_type"] == "marker": validate_understood_properties(shared_properties, bmad_parsed) - return cheetah.Marker(name=name, sanitize_name=sanitize_name) + return cheetah.Marker( + name=name, sanitize_name=sanitize_name, metadata=metadata + ) elif bmad_parsed["element_type"] == "monitor": validate_understood_properties(shared_properties + ["l"], bmad_parsed) if "l" in bmad_parsed: @@ -63,9 +70,12 @@ def convert_element( length=torch.tensor(bmad_parsed["l"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) else: - return cheetah.Marker(name=name, sanitize_name=sanitize_name) + return cheetah.Marker( + name=name, sanitize_name=sanitize_name, metadata=metadata + ) elif bmad_parsed["element_type"] == "instrument": validate_understood_properties(shared_properties + ["l"], bmad_parsed) if "l" in bmad_parsed: @@ -73,9 +83,12 @@ def convert_element( length=torch.tensor(bmad_parsed["l"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) else: - return cheetah.Marker(name=name, sanitize_name=sanitize_name) + return cheetah.Marker( + name=name, sanitize_name=sanitize_name, metadata=metadata + ) elif bmad_parsed["element_type"] == "pipe": validate_understood_properties( shared_properties + ["l", "descrip"], bmad_parsed @@ -84,6 +97,7 @@ def convert_element( length=torch.tensor(bmad_parsed["l"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "drift": validate_understood_properties( @@ -93,6 +107,7 @@ def convert_element( length=torch.tensor(bmad_parsed["l"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "hkicker": validate_understood_properties(shared_properties + ["kick"], bmad_parsed) @@ -101,6 +116,7 @@ def convert_element( angle=torch.tensor(bmad_parsed.get("kick", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "vkicker": validate_understood_properties(shared_properties + ["kick"], bmad_parsed) @@ -109,6 +125,7 @@ def convert_element( angle=torch.tensor(bmad_parsed.get("kick", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "sbend": validate_understood_properties( @@ -120,7 +137,7 @@ def convert_element( length=torch.tensor(bmad_parsed["l"], **factory_kwargs), gap=torch.tensor(2 * bmad_parsed.get("hgap", 0.0), **factory_kwargs), angle=torch.tensor(bmad_parsed.get("angle", 0.0), **factory_kwargs), - dipole_e1=torch.tensor(bmad_parsed["e1"], **factory_kwargs), + dipole_e1=torch.tensor(bmad_parsed.get("e1", 0.0), **factory_kwargs), dipole_e2=torch.tensor(bmad_parsed.get("e2", 0.0), **factory_kwargs), tilt=torch.tensor(bmad_parsed.get("ref_tilt", 0.0), **factory_kwargs), fringe_integral=torch.tensor( @@ -133,6 +150,7 @@ def convert_element( ), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "quadrupole": validate_understood_properties( @@ -144,6 +162,7 @@ def convert_element( tilt=torch.tensor(bmad_parsed.get("tilt", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "sextupole": validate_understood_properties( @@ -155,6 +174,7 @@ def convert_element( tilt=torch.tensor(bmad_parsed.get("tilt", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "solenoid": validate_understood_properties(shared_properties + ["l", "ks"], bmad_parsed) @@ -163,6 +183,7 @@ def convert_element( k=torch.tensor(bmad_parsed["ks"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "lcavity": validate_understood_properties( @@ -181,6 +202,21 @@ def convert_element( cavity_type=bmad_parsed["cavity_type"], name=name, sanitize_name=sanitize_name, + metadata=metadata, + ) + elif bmad_parsed["element_type"] == "crab_cavity": + validate_understood_properties( + shared_properties + ["l", "rf_frequency", "voltage", "phi0"], + bmad_parsed, + ) + return cheetah.TransverseDeflectingCavity( + length=torch.tensor(bmad_parsed["l"], **factory_kwargs), + voltage=torch.tensor(bmad_parsed.get("voltage", 0.0), **factory_kwargs), + phase=-(torch.tensor(bmad_parsed.get("phi0", 0.0), **factory_kwargs)), + frequency=torch.tensor(bmad_parsed["rf_frequency"], **factory_kwargs), + name=name, + sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "rcollimator": validate_understood_properties( @@ -210,6 +246,7 @@ def convert_element( ], name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "ecollimator": validate_understood_properties( @@ -239,6 +276,7 @@ def convert_element( ], name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "wiggler": validate_understood_properties( @@ -252,6 +290,7 @@ def convert_element( period=torch.tensor(bmad_parsed["l_period"], **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) elif bmad_parsed["element_type"] == "patch": # TODO: Does this need to be implemented in Cheetah in a more proper way? @@ -260,6 +299,7 @@ def convert_element( length=torch.tensor(bmad_parsed.get("l", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) else: warnings.warn( @@ -272,6 +312,7 @@ def convert_element( length=torch.tensor(bmad_parsed.get("l", 0.0), **factory_kwargs), name=name, sanitize_name=sanitize_name, + metadata=metadata, ) else: raise ValueError(f"Unknown Bmad element type for {name = }") # noqa: E202, E251 diff --git a/cheetah/converters/utils/fortran_namelist.py b/cheetah/converters/utils/fortran_namelist.py index ba3a187b0..bc8b4bbc6 100644 --- a/cheetah/converters/utils/fortran_namelist.py +++ b/cheetah/converters/utils/fortran_namelist.py @@ -28,13 +28,7 @@ ) LINE_DEFINITION_PATTERN = f"({ELEMENT_NAME_PATTERN})" + r"\s*\:\s*line\s*=\s*\((.*)\)" USE_LINE_PATTERN = r'use\s*\,\s*([a-z0-9_]+|"[a-z0-9_\-\.\:]+")' -OVERLAY_DEFINITION_PATTERN = ( - f"({ELEMENT_NAME_PATTERN})" r"\s*\:\s*overlay\s*=\s*\{(.*)\}\s*\,\s*var\s*=\s*" -) -OVERLAY_KNOT_BASED_PATTERN = ( - OVERLAY_DEFINITION_PATTERN + r"\{\s*([a-z0-9_]+)\s*\}\s*\,\s*x_knot\s*=\s*\{(.*)\}" -) -OVERLAY_EXPRESSION_BASED_PATTERN = OVERLAY_DEFINITION_PATTERN + r"\{(.*)\}\s*(\,.*)*" +CONTROL_DEFINITION_PATTERN = rf"({ELEMENT_NAME_PATTERN})\s*:\s*(?:overlay|group)\b.*" def read_clean_lines(lattice_file_path: Path) -> list[str]: @@ -143,6 +137,12 @@ def evaluate_expression(expression: str, context: dict) -> Any: except ValueError: pass + # Check against string literals enclosed in quotes + if (expression.startswith('"') and expression.endswith('"')) or ( + expression.startswith("'") and expression.endswith("'") + ): + return expression[1:-1] + # Check against allowed keywords if expression in ["open", "electron", "t", "f", "traveling_wave", "full"]: return expression @@ -168,28 +168,26 @@ def evaluate_expression(expression: str, context: dict) -> Any: return expression -def resolve_object_name_wildcard(wildcard_pattern: str, context: dict) -> list: +def resolve_object_name_wildcard(wildcard_pattern: str, context: dict) -> list[str]: """ - Return a list of object names that match the given wildcard pattern. - - :param wildcard_pattern: Wildcard pattern to match. - :param context: Dictionary of variables among which to search for matching object. - :return: List of object names that match the given wildcard pattern, both in terms - of name and element type. + Return a list of element names in context matching a name pattern and/or type + prefix. """ - object_type, object_name = wildcard_pattern.split("::") + if "::" in wildcard_pattern: + object_type, object_name = wildcard_pattern.split("::", maxsplit=1) + else: + object_type, object_name = None, wildcard_pattern pattern = object_name.replace("*", ".*").replace("%", ".") - name_matching_keys = [key for key in context.keys() if re.fullmatch(pattern, key)] - type_matching_keys = [ - key - for key in name_matching_keys - if isinstance(context[key], dict) - and "element_type" in context[key] - and context[key]["element_type"] == object_type - ] - return type_matching_keys + return [ + name + for name, element in context.items() + if isinstance(element, dict) + and "element_type" in element + and (object_type is None or element["element_type"] == object_type) + and re.fullmatch(pattern, name) + ] def assign_property(line: str, context: dict) -> dict: @@ -207,7 +205,7 @@ def assign_property(line: str, context: dict) -> dict: property_name = match.group(2).strip() property_expression = match.group(3).strip() # TODO: Evaluate expression first - if "*" in object_name or "%" in object_name: + if "*" in object_name or "%" in object_name or "::" in object_name: object_names = resolve_object_name_wildcard(object_name, context) else: object_names = [object_name] @@ -271,7 +269,7 @@ def define_element(line: str, context: dict) -> dict: for property_string in property_matches: property_string = property_string.strip() - property_name, property_expression = property_string.split("=") + property_name, property_expression = property_string.split("=", maxsplit=1) property_name = property_name.strip() property_expression = property_expression.strip() @@ -310,50 +308,6 @@ def define_line(line: str, context: dict) -> dict: return context -def define_overlay(line: str, context: dict) -> dict: - """ - Define an overlay in the context. - - :param line: Line of an overlay definition to be parsed. - :param context: Dictionary of variables to define the overlay in and from which to - read variables. - :return: Updated context. - """ - - expression_match = re.fullmatch(OVERLAY_EXPRESSION_BASED_PATTERN, line) - knot_match = re.fullmatch(OVERLAY_KNOT_BASED_PATTERN, line) - - if knot_match: - overlay_name = knot_match.group(1).strip() - overlay_definition = knot_match.group(2).strip() - overlay_variable = knot_match.group(3).strip() - overlay_x_knot = knot_match.group(4).strip() - - context[overlay_name] = { - "overlay_definition": overlay_definition, - "overlay_variable": overlay_variable, - "overlay_x_knot": overlay_x_knot, - } - elif expression_match: - overlay_name = expression_match.group(1).strip() - overlay_definition = expression_match.group(2).strip() - overlay_variables = expression_match.group(3).strip() - if expression_match.group(4) is not None: - overlay_parameters = expression_match.group(4).strip()[1:].strip() - else: - overlay_parameters = None - - context[overlay_name] = { - "overlay_definition": overlay_definition, - "overlay_variables": overlay_variables, - "overlay_parameters": overlay_parameters, - } - else: - raise ValueError(f"Overlay definition {line} not understood.") - - return context - - def parse_use_line(line: str, context: dict) -> dict: """ Parse a use line. @@ -410,8 +364,9 @@ def parse_lines(lines: str) -> dict: context = assign_variable(line, context) elif re.fullmatch(LINE_DEFINITION_PATTERN, line): context = define_line(line, context) - elif re.fullmatch(OVERLAY_DEFINITION_PATTERN, line): - context = define_overlay(line, context) + elif re.fullmatch(CONTROL_DEFINITION_PATTERN, line): + # Overlay and group definitions are control entries; skip for simplicity. + continue elif re.fullmatch(ELEMENT_DEFINITION_PATTERN, line): context = define_element(line, context) elif re.fullmatch(USE_LINE_PATTERN, line): diff --git a/cheetah/converters/utils/infix.py b/cheetah/converters/utils/infix.py index 34b86e8a0..38700574b 100644 --- a/cheetah/converters/utils/infix.py +++ b/cheetah/converters/utils/infix.py @@ -6,16 +6,18 @@ "-": {"precedence": 1, "inputs": 2, "func": lambda a, b: a - b}, "*": {"precedence": 2, "inputs": 2, "func": lambda a, b: a * b}, "/": {"precedence": 2, "inputs": 2, "func": lambda a, b: a / b}, - "^": {"precedence": 3, "inputs": 2, "func": lambda a, b: a**b}, - "sqrt": {"precedence": 4, "inputs": 1, "func": lambda a: math.sqrt(a)}, - "sin": {"precedence": 4, "inputs": 1, "func": lambda a: math.sin(a)}, - "asin": {"precedence": 4, "inputs": 1, "func": lambda a: math.asin(a)}, - "cos": {"precedence": 4, "inputs": 1, "func": lambda a: math.cos(a)}, - "acos": {"precedence": 4, "inputs": 1, "func": lambda a: math.acos(a)}, - "tan": {"precedence": 4, "inputs": 1, "func": lambda a: math.tan(a)}, - "atan": {"precedence": 4, "inputs": 1, "func": lambda a: math.atan(a)}, - "abs": {"precedence": 4, "inputs": 1, "func": lambda a: abs(a)}, - "log": {"precedence": 4, "inputs": 1, "func": lambda a: math.log(a)}, + "u+": {"precedence": 3, "inputs": 1, "func": lambda a: +a}, + "u-": {"precedence": 3, "inputs": 1, "func": lambda a: -a}, + "^": {"precedence": 4, "inputs": 2, "func": lambda a, b: a**b}, + "sqrt": {"precedence": 5, "inputs": 1, "func": lambda a: math.sqrt(a)}, + "sin": {"precedence": 5, "inputs": 1, "func": lambda a: math.sin(a)}, + "asin": {"precedence": 5, "inputs": 1, "func": lambda a: math.asin(a)}, + "cos": {"precedence": 5, "inputs": 1, "func": lambda a: math.cos(a)}, + "acos": {"precedence": 5, "inputs": 1, "func": lambda a: math.acos(a)}, + "tan": {"precedence": 5, "inputs": 1, "func": lambda a: math.tan(a)}, + "atan": {"precedence": 5, "inputs": 1, "func": lambda a: math.atan(a)}, + "abs": {"precedence": 5, "inputs": 1, "func": lambda a: abs(a)}, + "log": {"precedence": 5, "inputs": 1, "func": lambda a: math.log(a)}, } @@ -76,6 +78,7 @@ def _parse_expression(tokens: list[str]) -> dict: output = None stack = [] operator_stack = [] + expecting_value = True while len(tokens) > 0: token = tokens.pop(0) @@ -95,11 +98,19 @@ def _parse_expression(tokens: list[str]) -> dict: raise SyntaxError("Mismatched parentheses in expression") nested_ast = _parse_expression(nested_tokens) stack.append(nested_ast) + expecting_value = False elif token not in operators: stack.append({"value": float(token), "left": None, "right": None}) + expecting_value = False else: + if token in {"+", "-"} and expecting_value: + token = "u+" if token == "+" else "u-" + elif expecting_value and operators[token]["inputs"] == 2: + raise SyntaxError(f"Unexpected binary operator '{token}'") + while ( operator_stack + and not expecting_value and operators[operator_stack[-1]]["precedence"] >= operators[token]["precedence"] ): @@ -112,19 +123,13 @@ def _parse_expression(tokens: list[str]) -> dict: output = {"value": operator, "left": left, "right": right} stack.append(output) operator_stack.append(token) + expecting_value = True while operator_stack: - if operators[operator_stack[-1]]["inputs"] == 1 or ( - operator_stack[-1] == "-" and len(stack) == 1 - ): - # Handle unary functions or unary minus + if operators[operator_stack[-1]]["inputs"] == 1: + # Handle unary functions and sign operators. right = None - if operator_stack[-1] == "-" and len(stack) == 1: - # Unary minus, treat as 0 - x - left = {"value": 0, "left": None, "right": None} - right = stack.pop() - else: - left = stack.pop() + left = stack.pop() else: right = stack.pop() left = stack.pop() @@ -150,6 +155,21 @@ def _tokenize_expression(expression: str, context: dict) -> list[str]: for char in expression: if char.isspace() or char in "+-*/^()[]": + # Keep scientific-notation exponents in one token (e.g. 0.750e-3) + if ( + char in "+-" + and current_key is None + and current_token + and current_token[-1] in "eE" + ): + mantissa = current_token[:-1] + try: + float(mantissa) + current_token += char + continue + except ValueError: + pass + if current_token: if char == "]" and current_key is not None: # This will throw an index error if current key is invalid diff --git a/tests/resources/lcls/cu_hxr.lat.bmad b/tests/resources/lcls/cu_hxr.lat.bmad new file mode 100644 index 000000000..b5189ff1b --- /dev/null +++ b/tests/resources/lcls/cu_hxr.lat.bmad @@ -0,0 +1,111 @@ +! +! Reduced CU_HXR-like lattice fixture for tests. +! Main lattice file with companion overlay/fixer include. +! + +! ------------------------------------------------------------------------------ +! Global parameters and beginning conditions +! ------------------------------------------------------------------------------ + +parameter[custom_attribute1] = taylor::mat_und_k +parameter[custom_attribute2] = taylor::mat_und_l +parameter[geometry] = open +parameter[particle] = electron + +beginning[beta_a] = 1.29704949763868189E+001 +beginning[alpha_a] = -4.39664934865974200E+000 +beginning[beta_b] = 3.47990377067604273E-001 +beginning[alpha_b] = 2.93951797746668075E-001 +beginning[e_tot] = 6e6 +beginning[theta_position] = -35*pi/180 +beginning[z_position] = 3050.512000 - 1032.60052 +beginning[x_position] = 10.448934545 + +setsp = 0 +setcus = 0 +setda = 0 +sethxrss = 0 +setsxrss = 0 +setpepx = 0 +setcbxfel = 0 +setxl2wig = 0 +setxl2ss = 0 + +intgsx = 30.0 +intghx = 30.0 +cb = 1.0e10/c_light +e00 = 0.006 +in2m = 0.0254 + +! ------------------------------------------------------------------------------ +! Representative element inventory +! ------------------------------------------------------------------------------ + +dbmark80: marker +beggunb: marker +l0bbeg: marker +otr2: marker +yag03: marker +ws12: marker +ws24: marker +ws28144: marker +ws32: marker +ws32b: marker +bpm4: marker, superimpose, ref=qa01 +bpm5: marker, superimpose, ref=qa02 +bpm50: marker, superimpose, ref=q50q3 + +bq1 = 10 +brho = 10 + +qa01: quadrupole, l = 0.122, k1 = 10.0 +qa02: quadrupole, l = 0.122, k1 = 13.0 +qe01: quadrupole, l = 0.150, k1 = +bq1 +qe02: quadrupole, l = 0.150, k1 = -(10*bq1)/brho + +l0a: lcavity, l = 3.0, voltage = 4.0e7, phi0 = 0.0, rf_frequency = 2.856e9 +l0b: lcavity, l = 3.0, voltage = 8.0e7, phi0 = 0.0, rf_frequency = 2.856e9 + +d11o: drift, l = 2.4349 +d11oa: drift, l = 2.4349 +d_mid: drift, l = 0.2500 + +bx11: sbend, l = 0.2032, angle = -0.093575352547, e2 = -0.093575352547, hgap = 0.01, fint = 0.5 +bx12: sbend, l = 0.2032, angle = 0.093575352547, e1 = 0.093575352547, hgap = 0.01, fint = 0.5 +bx13: sbend, l = 0.2032, angle = 0.093575352547, e2 = 0.093575352547, hgap = 0.01, fint = 0.5 +bx14: sbend, l = 0.2032, angle = -0.093575352547, e1 = -0.093575352547, hgap = 0.01, fint = 0.5 + +tcxdg0: crab_cavity, type = "STCAV_X", rf_frequency = 2856 * 1e6, l = 20*in2m/2 + +! ------------------------------------------------------------------------------ +! Characteristic property assignments and aliases +! ------------------------------------------------------------------------------ + +qa01[alias] = "QUAD:IN20:121" +qa02[alias] = "QUAD:IN20:122" +l0a[alias] = "ACCL:IN20:300" +l0b[alias] = "ACCL:IN20:400" +yag03[alias] = "YAGS:IN20:131" + +lcavity::*[field_autoscale] = 1.0 +lcavity::*[cavity_type] = traveling_wave +sbend::*[fringe_type] = full + +qa01[k1] = 0.384840836193 + +l0a[voltage] = 4.0e7 +l0b[voltage] = 8.0e7 +lcavity::*[n_rf_steps] = 1000 +l0*[phi0] = 10 + +call, file = cu_hxr_overlays_and_fixers.bmad + +! ------------------------------------------------------------------------------ +! Line hierarchy +! ------------------------------------------------------------------------------ + +gunl0a: line = (beggunb, l0a, l0b, qa01, d_mid, qa02) +bc1: line = (bx11, d11o, bx12, d11oa, bx13, d11o, bx14) +cu_hxr: line = (gunl0a, bc1, yag03, tcxdg0) + +use, cu_hxr diff --git a/tests/resources/lcls/cu_hxr_overlays_and_fixers.bmad b/tests/resources/lcls/cu_hxr_overlays_and_fixers.bmad new file mode 100644 index 000000000..029fe9026 --- /dev/null +++ b/tests/resources/lcls/cu_hxr_overlays_and_fixers.bmad @@ -0,0 +1,44 @@ +! +! Companion include for CU_HXR test lattice. +! Contains representative overlay and fixer definitions. +! + +! ------------------------------------------------------------------------------ +! Minimal overlay examples +! ------------------------------------------------------------------------------ + +bc1_theta_default = -0.093575352547 +bc1_lp_default = 0.2032 +bc1_lp_drift_default = 2.4349 + +o_bx11: overlay = { + bx11[g]:sin(theta)/lp, + bx11[l]:lp*theta/sin(theta), + bx11[e2]:theta}, + var = {lp, theta}, + theta = bc1_theta_default, + lp = bc1_lp_default + +o_bc1: overlay = { + o_bx11[theta]:angle_deg*pi/180, + d11o[l]:d11o[l] + lp_drift*(1/cos(angle_deg*pi/180)-1/cos(bc1_theta_default)), + d11oa[l]:d11oa[l] + lp_drift*(1/cos(angle_deg*pi/180)-1/cos(bc1_theta_default))}, + var = {angle_deg, lp_drift}, + angle_deg = bc1_theta_default*180/pi, + lp_drift = bc1_lp_drift_default + +o_bc1_offset: overlay = {o_bc1[angle_deg]: {-5.0, 0.0, 5.0}}, + var = {offset}, + x_knot = {-0.2, 0.0, 0.2} +o_bc1_offset[offset] = 0.0 + +o_quad_fudge: overlay = {qa01[k1]:k1_scale * qa01[k1], qa02[k1]:k1_scale * qa02[k1]}, + var = {k1_scale}, + k1_scale = 1.0 + +! ------------------------------------------------------------------------------ +! Representative superimpose definitions +! ------------------------------------------------------------------------------ + +otr2: marker, superimpose, ref=qa02 + diff --git a/tests/test_bmad_conversion.py b/tests/test_bmad_conversion.py index 05abbdba3..db8ffc7ed 100644 --- a/tests/test_bmad_conversion.py +++ b/tests/test_bmad_conversion.py @@ -136,3 +136,26 @@ def test_default_dtype(default_torch_dtype): assert converted.q.k1.dtype == default_torch_dtype assert converted.s.length.dtype == default_torch_dtype assert converted.s.k2.dtype == default_torch_dtype + + +def test_cu_hxr_lcls_fixture_conversion(): + """Test converting the reduced split CU_HXR fixture into Cheetah.""" + file_path = "tests/resources/lcls/cu_hxr.lat.bmad" + + converted = cheetah.Segment.from_bmad(file_path, dtype=torch.float64) + flattened = converted.flattened() + + assert isinstance(converted, cheetah.Segment) + assert converted.name == "cu_hxr" + assert isinstance(flattened.bx11, cheetah.Dipole) + + assert flattened.qa01.k1.item() == pytest.approx(0.384840836193) + assert flattened.qa01.metadata["alias"] == "quad:in20:121" + + assert flattened.l0a.phase.item() == pytest.approx(-3600.0) + assert flattened.l0b.phase.item() == pytest.approx(-3600.0) + + assert isinstance(flattened.tcxdg0, cheetah.TransverseDeflectingCavity) + assert flattened.tcxdg0.metadata["type"] == "stcav_x" + assert flattened.tcxdg0.frequency.item() == pytest.approx(2.856e9) + assert flattened.tcxdg0.length.item() == pytest.approx(0.254) diff --git a/tests/test_fortran_namelist.py b/tests/test_fortran_namelist.py new file mode 100644 index 000000000..19d0f5307 --- /dev/null +++ b/tests/test_fortran_namelist.py @@ -0,0 +1,77 @@ +import pytest + +from cheetah.converters.utils.fortran_namelist import evaluate_expression, parse_lines + + +def test_evaluate_expression_with_context(): + """ + Test evaluating expressions with context variables, scientific notation, and unary + signs in the Fortran namelist parser. + """ + context = {"mc2": 0.511750} + + assert evaluate_expression("mc2+0.750e-3", context) == pytest.approx( + context["mc2"] + 0.750e-3 + ) + assert evaluate_expression("+mc2", context) == pytest.approx(context["mc2"]) + assert evaluate_expression("-mc2", context) == pytest.approx(-context["mc2"]) + + +def test_evaluate_quoted_strings(): + """ + Test evaluating quoted string literals in the Fortran namelist parser. + """ + assert evaluate_expression('"test_string"', {}) == "test_string" + assert evaluate_expression("'single_quoted'", {}) == "single_quoted" + + +def test_define_element_string_attributes(): + """ + Test that element string attributes such as alias and type are correctly parsed + and stored in the element dictionary without spurious quotes. + """ + lines = [ + 'q1: quadrupole, l = 0.2, alias = "q1_alias", type = "control_label", k1 = 1.0' + ] + + context = parse_lines(lines) + + q1 = context["q1"] + + assert q1["alias"] == "q1_alias" + assert q1["type"] == "control_label" + assert q1["k1"] == pytest.approx(1.0) + + +def test_typed_property_assignment(): + """ + Test that typed property assignments targeting specific elements (e.g. + `lcavity::l0a[voltage] = ...`) resolve and update the correct element in context. + """ + lines = [ + "l0a: lcavity, l = 3.0, voltage = 0.0, rf_frequency = 2.856e9", + "lcavity::l0a[voltage] = 4.0e7", + ] + + context = parse_lines(lines) + + assert "lcavity::l0a" not in context + assert context["l0a"]["voltage"] == pytest.approx(4.0e7) + + +def test_skip_control_definitions(): + """ + Test that overlay and group control definitions are skipped without causing + parsing errors or creating rogue elements. + """ + lines = [ + "q1: quadrupole, l = 0.2, k1 = 1.0", + "o_q1: overlay = {q1[k1]: scale * q1[k1]}, var = {scale}, scale = 1.0", + "g_all: group = {q1[k1]: 2.0}, var = {k1}", + ] + + context = parse_lines(lines) + + assert "q1" in context + assert "o_q1" not in context + assert "g_all" not in context diff --git a/tests/test_infix.py b/tests/test_infix.py index 80a88ea9f..97574aff6 100644 --- a/tests/test_infix.py +++ b/tests/test_infix.py @@ -1,3 +1,5 @@ +import math + import pytest from cheetah.converters.utils import infix @@ -95,6 +97,40 @@ def test_infix_expression_with_function_call(): assert infix.evaluate_expression(expression) == 8 +def test_infix_expression_with_unary_minus_before_function_call(): + """ + Test that an infix expression with a unary minus directly in front of a function + call is correctly evaluated. + """ + expression = "-sin(argw)*sqrt(kqwig)" + context = {"argw": math.pi / 2.0, "kqwig": 4.0} + + assert infix.evaluate_expression(expression, context) == pytest.approx(-2.0) + + +def test_infix_unary_sign_exponentiation_precedence(): + """ + Test that exponentiation has higher precedence than unary signs (e.g. -2^2 == -4). + """ + assert infix.evaluate_expression("-2^2") == -4.0 + assert infix.evaluate_expression("(-2)^2") == 4.0 + assert infix.evaluate_expression("2^-3") == pytest.approx(0.125) + + context = {"x": math.pi / 2.0} + assert infix.evaluate_expression("-sin(x)^2", context) == pytest.approx(-1.0) + assert infix.evaluate_expression("(-sin(x))^2", context) == pytest.approx(1.0) + + +def test_infix_scientific_notation(): + """ + Test that numbers in scientific notation with standard ('e'/'E') exponents are + correctly parsed and evaluated. + """ + expression = "1.0e-3 + 2.5e-3 - 0.5E-3" + + assert infix.evaluate_expression(expression) == pytest.approx(3.0e-3) + + def test_infix_expression_with_var_conflict(): """ Test that an infix expression with a variable name that conflicts with a function