fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
2.69 kB
"""Shared primitives: union-find, cluster -> representative selection, normalize."""
from __future__ import annotations
import re
from collections import defaultdict
from typing import Any, Callable, Iterable
class UnionFind:
"""Union-find with path compression and union-by-rank."""
def __init__(self, n: int):
self.parent = list(range(n))
self.rank = [0] * n
def find(self, x: int) -> int:
root = x
while self.parent[root] != root:
root = self.parent[root]
while self.parent[x] != root:
self.parent[x], x = root, self.parent[x]
return root
def union(self, x: int, y: int) -> None:
rx, ry = self.find(x), self.find(y)
if rx == ry:
return
if self.rank[rx] < self.rank[ry]:
rx, ry = ry, rx
self.parent[ry] = rx
if self.rank[rx] == self.rank[ry]:
self.rank[rx] += 1
def cluster_and_pick(
n: int,
pairs: Iterable[tuple[int, int]],
key_fn: Callable[[int], Any] | None = None,
) -> tuple[list[int], dict[int, list[int]]]:
"""Cluster items in [0, n) via `pairs`, pick one rep per cluster.
Returns (keep, clusters):
keep: sorted list of representative indices (one per cluster).
clusters: dict {rep_idx -> sorted list of ALL member indices in that
cluster, including the rep}. Singletons appear as {i: [i]}.
"""
uf = UnionFind(n)
for i, j in pairs:
uf.union(i, j)
members_by_root: dict[int, list[int]] = defaultdict(list)
for i in range(n):
members_by_root[uf.find(i)].append(i)
clusters: dict[int, list[int]] = {}
for members in members_by_root.values():
rep = min(members) if key_fn is None else min(members, key=key_fn)
clusters[rep] = sorted(members)
keep = sorted(clusters.keys())
return keep, clusters
def pick_representatives(
n: int,
pairs: Iterable[tuple[int, int]],
key_fn: Callable[[int], Any] | None = None,
) -> list[int]:
"""Cluster items in [0, n) via `pairs`, then keep one rep per cluster.
Representative is the cluster member with the smallest `key_fn(i)`. With
no key_fn, falls back to smallest index ("first seen"). Returned indices
are sorted ascending so caller can re-index in original order.
Thin wrapper over `cluster_and_pick` that drops the cluster-membership map.
"""
keep, _ = cluster_and_pick(n, pairs, key_fn)
return keep
_WHITESPACE_RE = re.compile(r"\s+")
def normalize_text(text: str) -> str:
"""Lowercase + collapse internal whitespace + strip."""
return _WHITESPACE_RE.sub(" ", text.lower()).strip()