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))