diff --git a/src/pubget/_vectorization.py b/src/pubget/_vectorization.py index b0db48c..418bd1c 100644 --- a/src/pubget/_vectorization.py +++ b/src/pubget/_vectorization.py @@ -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) @@ -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 diff --git a/src/pubget/_vocabulary.py b/src/pubget/_vocabulary.py index a1dc5a4..ac7ac91 100644 --- a/src/pubget/_vocabulary.py +++ b/src/pubget/_vocabulary.py @@ -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" ) diff --git a/tests/test_data_extraction.py b/tests/test_data_extraction.py index 536a2ac..36aa351 100644 --- a/tests/test_data_extraction.py +++ b/tests/test_data_extraction.py @@ -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 @@ -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): diff --git a/tests/test_labelbuddy.py b/tests/test_labelbuddy.py index ab89937..2716ad9 100644 --- a/tests/test_labelbuddy.py +++ b/tests/test_labelbuddy.py @@ -1,5 +1,6 @@ import json import math +import re from unittest.mock import Mock import pandas as pd @@ -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)