Skip to content
Draft
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
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -216,3 +216,5 @@ __marimo__/

# Streamlit
.streamlit/secrets.toml
data/

95 changes: 95 additions & 0 deletions datasets/amfv_datasets/chunking.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
""""""

from __future__ import annotations

import re

from .corpus import Chunk, Document

_HEADING = re.compile(r"^\s{0,3}(#{1,6}\s|\d+(\.\d+)*\s+\S|[A-Z][A-Z \-]{6,}$)")


def _segments(text: str) -> list[str]:
""""""
raw = re.split(r"\n\s*\n", text.replace("\r\n", "\n"))
return [s.strip() for s in raw if s.strip()]


def chunk_document(
doc: Document,
target_words: int = 350,
overlap_words: int = 40,
min_words: int = 20,
) -> list[Chunk]:
""""""
segments = _segments(doc.text)
chunks: list[Chunk] = []
buf: list[str] = []
buf_words = 0
ordinal = 0

def flush() -> None:
nonlocal buf, buf_words, ordinal
if not buf:
return
body = "\n\n".join(buf).strip()
if len(body.split()) >= min_words or not chunks:
chunks.append(
Chunk(
chunk_id=f"{doc.doc_id}::{ordinal}",
doc_id=doc.doc_id,
source=doc.source,
title=doc.title,
text=body,
ordinal=ordinal,
url=doc.url,
metadata=dict(doc.metadata),
)
)
ordinal += 1
buf, buf_words = [], 0

for seg in segments:
seg_words = len(seg.split())

if seg_words > target_words and not _HEADING.match(seg):
flush()
words = seg.split()
step = max(1, target_words - overlap_words)
for start in range(0, len(words), step):
window = words[start : start + target_words]
if len(window) < min_words and chunks:
break
chunks.append(
Chunk(
chunk_id=f"{doc.doc_id}::{ordinal}",
doc_id=doc.doc_id,
source=doc.source,
title=doc.title,
text=" ".join(window),
ordinal=ordinal,
url=doc.url,
metadata=dict(doc.metadata),
)
)
ordinal += 1
continue

if buf_words + seg_words > target_words and buf:
flush()

if overlap_words and chunks:
tail = chunks[-1].text.split()[-overlap_words:]
buf, buf_words = [" ".join(tail)], len(tail)
buf.append(seg)
buf_words += seg_words

flush()
return chunks


def chunk_documents(docs, **kwargs) -> list[Chunk]:
out: list[Chunk] = []
for doc in docs:
out.extend(chunk_document(doc, **kwargs))
return out
Empty file.
51 changes: 51 additions & 0 deletions datasets/amfv_datasets/cli/build_index.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
""""""

from __future__ import annotations

import argparse

from ..chunking import chunk_documents
from ..corpus import read_chunks, read_documents, write_jsonl


def main() -> None:
ap = argparse.ArgumentParser(description="Chunk documents and build a retrieval index")
ap.add_argument("--documents", help="documents.jsonl (skip if --chunks given)")
ap.add_argument("--chunks", help="prebuilt chunks.jsonl (skip chunking)")
ap.add_argument("--chunks-out", default=None, help="where to write chunks.jsonl")
ap.add_argument("--backend", choices=["bm25", "colbert"], required=True)
ap.add_argument("--index-dir", required=True)
ap.add_argument("--model", default=None, help="override ColBERT model (default LateOn)")
ap.add_argument("--target-words", type=int, default=350)
ap.add_argument("--overlap-words", type=int, default=40)
args = ap.parse_args()

if args.chunks:
chunks = read_chunks(args.chunks)
elif args.documents:
docs = list(read_documents(args.documents))
chunks = chunk_documents(docs, target_words=args.target_words, overlap_words=args.overlap_words)
out = args.chunks_out or f"{args.index_dir.rstrip('/')}.chunks.jsonl"
write_jsonl(out, chunks)
print(f"Chunked {len(docs)} docs -> {len(chunks)} chunks ({out})")
else:
ap.error("provide --documents or --chunks")

if args.backend == "bm25":
from ..retrieval.bm25 import BM25Retriever

r = BM25Retriever()
r.index(chunks)
r.save(args.index_dir)
else:
from ..retrieval.colbert import DEFAULT_MODEL, ColBERTRetriever

r = ColBERTRetriever(model_name=args.model or DEFAULT_MODEL, index_dir=args.index_dir)
r.index(chunks)
r.save(args.index_dir)

print(f"Built {args.backend} index over {len(chunks)} chunks -> {args.index_dir}")


if __name__ == "__main__":
main()
53 changes: 53 additions & 0 deletions datasets/amfv_datasets/cli/query.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
""""""

from __future__ import annotations

import argparse
import textwrap


def main() -> None:
ap = argparse.ArgumentParser(description="Query an AMFV retrieval index")
ap.add_argument("-q", "--query", required=True)
ap.add_argument("-k", type=int, default=5)
ap.add_argument("--bm25-dir")
ap.add_argument("--colbert-dir")
ap.add_argument(
"--weights",
type=float,
nargs="+",
default=None,
help="RRF weights, order: bm25 then colbert (only if both given)",
)
ap.add_argument("--fetch-k", type=int, default=50)
args = ap.parse_args()

retrievers = []
if args.bm25_dir:
from ..retrieval.bm25 import BM25Retriever

retrievers.append(BM25Retriever.load(args.bm25_dir))
if args.colbert_dir:
from ..retrieval.colbert import ColBERTRetriever

retrievers.append(ColBERTRetriever.load(args.colbert_dir))
if not retrievers:
ap.error("provide --bm25-dir and/or --colbert-dir")

if len(retrievers) == 1:
hits = retrievers[0].retrieve(args.query, k=args.k)
else:
from ..retrieval.hybrid import HybridRetriever

hits = HybridRetriever(retrievers, weights=args.weights, fetch_k=args.fetch_k).retrieve(args.query, k=args.k)

print(f"\nQuery: {args.query}\n" + "=" * 72)
for h in hits:
head = f"[{h.rank}] {h.score:.4f} {h.source} | {h.title or h.doc_id} ({h.chunk_id})"
print(head)
print(textwrap.indent(textwrap.shorten(h.text, width=320, placeholder=" …"), " "))
print("-" * 72)


if __name__ == "__main__":
main()
64 changes: 64 additions & 0 deletions datasets/amfv_datasets/corpus.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
""""""

from __future__ import annotations

import json
from collections.abc import Iterable, Iterator
from dataclasses import asdict, dataclass, field
from pathlib import Path


@dataclass(slots=True)
class Document:
""""""

doc_id: str
source: str
title: str
text: str
url: str = ""
metadata: dict = field(default_factory=dict)


@dataclass(slots=True)
class Chunk:
""""""

chunk_id: str
doc_id: str
source: str
title: str
text: str
ordinal: int
url: str = ""
metadata: dict = field(default_factory=dict)


def write_jsonl(path: str | Path, records: Iterable) -> int:
""""""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
n = 0
with path.open("w", encoding="utf-8") as fh:
for rec in records:
obj = rec if isinstance(rec, dict) else asdict(rec)
fh.write(json.dumps(obj, ensure_ascii=False) + "\n")
n += 1
return n


def _read_jsonl(path: str | Path) -> Iterator[dict]:
with Path(path).open(encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if line:
yield json.loads(line)


def read_documents(path: str | Path) -> Iterator[Document]:
for obj in _read_jsonl(path):
yield Document(**obj)


def read_chunks(path: str | Path) -> list[Chunk]:
return [Chunk(**obj) for obj in _read_jsonl(path)]
Empty file.
76 changes: 76 additions & 0 deletions datasets/amfv_datasets/ingest/nice.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
""""""

from __future__ import annotations

import argparse

from ..corpus import Document, write_jsonl

HF_DATASET = "epfl-llm/guidelines"


_TEXT_FIELDS = ("clean_text", "text", "content", "raw_text", "body")
_TITLE_FIELDS = ("title", "name", "heading")
_SOURCE_FIELDS = ("source", "dataset", "origin")
_URL_FIELDS = ("url", "link", "source_url")
_ID_FIELDS = ("id", "doc_id", "uuid")


def _pick(columns, candidates) -> str | None:
for c in candidates:
if c in columns:
return c
return None


def load_nice(limit: int | None = None, inspect: bool = False):
from datasets import load_dataset

ds = load_dataset(HF_DATASET, split="train")
columns = set(ds.column_names)
if inspect:
print("Columns:", sorted(columns))
print("Example:", {k: str(v)[:120] for k, v in ds[0].items()})

text_f = _pick(columns, _TEXT_FIELDS)
title_f = _pick(columns, _TITLE_FIELDS)
source_f = _pick(columns, _SOURCE_FIELDS)
url_f = _pick(columns, _URL_FIELDS)
id_f = _pick(columns, _ID_FIELDS)
if text_f is None:
raise RuntimeError(f"No text column found among {_TEXT_FIELDS}; columns are {sorted(columns)}")

if source_f is not None:
ds = ds.filter(lambda r: str(r[source_f]).strip().upper() == "NICE")

n = 0
for i, row in enumerate(ds):
text = (row.get(text_f) or "").strip()
if not text:
continue
yield Document(
doc_id=str(row.get(id_f, f"nice-{i}")),
source="NICE",
title=(row.get(title_f) or "").strip() if title_f else "",
text=text,
url=(row.get(url_f) or "").strip() if url_f else "",
metadata={"hf_dataset": HF_DATASET, "row": i},
)
n += 1
if limit is not None and n >= limit:
break


def main() -> None:
ap = argparse.ArgumentParser(description="Ingest NICE guidelines into documents.jsonl")
ap.add_argument("--out", default="data/nice/documents.jsonl")
ap.add_argument("--limit", type=int, default=None, help="cap document count (handy for smoke tests)")
ap.add_argument("--inspect", action="store_true", help="print dataset schema and exit-ish")
args = ap.parse_args()

written = write_jsonl(args.out, load_nice(limit=args.limit, inspect=args.inspect))
print(f"Wrote {written} NICE documents -> {args.out}")


if __name__ == "__main__":
main()
3 changes: 3 additions & 0 deletions datasets/amfv_datasets/retrieval/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from .base import Retriever, SearchHit, reciprocal_rank_fusion

__all__ = ["Retriever", "SearchHit", "reciprocal_rank_fusion"]
Loading