File size: 3,314 Bytes
5499d76
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
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")