diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index 3a75478..f25d462 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -309,6 +309,28 @@ def _split_by_first_symbol(symbol_seqs, min_subgroup=3): return viable +def _recursive_split(symbol_seqs, min_subgroup=3, max_depth=3, _depth=0): + """Recursively split by first symbol until sub-groups are uniform. + + Returns: + dict mapping "first1.first2..." → list of sequences (leaf groups). + """ + if _depth >= max_depth: + return {"": symbol_seqs} + + splits = _split_by_first_symbol(symbol_seqs, min_subgroup=min_subgroup) + if splits is None: + return {"": symbol_seqs} + + result = {} + for first_sym, sub_seqs in splits.items(): + sub_leaves = _recursive_split(sub_seqs, min_subgroup, max_depth, _depth + 1) + for suffix, leaf_seqs in sub_leaves.items(): + key = f"{first_sym}.{suffix}" if suffix else first_sym + result[key] = leaf_seqs + return result + + def _infer_group(label, group_seqs, group_files, project_root, min_coverage, prefer, kmax, N, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, split_mixed=False): """Infer grammar for one package group. Module-level for ProcessPoolExecutor.""" filtered = frequency_filter(group_seqs, min_coverage=min_coverage) @@ -331,29 +353,29 @@ def _infer_group(label, group_seqs, group_files, project_root, min_coverage, pre # Split mixed-pattern groups before CRX if split_mixed: - splits = _split_by_first_symbol(symbol_seqs, min_subgroup=min_methods) - if splits is not None: - # Infer each sub-group, pick the best - best_result = None - best_score = -1 - best_label = label + leaves = _recursive_split(symbol_seqs, min_subgroup=min_methods, max_depth=3) + if len(leaves) > 1: + # Infer each leaf, return ALL that pass + all_results = [] total_count = 0 - for first_sym, sub_seqs in splits.items(): - sub_label = f"{label} [{first_sym}]" - sub_result = infer_ensemble(sub_seqs, kmax=kmax, N=N, prefer=prefer, min_coverage=min_coverage, include_kore=include_kore, include_idregex=include_idregex, method=method) - if sub_result and sub_result.get('best') and sub_result['best'].get('grammar'): - g = sub_result['best']['grammar'] + for leaf_key, leaf_seqs in leaves.items(): + leaf_label = f"{label} [{leaf_key}]" if leaf_key else label + leaf_result = infer_ensemble(leaf_seqs, kmax=kmax, N=N, prefer=prefer, min_coverage=min_coverage, include_kore=include_kore, include_idregex=include_idregex, method=method) + total_count += len(leaf_seqs) + if leaf_result and leaf_result.get('best') and leaf_result['best'].get('grammar'): + g = leaf_result['best']['grammar'] ok, _ = validate_sore(g) if ok: - score = grammar_structure_score(g) if min_structure > 0 else sub_result['best']['mdl_score'] - if score > best_score: - best_score = score - best_result = sub_result - best_label = sub_label - total_count += len(sub_seqs) + if min_structure > 0 and grammar_structure_score(g) < min_structure: + continue + all_results.append((leaf_label, leaf_result, len(leaf_seqs))) - if best_result is not None: - meta = {"files": group_files, "imports": imports, "arg_patterns": arg_patterns, "packages": packages, "split": True, "n_splits": len(splits)} + if all_results: + # Return the best one, but store all in meta for later use + best_label, best_result, best_count = max(all_results, key=lambda x: grammar_structure_score(x[1]['best']['grammar'])) + meta = {"files": group_files, "imports": imports, "arg_patterns": arg_patterns, "packages": packages, + "split": True, "n_leaves": len(leaves), "n_grammars": len(all_results), + "all_grammars": [(l, r['best']['grammar'], grammar_structure_score(r['best']['grammar']), c) for l, r, c in all_results]} return (best_label, best_result, total_count, meta) # fall through to unsplit inference