viclickbait_gnn / src /preprocess.py
minhy112's picture
Upload viclickbait_gnn project
877049d verified
Raw
History Blame Contribute Delete
3.09 kB
import argparse
from pathlib import Path
import pandas as pd
from utils import clean_title, ensure_dir, normalize_label, save_json
def load_input(path: str | Path) -> pd.DataFrame:
path = Path(path)
if path.suffix.lower() == ".csv":
return pd.read_csv(path)
if path.suffix.lower() == ".jsonl":
return pd.read_json(path, lines=True)
raise ValueError(f"Unsupported input format: {path.suffix}")
def preprocess_frame(dataframe: pd.DataFrame, min_title_length: int) -> tuple[pd.DataFrame, dict]:
stats = {"before_rows": int(len(dataframe))}
if "title" not in dataframe.columns or "label" not in dataframe.columns:
raise ValueError("Input data must contain `title` and `label` columns.")
processed = dataframe.copy()
processed["title"] = processed["title"].map(clean_title)
processed = processed[processed["title"].str.len() >= min_title_length].copy()
stats["after_title_filter"] = int(len(processed))
processed["label_id"] = processed["label"].map(normalize_label)
processed["label"] = processed["label_id"].map({0: "non-clickbait", 1: "clickbait"})
processed = processed.drop_duplicates(subset=["title"]).reset_index(drop=True)
stats["after_deduplicate"] = int(len(processed))
processed["node_id"] = processed.index.astype(int)
preferred_columns = ["node_id", "id", "title", "label", "label_id", "source", "category", "publish_datetime", "url"]
existing_columns = [column for column in preferred_columns if column in processed.columns]
remaining_columns = [column for column in processed.columns if column not in existing_columns]
processed = processed[existing_columns + remaining_columns]
stats["label_distribution"] = processed["label"].value_counts().to_dict()
return processed, stats
def main() -> None:
parser = argparse.ArgumentParser(description="Preprocess ViClickbait dataset.")
parser.add_argument("--input", required=True, help="Path to raw CSV or JSONL.")
parser.add_argument("--output", required=True, help="Path to write processed CSV.")
parser.add_argument("--label_map", default=None, help="Optional label_map.json path.")
parser.add_argument("--stats_path", default=None, help="Optional preprocess_stats.json path.")
parser.add_argument("--min_title_length", type=int, default=4)
args = parser.parse_args()
frame = load_input(args.input)
processed, stats = preprocess_frame(frame, min_title_length=args.min_title_length)
output_path = Path(args.output)
ensure_dir(output_path.parent)
processed.to_csv(output_path, index=False)
label_map_path = Path(args.label_map) if args.label_map else output_path.parent / "label_map.json"
save_json({"label_to_id": {"non-clickbait": 0, "clickbait": 1}}, label_map_path)
stats_path = Path(args.stats_path) if args.stats_path else output_path.parent / "preprocess_stats.json"
save_json(stats, stats_path)
print(f"Saved processed data to {output_path}")
print(f"Rows after preprocessing: {len(processed)}")
if __name__ == "__main__":
main()