From 841f5efcf59cff31e9bae144db1f151910c603ee Mon Sep 17 00:00:00 2001 From: tobjend Date: Sun, 12 Jul 2026 17:37:36 +0200 Subject: [PATCH] feat: distributional clustering for better grouping (Phase 1) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- bex/distributional.py | 298 ++++++++++++++++++++++++++++++++ bex/tag_preprocessor/analyze.py | 53 ++++-- experiments/PHASE1_PLAN.md | 169 ++++++++++++++++++ tests/test_distributional.py | 240 +++++++++++++++++++++++++ 4 files changed, 745 insertions(+), 15 deletions(-) create mode 100644 bex/distributional.py create mode 100644 experiments/PHASE1_PLAN.md create mode 100644 tests/test_distributional.py diff --git a/bex/distributional.py b/bex/distributional.py new file mode 100644 index 0000000..834f0c5 --- /dev/null +++ b/bex/distributional.py @@ -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) diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index f2806a5..1ea1d97 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -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": diff --git a/experiments/PHASE1_PLAN.md b/experiments/PHASE1_PLAN.md new file mode 100644 index 0000000..7a3ad36 --- /dev/null +++ b/experiments/PHASE1_PLAN.md @@ -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? diff --git a/tests/test_distributional.py b/tests/test_distributional.py new file mode 100644 index 0000000..e1823ae --- /dev/null +++ b/tests/test_distributional.py @@ -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)