Skip to content
Open
Show file tree
Hide file tree
Changes from 12 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Comment thread
jank324 marked this conversation as resolved.
Outdated

### 🐆 Other

- Update openPMD dependency to renamed package (see #684) (@jank324)
Expand Down
49 changes: 45 additions & 4 deletions cheetah/converters/bmad.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]

Expand All @@ -55,27 +60,35 @@ 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:
return cheetah.Drift(
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:
return cheetah.Drift(
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
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -133,6 +150,7 @@ def convert_element(
),
name=name,
sanitize_name=sanitize_name,
metadata=metadata,
)
elif bmad_parsed["element_type"] == "quadrupole":
validate_understood_properties(
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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(
Expand All @@ -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(
Expand Down Expand Up @@ -210,6 +246,7 @@ def convert_element(
],
name=name,
sanitize_name=sanitize_name,
metadata=metadata,
)
elif bmad_parsed["element_type"] == "ecollimator":
validate_understood_properties(
Expand Down Expand Up @@ -239,6 +276,7 @@ def convert_element(
],
name=name,
sanitize_name=sanitize_name,
metadata=metadata,
)
elif bmad_parsed["element_type"] == "wiggler":
validate_understood_properties(
Expand All @@ -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?
Expand All @@ -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(
Expand All @@ -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
Expand Down
99 changes: 27 additions & 72 deletions cheetah/converters/utils/fortran_namelist.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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]
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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
Comment on lines -313 to -354

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'm not sure how good of an idea it is to remove this now, but it seems like currently this is dead code ... and who knows if it actually works. I don't remember.



def parse_use_line(line: str, context: dict) -> dict:
"""
Parse a use line.
Expand Down Expand Up @@ -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):
Expand Down
Loading
Loading