proofread-demo / solar_eval /core /dataset_loader.py
dev-strender's picture
Replace v24-era demo with v34 pipeline demo (engine-vendored bundle)
9c84f9d verified
Raw History Blame Contribute Delete
3.62 kB
import csv
import json
import logging
from pathlib import Path
from huggingface_hub import hf_hub_download, list_repo_files
from solar_eval.config import get_provider_config
logger = logging.getLogger(__name__)
class DatasetLoader:
"""Loads datasets from HuggingFace repos or local filesystem."""
def __init__(self, hf_token: str | None = None, base_dir: Path | None = None) -> None:
self.hf_token = hf_token or get_provider_config().hf_token or None
self.base_dir = base_dir
def download_file(self, repo: str, path: str, force: bool = False) -> Path:
"""Download a single file from HF repo. Returns local cache path."""
local_path = hf_hub_download(
repo_id=repo,
filename=path,
repo_type="dataset",
token=self.hf_token,
force_download=force,
)
return Path(local_path)
def _resolve_local_path(self, repo: str, path: str) -> Path:
"""Resolve a local file path from repo dir and dataset path."""
repo_path = Path(repo)
if repo_path.is_absolute():
return repo_path / path
if self.base_dir:
return self.base_dir / repo / path
return repo_path / path
def load_jsonl(
self, repo: str, path: str, source: str = "huggingface", force: bool = False
) -> list[dict]:
"""Load and parse a JSONL file from HuggingFace or local filesystem."""
if source == "local":
local_path = self._resolve_local_path(repo, path)
if not local_path.exists():
raise FileNotFoundError(f"Local dataset not found: {local_path}")
else:
local_path = self.download_file(repo, path, force=force)
records = []
with open(local_path, encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
records.append(json.loads(line))
logger.info(f"Loaded {len(records)} records from {local_path}")
return records
def load_csv(
self, repo: str, path: str, source: str = "huggingface", force: bool = False
) -> list[dict]:
"""Download and parse a CSV file.
Supports both DictReader (with headers) and headerless CSV.
For headerless CSV, columns are mapped to 'wrong' and 'correct' keys.
"""
if source == "local":
local_path = self._resolve_local_path(repo, path)
if not local_path.exists():
raise FileNotFoundError(f"Local CSV not found: {local_path}")
else:
local_path = self.download_file(repo, path, force=force)
with open(local_path, encoding="utf-8") as f:
# Peek first line to detect headers
first_line = f.readline().strip()
f.seek(0)
if first_line and ("wrong" in first_line or "before" in first_line):
reader = csv.DictReader(f)
records = list(reader)
else:
# Headerless CSV: treat as (wrong, correct) pairs
reader = csv.reader(f)
records = []
for row in reader:
if len(row) >= 2 and row[0].strip():
records.append({"wrong": row[0].strip(), "correct": row[1].strip()})
logger.info(f"Loaded {len(records)} records from {local_path}")
return records
def list_files(self, repo: str) -> list[str]:
"""List all files in HF repo."""
return list(list_repo_files(repo, repo_type="dataset", token=self.hf_token))