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