dev-strender's picture
Replace v24-era demo with v34 pipeline demo (engine-vendored bundle)
9c84f9d verified
Raw History Blame
2.68 kB
"""Dataset management commands."""
import click
from tabulate import tabulate
@click.group("datasets")
def datasets_group() -> None:
"""Manage evaluation datasets."""
@datasets_group.command("sync")
@click.argument("repo")
@click.option("--force", "-f", is_flag=True, help="Force re-download")
@click.pass_context
def sync_dataset(ctx: click.Context, repo: str, force: bool) -> None:
"""Sync a dataset from HuggingFace."""
cfg = ctx.obj["config"]
if cfg.is_remote:
from solar_eval.cli.client import EvalClient
client = EvalClient(cfg.remote_url, cfg.timeout)
click.secho(f"Syncing via server: {repo}...", fg="blue")
result = client.post("/api/datasets/sync", json={"repo": repo, "force": force})
if result.get("error"):
click.secho(f"Sync failed: {result['error']}", fg="red")
raise SystemExit(1)
click.secho(f"Synced {len(result.get('synced_files', []))} file(s)", fg="green")
else:
from solar_eval.core.dataset_loader import DatasetLoader
click.secho(f"Syncing locally: {repo}...", fg="blue")
loader = DatasetLoader()
files = loader.list_files(repo)
for f in files:
click.echo(f" Downloading: {f}")
loader.download_file(repo, f, force=force)
click.secho(f"Synced {len(files)} file(s)", fg="green")
@datasets_group.command("list")
@click.pass_context
def list_datasets(ctx: click.Context) -> None:
"""List available datasets."""
cfg = ctx.obj["config"]
if cfg.is_remote:
from solar_eval.cli.client import EvalClient
client = EvalClient(cfg.remote_url, cfg.timeout)
datasets = client.get("/api/datasets")
else:
# List datasets from project configs
from solar_eval.core.project_loader import load_all_project_configs
configs = load_all_project_configs(cfg.projects_dir, cfg.config_dirs)
seen = set()
datasets = []
for c in configs:
repo = c.get("dataset", {}).get("repo", "")
for t in c.get("tasks", []):
path = t.get("dataset_path", "")
key = f"{repo}:{path}"
if key not in seen:
seen.add(key)
datasets.append({"repo": repo, "path": path, "cached": False})
if not datasets:
click.secho("No datasets found.", fg="yellow")
return
rows = [[d.get("repo", ""), d.get("path", ""), d.get("num_samples", "-"),
"Yes" if d.get("cached") else "No"] for d in datasets]
click.secho(tabulate(rows, headers=["Repository", "Path", "Samples", "Cached"], tablefmt="simple"), fg="blue")