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.
298 lines
9 KiB
Python
298 lines
9 KiB
Python
"""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)
|