diff --git a/apsuite/loco/config.py b/apsuite/loco/config.py index 21bf77cfe..d58ed6c1f 100644 --- a/apsuite/loco/config.py +++ b/apsuite/loco/config.py @@ -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: @@ -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): """.""" @@ -727,6 +746,11 @@ def famname_skewquadset(self): 'SDP3', ] + @property + def famname_idskewquadset(self): + """.""" + return ['IDQS'] + class LOCOConfigBO(LOCOConfig): """Sirius Booster LOCO configuration.""" diff --git a/apsuite/loco/main.py b/apsuite/loco/main.py index 20b6319a9..31042989e 100644 --- a/apsuite/loco/main.py +++ b/apsuite/loco/main.py @@ -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 ) @@ -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): """."""