diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index 295a5b7..5956e34 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -204,6 +204,32 @@ def frequency_filter(sequences, min_coverage=0.2): return filtered +def _preprocess_file(fp): + """Preprocess one file. Module-level for ProcessPoolExecutor.""" + with open(fp) as f: + code = f.read() + sequences = [] + for method_seq in preprocess_by_method(fp, code): + if method_seq: + sequences.append(method_seq) + return (fp, sequences) + + +def _preprocess_files(file_paths): + """Preprocess multiple files in parallel.""" + sequences = [] + seq_files = [] + n_workers = os.cpu_count() + with ProcessPoolExecutor(max_workers=n_workers) as ex: + futures = {ex.submit(_preprocess_file, fp): fp for fp in file_paths} + for f in as_completed(futures): + fp, method_seqs = f.result() + for seq in method_seqs: + sequences.append(seq) + seq_files.append(fp) + return sequences, seq_files + + def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False): """Run full pipeline: preprocess → frequency filter → ensemble infer. @@ -211,16 +237,7 @@ def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAUL list of (label, ensemble_result_dict, method_count, meta) tuples. meta = {"files": set(paths), "imports": [lines], "arg_patterns": {...}, "packages": [...]}. """ - sequences = [] - seq_files = [] - for fp in file_paths: - with open(fp) as f: - code = f.read() - for method_seq in preprocess_by_method(fp, code): - if method_seq: - sequences.append(method_seq) - seq_files.append(fp) - + sequences, seq_files = _preprocess_files(file_paths) if not sequences: return [] @@ -259,17 +276,8 @@ def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFA Returns: list of (package_label, ensemble_result_dict, method_count, meta). """ - sequences = [] - seq_files = [] t0 = time.time() - for fp in file_paths: - with open(fp) as f: - code = f.read() - for method_seq in preprocess_by_method(fp, code): - if method_seq: - sequences.append(method_seq) - seq_files.append(fp) - + sequences, seq_files = _preprocess_files(file_paths) if not sequences: return [] _vprint(f"Preprocess: {len(sequences)} methods from {len(file_paths)} {extension} files ({time.time()-t0:.1f}s)") @@ -326,14 +334,7 @@ def infer(file_paths, extension, min_coverage=DEFAULT_COVERAGE, prefer=None, kma Returns: Ensemble result dict from infer_ensemble. """ - sequences = [] - for fp in file_paths: - with open(fp) as f: - code = f.read() - for method_seq in preprocess_by_method(fp, code): - if method_seq: - sequences.append(method_seq) - + sequences, _ = _preprocess_files(file_paths) sequences = frequency_filter(sequences, min_coverage=0.2) symbol_seqs = [[text for _, text, _ in seq] for seq in sequences]