Skip to content
Open
Show file tree
Hide file tree
Changes from all 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
30 changes: 10 additions & 20 deletions examples/cheetah_diag0_model.ipynb
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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()"
]
Expand Down Expand Up @@ -134,44 +125,43 @@
"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": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"display_name": "linac-simulation",
"language": "python",
"name": "python3"
},
Expand Down
108 changes: 108 additions & 0 deletions virtual_accelerator/cheetah/diag0.py
Original file line number Diff line number Diff line change
@@ -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
31 changes: 24 additions & 7 deletions virtual_accelerator/cheetah/transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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
)
Loading
Loading