Skip to content
Merged
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
19 changes: 19 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,25 @@ QUERY_FALLBACK_ENABLED=true
When disabled, the service runs only the original generated query. When enabled,
it may try later fallback candidates produced by the parser.

### Coreference preprocessing

Coreference resolution is disabled by default. To enable it, install the
optional `fastcoref` dependency/model and set:

```bash
COREFERENCE_ENABLED=true
COREFERENCE_MODEL=biu-nlp/f-coref
COREFERENCE_MIN_CONFIDENCE=0.65
```

The resolver runs before document chunking and leaves text unchanged if the
optional model is unavailable. Compare benchmark runs with:

```bash
python3 benchmark_parsers.py --suite-file data/benchmarks/stress25_v1.json --parsers canonical_pln --no-coreference
python3 benchmark_parsers.py --suite-file data/benchmarks/stress25_v1.json --parsers canonical_pln --coreference
```

## ConceptNet Background Knowledge

ConceptNet can be loaded as readonly background knowledge and indexed into the
Expand Down
14 changes: 14 additions & 0 deletions benchmark_parsers.py
Original file line number Diff line number Diff line change
Expand Up @@ -638,6 +638,16 @@ async def main() -> int:
help="Print per-case progress to stderr",
)
cli.add_argument("--quick", action="store_true", help="Run a reduced representative case set")
coref_group = cli.add_mutually_exclusive_group()
coref_group.add_argument(
"--coreference", dest="coreference", action="store_true",
help="Enable optional coreference preprocessing",
)
coref_group.add_argument(
"--no-coreference", dest="coreference", action="store_false",
help="Disable coreference preprocessing",
)
cli.set_defaults(coreference=None)
cli.add_argument(
"--output-dir",
default="data/benchmarks",
Expand All @@ -657,6 +667,9 @@ async def main() -> int:
)
args = cli.parse_args()

if args.coreference is not None:
os.environ["COREFERENCE_ENABLED"] = str(args.coreference).lower()

if args.suite_file:
suite_path = Path(args.suite_file)
suite_metadata, cases = _select_cases_from_file(suite_path, args.quick)
Expand All @@ -681,6 +694,7 @@ async def main() -> int:
payload: dict[str, object] = {
"run_id": run_id,
"conceptnet_enabled": False,
"coreference_enabled": bool(get_settings().coreference_enabled),
"mode": args.mode,
"suite": suite_label,
"suite_metadata": suite_metadata,
Expand Down
5 changes: 5 additions & 0 deletions config.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,11 @@ class Settings(BaseSettings):
parser_batch_sentences: int = 4
parser_batch_max_chars: int = 2000

# Optional document-level coreference preprocessing
coreference_enabled: bool = False
coreference_model: str = "biu-nlp/f-coref"
coreference_min_confidence: float = 0.65

# Reasoning
chaining_timeout: int = 30 # seconds before proof search is killed
chaining_max_steps: int = 100
Expand Down
124 changes: 124 additions & 0 deletions core/coreference.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,124 @@
"""Optional document-level coreference resolution."""

from __future__ import annotations

import logging
from dataclasses import dataclass, field
from typing import Any


logger = logging.getLogger(__name__)


@dataclass
class ResolvedDocument:
"""Original and parser-facing text plus optional resolution diagnostics."""

original: str
resolved: str
mentions: list[dict[str, Any]] = field(default_factory=list)


class CoreferenceResolver:
"""Resolve text with fastcoref when enabled, otherwise leave it unchanged."""

def __init__(
self,
enabled: bool = False,
model_name: str = "biu-nlp/f-coref",
min_confidence: float = 0.65,
) -> None:
self.enabled = enabled
self.model_name = model_name
self.min_confidence = min(1.0, max(0.0, min_confidence))
self._model: Any | None = None
self._load_failed = False

def resolve(self, text: str) -> ResolvedDocument:
"""Return resolved text, failing open to the original input."""
original = text or ""
if not self.enabled or not original.strip():
return ResolvedDocument(original=original, resolved=original)

try:
return self._resolve_with_fastcoref(original)
except Exception as exc: # Optional dependency/model must not break ingest.
if not self._load_failed:
logger.warning("Coreference resolver unavailable; using original text: %s", exc)
self._load_failed = True
return ResolvedDocument(original=original, resolved=original)

def _resolve_with_fastcoref(self, text: str) -> ResolvedDocument:
if self._model is None:
from fastcoref import FCoref
from fastcoref.coref_models.modeling_fcoref import FCorefModel

# Patch for compatibility with newer versions of transformers
FCorefModel.all_tied_weights_keys = {}

self._model = FCoref(model_name_or_path=self.model_name)

prediction = self._model.predict(texts=[text])[0]
clusters = prediction.get_clusters(as_strings=False)
replacements: list[tuple[int, int, str]] = []
mentions: list[dict[str, Any]] = []

pronouns = {
"he", "him", "his",
"she", "her", "hers",
"it", "its",
"they", "them", "their", "theirs"
}
possessives = {"his", "her", "hers", "its", "their", "theirs"}

for cluster_id, cluster in enumerate(clusters or []):
if len(cluster) < 2:
continue
antecedent_start, antecedent_end = self._span(cluster[0])
antecedent = text[antecedent_start:antecedent_end]
if not antecedent.strip():
continue
for span in cluster[1:]:
start, end = self._span(span)
mention = text[start:end]
mention_lower = mention.casefold()

# Restrict to exact pronoun matches to avoid breaking rigid grammar rules
if not mention.strip() or mention_lower == antecedent.casefold() or mention_lower not in pronouns:
continue

# Handle possessive pronoun replacement (e.g., 'his' -> 'John\'s')
replacement = antecedent
if mention_lower in possessives:
if not replacement.endswith("'s") and not replacement.endswith("'"):
replacement += "'" if replacement.endswith("s") else "'s"

replacements.append((start, end, replacement))
mentions.append(
{
"cluster_id": cluster_id,
"mention": mention,
"antecedent": replacement,
"start": start,
"end": end,
"confidence": 1.0,
}
)

resolved = self._replace_spans(text, replacements)
return ResolvedDocument(original=text, resolved=resolved, mentions=mentions)

@staticmethod
def _span(span: Any) -> tuple[int, int]:
if not isinstance(span, (list, tuple)) or len(span) != 2:
raise ValueError(f"Invalid coreference span: {span!r}")
start, end = int(span[0]), int(span[1])
if start < 0 or end <= start:
raise ValueError(f"Invalid coreference span: {span!r}")
return start, end

@staticmethod
def _replace_spans(text: str, replacements: list[tuple[int, int, str]]) -> str:
for start, end, replacement in sorted(replacements, reverse=True):
text = text[:start] + replacement + text[end:]
return text
12 changes: 10 additions & 2 deletions core/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

from config import get_settings
from core.chunker import Chunker
from core.coreference import CoreferenceResolver
from core.parser import SemanticParser
from core.reasoner import Reasoner
from core.answer_generator import AnswerGenerator
Expand All @@ -29,6 +30,11 @@ class PLNRAGService:
def __init__(self, parser: SemanticParser):
cfg = get_settings()
self._parser = parser
self._coreference = CoreferenceResolver(
enabled=cfg.coreference_enabled,
model_name=cfg.coreference_model,
min_confidence=cfg.coreference_min_confidence,
)
create_chunker = getattr(parser, "create_chunker", None)
self._chunker = create_chunker() if callable(create_chunker) else Chunker()
self._reasoner = Reasoner()
Expand Down Expand Up @@ -57,12 +63,14 @@ async def ingest_batch(self, texts: List[str]) -> List[IngestItemResult]:

def _ingest_single(self, text: str) -> IngestItemResult:
try:
resolved = self._coreference.resolve(text)
parser_text = resolved.resolved
all_atoms: List[str] = []
rejected: List[dict] = []
supports_batch_parse = (
self._parser.__class__.parse_batch is not SemanticParser.parse_batch
)
chunk_units = self._chunker.chunk(text)
chunk_units = self._chunker.chunk(parser_text)
chunk_count = len(chunk_units)
batch_count = 0
batch_sizes: List[int] = []
Expand All @@ -71,7 +79,7 @@ def _ingest_single(self, text: str) -> IngestItemResult:

if supports_batch_parse:
batches = self._chunker.batch_chunks(
text,
parser_text,
max_sentences=get_settings().parser_batch_sentences,
max_chars=get_settings().parser_batch_max_chars,
)
Expand Down
1 change: 1 addition & 0 deletions requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ dspy==3.1.3
litellm==1.95.0
openai>=1.0.0
spacy>=3.7.0
fastcoref>=2.1.6
faiss-cpu>=1.8.0
numpy>=1.26.0
networkx>=3.0.0
Expand Down