Skip to content
Draft
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
48 changes: 36 additions & 12 deletions apsuite/loco/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -459,6 +459,12 @@ def update_weight(self):

def update_quad_knobs(self, use_families):
"""."""
if self.quadrupoles_to_fit == []:
# empty list does not update knobs
# knobs must be set directly to the
# quad_indices_kl attribute
self.quad_indices_ksl = self.quad_indices_kl
return
if self.quadrupoles_to_fit is None:
self.quadrupoles_to_fit = self.famname_quadset
else:
Expand Down Expand Up @@ -514,22 +520,35 @@ def update_sext_knobs(self, use_families):

def update_skew_quad_knobs(self):
"""."""
if self.skew_quadrupoles_to_fit is None:
self.skew_quad_indices_ksl = self.respm.fam_data['QS']['index']
else:
if self.skew_quadrupoles_to_fit == []:
# empty list does not update knobs
# knobs must be set directly to the
# quad_indices_kl attribute
return
if self.skew_quadrupoles_to_fit is not None:
skewquadfit = set(self.skew_quadrupoles_to_fit)
skewquadall = set(self.famname_skewquadset)
skewquadall = skewquadall.union(set(self.famname_idskewquadset))
if not skewquadfit.issubset(skewquadall):
raise Exception('invalid skew quadrupole name used to fit!')
else:
self.skew_quad_indices_ksl = []
for fam_name in self.skew_quadrupoles_to_fit:
fam = self.respm.fam_data
self.skew_quad_indices_ksl += fam[fam_name]['index']
idx_all = _np.array(self.respm.fam_data['QS']['index']).ravel()
idx_sub = _np.array(self.skew_quad_indices_ksl).ravel()
self.skew_quad_indices_ksl = list(set(idx_sub) & set(idx_all))
self.skew_quad_indices_ksl.sort()

self.skew_quad_indices_ksl = []
for fam_name in self.skew_quadrupoles_to_fit:
fam = self.respm.fam_data
self.skew_quad_indices_ksl += fam[fam_name]['index']
idx_all = _np.array(self.respm.fam_data['QS']['index']).ravel()
idx_all = _np.concatenate((
idx_all,
_np.array(self.respm.fam_data['IDQS']['index']).ravel(),
))
idx_sub = _np.array(self.skew_quad_indices_ksl).ravel()
self.skew_quad_indices_ksl = list(set(idx_sub) & set(idx_all))
self.skew_quad_indices_ksl.sort()
self.skew_quad_indices_ksl = [
[idx] for idx in self.skew_quad_indices_ksl
]
return
self.skew_quad_indices_ksl = self.respm.fam_data['QS']['index']

def update_b1_knobs(self):
"""."""
Expand Down Expand Up @@ -727,6 +746,11 @@ def famname_skewquadset(self):
'SDP3',
]

@property
def famname_idskewquadset(self):
"""."""
return ['IDQS']


class LOCOConfigBO(LOCOConfig):
"""Sirius Booster LOCO configuration."""
Expand Down
26 changes: 19 additions & 7 deletions apsuite/loco/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -479,14 +479,25 @@ def save_jacobian(self):
_LOCOUtils.save_data('6d_KL_sextupoles', jloco_kl_sext)

if self.config.fit_skew_quadrupoles:
idx_qs = self.config.respm.fam_data['QS']['index']
sub_qs = self.config.respm.fam_data['QS']['subsection']
idx_qs_all = self.config.respm.fam_data['QS']['index']
idx_qs_all += self.config.respm.fam_data['IDQS']['index']
sub_qs_all = self.config.respm.fam_data['QS']['subsection']
sub_qs_all += self.config.respm.fam_data['IDQS']['subsection']
idx_qs = self.config.skew_quad_indices_ksl
selidx = []
for sel in self.config.skew_quad_indices_ksl:
selidx.append(idx_qs.index([sel]))
idx_qs = [idx_qs[idx] for idx in selidx]
sub_qs = [sub_qs[idx] for idx in selidx]

idx_flat, sub_flat = [], []
for idxs, sub in zip(idx_qs_all, sub_qs_all):
if isinstance(idxs, list):
idx_flat.extend(idxs)
sub_flat.extend([sub] * len(idxs))
else:
idx_flat.append(idxs)
sub_flat.append(sub)

sub_qs = []
for i in idx_qs:
sub_qs.append(sub_flat[idx_flat.index(i[0])])

jloco_ksl_skewquad = self.create_new_jacobian_dict(
self._jloco_ksl_skew_quad, idx_qs, sub_qs
)
Expand Down Expand Up @@ -946,6 +957,7 @@ def _create_output_vars(self):
self._girders_shift_inival + self._girders_shift_deltas
)
self.kldelta_history = self._kldelta_history
self.ksldelta_history = self._ksldelta_history

def clear_output_vars(self):
"""."""
Expand Down