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
11 changes: 9 additions & 2 deletions src/pubget/_vectorization.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,7 +165,6 @@ def _vectorize_articles(

Returns the pmcids and the mapping text field: csr matrix of features.
"""
articles.fillna("", inplace=True)
vectorized = {}
for field in _FIELDS:
vectorized[field] = vectorizer.transform(articles[field].values)
Expand All @@ -185,7 +184,15 @@ def _extract_word_counts(
).fit()
chunksize = 200
with open(corpus_file, encoding="utf-8") as corpus_fh:
all_chunks = pd.read_csv(corpus_fh, chunksize=chunksize)
# The text fields are read as strings and missing values as empty
# strings: pandas would otherwise give an empty field the float dtype,
# and the vectorizer expects text.
all_chunks = pd.read_csv(
corpus_fh,
chunksize=chunksize,
dtype=dict.fromkeys(_FIELDS, str),
keep_default_na=False,
)
vectorized_chunks = Parallel(n_jobs=n_jobs, verbose=8)(
delayed(_vectorize_articles)(chunk, vectorizer=vectorizer)
for chunk in all_chunks
Expand Down
12 changes: 10 additions & 2 deletions src/pubget/_vocabulary.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,13 +21,21 @@
_LOG = logging.getLogger(__name__)
_STEP_NAME = "extract_vocabulary"
_STEP_DESCRIPTION = "Extract vocabulary of word n-grams from text."
_TEXT_FIELDS = ("title", "keywords", "abstract", "body")


def _iter_corpus(corpus_fh: TextIO) -> Generator[str, None, None]:
"""Yield the concatenated text fields of articles one by one."""
n_articles = 0
for chunk in pd.read_csv(corpus_fh, chunksize=500):
chunk.fillna("", inplace=True)
# The text fields are read as strings and missing values as empty strings:
# pandas would otherwise give an empty field the float dtype, and refuse to
# store the empty strings in it.
for chunk in pd.read_csv(
corpus_fh,
chunksize=500,
dtype=dict.fromkeys(_TEXT_FIELDS, str),
keep_default_na=False,
):
text = chunk["title"].str.cat(
chunk.loc[:, ["keywords", "abstract", "body"]], sep="\n"
)
Expand Down
16 changes: 12 additions & 4 deletions tests/test_data_extraction.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,8 @@ def _check_extracted_data(data_dir, articles_with_coords_only):
assert text.shape == (n_articles, 5)
assert text.at[0, "body"].strip().startswith("The text")
coordinates = pd.read_csv(data_dir.joinpath("coordinates.csv"))
assert coordinates.shape == (12, 6)
# 5 articles with a 2-row coordinate table and one with a 10-row one
assert coordinates.shape == (20, 6)
authors = pd.read_csv(data_dir.joinpath("authors.csv"))
assert authors.shape == (n_authors, 3)
assert authors["pmcid"].nunique() == n_articles
Expand Down Expand Up @@ -104,9 +105,16 @@ def test_extract_data_to_csv_with_tables(tmp_path, articles_dir):
articles_dir, tmp_path.joinpath("extracted_data"), keep_tables=True
)
assert code == ExitCode.COMPLETED
text = pd.read_csv(data_dir.joinpath("text.csv"))
assert text.at[0, "body"].strip().startswith("The text")
assert "X\tY\tZ\n10\t20\t30\n" in text.at[0, "body"]
# rows are in the order in which the article directories are visited, so
# articles are looked up by pmcid rather than by position.
text = pd.read_csv(data_dir.joinpath("text.csv"), index_col="pmcid")
body = text.at[9054084, "body"]
assert body.strip().startswith("The text")
assert "X\tY\tZ\n10\t20\t30\n" in body
# a table with many columns is inserted in full, not truncated
wide_table_body = text.at[9056519, "body"]
assert "X\tY\tZ\tRegion\tActivation_Level" in wide_table_body
assert "-30\t-40\t-50\tBasal Ganglia" in wide_table_body


def test_extractor_failures(articles_dir, tmp_path, monkeypatch):
Expand Down
5 changes: 4 additions & 1 deletion tests/test_labelbuddy.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import json
import math
import re
from unittest.mock import Mock

import pandas as pd
Expand Down Expand Up @@ -65,7 +66,9 @@ def test_make_labelbuddy_documents(
) as f:
docs = [json.loads(doc_json) for doc_json in f]
assert len(docs) == expected_batch_size
assert all("Body\n The text of" in d["text"] for d in docs)
# the whitespace between the "Body" heading and the start of the body
# depends on how the article's XML is indented
assert all(re.search(r"# Body\s+The text of", d["text"]) for d in docs)
_check_batch_info(labelbuddy_dir)


Expand Down
Loading