58 nFolds=None, **kwargs):
59 """Configure the Normalizing Flow-based photon shower shape correction tool
60
61 nFolds: number of folds to use. If None (default), only 1 fold (fold 0) is used.
62 Otherwise must be >=1 and not exceed the NFolds value in the tool's config file.
63 """
64 acc = ComponentAccumulator()
65
66 from AthenaConfiguration.Enums import LHCPeriod
67
68 if not flags.Input.isMC:
69 raise RuntimeError("ElectronPhotonVariableNFCorrectionToolCfg: "
70 "NF correction tool should not be called for data"
71 )
72
73 isFullSim = flags.Sim.ISF.Simulator.isFullSim()
74 isRun3 = flags.GeoModel.Run is LHCPeriod.Run3
75 isRun2 = flags.GeoModel.Run is LHCPeriod.Run2
76
77 if isFullSim and isRun3:
78 default_conf = "EGammaVariableCorrection/NF_y_TUNE1/Run3FS/ElectronPhotonVariableNFCorrectionTool.conf"
79 elif isFullSim and isRun2:
80 default_conf = "EGammaVariableCorrection/NF_y_TUNE1/Run2FS/ElectronPhotonVariableNFCorrectionTool.conf"
81 elif not isFullSim and isRun3:
82 default_conf = "EGammaVariableCorrection/NF_y_TUNE1/Run3AF3/ElectronPhotonVariableNFCorrectionTool.conf"
83 elif not isFullSim and isRun2:
84
85 default_conf = "EGammaVariableCorrection/NF_y_TUNE1/Run3AF3/ElectronPhotonVariableNFCorrectionTool.conf"
86 else:
87 raise RuntimeError(
88 f"ElectronPhotonVariableNFCorrectionToolCfg: no NF correction config available for Run period {flags.GeoModel.Run} "
89 f"(isFullSim={isFullSim}). Only Run2 and Run3 are supported."
90 )
91
92 conf_key = kwargs.setdefault("ConfigFile", default_conf)
93
95 if not conf_file:
96 raise RuntimeError(f"PathResolver cannot find {conf_key}")
97
98
99 n_folds_config = None
100 pattern = None
101 with open(conf_file, 'r') as f:
102 for line in f:
103 line = line.strip()
104 if not line or line.startswith('#'):
105 continue
106 key, _, value = line.partition(':')
107 key = key.strip()
108 value = value.strip()
109 if key == 'NFolds':
110 n_folds_config = int(value)
111 elif key == 'ONNXnamePattern':
112 pattern = value
113
114 if n_folds_config is None or pattern is None:
115 raise RuntimeError(f'NFolds or ONNXnamePattern not found in config: {conf_file}')
116
117
118 if nFolds is None:
119 n_folds_used = 1
120 else:
121 if nFolds < 1:
122 raise ValueError(f'nFolds must be >= 1, got {nFolds}')
123 if nFolds > n_folds_config:
124 raise ValueError(
125 f'Requested nFolds={nFolds} exceeds NFolds={n_folds_config} available in config: {conf_file}'
126 )
127 n_folds_used = nFolds
128
129 kwargs.setdefault("NFoldsOverride", n_folds_used)
130
131
132 forward_tools = []
133 backward_tools = []
134 for i in range(n_folds_used):
135
136 fwd_session = CompFactory.AthOnnx.OnnxRuntimeSessionToolCPU(
137 f'NFCorrectionORTSessionToolForward_{i}',
138 ModelFileName=f'{pattern}_forward_{i}.onnx')
139 fwd_tool = CompFactory.AthOnnx.OnnxRuntimeInferenceTool(
140 f'NFCorrectionOnnxToolForward_{i}',
141 ORTSessionTool=fwd_session)
142 forward_tools.append(fwd_tool)
143
144 bwd_session = CompFactory.AthOnnx.OnnxRuntimeSessionToolCPU(
145 f'NFCorrectionORTSessionToolBackward_{i}',
146 ModelFileName=f'{pattern}_backward_{i}.onnx')
147 bwd_tool = CompFactory.AthOnnx.OnnxRuntimeInferenceTool(
148 f'NFCorrectionOnnxToolBackward_{i}',
149 ORTSessionTool=bwd_session)
150 backward_tools.append(bwd_tool)
151
152 kwargs.setdefault("OnnxInferenceToolsForward", forward_tools)
153 kwargs.setdefault("OnnxInferenceToolsBackward", backward_tools)
154
155 acc.setPrivateTools(
156 CompFactory.ElectronPhotonVariableNFCorrectionTool(name, **kwargs))
157 return acc
158
159
static std::string FindCalibFile(const std::string &logical_file_name)