feat: distributional clustering for better grouping (Phase 1)

Inspired by Crucio's distributional matrix (ICSE 2026).
Groups symbols by context similarity instead of first-symbol.

New module: bex/distributional.py
- extract_contexts(): get (left, right) context pairs
- build_distribution_matrix(): symbol × context counts
- cluster_symbols(): agglomerative clustering by similarity
- split_by_cluster(): split sequences by cluster membership
- distributional_split(): drop-in replacement for _split_by_first_symbol

CLI: --cluster-method distributional (opt-in, first-symbol is default)

Example: 'if' and 'while' both appear before 'return'
→ first-symbol: 3 groups (if, while, return)
→ distributional: 2 groups (if/while merged, return)

23 tests pass. Full suite: 257 tests pass.
This commit is contained in:
tobjend 2026-07-12 17:37:36 +02:00
parent 2da3a7161f
commit 841f5efcf5
4 changed files with 745 additions and 15 deletions

298
bex/distributional.py Normal file
View file

@ -0,0 +1,298 @@
"""Distributional clustering for behavioral sequences.
Inspired by Crucio's distributional matrix (ICSE 2026).
Groups symbols by context similarity for better sequence classification.
Instead of splitting groups by first symbol (which is crude), we:
1. Extract context pairs for each symbol (what appears before/after)
2. Build a distribution matrix (symbol × context)
3. Cluster symbols with similar distributions
4. Split sequences by cluster membership
This captures "symbols that appear in similar contexts are equivalent"
which is the core insight from Crucio's distributional learning.
"""
from collections import defaultdict, Counter
import math
def extract_contexts(sequences):
"""Extract context pairs for each symbol.
For each symbol in each sequence, record (left_context, right_context).
Left context is the symbol before, right context is the symbol after.
None represents start/end of sequence.
Args:
sequences: List of lists of symbols
Returns:
Dict mapping symbol list of (left, right) context pairs
"""
contexts = defaultdict(list)
for seq in sequences:
for i, sym in enumerate(seq):
left = seq[i-1] if i > 0 else None
right = seq[i+1] if i < len(seq) - 1 else None
contexts[sym].append((left, right))
return dict(contexts)
def build_distribution_matrix(sequences, min_occurrences=2):
"""Build distribution matrix from sequences.
Rows = symbols, Columns = unique contexts.
M[i,j] = count of symbol i in context j.
Args:
sequences: List of lists of symbols
min_occurrences: Minimum occurrences to include symbol
Returns:
symbols: List of symbols (rows)
contexts: List of unique contexts (columns)
matrix: 2D list of counts
"""
# Extract all contexts
sym_contexts = extract_contexts(sequences)
# Filter by minimum occurrences
symbols = [sym for sym, ctxs in sym_contexts.items()
if len(ctxs) >= min_occurrences]
if not symbols:
return [], [], []
# Collect all unique contexts
all_contexts = set()
for sym in symbols:
all_contexts.update(sym_contexts[sym])
contexts = sorted(all_contexts, key=lambda x: (str(x[0]), str(x[1])))
# Build matrix
matrix = []
for sym in symbols:
row = []
ctx_counts = Counter(sym_contexts[sym])
for ctx in contexts:
row.append(ctx_counts.get(ctx, 0))
matrix.append(row)
return symbols, contexts, matrix
def cosine_similarity(vec1, vec2):
"""Compute cosine similarity between two vectors."""
if not vec1 or not vec2:
return 0.0
dot = sum(a * b for a, b in zip(vec1, vec2))
norm1 = math.sqrt(sum(a * a for a in vec1))
norm2 = math.sqrt(sum(b * b for b in vec2))
if norm1 == 0 or norm2 == 0:
return 0.0
return dot / (norm1 * norm2)
def jaccard_similarity(set1, set2):
"""Compute Jaccard similarity between two sets."""
if not set1 and not set2:
return 1.0
if not set1 or not set2:
return 0.0
intersection = len(set1 & set2)
union = len(set1 | set2)
return intersection / union if union > 0 else 0.0
def cluster_symbols(symbols, matrix, threshold=0.5, method='cosine'):
"""Cluster symbols by distribution similarity.
Uses agglomerative clustering: start with each symbol in own cluster,
merge most similar pairs until no pair exceeds threshold.
Args:
symbols: List of symbols
matrix: Distribution matrix (rows = symbols)
threshold: Similarity threshold for merging
method: 'cosine' or 'jaccard'
Returns:
Dict mapping symbol cluster_id
"""
if not symbols:
return {}
n = len(symbols)
# Initialize: each symbol in its own cluster
clusters = {i: i for i in range(n)}
cluster_members = {i: [i] for i in range(n)}
# Compute similarity matrix
sim = {}
for i in range(n):
for j in range(i+1, n):
if method == 'cosine':
sim[(i,j)] = cosine_similarity(matrix[i], matrix[j])
else:
# Jaccard on non-zero positions
set_i = {k for k, v in enumerate(matrix[i]) if v > 0}
set_j = {k for k, v in enumerate(matrix[j]) if v > 0}
sim[(i,j)] = jaccard_similarity(set_i, set_j)
# Agglomerative clustering
while True:
# Find most similar pair in different clusters
best_sim = -1
best_pair = None
for i in range(n):
for j in range(i+1, n):
if clusters[i] != clusters[j]:
if sim.get((i,j), 0) > best_sim:
best_sim = sim[(i,j)]
best_pair = (i, j)
if best_pair is None or best_sim < threshold:
break
# Merge clusters
ci, cj = best_pair
cluster_i = clusters[ci]
cluster_j = clusters[cj]
# Move all cluster_j members to cluster_i
for idx in cluster_members[cluster_j]:
clusters[idx] = cluster_i
cluster_members[cluster_i].append(idx)
del cluster_members[cluster_j]
# Build result mapping
result = {}
for sym_idx, cluster_id in clusters.items():
result[symbols[sym_idx]] = cluster_id
return result
def split_by_cluster(sequences, clusters, min_cluster_size=2):
"""Split sequences by first symbol's cluster.
Instead of splitting by first symbol, split by which cluster
the first symbol belongs to. This groups sequences that start
with "distributionally equivalent" symbols.
Args:
sequences: List of lists of symbols
clusters: Dict mapping symbol cluster_id
min_cluster_size: Minimum sequences to keep a cluster
Returns:
Dict mapping cluster_id list of sequences
"""
groups = defaultdict(list)
for seq in sequences:
if not seq:
groups['__empty__'].append(seq)
continue
first_sym = seq[0]
cluster_id = clusters.get(first_sym, f'cluster_{first_sym}')
groups[cluster_id].append(seq)
# Filter small clusters
result = {}
for cluster_id, seqs in groups.items():
if len(seqs) >= min_cluster_size or cluster_id == '__empty__':
result[cluster_id] = seqs
return result
def split_by_distributional(sequences, threshold=0.5, min_cluster_size=2):
"""Split sequences by distributional clustering of first symbols.
High-level function that combines all steps:
1. Extract first symbols from sequences
2. Build distribution matrix for first symbols
3. Cluster by context similarity
4. Split sequences by cluster
Args:
sequences: List of lists of symbols
threshold: Similarity threshold for merging
min_cluster_size: Minimum sequences to keep a cluster
Returns:
Dict mapping cluster_id list of sequences
"""
if not sequences:
return {}
# Extract first symbols
first_symbols = [seq[0] for seq in sequences if seq]
if not first_symbols:
return {'__empty__': sequences}
# Get all sequences containing each first symbol
# to build context distribution
sym_sequences = defaultdict(list)
for seq in sequences:
if seq:
sym_sequences[seq[0]].append(seq)
# Build distribution matrix for first symbols
# Use their full context (not just first position)
symbols = list(sym_sequences.keys())
if len(symbols) <= 1:
# Only one symbol type, no need to split
return {symbols[0]: sequences} if symbols else {}
# Extract contexts for each first symbol
contexts = extract_contexts(sequences)
# Build distribution matrix
all_contexts = set()
for sym in symbols:
all_contexts.update(contexts.get(sym, []))
context_list = sorted(all_contexts, key=lambda x: (str(x[0]), str(x[1])))
matrix = []
for sym in symbols:
row = []
ctx_counts = Counter(contexts.get(sym, []))
for ctx in context_list:
row.append(ctx_counts.get(ctx, 0))
matrix.append(row)
# Cluster
clusters = cluster_symbols(symbols, matrix, threshold)
# Split
return split_by_cluster(sequences, clusters, min_cluster_size)
# Convenience function for pipeline integration
def distributional_split(sequences, threshold=0.5, min_cluster_size=2):
"""Distributional clustering split (drop-in replacement for first-symbol).
Use this as a drop-in replacement for _split_by_first_symbol():
groups = distributional_split(sequences, threshold=0.5)
Returns:
Dict mapping cluster_id list of sequences
"""
return split_by_distributional(sequences, threshold, min_cluster_size)

View file

@ -21,6 +21,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 validate_sore, grammar_structure_score
from bex.distributional import distributional_split
SUPPORTED_EXTENSIONS = {
".py", ".js", ".ts", ".kt", ".rb", ".go", ".rs", ".java", ".c", ".cpp",
@ -309,24 +310,37 @@ 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):
def _recursive_split(symbol_seqs, min_subgroup=3, max_depth=3, _depth=0, cluster_method='first-symbol'):
"""Recursively split by first symbol until sub-groups are uniform.
Args:
symbol_seqs: List of sequences to split
min_subgroup: Minimum size to keep a group
max_depth: Maximum recursion depth
_depth: Current depth (internal)
cluster_method: 'first-symbol' (fast) or 'distributional' (smarter clustering)
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}
if cluster_method == 'distributional':
# Use distributional clustering for smarter splitting
splits = distributional_split(symbol_seqs, threshold=0.5, min_cluster_size=min_subgroup)
if not splits:
return {"": symbol_seqs}
else:
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 cluster_id, sub_seqs in splits.items():
sub_leaves = _recursive_split(sub_seqs, min_subgroup, max_depth, _depth + 1, cluster_method)
for suffix, leaf_seqs in sub_leaves.items():
key = f"{first_sym}.{suffix}" if suffix else first_sym
key = f"{cluster_id}.{suffix}" if suffix else str(cluster_id)
result[key] = leaf_seqs
return result
@ -378,7 +392,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):
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'):
"""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,7 +414,7 @@ def _infer_group(label, group_seqs, group_files, project_root, min_coverage, pre
# Split mixed-pattern groups before CRX
if split_mixed:
leaves = _recursive_split(symbol_seqs, min_subgroup=min_methods, max_depth=3)
leaves = _recursive_split(symbol_seqs, min_subgroup=min_methods, max_depth=3, cluster_method=cluster_method)
if len(leaves) > 1:
# Infer each leaf, return ALL that pass
all_results = []
@ -471,7 +485,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):
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'):
"""Preprocess and group by package directory, infer per group.
Groups methods by their file's relative directory path, merging
@ -506,7 +520,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)
min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, split_mixed, idregex_refine, cluster_method)
futures[f] = label
for f in as_completed(futures):
@ -605,7 +619,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):
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'):
"""Reduce-style analysis: group by directory, then merge similar groups.
Uses Algorithm 4 (Reduce, TODS 2010) to merge directories with similar
@ -645,7 +659,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)
min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method)
futures[f] = label
done = 0
@ -658,7 +672,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):
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'):
"""iLocal-style analysis: extract (context, sequence) pairs, reduce, infer.
Instead of hard-coding directory as grouping key, this extracts contexts
@ -706,7 +720,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)
min_coverage, prefer, kmax, N, include_kore, include_idregex, method, min_methods, crx_method, min_structure, False, idregex_refine, cluster_method)
futures[f] = label
done = 0
@ -738,6 +752,7 @@ def analyze_directory(
reduce_threshold=0.15,
split_mixed=False,
idregex_refine=False,
cluster_method='first-symbol',
):
"""Scan a directory and run analysis for each language found.
@ -779,6 +794,7 @@ def analyze_directory(
min_structure=min_structure,
split_mixed=split_mixed,
idregex_refine=idregex_refine,
cluster_method=cluster_method,
)
elif slice == "reduce":
results[ext] = analyze_by_reduce(
@ -795,6 +811,7 @@ def analyze_directory(
min_structure=min_structure,
reduce_threshold=reduce_threshold,
idregex_refine=idregex_refine,
cluster_method=cluster_method,
)
elif slice == "ilocal":
results[ext] = analyze_by_ilocal(
@ -811,6 +828,7 @@ def analyze_directory(
min_structure=min_structure,
context_strategy=context_strategy,
idregex_refine=idregex_refine,
cluster_method=cluster_method,
)
else:
results[ext] = analyze_clusters(
@ -1012,6 +1030,10 @@ def _parse_args(argv=None):
"--split-mixed", action="store_true",
help="Split groups with mixed first symbols before CRX inference (produces tighter grammars)",
)
parser.add_argument(
"--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(
"--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",
@ -1044,6 +1066,7 @@ def main():
reduce_threshold=args.reduce_threshold,
split_mixed=args.split_mixed,
idregex_refine=args.idregex_refine,
cluster_method=args.cluster_method,
)
if args.json_flag or args.format == "json":

169
experiments/PHASE1_PLAN.md Normal file
View file

@ -0,0 +1,169 @@
# Phase 1: Distributional Clustering
## Goal
Replace `_split_by_first_symbol()` with smarter clustering based on context similarity.
This is inspired by Crucio's distributional matrix but adapted for behavioral sequences.
## Current Problem
Our current `_split_by_first_symbol()` is too crude:
```python
# Current: split by first symbol
groups[seq[0]].append(seq)
# Problem: "if return" and "return if" look different but might be equivalent
# if they appear in similar contexts
```
## Crucio's Insight
Symbols that appear in the same contexts are distributionally equivalent:
```
Symbol "a" appears in contexts: [_, b], [c, _], [_, _]
Symbol "b" appears in contexts: [a, _], [_, c], [_, _]
→ If "a" and "b" have same context distribution, they're equivalent
```
## Our Adaptation
### Step 1: Context Extraction
For each symbol in all sequences, extract context pairs:
```python
# Example sequences:
# ["if", "return", "if", "return"]
# ["return", "if", "return"]
# Contexts for "return":
# - ("if", "if") at position 1
# - ("if", None) at position 3
# - (None, "if") at position 0
# Contexts for "if":
# - (None, "return") at position 0
# - ("return", "return") at position 1
# - ("return", None) at position 2
```
### Step 2: Distribution Vector
For each symbol, create a distribution vector:
```python
# Symbol: "return"
# Context distribution: {("if", "if"): 1, ("if", None): 1, (None, "if"): 1}
# Symbol: "if"
# Context distribution: {(None, "return"): 1, ("return", "return"): 1, ("return", None): 1}
```
### Step 3: Similarity Measure
Compare context distributions using cosine similarity or Jaccard:
```python
def context_similarity(sym1_contexts, sym2_contexts):
# Compare the sets of contexts
# High similarity → symbols are distributionally equivalent
```
### Step 4: Clustering
Group symbols with high similarity:
```python
# Cluster 1: ["if", "while", "for"] (conditional contexts)
# Cluster 2: ["return", "yield"] (return contexts)
# Cluster 3: ["class", "def"] (definition contexts)
```
### Step 5: Split Sequences by Cluster
Replace first-symbol split with cluster-based split:
```python
# Instead of:
# groups[seq[0]].append(seq)
# Do:
# cluster = symbol_to_cluster[seq[0]]
# groups[cluster].append(seq)
```
## Implementation Plan
### File: `bex/distributional.py` (new)
```python
"""Distributional clustering for behavioral sequences.
Inspired by Crucio's distributional matrix (ICSE 2026).
Groups symbols by context similarity for better sequence classification.
"""
def extract_contexts(sequences):
"""Extract context pairs for each symbol."""
# For each symbol, collect (left_context, right_context) pairs
pass
def build_distribution_matrix(contexts):
"""Build distribution matrix from contexts."""
# Rows = symbols, Columns = unique contexts
# M[i,j] = count of symbol i in context j
pass
def cluster_symbols(distribution_matrix, threshold=0.7):
"""Cluster symbols by distribution similarity."""
# Use cosine similarity or Jaccard
# Return dict: symbol → cluster_id
pass
def split_by_cluster(sequences, clusters):
"""Split sequences by first symbol's cluster."""
# Instead of first-symbol, use cluster membership
pass
```
### File: `bex/tag_preprocessor/analyze.py` (modify)
Add `--cluster-method` flag:
```python
# New flag
parser.add_argument('--cluster-method',
choices=['first-symbol', 'distributional'],
default='first-symbol',
help='Method to split mixed groups')
# In _recursive_split():
if cluster_method == 'distributional':
clusters = cluster_symbols(sequences)
return split_by_cluster(sequences, clusters)
else:
return _split_by_first_symbol(sequences)
```
## Expected Benefits
1. **Better grouping**: Symbols with same context → same group
2. **More general**: Handles cases where first symbol varies
3. **Still fast**: Distributional clustering is O(n * k) where n = symbols, k = contexts
## Test Plan
1. **Unit tests**: `tests/test_distributional.py`
- Test context extraction
- Test distribution matrix
- Test clustering
- Test sequence splitting
2. **Integration test**: Run on RAGSAK with `--cluster-method distributional`
- Compare grammar count and quality vs first-symbol
3. **Evaluation metric**: `grammar_structure_score()` and `lang_size_score()`
- Higher structure = better grouping
- Tighter grammars = better patterns captured
## Questions to Answer
1. Does distributional clustering actually improve grammar quality?
2. What similarity threshold works best?
3. How much slower is it than first-symbol?
4. Does it help on flat bags specifically?

View file

@ -0,0 +1,240 @@
"""Tests for distributional clustering."""
import pytest
from bex.distributional import (
extract_contexts,
build_distribution_matrix,
cosine_similarity,
jaccard_similarity,
cluster_symbols,
split_by_cluster,
split_by_distributional,
distributional_split,
)
class TestExtractContexts:
"""Test context extraction."""
def test_simple_sequence(self):
seqs = [["if", "return", "if"]]
ctxs = extract_contexts(seqs)
# "if" at position 0: (None, "return")
# "return" at position 1: ("if", "if")
# "if" at position 2: ("return", None)
assert ctxs["if"] == [(None, "return"), ("return", None)]
assert ctxs["return"] == [("if", "if")]
def test_empty_sequence(self):
seqs = [[]]
ctxs = extract_contexts(seqs)
assert ctxs == {}
def test_single_symbol(self):
seqs = [["return"]]
ctxs = extract_contexts(seqs)
assert ctxs["return"] == [(None, None)]
def test_multiple_sequences(self):
seqs = [
["if", "return"], # if at pos 0: (None, return), return at pos 1: (if, None)
["return", "if"], # return at pos 0: (None, if), if at pos 1: (return, None)
]
ctxs = extract_contexts(seqs)
# if appears at: pos 0 in seq1, pos 1 in seq2
assert (None, "return") in ctxs["if"] # from seq1
assert ("return", None) in ctxs["if"] # from seq2
# return appears at: pos 1 in seq1, pos 0 in seq2
assert ("if", None) in ctxs["return"] # from seq1
assert (None, "if") in ctxs["return"] # from seq2
class TestBuildDistributionMatrix:
"""Test distribution matrix construction."""
def test_simple(self):
seqs = [
["if", "return", "if", "return"],
["return", "if", "return", "if"],
]
syms, ctxs, mat = build_distribution_matrix(seqs, min_occurrences=1)
assert len(syms) == 2
assert "if" in syms
assert "return" in syms
assert len(ctxs) > 0
assert len(mat) == 2
def test_filter_rare(self):
seqs = [
["if", "return"],
["rare"], # rare symbol
]
syms, ctxs, mat = build_distribution_matrix(seqs, min_occurrences=2)
# "rare" should be filtered out
assert "rare" not in syms
def test_empty(self):
syms, ctxs, mat = build_distribution_matrix([], min_occurrences=1)
assert syms == []
class TestSimilarity:
"""Test similarity measures."""
def test_cosine_identical(self):
assert cosine_similarity([1, 2, 3], [1, 2, 3]) == pytest.approx(1.0)
def test_cosine_orthogonal(self):
assert cosine_similarity([1, 0], [0, 1]) == pytest.approx(0.0)
def test_cosine_similar(self):
sim = cosine_similarity([1, 2, 3], [1, 2, 4])
assert sim > 0.9
def test_jaccard_identical(self):
assert jaccard_similarity({1, 2}, {1, 2}) == 1.0
def test_jaccard_disjoint(self):
assert jaccard_similarity({1}, {2}) == 0.0
def test_jaccard_partial(self):
assert jaccard_similarity({1, 2}, {2, 3}) == pytest.approx(1/3)
class TestClusterSymbols:
"""Test symbol clustering."""
def test_similar_symbols(self):
# Two symbols with identical context distributions
symbols = ["a", "b"]
matrix = [
[1, 2, 0], # a
[1, 2, 0], # b (same as a)
]
clusters = cluster_symbols(symbols, matrix, threshold=0.5)
# Should be in same cluster
assert clusters["a"] == clusters["b"]
def test_dissimilar_symbols(self):
# Two symbols with different contexts
symbols = ["a", "b"]
matrix = [
[1, 0, 0], # a
[0, 0, 1], # b (different from a)
]
clusters = cluster_symbols(symbols, matrix, threshold=0.5)
# Should be in different clusters
assert clusters["a"] != clusters["b"]
def test_empty(self):
clusters = cluster_symbols([], [], threshold=0.5)
assert clusters == {}
class TestSplitByCluster:
"""Test sequence splitting by cluster."""
def test_basic(self):
seqs = [
["if", "return"],
["if", "return"],
["return", "if"],
["return", "if"],
]
clusters = {"if": 0, "return": 1}
groups = split_by_cluster(seqs, clusters, min_cluster_size=1)
assert 0 in groups
assert 1 in groups
assert len(groups[0]) == 2
assert len(groups[1]) == 2
def test_empty_sequences(self):
seqs = [
[],
["if", "return"],
]
clusters = {"if": 0}
groups = split_by_cluster(seqs, clusters, min_cluster_size=1)
assert "__empty__" in groups
assert len(groups["__empty__"]) == 1
def test_filter_small(self):
seqs = [
["if", "return"],
["return", "if"],
]
clusters = {"if": 0, "return": 1}
groups = split_by_cluster(seqs, clusters, min_cluster_size=2)
# Both clusters have only 1 sequence, should be filtered
assert 0 not in groups
assert 1 not in groups
class TestDistributionalSplit:
"""Test high-level distributional split."""
def test_similar_first_symbols(self):
# "if" and "while" both appear before "return"
# They should cluster together
seqs = [
["if", "return", "if", "return"],
["if", "return"],
["while", "return", "while", "return"],
["while", "return"],
["return", "if"],
["return", "if"],
]
groups = split_by_distributional(seqs, threshold=0.3, min_cluster_size=2)
# Should have fewer groups than first-symbol split
# (if and while should merge)
assert len(groups) < 6
def test_different_first_symbols(self):
# "if" and "return" have different contexts
seqs = [
["if", "return", "if"],
["if", "return"],
["return", "if", "return"],
["return", "if"],
]
groups = split_by_distributional(seqs, threshold=0.5, min_cluster_size=2)
# Should still have 2 groups
assert len(groups) >= 2
def test_empty(self):
groups = split_by_distributional([], threshold=0.5)
assert groups == {}
class TestDropInReplacement:
"""Test that distributional_split is a drop-in replacement."""
def test_same_signature(self):
# Should work like _split_by_first_symbol but smarter
seqs = [
["if", "return", "if"],
["if", "return"],
["return", "if"],
["return", "if"],
]
groups = distributional_split(seqs, threshold=0.5)
# Should return dict of groups
assert isinstance(groups, dict)
assert len(groups) > 0
# All sequences should be in some group
total = sum(len(v) for v in groups.values())
assert total == len(seqs)