genmod / database_refresh /refresh_database.py
pbordescnil's picture
Revert "Refresh graph directly from Hub API"
5be9274
Raw
History Blame Contribute Delete
15.2 kB
#!/usr/bin/env python3
"""Build and publish a Neo4j dump from the public cfahlgren1/hub-stats dataset."""
from __future__ import annotations
import argparse
import json
import os
import shutil
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
import duckdb
from huggingface_hub import CommitOperationAdd, HfApi, hf_hub_download
SOURCE_REPO = "cfahlgren1/hub-stats"
DEFAULT_DUMP_REPO = "cnil/genmod-dump-neo4j"
PARQUET_REVISION = "refs/convert/parquet"
def log(message: str) -> None:
print(f"[genmod-refresh] {message}", flush=True)
def download_sources(work_dir: Path) -> tuple[Path, Path]:
cache_dir = work_dir / "hf-cache"
log(f"Downloading the model snapshot from {SOURCE_REPO}")
models = Path(
hf_hub_download(
repo_id=SOURCE_REPO,
repo_type="dataset",
revision=PARQUET_REVISION,
filename="models/train/0000.parquet",
cache_dir=cache_dir,
)
)
log(f"Downloading the dataset snapshot from {SOURCE_REPO}")
datasets = Path(
hf_hub_download(
repo_id=SOURCE_REPO,
repo_type="dataset",
revision=PARQUET_REVISION,
filename="datasets/train/0000.parquet",
cache_dir=cache_dir,
)
)
return models, datasets
def sql_path(path: Path) -> str:
return str(path).replace("'", "''")
def create_views(
connection: duckdb.DuckDBPyConnection,
models_path: Path,
datasets_path: Path,
max_models: int | None,
max_datasets: int | None,
) -> None:
model_limit = f" LIMIT {max_models}" if max_models else ""
dataset_limit = f" LIMIT {max_datasets}" if max_datasets else ""
connection.execute(
f"""
CREATE VIEW source_models AS
SELECT * EXCLUDE (_dedupe_rank)
FROM (
SELECT
*,
row_number() OVER (
PARTITION BY id
ORDER BY lastModified DESC NULLS LAST, _id DESC NULLS LAST
) AS _dedupe_rank
FROM read_parquet('{sql_path(models_path)}')
WHERE id IS NOT NULL AND trim(id) <> ''
)
WHERE _dedupe_rank = 1
{model_limit}
"""
)
connection.execute(
f"""
CREATE VIEW source_datasets AS
SELECT * EXCLUDE (_dedupe_rank)
FROM (
SELECT
*,
row_number() OVER (
PARTITION BY id
ORDER BY lastModified DESC NULLS LAST, _id DESC NULLS LAST
) AS _dedupe_rank
FROM read_parquet('{sql_path(datasets_path)}')
WHERE id IS NOT NULL AND trim(id) <> ''
)
WHERE _dedupe_rank = 1
{dataset_limit}
"""
)
connection.execute(
"""
CREATE VIEW base_model_edges AS
SELECT DISTINCT
base.id AS parent_id,
child.id AS child_id,
COALESCE(child.baseModels.relation, 'derived') AS relation_name
FROM source_models AS child,
UNNEST(child.baseModels.models) AS nested(base)
WHERE child.baseModels IS NOT NULL
AND base.id IS NOT NULL
AND trim(base.id) <> ''
AND child.id IS NOT NULL
"""
)
connection.execute(
"""
CREATE VIEW model_dataset_edges AS
SELECT DISTINCT
substr(tag, 9) AS dataset_id,
model.id AS model_id
FROM source_models AS model,
UNNEST(model.tags) AS nested(tag)
WHERE starts_with(tag, 'dataset:')
AND length(trim(substr(tag, 9))) > 0
AND model.id IS NOT NULL
"""
)
def export_csv(
connection: duckdb.DuckDBPyConnection,
output_dir: Path,
filename: str,
header: str,
query: str,
) -> Path:
path = output_dir / filename
header_path = output_dir / filename.replace(".csv", "-header.csv")
header_path.write_text(header + "\n", encoding="utf-8")
connection.execute(
f"""
COPY ({query})
TO '{sql_path(path)}'
(FORMAT CSV, HEADER false, DELIMITER ',', QUOTE '"', ESCAPE '"')
"""
)
log(f"Created {filename}")
return path
def prepare_csv_files(
models_path: Path,
datasets_path: Path,
output_dir: Path,
max_models: int | None = None,
max_datasets: int | None = None,
) -> dict[str, Path]:
output_dir.mkdir(parents=True, exist_ok=True)
database_path = output_dir / "refresh.duckdb"
connection = duckdb.connect(str(database_path))
connection.execute("SET preserve_insertion_order = false")
connection.execute("SET threads = 2")
create_views(connection, models_path, datasets_path, max_models, max_datasets)
files: dict[str, Path] = {}
files["models"] = export_csv(
connection,
output_dir,
"models.csv",
"modelId:ID(Model),name,downloads:long,task,createdAt,parameters,likes:long,license",
"""
WITH actual_models AS (
SELECT
id,
id AS name,
downloadsAllTime AS downloads,
pipeline_tag AS task,
CAST(createdAt AS VARCHAR) AS created_at,
CASE
WHEN safetensors.total >= 1000000000
THEN printf('%.1fB', safetensors.total / 1000000000.0)
WHEN safetensors.total >= 1000000
THEN printf('%.1fM', safetensors.total / 1000000.0)
WHEN safetensors.total >= 1000
THEN printf('%.1fK', safetensors.total / 1000.0)
WHEN safetensors.total IS NOT NULL
THEN CAST(safetensors.total AS VARCHAR)
END AS parameters,
likes,
json_extract_string(cardData, '$.license') AS license
FROM source_models
WHERE id IS NOT NULL AND trim(id) <> ''
),
missing_parents AS (
SELECT DISTINCT parent_id AS id
FROM base_model_edges
WHERE parent_id NOT IN (SELECT id FROM actual_models)
)
SELECT id, name, downloads, task, created_at, parameters, likes, license
FROM actual_models
UNION ALL
SELECT id, id, NULL, NULL, NULL, NULL, NULL, NULL
FROM missing_parents
""",
)
files["datasets"] = export_csv(
connection,
output_dir,
"datasets.csv",
"datasetId:ID(Dataset),name,downloads:long,createdAt_dataset",
"""
WITH actual_datasets AS (
SELECT
id,
id AS name,
downloadsAllTime AS downloads,
CAST(createdAt AS VARCHAR) AS created_at
FROM source_datasets
WHERE id IS NOT NULL AND trim(id) <> ''
),
missing_datasets AS (
SELECT DISTINCT dataset_id AS id
FROM model_dataset_edges
WHERE dataset_id NOT IN (SELECT id FROM actual_datasets)
)
SELECT id, name, downloads, created_at
FROM actual_datasets
UNION ALL
SELECT id, id, NULL, NULL
FROM missing_datasets
""",
)
files["authors"] = export_csv(
connection,
output_dir,
"authors.csv",
"authorId:ID(Author),name,type,followers:long",
"""
SELECT author, author, 'unknown', NULL
FROM (
SELECT author FROM source_models
UNION
SELECT author FROM source_datasets
)
WHERE author IS NOT NULL AND trim(author) <> ''
""",
)
files["base_model_edges"] = export_csv(
connection,
output_dir,
"base-model-edges.csv",
":START_ID(Model),:END_ID(Model),name",
"SELECT parent_id, child_id, relation_name FROM base_model_edges",
)
files["model_dataset_edges"] = export_csv(
connection,
output_dir,
"model-dataset-edges.csv",
":START_ID(Dataset),:END_ID(Model),name",
"""
SELECT dataset_id, model_id, 'A été utilisé dans ce modèle'
FROM model_dataset_edges
""",
)
files["author_model_edges"] = export_csv(
connection,
output_dir,
"author-model-edges.csv",
":START_ID(Author),:END_ID(Model),name",
"""
SELECT DISTINCT author, id, 'A publié'
FROM source_models
WHERE author IS NOT NULL AND trim(author) <> '' AND id IS NOT NULL
""",
)
files["author_dataset_edges"] = export_csv(
connection,
output_dir,
"author-dataset-edges.csv",
":START_ID(Author),:END_ID(Dataset),name",
"""
SELECT DISTINCT author, id, 'A publié'
FROM source_datasets
WHERE author IS NOT NULL AND trim(author) <> '' AND id IS NOT NULL
""",
)
connection.close()
database_path.unlink(missing_ok=True)
return files
def header_for(path: Path) -> Path:
return path.with_name(path.name.replace(".csv", "-header.csv"))
def build_dump(files: dict[str, Path], output_dir: Path, neo4j_admin: str) -> Path:
def group(name: str) -> str:
return f"{header_for(files[name])},{files[name]}"
command = [
neo4j_admin,
"database",
"import",
"full",
"neo4j",
"--overwrite-destination=true",
"--id-type=string",
"--threads=2",
"--verbose",
f"--nodes=Model={group('models')}",
f"--nodes=Dataset={group('datasets')}",
f"--nodes=Author={group('authors')}",
f"--relationships=USED_IN={group('base_model_edges')}",
f"--relationships=USED_IN={group('model_dataset_edges')}",
f"--relationships=POSTED={group('author_model_edges')}",
f"--relationships=POSTED={group('author_dataset_edges')}",
]
log("Building the offline Neo4j database")
subprocess.run(command, check=True)
dump_dir = output_dir / "dump"
dump_dir.mkdir(exist_ok=True)
log("Creating neo4j.dump")
subprocess.run(
[
neo4j_admin,
"database",
"dump",
"neo4j",
f"--to-path={dump_dir}",
"--overwrite-destination=true",
],
check=True,
)
return dump_dir / "neo4j.dump"
def write_metadata(
output_dir: Path,
source_revision: str,
model_count: int,
dataset_count: int,
) -> Path:
metadata = {
"built_at": datetime.now(timezone.utc).isoformat(),
"source_repo": SOURCE_REPO,
"source_revision": source_revision,
"model_count": model_count,
"dataset_count": dataset_count,
}
path = output_dir / "database_metadata.json"
path.write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8")
return path
def parquet_unique_id_count(path: Path) -> int:
connection = duckdb.connect()
count = connection.execute(
f"""
SELECT count(DISTINCT id)
FROM read_parquet('{sql_path(path)}')
WHERE id IS NOT NULL AND trim(id) <> ''
"""
).fetchone()[0]
connection.close()
return int(count)
def publish_dump(
api: HfApi,
dump_path: Path,
metadata_path: Path,
repo_id: str,
revision: str,
) -> None:
if revision != "main":
api.create_branch(
repo_id=repo_id,
repo_type="dataset",
branch=revision,
exist_ok=True,
)
log(f"Publishing the dump to {repo_id}@{revision}")
api.create_commit(
repo_id=repo_id,
repo_type="dataset",
revision=revision,
operations=[
CommitOperationAdd(
path_in_repo="neo4j.dump",
path_or_fileobj=str(dump_path),
),
CommitOperationAdd(
path_in_repo="database_metadata.json",
path_or_fileobj=str(metadata_path),
),
],
commit_message="Refresh Neo4j graph from cfahlgren1/hub-stats",
)
def restart_spaces(api: HfApi, space_ids: list[str]) -> None:
for space_id in space_ids:
log(f"Restarting Space {space_id}")
api.restart_space(repo_id=space_id)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--work-dir", type=Path, default=Path("/tmp/genmod-refresh"))
parser.add_argument("--dump-repo", default=os.getenv("NEO4J_DUMP_REPO", DEFAULT_DUMP_REPO))
parser.add_argument("--dump-revision", default=os.getenv("NEO4J_DUMP_REVISION", "main"))
parser.add_argument("--neo4j-admin", default=os.getenv("NEO4J_ADMIN", "neo4j-admin"))
parser.add_argument("--models-parquet", type=Path)
parser.add_argument("--datasets-parquet", type=Path)
parser.add_argument("--max-models", type=int)
parser.add_argument("--max-datasets", type=int)
parser.add_argument("--prepare-only", action="store_true")
parser.add_argument("--no-upload", action="store_true")
parser.add_argument("--keep-work-dir", action="store_true")
parser.add_argument("--restart-space", action="append", default=[])
return parser.parse_args()
def main() -> int:
args = parse_args()
if args.work_dir.exists() and not args.keep_work_dir:
shutil.rmtree(args.work_dir)
args.work_dir.mkdir(parents=True, exist_ok=True)
api = HfApi()
source_revision = api.dataset_info(SOURCE_REPO).sha
if bool(args.models_parquet) != bool(args.datasets_parquet):
raise SystemExit("Provide both --models-parquet and --datasets-parquet.")
if args.models_parquet:
models_path, datasets_path = args.models_parquet, args.datasets_parquet
else:
models_path, datasets_path = download_sources(args.work_dir)
csv_dir = args.work_dir / "csv"
files = prepare_csv_files(
models_path,
datasets_path,
csv_dir,
max_models=args.max_models,
max_datasets=args.max_datasets,
)
if args.prepare_only:
log(f"CSV preparation completed in {csv_dir}")
return 0
dump_path = build_dump(files, args.work_dir, args.neo4j_admin)
model_count = args.max_models or parquet_unique_id_count(models_path)
dataset_count = args.max_datasets or parquet_unique_id_count(datasets_path)
metadata_path = write_metadata(
args.work_dir,
source_revision,
model_count,
dataset_count,
)
if not args.no_upload:
if not os.getenv("HF_TOKEN"):
raise SystemExit("HF_TOKEN is required to upload the refreshed dump.")
publish_dump(
api,
dump_path,
metadata_path,
args.dump_repo,
args.dump_revision,
)
configured_spaces = [
value.strip()
for value in os.getenv("SPACES_TO_RESTART", "").split(",")
if value.strip()
]
restart_spaces(api, list(dict.fromkeys(configured_spaces + args.restart_space)))
log("Refresh completed successfully")
return 0
if __name__ == "__main__":
sys.exit(main())