10"""Convert ModeSelector PyTorch checkpoints to ONNX format for basf2 inference.
13 python3 convert_to_onnx.py --input-dir networks/ --output-dir onnx/
22from modeSelector
import config
23from modeSelector.training.train
import MultiClassNet
24from ROOT
import Belle2
27def convert_network_to_onnx(pt_path, onnx_path):
28 """Convert a ModeSelector .pt checkpoint to ONNX.
30 Input size and number of output classes are derived automatically
31 from the saved weights. The category network has 3 output classes;
32 the main network has N_INPUT_IDS + 3 = 139 output classes.
34 checkpoint = torch.load(pt_path, map_location=
'cpu', weights_only=
True)
35 state_dict = checkpoint[
'model_state_dict']
39 input_size = state_dict[
'network.0.weight'].shape[1]
40 last_layer_idx = max(int(k.split(
'.')[1])
for k
in state_dict
if k.endswith(
'.bias'))
41 num_labels = state_dict[f
'network.{last_layer_idx}.bias'].shape[0]
43 model = MultiClassNet(input_size, num_labels)
44 model.load_state_dict(state_dict)
48 model_with_softmax = nn.Sequential(model, nn.Softmax(dim=1))
49 model_with_softmax.eval()
51 dummy_input = torch.randn(1, input_size)
56 input_names=[
'input'],
57 output_names=[
'output'],
58 dynamic_axes={
'input': {0:
'batch_size'},
'output': {0:
'batch_size'}},
61 print(f
"Exported {pt_path} -> {onnx_path} (input={input_size}, labels={num_labels})")
62 _validate_onnx_export(model_with_softmax, onnx_path, input_size)
63 return input_size, num_labels, checkpoint.get(
'has_inputs')
66def _validate_onnx_export(model_with_softmax, onnx_path, input_size):
67 """Check that the exported ONNX model produces the same output as the PyTorch model.
69 Uses a batch of 4 inputs from a fixed seed to exercise the dynamic batch
70 dimension reproducibly. Tolerances are set for float32: torch and onnxruntime
71 accumulate the matmuls in a different order, which shifts individual softmax
72 probabilities of the 139-class main network by a few 1e-6. The numpy defaults
73 (atol=1e-8) are meant for float64 and reject those runs at random.
74 The predicted class is compared separately, since that is what the module uses.
75 Raises RuntimeError if outputs do not match.
77 import onnxruntime
as ort
79 generator = torch.Generator().manual_seed(0)
80 dummy = torch.randn(4, input_size, generator=generator)
82 pt_out = model_with_softmax(dummy).numpy()
83 session = ort.InferenceSession(onnx_path)
84 onnx_out = session.run([
'output'], {
'input': dummy.numpy()})[0]
85 if not np.allclose(pt_out, onnx_out, rtol=1e-4, atol=1e-6):
86 max_diff = float(np.abs(pt_out - onnx_out).max())
88 'ONNX validation failed for ' + onnx_path
89 +
': outputs do not match PyTorch model (max abs diff '
92 if not (pt_out.argmax(axis=1) == onnx_out.argmax(axis=1)).all():
94 'ONNX validation failed for ' + onnx_path +
': predicted class does not match PyTorch model'
96 print(f
"ONNX validation passed for {onnx_path}")
99def build_variable_names(has_inputs, input_size, is_main):
100 """Build the weightfile variable list describing this model's inputs.
102 The names are not resolved through the VariableManager. ModeSelectorModule fills
103 the feature vector manually. They are used to record which raw feature indices the
104 model was trained on, so the payload carries its own feature selection and a
105 retraining can change it without a software release.
109 'Checkpoint does not store has_inputs, cannot record the feature selection '
110 'in the weightfile. Retrain or export with a checkpoint written by train.py.'
112 names = [config.FEATURE_VAR_PREFIX + str(i)
for i
in has_inputs]
114 names += list(config.MAIN_EXTRA_VARS)
115 if len(names) != input_size:
117 'Feature selection does not match the model: built ' + str(len(names))
118 +
' variable names but the network takes ' + str(input_size) +
' inputs.'
123def package_as_mva_weightfile(onnx_path, root_path, variables, n_classes, identifier):
124 """Wrap an ONNX file in a basf2 MVA weightfile, saved as .root file.
126 The variable list records the raw feature indices the model was trained on. The contract
127 version is added as an extra weightfile element, and the identifier holds the training
128 name. All three are read back at inference time.
130 from basf2_mva_util
import create_onnx_mva_weightfile
131 wf = create_onnx_mva_weightfile(
135 identifier=identifier,
137 wf.addContractVersion(config.MODEL_CONTRACT_VERSION)
139 print(f
"Packaged {onnx_path} -> {root_path} ({len(variables)} features, {n_classes} classes)")
140 print(f
" identifier: {identifier}")
143def add_onnx_payloads(cat_model, main_model, cat_payload_name, main_payload_name, iov):
144 """Copy the ModeSelector MVA weightfiles into localdb/database.txt."""
147 if not database.addPayload(cat_payload_name, cat_model, iov):
149 'Failed to add category payload '
154 if not database.addPayload(main_payload_name, main_model, iov):
156 'Failed to add main payload '
162 print(
'Created localdb/database.txt with ModeSelector payloads')
166 """Convert net_category.pt and net_main.pt to ONNX, package as MVA weightfiles, and optionally upload to localdb."""
167 parser = argparse.ArgumentParser(description=
'Convert ModeSelector networks to ONNX')
168 parser.add_argument(
'--input-dir', required=
True,
169 help=
'Directory containing net_category.pt and net_main.pt')
170 parser.add_argument(
'--output-dir', required=
True,
171 help=
'Directory for output ONNX files')
172 parser.add_argument(
'--add-payloads', action=
'store_true',
173 help=
'Copy the exported ONNX files into localdb/database.txt')
174 parser.add_argument(
'--cat-payload-name', default=
None,
175 help=
'Payload name for the category model (default: derived from the contract version)')
176 parser.add_argument(
'--main-payload-name', default=
None,
177 help=
'Payload name for the main model (default: derived from the contract version)')
178 parser.add_argument(
'--identifier', default=
'unspecified',
179 help=
'Name of this training, stored as the weightfile identifier and '
180 'logged at inference time (for example sample and iteration)')
181 parser.add_argument(
'--first-exp', type=int, default=0,
182 help=
'First experiment of the interval of validity')
183 parser.add_argument(
'--first-run', type=int, default=0,
184 help=
'First run of the interval of validity')
185 parser.add_argument(
'--final-exp', type=int, default=-1,
186 help=
'Final experiment of the interval of validity')
187 parser.add_argument(
'--final-run', type=int, default=-1,
188 help=
'Final run of the interval of validity')
189 args = parser.parse_args()
191 print(f
"Contract version {config.MODEL_CONTRACT_VERSION}")
194 (
'net_category.pt',
'modeSelector_cat.onnx',
'modeSelector_cat.root',
False),
195 (
'net_main.pt',
'modeSelector_main.onnx',
'modeSelector_main.root',
True),
200 error = config.contract_consistency_error()
202 raise RuntimeError(
'Cannot export: ' + error)
206 if args.cat_payload_name
is None:
207 args.cat_payload_name = config.DEFAULT_CAT_PAYLOAD
208 if args.main_payload_name
is None:
209 args.main_payload_name = config.DEFAULT_MAIN_PAYLOAD
210 if args.add_payloads:
211 for given, default
in ((args.cat_payload_name, config.DEFAULT_CAT_PAYLOAD),
212 (args.main_payload_name, config.DEFAULT_MAIN_PAYLOAD)):
214 print(
'Warning: exporting payload ' + given +
' instead of ' + default +
'. Releases '
215 'implementing contract version ' + str(config.MODEL_CONTRACT_VERSION) +
' request '
216 + default +
', so they will not load this payload.')
221 localdb = os.path.join(
'localdb',
'database.txt')
222 if args.add_payloads
and os.path.exists(localdb):
224 os.path.abspath(localdb) +
' already exists. Adding payloads would append to it '
225 'instead of replacing it, leaving stale entries with overlapping iovs. Move or '
226 'delete the localdb directory and run again.'
229 missing = [pt_name
for pt_name, _onnx, _root, _is_main
in model_specs
230 if not os.path.exists(os.path.join(args.input_dir, pt_name))]
231 if len(missing) == len(model_specs):
233 'No checkpoints found in --input-dir ' + args.input_dir +
' (looked for '
234 +
', '.join(missing) +
'). Point --input-dir at the directory holding the '
238 os.makedirs(args.output_dir, exist_ok=
True)
242 for pt_name, onnx_name, root_name, is_main
in model_specs:
243 pt_path = os.path.join(args.input_dir, pt_name)
244 onnx_path = os.path.join(args.output_dir, onnx_name)
245 root_path = os.path.join(args.output_dir, root_name)
246 if os.path.exists(pt_path):
247 input_size, num_labels, has_inputs = convert_network_to_onnx(pt_path, onnx_path)
248 variables = build_variable_names(has_inputs, input_size, is_main)
249 package_as_mva_weightfile(onnx_path, root_path, variables, num_labels, args.identifier)
250 exported_root[root_name] = root_path
251 selections[pt_name] = list(has_inputs)
253 print(f
"Warning: {pt_path} not found")
257 if len(selections) == 2
and selections[
'net_category.pt'] != selections[
'net_main.pt']:
259 'net_category.pt and net_main.pt were trained on different feature selections. '
260 'Re-export from a matching pair of checkpoints.'
263 if selections
and sorted(selections.values())[0] != list(config.HAS_INPUTS):
264 print(
'Note: the exported feature selection differs from config.HAS_INPUTS. The payload '
265 'carries its own selection, so this is only a problem for training.')
267 if args.add_payloads:
269 name
for name
in (
'modeSelector_cat.root',
'modeSelector_main.root')
270 if name
not in exported_root
274 'Cannot add payloads because these weightfiles were not created: '
275 +
', '.join(missing_files)
285 exported_root[
'modeSelector_cat.root'],
286 exported_root[
'modeSelector_main.root'],
287 args.cat_payload_name,
288 args.main_payload_name,
293if __name__ ==
'__main__':
A class that describes the interval of experiments/runs for which an object in the database is valid.
static Database & Instance()
Instance of a singleton Database.