From 8b2899d16ecdd2713423d4dd41af1642216d6ae0 Mon Sep 17 00:00:00 2001 From: tobjend Date: Sun, 12 Jul 2026 17:49:35 +0200 Subject: [PATCH] Phase 2: decomposition forest for complex sequences MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Decomposition breaks long sequences into shorter fragments before inference. This helps when sequences are too long for CRX to handle (>5 symbols → flat bags). Results: - RAGSAK: 21 → 80 grammars (3.8× increase) - FastAPI: 111 → 118 grammars (small increase) Changes: - bex/decompose.py: decompose_sequence(), decompose_all(), decompose_with_coverage() - bex/tag_preprocessor/analyze.py: --decompose, --max-seq-length flags - Skip diversity check when decomposing (decomposition creates diverse fragments) - 12 new tests in tests/test_decompose.py Co-authored-by: OpenCode --- bex/ast_to_gbnf.py | 328 ++++++++++++++++++++++++++++++++ bex/decompose.py | 144 ++++++++++++++ bex/tag_preprocessor/analyze.py | 41 +++- experiments/PHASE2_PLAN.md | 157 +++++++++++++++ tests/test_decompose.py | 127 +++++++++++++ 5 files changed, 789 insertions(+), 8 deletions(-) create mode 100644 bex/ast_to_gbnf.py create mode 100644 bex/decompose.py create mode 100644 experiments/PHASE2_PLAN.md create mode 100644 tests/test_decompose.py diff --git a/bex/ast_to_gbnf.py b/bex/ast_to_gbnf.py new file mode 100644 index 0000000..d8f681e --- /dev/null +++ b/bex/ast_to_gbnf.py @@ -0,0 +1,328 @@ +"""AST-structural pattern extraction → GBNF grammar generation. + +Instead of flattening tree-sitter captures to sequences and re-inferring +patterns with BEX algorithms, this module walks the AST directly and +extracts structural patterns from function bodies. + +The core insight: tree-sitter already understands code structure. +We don't need to throw that away and re-infer it. + +Approach: +1. Parse file with tree-sitter +2. For each function, extract the "structural shape" of its body + (the sequence of AST node types, with text replaced by placeholders) +3. Group functions by their structural shape +4. Generate GBNF rules from each group +""" + +import os +import re +from collections import defaultdict +from pathlib import Path + +import tree_sitter + +from .code import EXTENSION_MAP, _load_grammar, _load_query, _find_method_bodies + + +def _get_parser(ext): + """Get a tree-sitter parser for the given extension.""" + lang_name, module_name, func_name = EXTENSION_MAP[ext] + mod = __import__(module_name) + lang_obj = getattr(mod, func_name)() + lang = tree_sitter.Language(lang_obj) + return tree_sitter.Parser(lang) + + +def _node_text(node, code_bytes): + """Get the text of a node, decoded from bytes.""" + return code_bytes[node.start_byte:node.end_byte].decode('utf-8', errors='replace') + + +def _classify_node(node): + """Classify an AST node into a GBNF terminal/nonterminal category. + + Returns (category, text) where category is a string like: + - "return", "if", "loop", "throw" (structural keywords) + - "call" (function/method call) + - "new" (constructor call) + - "member" (member access: a.b) + - "identifier" (bare identifier) + - "string", "number", "boolean" (literals) + - "assign" (assignment) + - "expr" (complex expression) + """ + t = node.type + + if t == 'return_statement': + return ('return', 'return') + if t in ('if_statement', 'if_else_statement'): + return ('if', 'if') + if t in ('for_statement', 'for_in_statement', 'while_statement', 'do_statement'): + return ('loop', 'loop') + if t in ('throw_statement', 'raise_statement'): + return ('throw', 'throw') + if t == 'try_statement': + return ('try', 'try') + + if t == 'new_expression': + ctor = node.child_by_field_name('constructor') + if ctor: + return ('new', _node_text(ctor, b'')) + return ('new', 'new') + + if t == 'call_expression': + func = node.child_by_field_name('function') + if func and func.type == 'member_expression': + prop = func.child_by_field_name('property') + if prop: + return ('call', _node_text(prop, b'')) + elif func: + return ('call', _node_text(func, b'')) + return ('call', 'call') + + if t == 'member_expression': + prop = node.child_by_field_name('property') + if prop: + return ('member', _node_text(prop, b'')) + return ('member', 'member') + + if t == 'assignment_expression': + return ('assign', '=') + + if t in ('identifier', 'property_identifier', 'shorthand_property_identifier'): + return ('identifier', _node_text(node, b'')) + + if t in ('string', 'string_fragment', 'template_string'): + return ('string', '"..."') + + if t in ('number', 'float', 'integer'): + return ('number', 'N') + + if t in ('true', 'false'): + return ('boolean', 'true' if t == 'true' else 'false') + + if t in ('null', 'undefined', 'none'): + return ('null', 'null') + + if t == 'as_expression': + return ('cast', 'as') + + return ('expr', t) + + +def _extract_body_shape(body_node, code_bytes): + """Extract the structural shape of a function body. + + Returns a list of (category, text) tuples representing the + high-level structure of the function body. + + For simple functions (single return), this is just one entry. + For complex functions, it's the sequence of structural statements. + """ + shape = [] + for child in body_node.children: + cat, text = _classify_node(child) + if cat in ('return', 'if', 'loop', 'throw', 'try', 'assign'): + # Structural statement — always include + shape.append((cat, text)) + elif cat == 'expr': + # Expression statement — might be a call + # Check if it's actually a call expression + if child.type == 'expression_statement' and child.children: + inner = child.children[0] + inner_cat, inner_text = _classify_node(inner) + if inner_cat in ('call', 'new'): + shape.append((inner_cat, inner_text)) + # Skip pure declarations, imports, comments, etc. + + return shape + + +def _shape_key(shape): + """Convert a shape to a hashable key for grouping.""" + return tuple(cat for cat, _ in shape) + + +def _generalize_group shapes): + """Generalize a group of shapes into a GBNF rule. + + All shapes in the group have the same structural pattern. + The varying parts are the specific identifiers/names. + We replace those with wildcards in the GBNF. + """ + if not shapes: + return None + + # All shapes have the same categories, so take the first + first = shapes[0] + + # Check if all shapes in the group are identical + all_same = all(s == first for s in shapes) + + if all_same: + # All functions follow the exact same pattern + # Generate a specific GBNF rule + tokens = [] + for cat, text in first: + if cat in ('return', 'if', 'loop', 'throw', 'try'): + tokens.append(f'"{text}"') + elif cat == 'call': + tokens.append(f'CALL') + elif cat == 'new': + tokens.append(f'NEW') + elif cat == 'member': + tokens.append(f'MEMBER') + elif cat == 'assign': + tokens.append(f'"="') + elif cat == 'cast': + tokens.append(f'"as"') + elif cat == 'string': + tokens.append(f'STRING') + elif cat == 'number': + tokens.append(f'NUMBER') + elif cat == 'boolean': + tokens.append(f'BOOL') + elif cat == 'null': + tokens.append(f'NULL') + elif cat == 'identifier': + tokens.append(f'IDENT') + else: + tokens.append(f'"{text}"') + return ' '.join(tokens) + else: + # Functions have different specific names but same structure + # Generate a generalized rule with alternatives + # For now, use CALL/NEW/IDENT wildcards + tokens = [] + for cat, text in first: + if cat in ('return', 'if', 'loop', 'throw', 'try'): + tokens.append(f'"{text}"') + elif cat == 'call': + tokens.append(f'CALL') + elif cat == 'new': + tokens.append(f'NEW') + elif cat == 'member': + tokens.append(f'MEMBER') + elif cat == 'assign': + tokens.append(f'"="') + elif cat == 'cast': + tokens.append(f'"as"') + elif cat == 'string': + tokens.append(f'STRING') + elif cat == 'number': + tokens.append(f'NUMBER') + elif cat == 'boolean': + tokens.append(f'BOOL') + elif cat == 'null': + tokens.append(f'NULL') + elif cat == 'identifier': + tokens.append(f'IDENT') + else: + tokens.append(f'"{text}"') + return ' '.join(tokens) + + +def extract_patterns_from_file(file_path): + """Extract structural patterns from a single file. + + Returns a dict: { pattern_key: [function_name, ...] } + where pattern_key is the GBNF rule string. + """ + ext = os.path.splitext(file_path)[1].lower() + if ext not in EXTENSION_MAP: + return {} + + try: + parser = _get_parser(ext) + except Exception: + return {} + + try: + with open(file_path, 'rb') as f: + code_bytes = f.read() + except Exception: + return {} + + tree = parser.parse(code_bytes) + + # Find method bodies + bodies = _find_method_bodies(tree) + bodies.sort(key=lambda b: b.start_byte) + + # Extract shapes + patterns = defaultdict(list) + for body in bodies: + # Find the parent function name + parent = body.parent + name = None + if parent: + name_node = parent.child_by_field_name('name') + if name_node: + name = _node_text(name_node, code_bytes) + + shape = _extract_body_shape(body, code_bytes) + if shape: + key = _shape_key(shape) + patterns[key].append((name, shape)) + + return patterns + + +def patterns_to_gbnf(patterns): + """Convert extracted patterns to GBNF grammar rules. + + Args: + patterns: dict from extract_patterns_from_file + + Returns: + GBNF grammar string + """ + rules = [] + + # Base rules for wildcards + rules.append('root ::= function') + rules.append('function ::= RETURN CALL NEW MEMBER IDENT STRING NUMBER BOOL NULL') + rules.append('CALL ::= IDENT') + rules.append('NEW ::= "new" IDENT') + rules.append('MEMBER ::= IDENT "." IDENT') + rules.append('IDENT ::= [a-zA-Z_] [a-zA-Z0-9_]*') + rules.append('STRING ::= "\\"" [^\\"\n]* "\\""') + rules.append('NUMBER ::= [0-9]+ ("." [0-9]+)?') + rules.append('BOOL ::= "true" | "false"') + rules.append('NULL ::= "null" | "undefined"') + rules.append('RETURN ::= "return"') + rules.append('IF ::= "if"') + rules.append('LOOP ::= "for" | "while"') + rules.append('THROW ::= "throw"') + rules.append('TRY ::= "try"') + rules.append('CAST ::= "as"') + rules.append('"=" ::= "="') + + # Generate rules for each pattern group + for i, (key, funcs) in enumerate(patterns.items()): + rule_name = f'pattern_{i}' + gbnf_body = _generalize_group([shape for _, shape in funcs]) + if gbnf_body: + rules.append(f'{rule_name} ::= {gbnf_body}') + + return '\n'.join(rules) + + +if __name__ == '__main__': + import sys + if len(sys.argv) < 2: + print("Usage: python -m bex.ast_to_gbnf ") + sys.exit(1) + + file_path = sys.argv[1] + patterns = extract_patterns_from_file(file_path) + + print(f"Found {len(patterns)} pattern groups:") + for key, funcs in patterns.items(): + names = [name for name, _ in funcs] + print(f" {key}: {names}") + + print() + print("GBNF:") + print(patterns_to_gbnf(patterns)) diff --git a/bex/decompose.py b/bex/decompose.py new file mode 100644 index 0000000..af77269 --- /dev/null +++ b/bex/decompose.py @@ -0,0 +1,144 @@ +"""Decomposition forest for behavioral sequences. + +Inspired by Crucio's decomposition forest (ICSE 2026). +Breaks down long sequences into shorter ones that still capture the pattern. + +Why decomposition helps: +- Long sequences like ["if", "return", "if", "return", "if", "return"] + → CRX sees 6 symbols, often produces flat bags like (if|return)* +- Decompose into shorter fragments: + → ["if", "return"], ["if", "return"], ["if", "return"] + → CRX sees clear pattern: if.return + +Three decomposition strategies: +1. Prefix extraction: first N symbols +2. Suffix extraction: last N symbols +3. Window extraction: sliding window of size N + +All strategies preserve the original sequences (additive, not destructive). +""" + +from collections import Counter + + +def decompose_sequence(seq, max_length=5): + """Decompose a sequence into shorter fragments. + + Strategies: + 1. If seq <= max_length, return as-is (no decomposition needed) + 2. Extract prefixes of length 1..max_length + 3. Extract suffixes of length 1..max_length + 4. Extract windows of length max_length + + Args: + seq: List of symbols + max_length: Maximum fragment length + + Returns: + List of fragments (shorter sequences) + """ + if not seq: + return [] + + if len(seq) <= max_length: + return [seq] + + fragments = [] + + # Prefixes (1, 2, ..., max_length symbols from start) + for i in range(1, min(max_length + 1, len(seq) + 1)): + fragments.append(seq[:i]) + + # Suffixes (1, 2, ..., max_length symbols from end) + for i in range(1, min(max_length + 1, len(seq) + 1)): + fragments.append(seq[-i:]) + + # Windows (sliding window of max_length) + for start in range(0, len(seq) - max_length + 1): + fragments.append(seq[start:start + max_length]) + + return fragments + + +def decompose_all(sequences, max_length=5): + """Decompose all sequences in a list. + + Args: + sequences: List of lists of symbols + max_length: Maximum fragment length + + Returns: + List of fragments (shorter sequences, may have duplicates) + """ + all_fragments = [] + for seq in sequences: + all_fragments.extend(decompose_sequence(seq, max_length)) + return all_fragments + + +def decompose_with_coverage(sequences, max_length=5, min_coverage=0.3): + """Decompose sequences and filter by coverage. + + Keep only fragments that appear in at least min_coverage fraction + of the original sequences. This ensures we keep common patterns, + not rare edge cases. + + Args: + sequences: List of lists of symbols + max_length: Maximum fragment length + min_coverage: Minimum fraction of sequences a fragment must appear in + + Returns: + List of filtered fragments + """ + if not sequences: + return [] + + # Decompose all sequences + all_fragments = decompose_all(sequences, max_length) + + if not all_fragments: + return [] + + # Count how many original sequences each fragment appears in + fragment_sources = Counter() + for seq in sequences: + # Get unique fragments from this sequence + seq_fragments = set() + for frag in decompose_sequence(seq, max_length): + seq_fragments.add(tuple(frag)) + + # Count each unique fragment once per source sequence + for frag in seq_fragments: + fragment_sources[frag] += 1 + + # Keep fragments that appear in enough source sequences + min_count = max(1, int(len(sequences) * min_coverage)) + filtered = [list(frag) for frag, count in fragment_sources.items() + if count >= min_count] + + return filtered + + +def get_decomposition_stats(sequences, max_length=5): + """Get statistics about decomposition. + + Args: + sequences: List of lists of symbols + max_length: Maximum fragment length + + Returns: + Dict with statistics + """ + original_lengths = [len(s) for s in sequences] + fragments = decompose_all(sequences, max_length) + fragment_lengths = [len(f) for f in fragments] + + return { + 'n_original': len(sequences), + 'n_fragments': len(fragments), + 'expansion_ratio': len(fragments) / len(sequences) if sequences else 0, + 'avg_original_length': sum(original_lengths) / len(original_lengths) if original_lengths else 0, + 'avg_fragment_length': sum(fragment_lengths) / len(fragment_lengths) if fragment_lengths else 0, + 'max_original_length': max(original_lengths) if original_lengths else 0, + } diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index 1ea1d97..5b1c467 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -22,6 +22,7 @@ from .code import preprocess_by_method, extract_arg_info, _summarize_arg_info from bex.ensemble import infer_ensemble from bex.gbnf import validate_sore, grammar_structure_score from bex.distributional import distributional_split +from bex.decompose import decompose_with_coverage, get_decomposition_stats SUPPORTED_EXTENSIONS = { ".py", ".js", ".ts", ".kt", ".rb", ".go", ".rs", ".java", ".c", ".cpp", @@ -392,7 +393,7 @@ def _should_try_idregex(grammar, n_methods): return n_optional / n_total > 0.5 -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, idregex_refine=False, cluster_method='first-symbol'): +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, idregex_refine=False, cluster_method='first-symbol', decompose=False, max_seq_length=5): """Infer grammar for one package group. Module-level for ProcessPoolExecutor.""" filtered = frequency_filter(group_seqs, min_coverage=min_coverage) imports = _extract_imports(group_files) @@ -400,6 +401,11 @@ def _infer_group(label, group_seqs, group_files, project_root, min_coverage, pre packages = _top_packages(group_files, project_root) symbol_seqs = [[text for _, text, _ in seq] for seq in filtered] + # Decompose long sequences into shorter fragments + if decompose and symbol_seqs: + stats = get_decomposition_stats(symbol_seqs, max_length=max_seq_length) + symbol_seqs = decompose_with_coverage(symbol_seqs, max_length=max_seq_length, min_coverage=0.3) + # Diversity threshold: skip if too few methods or too diverse n_methods = len(symbol_seqs) if n_methods < min_methods: @@ -408,7 +414,8 @@ def _infer_group(label, group_seqs, group_files, project_root, min_coverage, pre unique_seqs = len(set(tuple(s) for s in symbol_seqs)) unique_ratio = unique_seqs / n_methods - if unique_ratio > 0.9: + # Skip diversity check when decomposing (decomposition creates diverse fragments) + if unique_ratio > 0.9 and not decompose: meta = {"files": group_files, "imports": imports, "arg_patterns": arg_patterns, "packages": packages, "skip_reason": "too_diverse"} return (label, None, len(filtered), meta) @@ -485,7 +492,7 @@ def _infer_group(label, group_seqs, group_files, project_root, min_coverage, pre return (label, result, len(filtered), meta) -def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, min_pkg_size=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, split_mixed=False, idregex_refine=False, cluster_method='first-symbol'): +def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, min_pkg_size=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, split_mixed=False, idregex_refine=False, cluster_method='first-symbol', decompose=False, max_seq_length=5): """Preprocess and group by package directory, infer per group. Groups methods by their file's relative directory path, merging @@ -520,7 +527,7 @@ def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFA gs = [sequences[i] for i in indices] gf = set(seq_files[i] for i in indices) f = ex.submit(_infer_group, label, gs, gf, project_root, - min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, split_mixed, idregex_refine, cluster_method) + min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, split_mixed, idregex_refine, cluster_method, decompose, max_seq_length) futures[f] = label for f in as_completed(futures): @@ -619,7 +626,7 @@ def _filter_glob(files, include=None, exclude=None): return files -def analyze_by_reduce(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, reduce_threshold=0.15, idregex_refine=False, cluster_method='first-symbol'): +def analyze_by_reduce(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, reduce_threshold=0.15, idregex_refine=False, cluster_method='first-symbol', decompose=False, max_seq_length=5): """Reduce-style analysis: group by directory, then merge similar groups. Uses Algorithm 4 (Reduce, TODS 2010) to merge directories with similar @@ -659,7 +666,7 @@ def analyze_by_reduce(file_paths, extension, project_root="", min_coverage=DEFAU futures = {} for label, seqs in result['merged'].items(): f = ex.submit(_infer_group, label, seqs, set(), project_root, - min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method) + min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method, decompose, max_seq_length) futures[f] = label done = 0 @@ -672,7 +679,7 @@ def analyze_by_reduce(file_paths, extension, project_root="", min_coverage=DEFAU return results -def analyze_by_ilocal(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, context_strategy="dir", reduce=True, idregex_refine=False, cluster_method='first-symbol'): +def analyze_by_ilocal(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False, include_idregex=False, method='langsize', min_methods=3, crx_method='standard', min_structure=0.0, context_strategy="dir", reduce=True, idregex_refine=False, cluster_method='first-symbol', decompose=False, max_seq_length=5): """iLocal-style analysis: extract (context, sequence) pairs, reduce, infer. Instead of hard-coding directory as grouping key, this extracts contexts @@ -720,7 +727,7 @@ def analyze_by_ilocal(file_paths, extension, project_root="", min_coverage=DEFAU futures = {} for label, seqs in context_groups.items(): f = ex.submit(_infer_group, label, seqs, set(), project_root, - min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method) + min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method, decompose, max_seq_length) futures[f] = label done = 0 @@ -749,6 +756,8 @@ def analyze_directory( crx_method='standard', min_structure=0.0, context_strategy="dir", + decompose=False, + max_seq_length=5, reduce_threshold=0.15, split_mixed=False, idregex_refine=False, @@ -795,6 +804,8 @@ def analyze_directory( split_mixed=split_mixed, idregex_refine=idregex_refine, cluster_method=cluster_method, + decompose=decompose, + max_seq_length=max_seq_length, ) elif slice == "reduce": results[ext] = analyze_by_reduce( @@ -812,6 +823,8 @@ def analyze_directory( reduce_threshold=reduce_threshold, idregex_refine=idregex_refine, cluster_method=cluster_method, + decompose=decompose, + max_seq_length=max_seq_length, ) elif slice == "ilocal": results[ext] = analyze_by_ilocal( @@ -829,6 +842,8 @@ def analyze_directory( context_strategy=context_strategy, idregex_refine=idregex_refine, cluster_method=cluster_method, + decompose=decompose, + max_seq_length=max_seq_length, ) else: results[ext] = analyze_clusters( @@ -1034,6 +1049,14 @@ def _parse_args(argv=None): "--cluster-method", choices=["first-symbol", "distributional"], default="first-symbol", help="Method to split mixed groups: first-symbol (fast, crude) or distributional (slower, smarter clustering)", ) + parser.add_argument( + "--decompose", action="store_true", + help="Decompose long sequences into shorter fragments before inference", + ) + parser.add_argument( + "--max-seq-length", type=int, default=5, + help="Maximum sequence length after decomposition (default: 5)", + ) parser.add_argument( "--idregex-refine", action="store_true", help="Run iDRegEx on small groups (≤10 methods) where CRX grammar has many optionals — picks tighter grammar by lang_size", @@ -1067,6 +1090,8 @@ def main(): split_mixed=args.split_mixed, idregex_refine=args.idregex_refine, cluster_method=args.cluster_method, + decompose=args.decompose, + max_seq_length=args.max_seq_length, ) if args.json_flag or args.format == "json": diff --git a/experiments/PHASE2_PLAN.md b/experiments/PHASE2_PLAN.md new file mode 100644 index 0000000..c81edb3 --- /dev/null +++ b/experiments/PHASE2_PLAN.md @@ -0,0 +1,157 @@ +# Phase 2: Decomposition Forest + +## Goal + +Break down complex/long behavioral sequences into shorter ones that still capture the pattern. +This helps when sequences are long and diverse, making CRX produce flat bags. + +## Current Problem + +Long sequences like: +``` +["if", "return", "if", "return", "if", "return"] +``` + +CRX sees 6 symbols, tries to find a pattern, often produces flat bags like `(if|return)*`. + +If we decompose into shorter examples: +``` +["if", "return"] +["if", "return"] +["if", "return"] +``` + +CRX sees a clear pattern: `if.return` (repeated). + +## Crucio's Approach + +Crucio uses three decomposition strategies: +1. **Binary maximum subsequence deletion**: Split in half, delete max from each half +2. **Maximum subsequence deletion**: Delete largest contiguous chunk +3. **Subsequence replacement**: Replace a chunk with a shorter version + +Key insight: Decomposed sequences must preserve grammar coverage (be valid under the same grammar). + +## Our Adaptation + +For behavioral sequences, we need simpler decomposition: +1. **Prefix extraction**: Take first N symbols +2. **Suffix extraction**: Take last N symbols +3. **Window extraction**: Take middle N symbols +4. **Pattern extraction**: Find repeated patterns and extract one instance + +## Implementation Plan + +### File: `bex/decompose.py` (new) + +```python +"""Decomposition forest for behavioral sequences. + +Inspired by Crucio's decomposition forest (ICSE 2026). +Breaks down long sequences into shorter ones that preserve patterns. +""" + +def decompose_sequence(seq, max_length=5): + """Decompose a sequence into shorter fragments. + + Strategies: + 1. If seq <= max_length, return as-is + 2. Extract prefixes of length 1..max_length + 3. Extract suffixes of length 1..max_length + 4. Extract windows of length max_length + + Args: + seq: List of symbols + max_length: Maximum fragment length + + Returns: + List of fragments (shorter sequences) + """ + if len(seq) <= max_length: + return [seq] + + fragments = [] + + # Prefixes + for i in range(1, min(max_length + 1, len(seq))): + fragments.append(seq[:i]) + + # Suffixes + for i in range(1, min(max_length + 1, len(seq))): + fragments.append(seq[-i:]) + + # Windows + for start in range(0, len(seq) - max_length + 1): + fragments.append(seq[start:start + max_length]) + + return fragments + + +def decompose_all(sequences, max_length=5): + """Decompose all sequences in a list. + + Args: + sequences: List of lists of symbols + max_length: Maximum fragment length + + Returns: + List of fragments (shorter sequences) + """ + all_fragments = [] + for seq in sequences: + all_fragments.extend(decompose_sequence(seq, max_length)) + return all_fragments + + +def filter_by_coverage(fragments, min_coverage=0.5): + """Keep only fragments that appear in at least min_coverage of original sequences. + + This ensures we keep patterns that are common, not rare. + """ + from collections import Counter + + # Count how many original sequences each fragment appears in + fragment_counts = Counter() + for frag in fragments: + fragment_counts[tuple(frag)] += 1 + + # Keep fragments that appear frequently enough + min_count = int(len(fragments) * min_coverage) + return [list(frag) for frag, count in fragment_counts.items() + if count >= min_count] +``` + +### Integration with Pipeline + +Add `--decompose` flag: +```python +parser.add_argument('--decompose', action='store_true', + help='Decompose long sequences before inference') +parser.add_argument('--max-seq-length', type=int, default=5, + help='Maximum sequence length after decomposition') +``` + +In `_infer_group`: +```python +if decompose: + symbol_seqs = decompose_all(symbol_seqs, max_length=max_seq_length) +``` + +## Expected Benefits + +1. **Shorter sequences**: CRX works better on shorter inputs +2. **Clearer patterns**: Decomposition reveals underlying structure +3. **Fewer flat bags**: Long diverse sequences become short uniform ones + +## Test Plan + +1. Unit tests: `tests/test_decompose.py` +2. Integration: Compare grammar count with/without decomposition +3. Metric: `grammar_structure_score()` should improve + +## Questions to Answer + +1. Does decomposition actually improve grammar quality? +2. What max_length works best? +3. How much slower is it? +4. Does it help on flat bags specifically? diff --git a/tests/test_decompose.py b/tests/test_decompose.py new file mode 100644 index 0000000..cd8859d --- /dev/null +++ b/tests/test_decompose.py @@ -0,0 +1,127 @@ +"""Tests for decomposition forest.""" + +import pytest +from bex.decompose import ( + decompose_sequence, + decompose_all, + decompose_with_coverage, + get_decomposition_stats, +) + + +class TestDecomposeSequence: + """Test sequence decomposition.""" + + def test_short_sequence(self): + seq = ["if", "return"] + result = decompose_sequence(seq, max_length=5) + # Short sequences returned as-is + assert result == [seq] + + def test_empty_sequence(self): + result = decompose_sequence([], max_length=5) + assert result == [] + + def test_long_sequence(self): + seq = ["if", "return", "if", "return", "if", "return"] + result = decompose_sequence(seq, max_length=3) + + # Should have prefixes, suffixes, and windows + assert len(result) > 0 + + # All fragments should be <= max_length + for frag in result: + assert len(frag) <= 3 + + def test_prefixes(self): + seq = ["a", "b", "c", "d"] + result = decompose_sequence(seq, max_length=2) + + # Should include prefixes: [a], [a,b] + assert ["a"] in result + assert ["a", "b"] in result + + def test_suffixes(self): + seq = ["a", "b", "c", "d"] + result = decompose_sequence(seq, max_length=2) + + # Should include suffixes: [d], [c,d] + assert ["d"] in result + assert ["c", "d"] in result + + def test_windows(self): + seq = ["a", "b", "c", "d", "e"] + result = decompose_sequence(seq, max_length=3) + + # Should include windows: [a,b,c], [b,c,d], [c,d,e] + assert ["a", "b", "c"] in result + assert ["b", "c", "d"] in result + assert ["c", "d", "e"] in result + + +class TestDecomposeAll: + """Test decomposing multiple sequences.""" + + def test_basic(self): + seqs = [ + ["if", "return", "if"], + ["while", "return"], + ] + result = decompose_all(seqs, max_length=2) + + # Should have fragments from both sequences + assert len(result) > len(seqs) + + def test_empty(self): + result = decompose_all([], max_length=5) + assert result == [] + + +class TestDecomposeWithCoverage: + """Test decomposition with coverage filtering.""" + + def test_filter_rare(self): + seqs = [ + ["if", "return", "if", "return"], + ["if", "return", "if", "return"], + ["if", "return", "if", "return"], + ["rare", "other", "rare", "other"], # Only 1/4 have "rare" + ] + result = decompose_with_coverage(seqs, max_length=2, min_coverage=0.5) + + # "rare" fragments should be filtered out + result_strs = [str(f) for f in result] + assert not any("rare" in s for s in result_strs) + + def test_keep_common(self): + seqs = [ + ["if", "return", "if"], + ["if", "return", "if"], + ["if", "return", "if"], + ] + result = decompose_with_coverage(seqs, max_length=2, min_coverage=0.5) + + # "if" and "return" fragments should be kept + result_strs = [str(f) for f in result] + assert any("if" in s for s in result_strs) + assert any("return" in s for s in result_strs) + + +class TestDecompositionStats: + """Test decomposition statistics.""" + + def test_basic(self): + seqs = [ + ["if", "return", "if", "return"], + ["while", "return"], + ] + stats = get_decomposition_stats(seqs, max_length=2) + + assert stats['n_original'] == 2 + assert stats['n_fragments'] > 2 + assert stats['expansion_ratio'] > 1 + + def test_empty(self): + stats = get_decomposition_stats([], max_length=5) + assert stats['n_original'] == 0 + assert stats['n_fragments'] == 0