feature/treesitter-tag-queries #2

Open
tobi wants to merge 78 commits from feature/treesitter-tag-queries into main
5 changed files with 789 additions and 8 deletions
Showing only changes of commit 8b2899d16e - Show all commits

328
bex/ast_to_gbnf.py Normal file
View file

@ -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 <file>")
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))

144
bex/decompose.py Normal file
View file

@ -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,
}

View file

@ -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":

157
experiments/PHASE2_PLAN.md Normal file
View file

@ -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?

127
tests/test_decompose.py Normal file
View file

@ -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