Skip to content
This repository was archived by the owner on Dec 16, 2022. It is now read-only.

Commit a8a3486

Browse files
authored
update conll (#298)
1 parent 54de9d6 commit a8a3486

3 files changed

Lines changed: 13 additions & 7 deletions

File tree

CHANGELOG.md

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
99

1010
### Added
1111

12-
- Added some additional `__init__()` parameters to the `T5` model in `allennlp_models.generation` for customizing
12+
- Added some additional `__init__()` parameters to the `T5` model in `allennlp_models.generation` for customizing.
1313
beam search and other options.
1414
- Added a configuration file for fine-tuning `t5-11b` on CCN-DM (requires at least 8 GPUs).
1515
- Added a configuration to train on the PIQA dataset with AllenNLP Tango.
@@ -18,8 +18,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1818

1919
### Fixed
2020

21-
- Fixed tests for Spacy versions greater than 3.1
22-
- Fixed the last step decoding when training CopyNet
21+
- Fixed tests for Spacy versions greater than 3.1.
22+
- Fixed the last step decoding when training CopyNet.
23+
- Allow singleton clusters in `ConllCorefScores`.
2324

2425
### Changed
2526

allennlp_models/coref/metrics/conll_coref_scores.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,9 @@ class ConllCorefScores(Metric):
1616

1717
supports_distributed = True
1818

19-
def __init__(self) -> None:
19+
def __init__(self, allow_singletons=False) -> None:
2020
self.scorers = [Scorer(m) for m in (Scorer.muc, Scorer.b_cubed, Scorer.ceafe)]
21+
self.allow_singletons = allow_singletons
2122

2223
@overrides
2324
def __call__(
@@ -56,7 +57,7 @@ def __call__(
5657
for i, metadata in enumerate(metadata_list):
5758
gold_clusters, mention_to_gold = self.get_gold_clusters(metadata["clusters"])
5859
predicted_clusters, mention_to_predicted = self.get_predicted_clusters(
59-
top_spans[i], antecedent_indices[i], predicted_antecedents[i]
60+
top_spans[i], antecedent_indices[i], predicted_antecedents[i], self.allow_singletons
6061
)
6162
for scorer in self.scorers:
6263
scorer.update(
@@ -91,6 +92,7 @@ def get_predicted_clusters(
9192
top_spans: torch.Tensor, # (num_spans, 2)
9293
antecedent_indices: torch.Tensor, # (num_spans, num_antecedents)
9394
predicted_antecedents: torch.Tensor, # (num_spans,)
95+
allow_singletons: bool,
9496
) -> Tuple[
9597
List[Tuple[Tuple[int, int], ...]], Dict[Tuple[int, int], Tuple[Tuple[int, int], ...]]
9698
]:
@@ -104,7 +106,10 @@ def get_predicted_clusters(
104106
# Find predicted index in the antecedent spans.
105107
predicted_index = antecedent_indices[i, predicted_antecedent]
106108
# Must be a previous span.
107-
assert i > predicted_index
109+
if allow_singletons:
110+
assert i >= predicted_index
111+
else:
112+
assert i > predicted_index
108113
antecedent_span: Tuple[int, int] = tuple( # type: ignore
109114
top_spans[predicted_index].tolist()
110115
)

tests/coref/metrics/conll_coref_scores_test.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@ def test_get_predicted_clusters(self, device: str):
1717
antecedent_indices = torch.tensor([[-1, -1, -1], [0, -1, -1], [0, 1, -1]], device=device)
1818
predicted_antecedents = torch.tensor([-1, -1, 1], device=device)
1919
clusters, mention_to_cluster = ConllCorefScores.get_predicted_clusters(
20-
top_spans, antecedent_indices, predicted_antecedents
20+
top_spans, antecedent_indices, predicted_antecedents, allow_singletons=False
2121
)
2222
assert len(clusters) == 1
2323
assert set(clusters[0]) == {(4, 6), (8, 9)}

0 commit comments

Comments
 (0)