Belle II Software light-2609-luna
convert_training_inputs.py
1#!/usr/bin/env python3
2"""
3Convert produceTrainingInputs.py ROOT outputs into compact .npz shards for train.py.
4
5Reading the ROOT training inputs directly is the dominant cost of a training run: the
6'events' tree carries one branch per feature (1644 ms_feat_* branches, ~1667 total) split
7into ~50k small baskets per file, which costs ~11 s per file even with an explicit branch
8list. On the full v7 sample (14364 files, 320 GB) that is ~44 core-hours -- paid again for
9every train.py invocation, i.e. twice per submission since the category and main networks
10are trained in separate processes.
11
12The features are only ~1.6% dense, so the same data stored as CSR is ~211 bytes/event:
13the full sample becomes roughly 20 GB of .npz that loads in minutes. This script does the
14conversion once, parallel over files within a job and over input directories across
15HTCondor jobs.
16
17No selection is applied here -- shards hold exactly what _load_root_training_file()
18returns, so --fraction/--cont_fraction/--sigprob_thresh stay train-time knobs and do not
19require reconverting.
20
21Usage:
22 # one shard from one gbasf2 dataset directory (the usual HTCondor job)
23 python3 convert_training_inputs.py \
24 --input '/path/to/ModeSelector_v7/ModeSelector_v7_ccbar_1/**/*.root' \
25 --output /path/to/converted \
26 --name ModeSelector_v7_ccbar_1
27
28 # everything in one go, split into shards of 200 input files
29 python3 convert_training_inputs.py \
30 --input '/path/to/ModeSelector_v7/**/*.root' \
31 --output /path/to/converted \
32 --name ModeSelector_v7 --files_per_shard 200
33
34Then train on the result:
35 python3 train.py --input '/path/to/converted/*.npz' --network category --use_sparse
36"""
37
38import argparse
39import glob
40import os
41import sys
42import time
43from concurrent.futures import ProcessPoolExecutor
44
45from tqdm import tqdm
46
47try:
48 from modeSelector.training.train import concat_file_data, load_training_file, save_npz_shard
49except Exception:
50 # Same reason (and same deliberate breadth) as the config import in train.py: the
51 # modeSelector package __init__ pulls in basf2 and ROOT, which a standalone venv does
52 # not have, and pybasf2 fails with SystemError rather than ImportError when a basf2
53 # PYTHONPATH is inherited. train.py sits next to this file, so import it by path.
54 sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
55 from train import concat_file_data, load_training_file, save_npz_shard
56
57
58def _load_one(input_file):
59 """
60 Load one input file, returning (data, error) so a single bad file cannot abort a job.
61
62 Parameters:
63 input_file (str): Path to a produceTrainingInputs.py ROOT file.
64
65 Returns:
66 tuple: (data, error). On success `data` is the per-file dict from
67 _load_root_training_file() and `error` is None; on failure `data` is None and
68 `error` is a message naming the file and the exception.
69 """
70 try:
71 return load_training_file(input_file), None
72 except Exception as exc:
73 return None, f"{input_file}: failed to read input file ({exc})"
74
75
76def convert_shard(input_files, out_path, n_workers):
77 """
78 Load `input_files` in parallel and write them as a single .npz shard.
79
80 Parameters:
81 input_files (list of str): Input files making up this shard.
82 out_path (str): Output .npz path.
83 n_workers (int): Number of loader processes.
84
85 Returns:
86 tuple: (n_events, n_failed) -- events written and inputs skipped as unreadable.
87 """
88 data_list = []
89 errors = []
90
91 with ProcessPoolExecutor(max_workers=n_workers) as executor:
92 results = executor.map(_load_one, input_files)
93 for data, error in tqdm(results, total=len(input_files),
94 desc=os.path.basename(out_path)):
95 if error is not None:
96 errors.append(error)
97 continue
98 data_list.append(data)
99
100 for err in errors[:20]:
101 print(f" WARNING: {err}")
102 if len(errors) > 20:
103 print(f" WARNING: ... and {len(errors) - 20} more unreadable files")
104
105 if not data_list:
106 raise ValueError(f"{out_path}: all {len(input_files)} input files failed to load.")
107
108 merged = concat_file_data(data_list)
109 del data_list
110
111 os.makedirs(os.path.dirname(os.path.abspath(out_path)), exist_ok=True)
112 save_npz_shard(out_path, merged)
113
114 return merged['features'].shape[0], len(errors)
115
116
117def main():
118 """Parse command-line arguments and convert the requested inputs."""
119 parser = argparse.ArgumentParser(
120 description='Convert ModeSelector ROOT training inputs to .npz shards')
121 parser.add_argument('--input', required=True, nargs='+',
122 help='One or more ROOT paths (shell glob or space-separated list). '
123 'Quote patterns containing ** so glob expands them recursively.')
124 parser.add_argument('--output', required=True,
125 help='Output directory for the .npz shard(s)')
126 parser.add_argument('--name', default='shard',
127 help='Base name for the output shard(s) (default: shard)')
128 parser.add_argument('--files_per_shard', type=int, default=None,
129 help='Split the inputs into shards of this many files '
130 '(default: one shard for all inputs)')
131 parser.add_argument('--num_workers', type=int, default=None,
132 help='Loader processes (default: min(16, cpu_count()))')
133 parser.add_argument('--overwrite', action='store_true',
134 help='Rewrite shards that already exist (default: skip them)')
135
136 args = parser.parse_args()
137
138 expanded = []
139 for pattern in args.input:
140 matches = glob.glob(pattern, recursive=True)
141 expanded.extend(matches if matches else [pattern])
142 input_files = sorted(expanded)
143
144 if not input_files:
145 raise ValueError("No input files found.")
146
147 n_workers = args.num_workers
148 if n_workers is None:
149 n_workers = min(16, os.cpu_count() or 1)
150 n_workers = max(1, min(n_workers, len(input_files)))
151
152 files_per_shard = args.files_per_shard or len(input_files)
153 shards = [input_files[i:i + files_per_shard]
154 for i in range(0, len(input_files), files_per_shard)]
155
156 print(f"Input files: {len(input_files)}")
157 print(f"Shards to write: {len(shards)}")
158 print(f"Loader workers: {n_workers}")
159 print(f"Output dir: {args.output}")
160
161 os.makedirs(args.output, exist_ok=True)
162
163 total_events = 0
164 total_failed = 0
165 t_start = time.time()
166
167 for i, shard_files in enumerate(shards):
168 suffix = '' if len(shards) == 1 else f"_{i:04d}"
169 out_path = os.path.join(args.output, f"{args.name}{suffix}.npz")
170
171 if os.path.exists(out_path) and not args.overwrite:
172 print(f"\nSkipping existing shard {out_path} (use --overwrite to rewrite)")
173 continue
174
175 print(f"\n[{i + 1}/{len(shards)}] {out_path} <- {len(shard_files)} files")
176 t0 = time.time()
177 n_events, n_failed = convert_shard(shard_files, out_path, n_workers)
178 dt = time.time() - t0
179 size_mb = os.path.getsize(out_path) / 1024**2
180 print(f" {n_events} events, {size_mb:.1f} MB, {dt:.0f} s "
181 f"({dt / max(1, len(shard_files)):.1f} s/file)")
182 total_events += n_events
183 total_failed += n_failed
184
185 dt = time.time() - t_start
186 print(f"\nDone: {total_events} events in {len(shards)} shard(s), "
187 f"{total_failed} unreadable file(s), {dt / 60:.1f} min")
188
189 return 0
190
191
192if __name__ == '__main__':
193 sys.exit(main())
convert_shard(input_files, out_path, n_workers)
Definition main.py:1