Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 32 additions & 21 deletions spacy_llm/tasks/rel/util.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
import re
import warnings
from typing import Iterable, List, Optional, Tuple
from typing import Iterable, List, Optional

from spacy import Vocab
from spacy.tokens import Doc, Span
from spacy.tokens import Doc
from spacy.training import Example

from ...compat import Self
Expand Down Expand Up @@ -48,34 +48,45 @@ def to_doc(self) -> Doc:
for i, word in enumerate(doc_words)
]
doc = Doc(words=doc_words, spaces=doc_spaces, vocab=Vocab(strings=doc_words))
char_offsets = _map_char_offsets(self.text, doc.text)

# Set entities after finding correct indices.
conv_ent_indices: List[Tuple[int, int]] = []
if len(self.ents):
ent_idx = 0
for token in doc:
if token.idx == self.ents[ent_idx].start_char:
conv_ent_indices.append((token.i, -1))
if token.idx + len(token.text) == self.ents[ent_idx].end_char:
conv_ent_indices[-1] = (conv_ent_indices[-1][0], token.i + 1)
ent_idx += 1
if ent_idx == len(self.ents):
break

# Set entities using offsets from the original, unnormalized text.
doc.ents = [
Span( # noqa: E731
doc=doc,
start=ent_idx[0],
end=ent_idx[1],
label=self.ents[i].label,
doc.char_span(
char_offsets[entity.start_char],
char_offsets[entity.end_char],
label=entity.label,
)
for i, ent_idx in enumerate(conv_ent_indices)
for entity in self.ents
]
doc.user_data["rel"] = self.relations

return doc


def _map_char_offsets(source: str, target: str) -> List[int]:
"""Map source character boundaries to a target with inserted spaces."""
offsets = [0] * (len(source) + 1)
source_index = 0
target_index = 0

while source_index < len(source):
while (
target_index < len(target)
and target[target_index] == " "
and target[target_index] != source[source_index]
):
target_index += 1
if target_index >= len(target) or target[target_index] != source[source_index]:
raise ValueError("Unable to map normalized text offsets")
offsets[source_index] = target_index
source_index += 1
target_index += 1

offsets[len(source)] = target_index
return offsets


def reduce_shards_to_doc(task: RELTask, shards: Iterable[Doc]) -> Doc:
"""Reduces shards to docs for RELTask.
task (RELTask): Task.
Expand Down
14 changes: 14 additions & 0 deletions tests/test_rel_util.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
from spacy_llm.tasks.rel.items import EntityItem
from spacy_llm.tasks.rel.util import RELExample


def test_to_doc_maps_entity_offsets_after_punctuation_normalization():
example = RELExample(
text="Alice, Bob works.",
ents=[EntityItem(start_char=7, end_char=10, label="PERSON")],
relations=[],
)

doc = example.to_doc()

assert [(ent.text, ent.label_) for ent in doc.ents] == [("Bob", "PERSON")]