diff --git a/.github/workflows/python-publish.yml b/.github/workflows/python-publish.yml new file mode 100644 index 0000000..bc0dc34 --- /dev/null +++ b/.github/workflows/python-publish.yml @@ -0,0 +1,43 @@ +# This workflow will upload a Python package using uv when a release is published +# Adapted from: https://docs.astral.sh/uv/guides/integration/github/#publishing-to-pypi + +name: Upload Python Package + +on: + # normal behavior: run when a new release is published + release: + types: [published] + # allow running manually on main (restriction within job) + workflow_dispatch: + +concurrency: + # Cancel existing job(s) for workflow when a new one is queued + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + pypi-publish: + name: Upload release to PyPI + runs-on: ubuntu-latest + environment: + name: pypi + url: https://pypi.org/project/spanerr/ + permissions: + id-token: write # IMPORTANT: this permission is mandatory for trusted publishing + contents: read + if: github.event_name == 'release' || github.ref_name == 'main' + steps: + - name: Checkout repository + uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6.0.3 + with: + persist-credentials: false + - name: Install uv + uses: astral-sh/setup-uv@08807647e7069bb48b6ef5acd8ec9567f424441b # v.8.2.0 + with: + enable-cache: false + - name: Install Python 3.12 + run: uv python install 3.12 + - name: Build package + run: uv build + - name: Publish package + run: uv publish diff --git a/README.md b/README.md index ebdb448..a408663 100644 --- a/README.md +++ b/README.md @@ -1 +1,101 @@ # spanerr + +`spanerr` is a Python library for evaluating span-level text annotations using customizable alignment and scoring strategies. +`spanerr` operates explicitly over span text boundaries (i.e., text indices) rather than over the annotated text itself. + +Many span-level annotation tasks diverge significantly enough from named-entity recognition (e.g., long text spans, large label set) that the typical formulations for the evaluation metrics of precision, recall, and $F_1$ scores become insuitable. +`spanerr` addresses this issue by not only supporting customized scoring of (partial) span matches, but also customizing how the spans within document-level annotation sets are aligned for evaluation. +`spanerr` is designed for maximal flexibility generally leaving it to the user to determine what assumptions and restrictions are required in their use case. + +[![unit tests](https://github.com/Princeton-CDH/spanerr/actions/workflows/unit-tests.yml/badge.svg)](https://github.com/Princeton-CDH/spanerr/actions/workflows/unit-tests.yml) +[![codecov](https://codecov.io/gh/Princeton-CDH/spanerr/graph/badge.svg?token=Wd3vZ38Bxz)](https://codecov.io/gh/Princeton-CDH/spanerr) + +## Basic Usage + +### Installation + +Use pip to install as a Python package directly from GitHub. +Use a branch or tag name, e.g. `@develop` or `@0.1.0` if you need to install a specific version + +```sh +pip install git+https://github.com/Princeton-CDH/spanerr.git#egg=spanerr +``` + +### Core Data Types + +`spanerr` has three first-class objects: + +- `Span`: An individual span annotation. +- `DocSpans`: A set of span annotations for a document. +- `SpanAlignment`: A set of aligned span annotations (reference, system) within a single document. + +All three of these data types are immutable, but `SpanAlignment` does not currently support hashing. + +#### Binarization + +`spanerr` provides functionality for "removing" span labels from `Span` and `DocSpans` by setting them to a default label (empty string). +For `DocSpans`, overlapping spans will be merged and optionally neighboring spans can be merged. + +#### Loading from dictionaries + +`Spans` can be loaded from dictionaries with the following fields: + +- `start` (int): starting text index (inclusive) +- `end` (int): ending text index (exclusive) +- `label` (str): optional span label (defaults to empty string) + +`DocSpans` can be loaded from dictionaries with the following fields: + +- `doc_id` (str): optional document id (defaults to empty string) +- `spans` (list[dict]): list of spans (in dictionary form, see above) + +### Core Functionality + +In `spanerr` there are two core components to evaluating span annotations: (1) how spans are aligned and (2) how aligned spans are scored. +`spanerr` is intentionally designed so that these two components can be heavily customized. + +#### Aligning span annotations + +An alignment strategy is represented as a function that takes two sets of annotations (`DocSpans`) as input and returns an alignment (`SpanAlignment`) which will then be used for scoring. +There are little restrictions on the alignments themselves: the resulting `SpanAlignment` may contain transformed versions of the input `DocSpans` and no restrictions are made on its mapping between reference and system spans. +The idea is to allow for the creation of whatever alignment is useful for scoring. + +The following alignment strategies are provided in `spanerr.align`: + +- Select First : Select the first (sequential) matching system span for each reference span. + By default, spans match if they overlap and have the same label. +- Select Best : Select the best matching system span for each reference span. + By default, given spans that overlap and have the same label, the best match is the span pair with the highest jaccard similarity. +- Corppa : The alignment strategy used by [`corppa`](https://github.com/Princeton-CDH/corppa). + See `corppa`'s [evaluation documentation](https://github.com/Princeton-CDH/corppa/tree/main/src/corppa/poetry_detection/evaluation) for more detail. + +For additional flexibility, alignment strategies may take additional inputs to further customize their behavior (e.g., use different span matching and span scoring strategies), but these will generally need to be set to a specific value (e.g., via a lambda function) before they can be used within `spanerr`'s evaluation workflow. +See the `spanerr.align.construct_aligner` for an example. + +### Scoring Aligmnents + +Scoring an alignment (`SpanAlignment`) entails computing an alignment's *relevance score* (i.e., true positive, numerator for precision and recall). +So, a scoring strategy is represented as a function that takes an alignment (`SpanAlignment`) as input and returns its score (`float`). +This allows for both alignment-independent scoring strategies in which scoring depends solely on the individual scores of each reference-system span pair as well as those that don't. + +Currently, `spanerr` provides the building blocks for constructing alignment-independent strategies using `spanerr.eval.relevance_score` and `spanerr.span_utils.composite_match_score`. +See `spanerr.compute_metrics.get_scorer` for an example of constructing scoring strategy functions. + +### Span Annotations File Format + +A set of span annotations can be loaded into `spanerr` by providing a JSONL file with each line corresponding to a different document's annotations (i.e. `DocSpans`). +For examples see the files in `tests/test_data`. + +### Scripts + +Installing `spanerr` currently provides access to the following command line script: + +- `spanerr-metrics`: For calculating entity- or document-level aggregated precision, recall, and F-1 scores for given reference and system span annotation sets. + (Corresponds to `spanerr.compute_metrics.py`) + +## License + +This project is licensed under the [Apache 2.0 License](LICENSE) + +(c)2026 Trustees of Princeton University. +Permission granted for non-commercial distribution online under a standard Open Source license. diff --git a/pyproject.toml b/pyproject.toml index f5b66a8..8bb4019 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,6 +42,9 @@ test = [ "pytest-cov>=7.1.0", ] +[project.scripts] +spanerr-metrics="spanerr.compute_metrics:main" + [tool.uv] exclude-newer= "1 week" diff --git a/src/spanerr/align.py b/src/spanerr/align.py index 01c759c..8f242c9 100644 --- a/src/spanerr/align.py +++ b/src/spanerr/align.py @@ -25,13 +25,14 @@ def select_first_match( ref: DocSpans, sys: DocSpans, - is_match: CheckSpanPair, + is_match: CheckSpanPair = partial_overlap, exclusive: bool = True, ) -> SpanAlignment: """ Builds a span alignment using a select first match strategy. Each reference - span is aligned with the first matching system span as determined by the provided - `is_match` method. By default, alignments are exclusive. + span is aligned with the first matching system span. By default, a match + corresponds to spans that overlap and have the same label. By default, + alignments are exclusive. """ align_map = {} sys_span_pool = dict.fromkeys(sys.spans) @@ -50,14 +51,16 @@ def select_first_match( def select_best_match( ref: DocSpans, sys: DocSpans, - is_match: CheckSpanPair, - score_match: ScoreSpanPair, + is_match: CheckSpanPair = partial_overlap, + score_match: ScoreSpanPair = Span.jaccard, exclusive: bool = True, ) -> SpanAlignment: """ Builds a span alignment using a greedy select best match strategy. Each reference - span is aligned with its best matching system span as defined by the provided - `is_match` and `score_match` methods. By default, alignments are exclusive. + span is aligned with its best matching system span. By default, a match + corresponds to spans that overlap and have the same label and the best match + corresponds to the match with the highest jaccard similarity. + By default, alignments are exclusive. """ mapping = {} sys_span_pool = dict.fromkeys(sys.spans) @@ -142,52 +145,34 @@ def construct_aligner( - select_best: corresponds to select_best_match - corppa: corresponds to corppa_align """ + ## Validate inputs and get alignment method + align_method = None match strategy: case "select_first": # Validate input parameters - if is_match is None: - raise ValueError(f"Strategy {strategy} requires is_match parameter") if score_match is not None: raise ValueError( f"Strategy {strategy} does not use score_match parameter" ) - # Construct aligner - if exclusive is None: - return lambda r, s: select_first_match(r, s, is_match) - else: - return lambda r, s: select_first_match( - r, s, is_match, exclusive=exclusive - ) + align_method = select_first_match case "select_best": - # Validate input parameters - if is_match is None or score_match is None: - raise ValueError( - f"Strategy {strategy} requires is_match and score_match parameters" - ) - # Construct aligner - if exclusive is None: - return lambda r, s: select_best_match(r, s, is_match, score_match) - else: - return lambda r, s: select_best_match( - r, s, is_match, score_match, exclusive=exclusive - ) - + align_method = select_best_match case "corppa": # Validate input parameters if exclusive is not None: raise ValueError( f"Strategy {strategy} does not use exclusive parameter" ) - # Construct aligner - if is_match is not None and score_match is not None: - return lambda r, s: corppa_align( - r, s, is_match=is_match, score_match=score_match - ) - elif is_match is not None: - return lambda r, s: corppa_align(r, s, is_match=is_match) - elif score_match is not None: - return lambda r, s: corppa_align(r, s, score_match=score_match) - else: - return lambda r, s: corppa_align(r, s) + align_method = corppa_align case _: raise ValueError(f"Unknown alignment strategy: {strategy}") + # Determine optional args + options = {} + if is_match is not None: + options["is_match"] = is_match + if score_match is not None: + options["score_match"] = score_match + if exclusive is not None: + options["exclusive"] = exclusive + # Construct aligner + return lambda r, s: align_method(r, s, **options) # ty: ignore[invalid-argument-type] diff --git a/src/spanerr/compute_metrics.py b/src/spanerr/compute_metrics.py index ab3b8c2..3e623ef 100644 --- a/src/spanerr/compute_metrics.py +++ b/src/spanerr/compute_metrics.py @@ -31,6 +31,9 @@ """ import argparse +import csv +import os +import sys from collections.abc import Callable, Iterator from pathlib import Path @@ -106,6 +109,7 @@ def get_span_alignments( sys_file: Path, aligner: AlignSpans, ignore_unmatched: bool = False, + binarize: bool = False, ) -> Iterator[SpanAlignment]: """ Yields document-level span alignments given two sets of span annotations (JSONL). @@ -122,6 +126,8 @@ def get_span_alignments( sys_annos = {} for sys_dict in orjsonl.stream(sys_file): anno = DocSpans.from_dict(sys_dict) # ty: ignore[invalid-argument-type] + if binarize: + anno = anno.binarize() doc_id = anno.doc_id # Validate system annotations by checking for duplicate doc ids if doc_id in sys_annos: @@ -132,6 +138,8 @@ def get_span_alignments( ref_doc_ids = set() # for tracking encountered reference doc ids for ref_json in orjsonl.stream(ref_file): ref_anno = DocSpans.from_dict(ref_json) # ty: ignore[invalid-argument-type] + if binarize: + ref_anno = ref_anno.binarize() doc_id = ref_anno.doc_id # Validate reference annotations by checking for duplicate doc ids if doc_id in ref_doc_ids: @@ -158,6 +166,7 @@ def compute_entity_metrics( alignments: Iterator[SpanAlignment], scorer: ScoreAlignment, beta: float = 1, + save_intmd: Path | None = None, show_progress: bool = True, ) -> dict[str, float]: """ @@ -178,19 +187,37 @@ def compute_entity_metrics( total_sys_spans = 0 total_ref_spans = 0 progress = tqdm(alignments, desc="Scoring alignments", disable=not show_progress) - for a in progress: - n_docs += 1 - rel_score = scorer(a) - n_sys_spans = len(a.sys.spans) - n_ref_spans = len(a.ref.spans) - if show_progress: - tqdm.write( - f" * {a.ref.doc_id}: relevance = {rel_score:.4g} | " - f"{n_ref_spans} ref spans | {n_sys_spans} sys spans" - ) - total_relevance += rel_score - total_sys_spans += n_sys_spans - total_ref_spans += n_ref_spans + # Set intermediate results path to os.devnull when unset for type checking + intmd_path = save_intmd if save_intmd else Path(os.devnull) + with intmd_path.open(mode="w", newline="") as intmd_file: + # Optionally set up intermdiate results CSV + if save_intmd: + fields = ["doc_id", "n_ref", "n_sys", "relevance"] + intmd_writer = csv.DictWriter(intmd_file, fields) + intmd_writer.writeheader() + for a in progress: + n_docs += 1 + rel_score = scorer(a) + n_sys_spans = len(a.sys.spans) + n_ref_spans = len(a.ref.spans) + # Optionally save intermediate results + if save_intmd: + intmd_writer.writerow( + { + "doc_id": a.ref.doc_id, + "n_ref": n_ref_spans, + "n_sys": n_sys_spans, + "relevance": rel_score, + } + ) + if show_progress: + tqdm.write( + f" * {a.ref.doc_id}: relevance = {rel_score:.4g} | " + f"{n_ref_spans} ref spans | {n_sys_spans} sys spans" + ) + total_relevance += rel_score + total_sys_spans += n_sys_spans + total_ref_spans += n_ref_spans # Raise error if iterator contains no alignments if n_docs == 0: raise ValueError("Found no alignments to score") @@ -210,6 +237,7 @@ def compute_document_metrics( alignments: Iterator[SpanAlignment], scorer: ScoreAlignment, beta: float = 1, + save_intmd: Path | None = None, show_progress: bool = True, ) -> dict[str, float]: """ @@ -230,22 +258,55 @@ def compute_document_metrics( cumulative_recall = 0 cumulative_fscore = 0 progress = tqdm(alignments, desc="Scoring alignments", disable=not show_progress) - for a in progress: - n_docs += 1 - relevance = scorer(a) - # Compute and accumulate precision and recall - doc_precision = precision(len(a.sys.spans), relevance) - doc_recall = recall(len(a.ref.spans), relevance) - doc_fscore = f_beta(beta, doc_precision, doc_recall) - if show_progress: - tqdm.write( - f" * {a.ref.doc_id}: precision = {doc_precision:.4g} | " - f"recall = {doc_recall:.4g} | F-{beta} = {doc_fscore:.4g}" - ) - # Add to running totals - cumulative_precision += doc_precision - cumulative_recall += doc_recall - cumulative_fscore += doc_fscore + # Set intermediate results path to os.devnull when unset for type checking + intmd_path = save_intmd if save_intmd else Path(os.devnull) + with intmd_path.open(mode="w", newline="") as intmd_file: + # Optionally set up intermdiate results CSV + if save_intmd: + fields = [ + "doc_id", + "n_ref", + "n_sys", + "relevance", + "precision", + "recall", + f"f-{beta}", + ] + intmd_writer = csv.DictWriter(intmd_file, fields) + intmd_writer.writeheader() + + for a in progress: + n_docs += 1 + relevance = scorer(a) + n_sys_spans = len(a.sys.spans) + n_ref_spans = len(a.ref.spans) + # Compute and accumulate precision and recall + doc_precision = precision(n_sys_spans, relevance) + doc_recall = recall(n_ref_spans, relevance) + doc_fscore = f_beta(beta, doc_precision, doc_recall) + # Optionally save intermediate results + if save_intmd: + intmd_writer.writerow( + { + "doc_id": a.ref.doc_id, + "n_ref": n_ref_spans, + "n_sys": n_sys_spans, + "relevance": relevance, + "precision": doc_precision, + "recall": doc_recall, + f"f-{beta}": doc_fscore, + } + ) + if show_progress: + tqdm.write( + f" * {a.ref.doc_id}: precision = {doc_precision:.4g} | " + f"recall = {doc_recall:.4g} | F-{beta} = {doc_fscore:.4g}" + ) + # Add to running totals + cumulative_precision += doc_precision + cumulative_recall += doc_recall + cumulative_fscore += doc_fscore + # Raise error if iterator contains no alignments if n_docs == 0: raise ValueError("Found no alignments to score") @@ -268,6 +329,8 @@ def compute_macro_metrics( aligner: AlignSpans, scorer: ScoreAlignment, beta: float = 1, + binarize: bool = False, + save_intmd: Path | None = None, show_progress: bool = True, ) -> dict[str, float]: """ @@ -281,16 +344,24 @@ def compute_macro_metrics( - macro F-score: float """ # Step 1: Span Alignments - alignments = get_span_alignments(ref_jsonl, sys_jsonl, aligner) + alignments = get_span_alignments(ref_jsonl, sys_jsonl, aligner, binarize=binarize) # Step 2: Compute macro metrics match macro_level: case "entity": return compute_entity_metrics( - alignments, scorer, beta=beta, show_progress=show_progress + alignments, + scorer, + beta=beta, + save_intmd=save_intmd, + show_progress=show_progress, ) case "document": return compute_document_metrics( - alignments, scorer, beta=beta, show_progress=show_progress + alignments, + scorer, + beta=beta, + save_intmd=save_intmd, + show_progress=show_progress, ) case _: raise ValueError(f"Unsupported macro level: {macro_level}") @@ -331,6 +402,16 @@ def main(): help="Strategy for scoring partial matches", ) # Optional arguments + parser.add_argument( + "--binary", + help='Collapse labels into single "True" label', + action="store_true", + ) + parser.add_argument( + "--save-intermediate", + help="Filename where intermediate (document-level) results should be written (CSV file)", + type=Path, + ) parser.add_argument( "--progress", help="Show progress", @@ -338,6 +419,16 @@ def main(): default=True, ) args = parser.parse_args() + + # Validate intermediate results file if specified + intmd_file = args.save_intermediate + if intmd_file and intmd_file.is_file(): + print( + f"Intermediate results file {intmd_file} already exists. Not overwriting.", + file=sys.stderr, + ) + sys.exit(1) + # Compute aggregated evaluation metrics results = compute_macro_metrics( args.ref_jsonl, @@ -345,6 +436,8 @@ def main(): args.macro_level, get_aligner(args.alignment_method), get_scorer(args.scoring_method), + binarize=args.binary, + save_intmd=intmd_file, show_progress=args.progress, ) if args.progress: diff --git a/src/spanerr/core.py b/src/spanerr/core.py index 840d9f7..4a3ec9a 100644 --- a/src/spanerr/core.py +++ b/src/spanerr/core.py @@ -93,12 +93,18 @@ def overlap_factor(self, other: Self) -> float: overlap = self.overlap_length(other) return overlap / max(len(self), len(other)) + def relabel(self, label: str) -> Self: + """ + Returns the relabeled version of this span. + """ + return self.__class__(self.start, self.end, label) + def binarize(self) -> Self: """ Returns the "binarized" version of this span (i.e., sets label to default) """ default_label = self.__class__.label - return self.__class__(self.start, self.end, default_label) + return self.relabel(default_label) def merge(self, other: Self) -> Self: """ diff --git a/tests/test_align.py b/tests/test_align.py index fd690d4..43727dd 100644 --- a/tests/test_align.py +++ b/tests/test_align.py @@ -197,41 +197,31 @@ def test_construct_aligner(mock_first, mock_best, mock_corppa): construct_aligner(strategy) # select_first strategy = "select_first" - ## Missing is_match parameter - err_msg = "Strategy select_first requires is_match parameter" - with pytest.raises(ValueError, match=err_msg): - construct_aligner(strategy) ## Includes extra score_match parameter err_msg = "Strategy select_first does not use score_match parameter" with pytest.raises(ValueError, match=err_msg): construct_aligner(strategy, is_match="test", score_match="score") ## Default - aligner = construct_aligner(strategy, is_match="test") + aligner = construct_aligner(strategy) assert callable(aligner) _ = aligner("ref_span", "sys_span") - mock_first.assert_called_once_with("ref_span", "sys_span", "test") - ## Set optional exclusive flag + mock_first.assert_called_once_with("ref_span", "sys_span") + ## Set optional args mock_first.reset_mock() aligner = construct_aligner(strategy, is_match="test", exclusive="flag") assert callable(aligner) _ = aligner("ref_span", "sys_span") - mock_first.assert_called_once_with("ref_span", "sys_span", "test", exclusive="flag") + mock_first.assert_called_once_with( + "ref_span", "sys_span", is_match="test", exclusive="flag" + ) # select_best strategy = "select_best" - ## Missing required input parameters - err_msg = "Strategy select_best requires is_match and score_match parameters" - with pytest.raises(ValueError, match=err_msg): - construct_aligner(strategy) - with pytest.raises(ValueError, match=err_msg): - construct_aligner(strategy, is_match="test") - with pytest.raises(ValueError, match=err_msg): - construct_aligner(strategy, score_match="score") ## Default aligner = construct_aligner(strategy, is_match="test", score_match="score") assert callable(aligner) _ = aligner("ref_span", "sys_span") mock_best.assert_called_once_with("ref_span", "sys_span", "test", "score") - ## Set optional exclusive flag + ## Set optional args mock_best.reset_mock() aligner = construct_aligner( strategy, is_match="test", score_match="score", exclusive="flag" @@ -239,7 +229,7 @@ def test_construct_aligner(mock_first, mock_best, mock_corppa): assert callable(aligner) _ = aligner("ref_span", "sys_span") mock_best.assert_called_once_with( - "ref_span", "sys_span", "test", "score", exclusive="flag" + "ref_span", "sys_span", is_match="test", score_match="score", exclusive="flag" ) # corppa strategy = "corppa" @@ -252,19 +242,7 @@ def test_construct_aligner(mock_first, mock_best, mock_corppa): assert callable(aligner) _ = aligner("ref_span", "sys_span") mock_corppa.assert_called_once_with("ref_span", "sys_span") - ## Set optional is_match parameter - mock_corppa.reset_mock() - aligner = construct_aligner(strategy, is_match="test") - assert callable(aligner) - _ = aligner("ref_span", "sys_span") - mock_corppa.assert_called_once_with("ref_span", "sys_span", is_match="test") - ## Set optional score_match parameter - mock_corppa.reset_mock() - aligner = construct_aligner(strategy, score_match="score") - assert callable(aligner) - _ = aligner("ref_span", "sys_span") - mock_corppa.assert_called_once_with("ref_span", "sys_span", score_match="score") - ## Set both optional parameters + ## Set optional args mock_corppa.reset_mock() aligner = construct_aligner(strategy, is_match="test", score_match="score") assert callable(aligner) diff --git a/tests/test_compute_metrics.py b/tests/test_compute_metrics.py index 8a996c9..a33fc3f 100644 --- a/tests/test_compute_metrics.py +++ b/tests/test_compute_metrics.py @@ -222,6 +222,33 @@ def test_get_span_alignments(mock_alignment, tmp_path): mock_aligner.assert_called_once_with(ref_docspans[1], sys_docspans[1]) mock_alignment.assert_not_called() + # Binarize labels + mock_aligner.reset_mock() + mock_alignment.reset_mock() + ref_labeled_lines = [ + '{"doc_id":"a","spans":[{"start":1,"end":2,"label":"a"}]}', + '{"doc_id":"b","spans":[{"start":3,"end":4,"label":"b"}]}', + ] + sys_labeled_lines = [ + '{"doc_id":"a","spans":[{"start":4,"end":5,"label":"c"}]}', + '{"doc_id":"b","spans":[{"start":6,"end":7,"label":"d"}]}', + ] + ref_labeled_jsonl = tmp_path / "ref_labeled.jsonl" + sys_labeled_jsonl = tmp_path / "sys_labeled.jsonl" + ref_labeled_jsonl.write_text("\n".join(ref_labeled_lines) + "\n") + sys_labeled_jsonl.write_text("\n".join(sys_labeled_lines) + "\n") + ref_docspans = [DocSpans.from_dict(json.loads(l)).binarize() for l in ref_lines] + sys_docspans = [DocSpans.from_dict(json.loads(l)).binarize() for l in sys_lines] + result = get_span_alignments(ref_jsonl, sys_jsonl, mock_aligner, binarize=True) + assert set(result) == {"a", "b"} + assert mock_aligner.call_count == 2 + expected_calls = [ + call(ref_docspans[0], sys_docspans[0].binarize()), + call(ref_docspans[1], sys_docspans[1].binarize()), + ] + mock_aligner.assert_has_calls(expected_calls, any_order=True) + mock_alignment.assert_not_called() + @pytest.mark.parametrize( "jsonl_text,dup_id", @@ -307,6 +334,38 @@ def test_compute_entity_metrics(mock_precision, mock_recall, mock_fscore): mock_fscore.assert_called_once_with(0.5, "p", "r") +@patch("spanerr.compute_metrics.f_beta", autospec=True, return_value="f") +@patch("spanerr.compute_metrics.recall", autospec=True, return_value="r") +@patch("spanerr.compute_metrics.precision", autospec=True, return_value="p") +def test_compute_entity_metrics_save_intmd( + mock_precision, mock_recall, mock_fscore, tmp_path +): + intmd_results = tmp_path / "intmd_results.csv" + + mock_scorer = Mock(autospec=ScoreAlignment, side_effect=[1, 0.2, 0.3]) + alignments = [ + SpanAlignment(DocSpans("a", []), DocSpans("a", []), {}), + SpanAlignment( + DocSpans("b", [Span(1, 2), Span(3, 4)]), DocSpans("b", [Span(4, 5)]), {} + ), + SpanAlignment( + DocSpans("c", [Span(6, 7)]), + DocSpans("c", [Span(7, 8), Span(8, 9), Span(9, 10)]), + {}, + ), + ] + expected_lines = [ + "doc_id,n_ref,n_sys,relevance", + "a,0,0,1", + f"b,2,1,{0.2}", + f"c,1,3,{0.3}", + ] + _ = compute_entity_metrics( + alignments, mock_scorer, save_intmd=intmd_results, show_progress=False + ) + assert intmd_results.read_text() == "\n".join(expected_lines) + "\n" + + @patch("spanerr.compute_metrics.f_beta", autospec=True) @patch("spanerr.compute_metrics.recall", autospec=True) @patch("spanerr.compute_metrics.precision", autospec=True) @@ -384,6 +443,41 @@ def test_document_entity_metrics(mock_precision, mock_recall, mock_fscore): ) +@patch("spanerr.compute_metrics.f_beta", autospec=True) +@patch("spanerr.compute_metrics.recall", autospec=True) +@patch("spanerr.compute_metrics.precision", autospec=True) +def test_compute_document_metrics_save_intmd( + mock_precision, mock_recall, mock_fscore, tmp_path +): + intmd_results = tmp_path / "intmd_results.csv" + + mock_scorer = Mock(autospec=ScoreAlignment, side_effect=[1, 0.2, 0.3]) + mock_precision.side_effect = [0.1, 0.2, 0.3] + mock_recall.side_effect = [0.4, 0.5, 0.6] + mock_fscore.side_effect = [0, 1, 0.5] + alignments = [ + SpanAlignment(DocSpans("a", []), DocSpans("a", []), {}), + SpanAlignment( + DocSpans("b", [Span(1, 2), Span(3, 4)]), DocSpans("b", [Span(4, 5)]), {} + ), + SpanAlignment( + DocSpans("c", [Span(6, 7)]), + DocSpans("c", [Span(7, 8), Span(8, 9), Span(9, 10)]), + {}, + ), + ] + expected_lines = [ + "doc_id,n_ref,n_sys,relevance,precision,recall,f-1", + f"a,0,0,1,{0.1},{0.4},0", + f"b,2,1,{0.2},{0.2},{0.5},1", + f"c,1,3,{0.3},{0.3},{0.6},{0.5}", + ] + _ = compute_document_metrics( + alignments, mock_scorer, save_intmd=intmd_results, show_progress=False + ) + assert intmd_results.read_text() == "\n".join(expected_lines) + "\n" + + @patch( "spanerr.compute_metrics.compute_document_metrics", autospec=True, @@ -403,7 +497,7 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me # Unknown macro level with pytest.raises(ValueError, match="Unsupported macro level: unknown"): compute_macro_metrics("ref", "sys", "unknown", "aligner", "scorer") - mock_alignments.assert_called_once_with("ref", "sys", "aligner") + mock_alignments.assert_called_once_with("ref", "sys", "aligner", binarize=False) mock_entity_metrics.assert_not_called() mock_doc_metrics.assert_not_called() @@ -411,20 +505,28 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me mock_alignments.reset_mock() result = compute_macro_metrics("ref", "sys", "entity", "aligner", "scorer") assert result == "entity metrics" - mock_alignments.assert_called_once_with("ref", "sys", "aligner") + mock_alignments.assert_called_once_with("ref", "sys", "aligner", binarize=False) mock_entity_metrics.assert_called_once_with( - "alignments", "scorer", beta=1, show_progress=True + "alignments", "scorer", beta=1, save_intmd=None, show_progress=True ) mock_doc_metrics.assert_not_called() ## Setting optional parameters mock_alignments.reset_mock() mock_entity_metrics.reset_mock() assert compute_macro_metrics( - "ref", "sys", "entity", "aligner", "scorer", beta="float", show_progress="bool" + "ref", + "sys", + "entity", + "aligner", + "scorer", + beta="float", + binarize="flag", + save_intmd="path", + show_progress="bool", ) - mock_alignments.assert_called_once_with("ref", "sys", "aligner") + mock_alignments.assert_called_once_with("ref", "sys", "aligner", binarize="flag") mock_entity_metrics.assert_called_once_with( - "alignments", "scorer", beta="float", show_progress="bool" + "alignments", "scorer", beta="float", save_intmd="path", show_progress="bool" ) mock_doc_metrics.assert_not_called() @@ -433,9 +535,9 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me mock_entity_metrics.reset_mock() result = compute_macro_metrics("ref", "sys", "document", "aligner", "scorer") assert result == "doc metrics" - mock_alignments.assert_called_once_with("ref", "sys", "aligner") + mock_alignments.assert_called_once_with("ref", "sys", "aligner", binarize=False) mock_doc_metrics.assert_called_once_with( - "alignments", "scorer", beta=1, show_progress=True + "alignments", "scorer", beta=1, save_intmd=None, show_progress=True ) mock_entity_metrics.assert_not_called() ## Setting optional parameters @@ -448,11 +550,13 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me "aligner", "scorer", beta="float", + binarize="flag", + save_intmd="path", show_progress="bool", ) - mock_alignments.assert_called_once_with("ref", "sys", "aligner") + mock_alignments.assert_called_once_with("ref", "sys", "aligner", binarize="flag") mock_doc_metrics.assert_called_once_with( - "alignments", "scorer", beta="float", show_progress="bool" + "alignments", "scorer", beta="float", save_intmd="path", show_progress="bool" ) mock_entity_metrics.assert_not_called() @@ -460,7 +564,7 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me @pytest.mark.parametrize( "cli_args,call_params", [ - # all required params, default progress behavior + # all required params, with defaults for optional [ [ "compute_metrics.py", @@ -478,7 +582,29 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me "select_first", "overlap_factor", ], - {"show_progress": True}, + {"binarize": False, "save_intmd": None, "show_progress": True}, + ), + ], + # set binary flag + [ + [ + "compute_metrics.py", + "ref.jsonl", + "sys.jsonl", + "document", + "corppa", + "jaccard", + "--binary", + ], + ( + [ + Path("ref.jsonl"), + Path("sys.jsonl"), + "document", + "corppa", + "jaccard", + ], + {"binarize": True, "save_intmd": None, "show_progress": True}, ), ], # disable progress @@ -500,7 +626,7 @@ def test_compute_macro_metrics(mock_alignments, mock_entity_metrics, mock_doc_me "select_best", "jaccard", ], - {"show_progress": False}, + {"binarize": False, "save_intmd": None, "show_progress": False}, ), ], ], @@ -537,3 +663,48 @@ def test_main(mock_metrics, mock_aligner, mock_scorer, cli_args, call_params, ca ) progress_pfx = "\n" if kwargs["show_progress"] else "" assert captured.out == f"{progress_pfx}{expected_reporting}\n" + + +@patch("spanerr.compute_metrics.get_scorer", return_value="scorer") +@patch("spanerr.compute_metrics.get_aligner", return_value="aligner") +@patch("spanerr.compute_metrics.compute_macro_metrics") +def test_main_save_intmd(mock_metrics, mock_aligner, mock_scorer, capsys, tmp_path): + mock_metrics.return_value = { + "n_docs": 0, + "precision": 0.1, + "recall": 0.1, + "f-score": 0.1, + } + intmd_file = tmp_path / "intmd_results.csv" + cli_args = [ + "compute_metrics.py", + "ref.jsonl", + "sys.jsonl", + "document", + "select_best", + "jaccard", + "--save-intermediate", + str(intmd_file), + ] + + # Save intermediate results (specified file does not exist) + ## patch in test args for argpars to parse + with patch("sys.argv", cli_args): + main() + mock_aligner.assert_called_once_with("select_best") + mock_scorer.assert_called_once_with("jaccard") + # Swap final args for the expected return values (based on patching) + args = [Path("ref.jsonl"), Path("sys.jsonl"), "document", "aligner", "scorer"] + kwargs = {"binarize": False, "save_intmd": intmd_file, "show_progress": True} + mock_metrics.assert_called_once_with(*args, **kwargs) + + # Raises ValueError if file already exists + intmd_file.write_text("some text") + with patch("sys.argv", cli_args): + with pytest.raises(SystemExit) as execinfo: + main() + assert execinfo.value.code == 1 + captured = capsys.readouterr() + assert captured.err.startswith( + f"Intermediate results file {intmd_file} already exists" + ) diff --git a/tests/test_core.py b/tests/test_core.py index 6a6d7f5..2d2ba51 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -152,6 +152,10 @@ def test_overlap_factor(self, mock_overlap_length): span_b = Span(2, 8) assert span_a.overlap_factor(span_b) == 3 / 6 + def test_relabel(self): + assert Span(1, 4, "label").relabel("new") == Span(1, 4, "new") + assert Span(3, 5).relabel("other") == Span(3, 5, "other") + def test_binarize(self): assert Span(1, 4, "label").binarize() == Span(1, 4, "") assert Span(3, 5).binarize() == Span(3, 5) diff --git a/uv.lock b/uv.lock index 3a8ea00..b7103de 100644 --- a/uv.lock +++ b/uv.lock @@ -3,7 +3,7 @@ revision = 3 requires-python = ">=3.12" [options] -exclude-newer = "2026-09-01T18:06:14.572133Z" +exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values. exclude-newer-span = "P1W" [[package]]