Belle II Software light-2609-luna
convert_to_onnx.py
1#!/usr/bin/env python3
2
9
10"""Convert ModeSelector PyTorch checkpoints to ONNX format for basf2 inference.
11
12Usage:
13 python3 convert_to_onnx.py --input-dir networks/ --output-dir onnx/
14"""
15
16import argparse
17import os
18
19import numpy as np
20import torch
21import torch.nn as nn
22from modeSelector import config
23from modeSelector.training.train import MultiClassNet
24from ROOT import Belle2
25
26
27def convert_network_to_onnx(pt_path, onnx_path):
28 """Convert a ModeSelector .pt checkpoint to ONNX.
29
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.
33 """
34 checkpoint = torch.load(pt_path, map_location='cpu', weights_only=True)
35 state_dict = checkpoint['model_state_dict']
36
37 # Derive dimensions from the saved weights.
38 # The last linear layer index depends on architecture depth, so find it dynamically.
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]
42
43 model = MultiClassNet(input_size, num_labels)
44 model.load_state_dict(state_dict)
45 model.eval()
46
47 # Wrap with softmax: train.py outputs raw logits, basf2 inference expects probabilities
48 model_with_softmax = nn.Sequential(model, nn.Softmax(dim=1))
49 model_with_softmax.eval()
50
51 dummy_input = torch.randn(1, input_size)
52 torch.onnx.export(
53 model_with_softmax,
54 dummy_input,
55 onnx_path,
56 input_names=['input'],
57 output_names=['output'],
58 dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}},
59 opset_version=14,
60 )
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')
64
65
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.
68
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.
76 """
77 import onnxruntime as ort
78
79 generator = torch.Generator().manual_seed(0)
80 dummy = torch.randn(4, input_size, generator=generator)
81 with torch.no_grad():
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())
87 raise RuntimeError(
88 'ONNX validation failed for ' + onnx_path
89 + ': outputs do not match PyTorch model (max abs diff '
90 + f'{max_diff:.3e})'
91 )
92 if not (pt_out.argmax(axis=1) == onnx_out.argmax(axis=1)).all():
93 raise RuntimeError(
94 'ONNX validation failed for ' + onnx_path + ': predicted class does not match PyTorch model'
95 )
96 print(f"ONNX validation passed for {onnx_path}")
97
98
99def build_variable_names(has_inputs, input_size, is_main):
100 """Build the weightfile variable list describing this model's inputs.
101
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.
106 """
107 if not has_inputs:
108 raise RuntimeError(
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.'
111 )
112 names = [config.FEATURE_VAR_PREFIX + str(i) for i in has_inputs]
113 if is_main:
114 names += list(config.MAIN_EXTRA_VARS)
115 if len(names) != input_size:
116 raise RuntimeError(
117 'Feature selection does not match the model: built ' + str(len(names))
118 + ' variable names but the network takes ' + str(input_size) + ' inputs.'
119 )
120 return names
121
122
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.
125
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.
129 """
130 from basf2_mva_util import create_onnx_mva_weightfile
131 wf = create_onnx_mva_weightfile(
132 onnx_path,
133 variables=variables,
134 nClasses=n_classes,
135 identifier=identifier,
136 )
137 wf.addContractVersion(config.MODEL_CONTRACT_VERSION)
138 wf.save(root_path)
139 print(f"Packaged {onnx_path} -> {root_path} ({len(variables)} features, {n_classes} classes)")
140 print(f" identifier: {identifier}")
141
142
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."""
145 database = Belle2.Database.Instance()
146
147 if not database.addPayload(cat_payload_name, cat_model, iov):
148 raise RuntimeError(
149 'Failed to add category payload '
150 + cat_payload_name
151 + ' from '
152 + cat_model
153 )
154 if not database.addPayload(main_payload_name, main_model, iov):
155 raise RuntimeError(
156 'Failed to add main payload '
157 + main_payload_name
158 + ' from '
159 + main_model
160 )
161
162 print('Created localdb/database.txt with ModeSelector payloads')
163
164
165def main():
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()
190
191 print(f"Contract version {config.MODEL_CONTRACT_VERSION}")
192
193 model_specs = [
194 ('net_category.pt', 'modeSelector_cat.onnx', 'modeSelector_cat.root', False),
195 ('net_main.pt', 'modeSelector_main.onnx', 'modeSelector_main.root', True),
196 ]
197
198 # A payload records the contract version of the code that exported it, so refuse to
199 # export from code whose contract configuration is inconsistent.
200 error = config.contract_consistency_error()
201 if error:
202 raise RuntimeError('Cannot export: ' + error)
203
204 # The default names encode the contract version of this code, and releases request
205 # exactly that name. Exporting under another name bypasses that link.
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)):
213 if given != default:
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.')
217
218 # Belle2.Database.addPayload() appends to an existing local database rather than
219 # replacing it, which would leave two iovs per payload name covering the same range.
220 # Checked before the conversion so the run fails immediately rather than at the end.
221 localdb = os.path.join('localdb', 'database.txt')
222 if args.add_payloads and os.path.exists(localdb):
223 raise RuntimeError(
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.'
227 )
228
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):
232 raise RuntimeError(
233 'No checkpoints found in --input-dir ' + args.input_dir + ' (looked for '
234 + ', '.join(missing) + '). Point --input-dir at the directory holding the '
235 'trained .pt files.'
236 )
237
238 os.makedirs(args.output_dir, exist_ok=True)
239
240 exported_root = {}
241 selections = {}
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)
252 else:
253 print(f"Warning: {pt_path} not found")
254
255 # Both networks read the same raw features, so a disagreement means the two
256 # checkpoints come from different trainings and must not be paired in one payload set.
257 if len(selections) == 2 and selections['net_category.pt'] != selections['net_main.pt']:
258 raise RuntimeError(
259 'net_category.pt and net_main.pt were trained on different feature selections. '
260 'Re-export from a matching pair of checkpoints.'
261 )
262
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.')
266
267 if args.add_payloads:
268 missing_files = [
269 name for name in ('modeSelector_cat.root', 'modeSelector_main.root')
270 if name not in exported_root
271 ]
272 if missing_files:
273 raise RuntimeError(
274 'Cannot add payloads because these weightfiles were not created: '
275 + ', '.join(missing_files)
276 )
277
279 args.first_exp,
280 args.first_run,
281 args.final_exp,
282 args.final_run,
283 )
284 add_onnx_payloads(
285 exported_root['modeSelector_cat.root'],
286 exported_root['modeSelector_main.root'],
287 args.cat_payload_name,
288 args.main_payload_name,
289 iov,
290 )
291
292
293if __name__ == '__main__':
294 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.
Definition Database.cc:42
Definition main.py:1