article_classifier / src /data_utils.py
asriel14's picture
Upload 4 files
5499d76 verified
Raw History Blame Contribute Delete
3.31 kB
from __future__ import annotations
import json
from pathlib import Path
import pandas as pd
TEXT_CANDIDATES = {
"title": ["title", "headline", "paper_title", "name"],
"abstract": ["abstract", "summary", "description"],
"label": ["label", "tag", "category", "primary_category", "topic"],
}
def _find_column(df: pd.DataFrame, candidates: list[str]) -> str | None:
lower_map = {c.lower(): c for c in df.columns}
for candidate in candidates:
if candidate.lower() in lower_map:
return lower_map[candidate.lower()]
return None
def normalize_whitespace(text: str) -> str:
return " ".join(str(text).strip().split())
def combine_text(title: str, abstract: str) -> str:
title = normalize_whitespace(title)
abstract = normalize_whitespace(abstract)
if title and abstract:
return f"[TITLE] {title} [ABSTRACT] {abstract}"
if title:
return f"[TITLE] {title}"
if abstract:
return f"[ABSTRACT] {abstract}"
return ""
def read_table(path: str | Path) -> pd.DataFrame:
path = Path(path)
suffix = path.suffix.lower()
if suffix == ".csv":
return pd.read_csv(path)
if suffix in {".jsonl", ".json"}:
return pd.read_json(path, lines=suffix == ".jsonl")
if suffix == ".parquet":
return pd.read_parquet(path)
raise ValueError(f"Unsupported file type: {suffix}")
def load_dataset_frame(path: str | Path, text_cols: list[str] | None = None, label_col: str | None = None) -> pd.DataFrame:
df = read_table(path).copy()
title_col = None
abstract_col = None
if text_cols:
if len(text_cols) == 1:
title_col = text_cols[0]
elif len(text_cols) >= 2:
title_col, abstract_col = text_cols[:2]
else:
title_col = _find_column(df, TEXT_CANDIDATES["title"])
abstract_col = _find_column(df, TEXT_CANDIDATES["abstract"])
if not title_col and not abstract_col:
raise ValueError("Could not identify title/abstract columns automatically.")
if label_col is None:
label_col = _find_column(df, TEXT_CANDIDATES["label"])
if label_col is None:
raise ValueError("Could not identify label column automatically.")
if title_col is None:
df["title"] = ""
else:
df["title"] = df[title_col].fillna("").astype(str)
if abstract_col is None:
df["abstract"] = ""
else:
df["abstract"] = df[abstract_col].fillna("").astype(str)
df["label"] = df[label_col].fillna("").astype(str)
df["text"] = [combine_text(t, a) for t, a in zip(df["title"], df["abstract"])]
df["text"] = df["text"].map(normalize_whitespace)
df["label"] = df["label"].map(normalize_whitespace)
df = df[(df["text"].str.len() > 0) & (df["label"].str.len() > 0)].reset_index(drop=True)
return df[["title", "abstract", "text", "label"]]
def filter_rare_classes(df: pd.DataFrame, min_examples_per_class: int) -> pd.DataFrame:
counts = df["label"].value_counts()
keep_labels = counts[counts >= min_examples_per_class].index
return df[df["label"].isin(keep_labels)].reset_index(drop=True)
def save_json(data: dict, path: str | Path) -> None:
path = Path(path)
path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")