diff --git a/bex/gbnf.py b/bex/gbnf.py index db767a4..dab45a0 100644 --- a/bex/gbnf.py +++ b/bex/gbnf.py @@ -224,3 +224,98 @@ def grammar_noise_ratio(node, noise_tokens=None): return grammar_noise_ratio(node.child, noise_tokens) return 0, 0 + + +def grammar_quality_score(node): + """Score grammar quality (0.0 = useless, 1.0 = excellent). + + Criteria: + - Has ordering (not a pure bag): +0.3 + - Has multiple alternation groups: +0.2 per group (max 0.4) + - Has enough symbols (>=3 domain tokens): +0.2 + - Not too short (>=2 concat parts): +0.1 + """ + from .grammar import Concat, Alt, Optional, Plus, Star + + if node is None or isinstance(node, (Epsilon, Empty)): + return 0.0 + + score = 0.0 + + # Check for ordering (Concat with multiple parts) + if isinstance(node, Concat) and len(node.parts) >= 2: + score += 0.3 + + # Check for alternation groups + n_groups = _count_alt_groups(node) + score += min(0.4, n_groups * 0.2) + + # Check symbol count + n_symbols = _count_symbols(node) + if n_symbols >= 3: + score += 0.2 + + # Check concat depth + n_concat = _count_concat_parts(node) + if n_concat >= 2: + score += 0.1 + + return min(1.0, score) + + +def _count_alt_groups(node): + """Count alternation groups in AST.""" + from .grammar import Concat, Alt, Optional, Plus, Star + + if node is None or isinstance(node, (Epsilon, Empty, Symbol)): + return 0 + if isinstance(node, Alt): + return 1 + sum(_count_alt_groups(p) for p in node.parts) + if isinstance(node, Concat): + return sum(_count_alt_groups(p) for p in node.parts) + if isinstance(node, (Plus, Optional, Star)): + return _count_alt_groups(node.child) + return 0 + + +def _count_symbols(node): + """Count total symbols in AST.""" + from .grammar import Concat, Alt, Optional, Plus, Star + + if node is None or isinstance(node, (Epsilon, Empty)): + return 0 + if isinstance(node, Symbol): + return 1 + if isinstance(node, Concat): + return sum(_count_symbols(p) for p in node.parts) + if isinstance(node, Alt): + return sum(_count_symbols(p) for p in node.parts) + if isinstance(node, (Plus, Optional, Star)): + return _count_symbols(node.child) + return 0 + + +def _count_concat_parts(node): + """Count top-level concat parts.""" + from .grammar import Concat, Optional, Plus, Star + + if isinstance(node, Concat): + return len(node.parts) + if isinstance(node, (Plus, Optional, Star)): + return _count_concat_parts(node.child) + return 1 + + +def is_useful_grammar(node, min_quality=0.3): + """Check if grammar is useful for LLM constraining. + + A grammar is useful if: + 1. Not empty after noise filtering + 2. Has some structure (not a pure bag) + 3. Has enough symbols to be constraining + """ + if node is None or isinstance(node, (Epsilon, Empty)): + return False + + quality = grammar_quality_score(node) + return quality >= min_quality diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index aef7816..8b48627 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -20,7 +20,7 @@ import pathspec from .code import preprocess_by_method, extract_arg_info, _summarize_arg_info from bex.ensemble import infer_ensemble -from bex.gbnf import grammar_structure_score, to_gbnf, filter_noise, grammar_noise_ratio +from bex.gbnf import grammar_structure_score, to_gbnf, filter_noise, grammar_noise_ratio, is_useful_grammar, grammar_quality_score from bex.grammar import Empty from bex.distributional import distributional_split from bex.decompose import decompose_with_coverage, get_decomposition_stats @@ -878,7 +878,7 @@ def analyze_directory( return results -def _build_json_output(results): +def _build_json_output(results, min_quality=0.3): """Convert results dict to a compact JSON structure for prompt injection.""" output = [] for ext, clusters in results.items(): @@ -896,14 +896,19 @@ def _build_json_output(results): filtered_grammar = filter_noise(grammar) if filtered_grammar and not isinstance(filtered_grammar, Empty): n_noise, n_total = grammar_noise_ratio(grammar) + quality = grammar_quality_score(filtered_grammar) entry["grammar"] = to_gbnf(filtered_grammar) entry["grammar_clean"] = to_gbnf(filtered_grammar) entry["noise_ratio"] = round(n_noise / n_total, 2) if n_total > 0 else 1.0 entry["symbols_before"] = n_total entry["symbols_after"] = n_total - n_noise + entry["quality"] = round(quality, 2) + entry["useful"] = quality >= min_quality else: entry["grammar"] = to_gbnf(grammar) entry["noise_ratio"] = 1.0 + entry["quality"] = 0.0 + entry["useful"] = False entry["algorithm"] = result["best"]["algorithm"] entry["mdl_score"] = round(result['best']['mdl_score'], 1) entry["imports"] = meta.get("imports", []) @@ -914,11 +919,11 @@ def _build_json_output(results): return json.dumps(output, indent=2) -def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filter_grammar_noise=True): +def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filter_grammar_noise=True, min_quality=0.3): """Build YAML output grouped by top-level module, sorted by MDL. Filters out (other), no-grammar groups, groups above max_mdl, - and groups below min_structure. + groups below min_structure, and low-quality grammars. Returns YAML string. """ import yaml @@ -952,6 +957,10 @@ def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filt else: continue # Skip grammars that become empty after filtering + # Apply quality gate + if not is_useful_grammar(grammar, min_quality): + continue + # Extract top-level module from package path parts = label.replace(os.sep, "/").split("/") module = parts[0] if len(parts) > 1 else "(root)" @@ -964,6 +973,7 @@ def _build_yaml_output(results, dir_path, max_mdl=500.0, min_structure=0.0, filt "score": round(best.get("mdl_score", 0), 3), "algorithm": best["algorithm"], "mdl": round(best["mdl_score"], 1), + "quality": round(grammar_quality_score(grammar), 2), } modules.setdefault(module, []).append(entry)