48 from modeSelector.training.train
import concat_file_data, load_training_file, save_npz_shard
78 Load `input_files` in parallel and write them as a single .npz shard.
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.
86 tuple: (n_events, n_failed) -- events written and inputs skipped as unreadable.
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)):
98 data_list.append(data)
100 for err
in errors[:20]:
101 print(f
" WARNING: {err}")
103 print(f
" WARNING: ... and {len(errors) - 20} more unreadable files")
106 raise ValueError(f
"{out_path}: all {len(input_files)} input files failed to load.")
108 merged = concat_file_data(data_list)
111 os.makedirs(os.path.dirname(os.path.abspath(out_path)), exist_ok=
True)
112 save_npz_shard(out_path, merged)
114 return merged[
'features'].shape[0], len(errors)
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)')
136 args = parser.parse_args()
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)
145 raise ValueError(
"No input files found.")
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)))
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)]
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}")
161 os.makedirs(args.output, exist_ok=
True)
165 t_start = time.time()
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")
171 if os.path.exists(out_path)
and not args.overwrite:
172 print(f
"\nSkipping existing shard {out_path} (use --overwrite to rewrite)")
175 print(f
"\n[{i + 1}/{len(shards)}] {out_path} <- {len(shard_files)} files")
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
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")