diff --git a/tools/README.rst b/tools/README.rst index bbc0e80..75b0b31 100644 --- a/tools/README.rst +++ b/tools/README.rst @@ -67,6 +67,12 @@ pool-keywords Deterministically pool the keywords of the two tagging runs into the clustering step's input, per the project's handoff contract. Check ``pop-fem-audit-tools pool-keywords -h`` for complete instructions on its usage. +cluster-keywords +---------------- + +Deterministically build the coding vocabulary from the pooled keywords by sentence-embedding and clustering them. Requires the optional ``cluster`` dependency group. Check ``pop-fem-audit-tools cluster-keywords -h`` for complete instructions on its usage. + + Copyright ========= diff --git a/tools/docs/source/pop_fem_audit_tools.commands.rst b/tools/docs/source/pop_fem_audit_tools.commands.rst index 0a13fcf..f187082 100644 --- a/tools/docs/source/pop_fem_audit_tools.commands.rst +++ b/tools/docs/source/pop_fem_audit_tools.commands.rst @@ -12,6 +12,14 @@ pop\_fem\_audit\_tools.commands.build\_db module :show-inheritance: :undoc-members: +pop\_fem\_audit\_tools.commands.cluster\_keywords module +-------------------------------------------------------- + +.. automodule:: pop_fem_audit_tools.commands.cluster_keywords + :members: + :show-inheritance: + :undoc-members: + pop\_fem\_audit\_tools.commands.export\_llm\_input module --------------------------------------------------------- diff --git a/tools/pyproject.toml b/tools/pyproject.toml index 3252f0a..b9c0c5f 100644 --- a/tools/pyproject.toml +++ b/tools/pyproject.toml @@ -42,6 +42,13 @@ dependencies = [ "anthropic", ] +[project.optional-dependencies] +cluster = [ + "torch", + "sentence-transformers", + "scikit-learn", +] + [project.scripts] pop-fem-audit-tools = "pop_fem_audit_tools.__main__:main" diff --git a/tools/src/pop_fem_audit_tools/__main__.py b/tools/src/pop_fem_audit_tools/__main__.py index 828c788..0df8131 100644 --- a/tools/src/pop_fem_audit_tools/__main__.py +++ b/tools/src/pop_fem_audit_tools/__main__.py @@ -17,6 +17,7 @@ from types import ModuleType from .commands import ( build_db_command, + cluster_keywords_command, export_llm_input_command, fetch_artists_command, fetch_lyrics_command, @@ -29,6 +30,7 @@ MODULE_PROG: str = "python -m pop_fem_audit_tools" SUBCOMMANDS: dict[str, Callable[[list[str] | None], int]] = { "build-db": build_db_command, + "cluster-keywords": cluster_keywords_command, "export-llm-input": export_llm_input_command, "fetch-artists": fetch_artists_command, "fetch-lyrics": fetch_lyrics_command, diff --git a/tools/src/pop_fem_audit_tools/commands/__init__.py b/tools/src/pop_fem_audit_tools/commands/__init__.py index 4f39df8..dbb0e96 100644 --- a/tools/src/pop_fem_audit_tools/commands/__init__.py +++ b/tools/src/pop_fem_audit_tools/commands/__init__.py @@ -4,6 +4,7 @@ # imacat@mail.imacat.idv.tw (imacat), 2026/8/4 """The registry of the CLI subcommands.""" from .build_db import main as build_db_command +from .cluster_keywords import main as cluster_keywords_command from .export_llm_input import main as export_llm_input_command from .fetch_artists import main as fetch_artists_command from .fetch_lyrics import main as fetch_lyrics_command diff --git a/tools/src/pop_fem_audit_tools/commands/cluster_keywords.py b/tools/src/pop_fem_audit_tools/commands/cluster_keywords.py new file mode 100644 index 0000000..24b13f4 --- /dev/null +++ b/tools/src/pop_fem_audit_tools/commands/cluster_keywords.py @@ -0,0 +1,270 @@ +# Tools for A Feminist Audit of Pop Music. +# Copyright 2026 imacat. All rights reserved. +# Authors: +# imacat@mail.imacat.idv.tw (imacat), 2026/8/5 +"""The deterministic clusterer of the pooled keywords. + +Builds the coding vocabulary from the pooled keyword list, given as +the first positional command-line argument, by sentence-embedding +every keyword and clustering the embeddings: the group membership, +given as the second positional argument, and the group name +vocabulary, given as the third positional argument, are written as +plain files. The step is fully deterministic; no LLM call is +made. +""" +import argparse +import csv +import sys +import time +from pathlib import Path +from typing import Any + +from ..utils import format_duration + +MODEL: str = "sentence-transformers/all-mpnet-base-v2" +DEFAULT_CLUSTERS: int = 50 +CLUSTER_EXTRA_MESSAGE: str = ( + "cluster-keywords requires the optional \"cluster\"" + " dependency group; install it with" + " pip install -e \"tools/[cluster]\"") +"""The error message shown when the heavy clustering dependencies +are not installed.""" + + +def parse_args(argv: list[str] | None) -> argparse.Namespace: + """Parse the command-line arguments. + + :param argv: The command-line arguments, or None for + ``sys.argv``. + :return: The parsed arguments. + """ + parser: argparse.ArgumentParser = argparse.ArgumentParser( + description="Build the coding vocabulary by clustering" + " the sentence embeddings of the pooled" + " keywords.") + parser.add_argument( + "keywords_txt", type=Path, + help="the pooled keyword list, one keyword per line") + parser.add_argument( + "groups_csv", type=Path, + help="the group membership CSV output file") + parser.add_argument( + "vocabulary_txt", type=Path, + help="the group name vocabulary text output file") + parser.add_argument( + "--model", default=MODEL, + help=f"the sentence embedding model (default \"{MODEL}\")") + parser.add_argument( + "--revision", default=None, + help="the model revision to pin (default: unpinned)") + parser.add_argument( + "--clusters", type=int, default=DEFAULT_CLUSTERS, + help=f"the number of clusters (default {DEFAULT_CLUSTERS})") + return parser.parse_args(argv) + + +def load_keywords(path: Path) -> list[str]: + """Load and validate the pooled keyword list. + + :param path: The path of the pooled keyword text file, one + keyword per line. + :return: The keywords, in file order. + :raises OSError: When the file cannot be read. + :raises ValueError: When the file has no keyword, or a + keyword is duplicated. + """ + text: str = path.read_text(encoding="utf-8") + lines: list[str] = text.split("\n") + if len(lines) > 0 and lines[-1] == "": + lines = lines[:-1] + if len(lines) == 0: + raise ValueError(f"{path}: no keywords") + seen: set[str] = set() + keyword: str + for keyword in lines: + if keyword in seen: + raise ValueError( + f"{path}: duplicate keyword \"{keyword}\"") + seen.add(keyword) + return lines + + +def encode_keywords(keywords: list[str], model_name: str, + revision: str | None) -> Any: + """Encode the keywords into L2-normalized sentence embeddings. + + :param keywords: The keywords to encode. + :param model_name: The sentence embedding model name. + :param revision: The model revision to pin, or None to use + the model's default revision. + :return: The float32 embeddings, one row per keyword, in the + given order. + :raises RuntimeError: When the optional clustering + dependencies are not installed. + """ + try: + import numpy as np + from sentence_transformers import SentenceTransformer + except ImportError as error: + raise RuntimeError(CLUSTER_EXTRA_MESSAGE) from error + kwargs: dict[str, Any] = {} + if revision is not None: + kwargs["revision"] = revision + model: Any = SentenceTransformer( + model_name, device="cpu", **kwargs) + texts: list[str] = [x.replace("-", " ") for x in keywords] + embeddings: Any = model.encode( + texts, normalize_embeddings=True) + return np.asarray(embeddings, dtype=np.float32) + + +def cluster_embeddings(embeddings: Any, n_clusters: int) -> Any: + """Cluster the embeddings with ward-linkage agglomeration. + + :param embeddings: The float32 embeddings, one row per + keyword. + :param n_clusters: The number of clusters to form. + :return: The cluster label of each embedding, in the given + order. + :raises RuntimeError: When the optional clustering + dependencies are not installed. + """ + try: + from sklearn.cluster import AgglomerativeClustering + except ImportError as error: + raise RuntimeError(CLUSTER_EXTRA_MESSAGE) from error + clustering: Any = AgglomerativeClustering( + n_clusters=n_clusters, linkage="ward") + return clustering.fit_predict(embeddings) + + +def build_groups(keywords: list[str], embeddings: Any, + labels: Any) -> dict[str, list[str]]: + """Group the keywords by cluster label, named by their medoid. + + The group name is its medoid: the member whose embedding has + the highest dot product with the cluster's mean vector + re-normalized to unit length. Ties break toward the + lexicographically smallest member. + + :param keywords: The keywords, in embedding row order. + :param embeddings: The float32 embeddings, one row per + keyword. + :param labels: The cluster label of each keyword, in the same + order. + :return: The keyword members of every group, keyed by the + group's medoid name. + :raises RuntimeError: When the optional clustering + dependencies are not installed. + :raises ValueError: When two clusters yield the same medoid + name. + """ + try: + import numpy as np + except ImportError as error: + raise RuntimeError(CLUSTER_EXTRA_MESSAGE) from error + clusters: dict[int, list[int]] = {} + index: int + label: int + for index, label in enumerate(labels): + clusters.setdefault(int(label), []).append(index) + groups: dict[str, list[str]] = {} + indices: list[int] + for indices in clusters.values(): + members: list[str] = [keywords[x] for x in indices] + vectors: Any = embeddings[indices] + mean_vector: Any = vectors.mean(axis=0) + norm: float = float(np.linalg.norm(mean_vector)) + direction: Any = ( + mean_vector / norm if norm > 0 else mean_vector) + scores: Any = vectors @ direction + best_score: float = float(scores.max()) + medoid: str = min( + member for member, score in zip(members, scores) + if float(score) == best_score) + if medoid in groups: + raise ValueError( + f"duplicate medoid group name \"{medoid}\"") + groups[medoid] = members + return groups + + +def write_groups(path: Path, groups: dict[str, list[str]]) -> None: + """Write the group membership CSV file. + + Writes a CSV file with the header row ``Group,Keyword``, one + row per member keyword. Rows are sorted by group name + lexicographically, then by keyword lexicographically. + + :param path: The path of the group membership CSV file to + write. + :param groups: The keyword members of every group, keyed by + the group's medoid name. + :return: None. + :raises OSError: When the file cannot be written. + """ + group: str + with open(path, "w", encoding="utf-8", newline="") as file: + writer: Any = csv.writer(file) + writer.writerow(["Group", "Keyword"]) + for group in sorted(groups.keys()): + keyword: str + for keyword in sorted(groups[group]): + writer.writerow([group, keyword]) + + +def write_vocabulary(path: Path, + groups: dict[str, list[str]]) -> None: + """Write the group name vocabulary text file. + + Writes a plain text file, one group name per line, + lexicographically sorted, UTF-8, LF line endings, with a + trailing newline. + + :param path: The path of the vocabulary text file to write. + :param groups: The keyword members of every group, keyed by + the group's medoid name. + :return: None. + :raises OSError: When the file cannot be written. + """ + path.write_text( + "".join(f"{x}\n" for x in sorted(groups.keys())), + encoding="utf-8") + + +def main(argv: list[str] | None = None) -> int: + """Cluster the pooled keywords into the coding vocabulary. + + Writes the group membership CSV file and the group name + vocabulary text file. + + :param argv: The command-line arguments, or None for + ``sys.argv``. + :return: The exit status: 0 on success, non-zero on failure. + """ + started: float = time.monotonic() + args: argparse.Namespace = parse_args(argv) + try: + keywords: list[str] = load_keywords(args.keywords_txt) + except (OSError, ValueError) as error: + print(f"error: {error}", file=sys.stderr) + return 1 + try: + embeddings: Any = encode_keywords( + keywords, args.model, args.revision) + labels: Any = cluster_embeddings(embeddings, args.clusters) + groups: dict[str, list[str]] = build_groups( + keywords, embeddings, labels) + except (RuntimeError, ValueError) as error: + print(f"error: {error}", file=sys.stderr) + return 1 + args.groups_csv.parent.mkdir(parents=True, exist_ok=True) + args.vocabulary_txt.parent.mkdir(parents=True, exist_ok=True) + write_groups(args.groups_csv, groups) + write_vocabulary(args.vocabulary_txt, groups) + elapsed: str = format_duration(time.monotonic() - started) + print( + f"done: {len(keywords)} keywords clustered into" + f" {len(groups)} groups. {elapsed} elapsed.", + file=sys.stderr) + return 0 diff --git a/tools/tests/test_cluster_keywords.py b/tools/tests/test_cluster_keywords.py new file mode 100644 index 0000000..6d4bbc0 --- /dev/null +++ b/tools/tests/test_cluster_keywords.py @@ -0,0 +1,207 @@ +# Tools for A Feminist Audit of Pop Music. +# Copyright 2026 imacat. All rights reserved. +# Authors: +# imacat@mail.imacat.idv.tw (imacat), 2026/8/5 +"""Unit tests for the keyword clusterer module.""" +import csv +import io +import tempfile +import unittest +from contextlib import redirect_stderr +from pathlib import Path +from typing import Any +from unittest import mock + +import numpy as np + +from pop_fem_audit_tools.commands import cluster_keywords + + +type Vectors = dict[str, tuple[float, float]] +"""A fixed 2D embedding, keyed by keyword.""" + + +class TestClusterKeywords(unittest.TestCase): + """Test cases for the keyword clusterer.""" + + def setUp(self) -> None: + """Create a temporary directory for the output files.""" + tmp: tempfile.TemporaryDirectory[str] \ + = tempfile.TemporaryDirectory() + self.addCleanup(tmp.cleanup) + self.__dir: Path = Path(tmp.name) + self.__keywords_txt: Path = self.__dir / "keywords.txt" + self.__groups_csv: Path = self.__dir / "groups.csv" + self.__vocabulary_txt: Path = self.__dir / "vocabulary.txt" + + @staticmethod + def __two_cluster_vectors() -> Vectors: + """Build two well-separated, exactly medoid-determined + clusters of three unit vectors each. + + Each cluster is three points symmetric around a central + angle on the unit circle, so the point at the exact + central angle is uniquely closest to the cluster's + renormalized mean direction. + + :return: The fixed embedding of every keyword. + """ + return { + "a-left": (0.9396926, -0.3420201), + "a-center": (1.0, 0.0), + "a-right": (0.9396926, 0.3420201), + "b-north": (-0.9396926, 0.3420201), + "b-middle": (-1.0, 0.0), + "b-south": (-0.9396926, -0.3420201), + } + + def __write_keywords(self, keywords: list[str]) -> None: + """Write the pooled keyword input file. + + :param keywords: The keywords, one per line. + :return: None. + """ + self.__keywords_txt.write_text( + "".join(f"{x}\n" for x in keywords), encoding="utf-8") + + @staticmethod + def __fake_encode(vectors: Vectors) -> Any: + """Build a test double for :func:`encode_keywords`. + + :param vectors: The fixed 2D embedding of every keyword + the double may be asked to encode. + :return: A callable with the same signature as + :func:`encode_keywords`, returning the fixed + embeddings in the requested keyword order. + """ + def fake(keywords: list[str], model_name: str, + revision: str | None) -> Any: + """Return the fixed embeddings of the given keywords. + + :param keywords: The keywords to "encode". + :param model_name: Unused; part of the seam contract. + :param revision: Unused; part of the seam contract. + :return: The fixed float32 embeddings, in order. + """ + return np.asarray( + [vectors[x] for x in keywords], dtype=np.float32) + return fake + + def __run_cluster(self, extra_args: list[str] | None = None, + vectors: Vectors | None = None, + ) -> tuple[int, str]: + """Run the clusterer with a fake encoder and captured + standard error. + + :param extra_args: Extra command-line arguments appended + after the three positional arguments. + :param vectors: The fixed embedding to encode with; the + two-cluster fixture is used when None. + :return: A tuple of the exit status and the standard + error. + """ + argv: list[str] = [ + str(self.__keywords_txt), str(self.__groups_csv), + str(self.__vocabulary_txt)] + argv.extend(extra_args or []) + fake: Any = self.__fake_encode( + vectors if vectors is not None + else self.__two_cluster_vectors()) + stderr: io.StringIO = io.StringIO() + with mock.patch.object( + cluster_keywords, "encode_keywords", fake), \ + redirect_stderr(stderr): + status: int = cluster_keywords.main( + argv + ["--clusters", "2"]) + return status, stderr.getvalue() + + def __read_groups(self) -> list[list[str]]: + """Read the group membership CSV file. + + :return: All rows, including the header row, in file + order. + """ + with open(self.__groups_csv, encoding="utf-8", + newline="") as file: + return list(csv.reader(file)) + + def __read_vocabulary(self) -> list[str]: + """Read the vocabulary text file. + + :return: The group names, one per line, with the trailing + empty line from the final newline removed. + """ + lines: list[str] = self.__vocabulary_txt.read_text( + encoding="utf-8").split("\n") + self.assertEqual(lines[-1], "") + return lines[:-1] + + def test_groups_csv_header_and_ordering(self) -> None: + """Test the header row and the group/keyword ordering of + the group membership CSV file.""" + self.__write_keywords([ + "a-left", "a-center", "a-right", + "b-north", "b-middle", "b-south"]) + status: int + status, _ = self.__run_cluster() + self.assertEqual(status, 0) + rows: list[list[str]] = self.__read_groups() + self.assertEqual(rows[0], ["Group", "Keyword"]) + self.assertEqual(rows[1:], [ + ["a-center", "a-center"], + ["a-center", "a-left"], + ["a-center", "a-right"], + ["b-middle", "b-middle"], + ["b-middle", "b-north"], + ["b-middle", "b-south"], + ]) + + def test_every_keyword_appears_exactly_once(self) -> None: + """Test that every input keyword appears in exactly one + row of the group membership CSV file.""" + keywords: list[str] = [ + "a-left", "a-center", "a-right", + "b-north", "b-middle", "b-south"] + self.__write_keywords(keywords) + status: int + status, _ = self.__run_cluster() + self.assertEqual(status, 0) + rows: list[list[str]] = self.__read_groups()[1:] + self.assertEqual( + sorted(x[1] for x in rows), sorted(keywords)) + + def test_vocabulary_file_sorted_medoids(self) -> None: + """Test that the vocabulary file holds the sorted medoid + group names.""" + self.__write_keywords([ + "a-left", "a-center", "a-right", + "b-north", "b-middle", "b-south"]) + status: int + status, _ = self.__run_cluster() + self.assertEqual(status, 0) + self.assertEqual( + self.__read_vocabulary(), ["a-center", "b-middle"]) + + def test_duplicate_keyword_rejected(self) -> None: + """Test that a duplicate keyword line fails the run + without writing any output file.""" + self.__write_keywords(["shared", "shared"]) + status: int + stderr: str + status, stderr = self.__run_cluster() + self.assertEqual(status, 1) + self.assertIn("duplicate keyword", stderr) + self.assertFalse(self.__groups_csv.exists()) + self.assertFalse(self.__vocabulary_txt.exists()) + + def test_empty_input_rejected(self) -> None: + """Test that an empty keyword file fails the run without + writing any output file.""" + self.__keywords_txt.write_text("", encoding="utf-8") + status: int + stderr: str + status, stderr = self.__run_cluster() + self.assertEqual(status, 1) + self.assertIn("no keywords", stderr) + self.assertFalse(self.__groups_csv.exists()) + self.assertFalse(self.__vocabulary_txt.exists())