Efradeca's picture
Upload folder using huggingface_hub
3e77c56 verified
Raw
History Blame Contribute Delete
4.48 kB
"""Download + verify the Geo-FNO Elasticity dataset.
Source: the Geo-FNO Google Drive collection (license MIT, ``neuraloperator/Geo-FNO``).
IMPORTANT: the canonical Drive *folder*
``https://drive.google.com/drive/folders/1YBuaoTdOSr_qzaow-G-iwvbUI7fiUzu8`` is the **entire**
neuraloperator dataset collection (~8 GB: airfoil, car-cfd, channel-shocks, elasticity, ...).
For this project we only need three files (~47 MB total), so we fetch them **by file id**
rather than downloading the whole folder. The ids below were resolved once from that folder
(``gdown.download_folder(..., skip_download=True)``) — see ``docs/RECONCILIATION.md``.
Usage:
python -m stress_operator.data.download --out data
python -m stress_operator.data.download --out data --report-only # just verify + print shapes
"""
from __future__ import annotations
import argparse
import hashlib
import os
import sys
import numpy as np
# Resolved file ids for elasticity/Meshes/* inside the Geo-FNO Drive collection.
DRIVE_FILE_IDS = {
"Random_UnitCell_sigma_10.npy": "1Ia5izgUum-IQLdO6PW70HO8AdAqA_IVb",
"Random_UnitCell_XY_10.npy": "1I-fO-RsFvD3nqBuFrg67R0yqTFdD_gpA",
"Random_UnitCell_rr_10.npy": "1Pjliqhxegoe5VpoLrpBa9n3P4gX9MfTt",
}
# Mirror Transolver's expected layout: <out>/elasticity/Meshes/<file>.
SUBDIR = os.path.join("elasticity", "Meshes")
REQUIRED_FILES = ["Random_UnitCell_sigma_10.npy", "Random_UnitCell_XY_10.npy"]
# Authoritative raw shapes (sample axis LAST). See docs/RECONCILIATION.md.
EXPECTED_RAW_SHAPES = {
"Random_UnitCell_sigma_10.npy": (972, 2000),
"Random_UnitCell_XY_10.npy": (972, 2, 2000),
"Random_UnitCell_rr_10.npy": (42, 2000),
}
def _find(data_dir: str, name: str):
"""Locate ``name`` anywhere under ``data_dir`` (handles nested layouts)."""
direct = os.path.join(data_dir, name)
if os.path.isfile(direct):
return direct
for root, _dirs, files in os.walk(data_dir):
if name in files:
return os.path.join(root, name)
return None
def have_all(data_dir: str) -> bool:
return all(_find(data_dir, f) is not None for f in REQUIRED_FILES)
def download(out_dir: str, force: bool = False) -> None:
import gdown
dest_dir = os.path.join(out_dir, SUBDIR)
os.makedirs(dest_dir, exist_ok=True)
for name, file_id in DRIVE_FILE_IDS.items():
existing = _find(out_dir, name)
if existing is not None and not force:
print(f"[download] {name}: already present ({existing}); skipping.")
continue
out_path = os.path.join(dest_dir, name)
print(f"[download] fetching {name} (id={file_id}) -> {out_path}")
gdown.download(id=file_id, output=out_path, quiet=False)
def _md5(path: str, chunk: int = 1 << 20) -> str:
h = hashlib.md5()
with open(path, "rb") as f:
for block in iter(lambda: f.read(chunk), b""):
h.update(block)
return h.hexdigest()
def report(data_dir: str) -> int:
"""Print shape/dtype/stats and check against expected raw shapes. Returns exit code."""
ok = True
print(f"[report] scanning {data_dir!r}")
for name in DRIVE_FILE_IDS:
path = _find(data_dir, name)
required = name in REQUIRED_FILES
if path is None:
print(f" - {name}: {'MISSING (required)' if required else 'missing (optional)'}")
if required:
ok = False
continue
arr = np.load(path)
exp = EXPECTED_RAW_SHAPES.get(name)
flag = "" if (exp is None or tuple(arr.shape) == tuple(exp)) else f" <-- WARNING: expected {exp}"
print(
f" - {name}: shape={tuple(arr.shape)} dtype={arr.dtype} "
f"min={float(arr.min()):.4g} max={float(arr.max()):.4g} md5={_md5(path)[:8]}{flag}"
)
print("[report] OK" if ok else "[report] FAILED: required files missing")
return 0 if ok else 1
def main() -> int:
ap = argparse.ArgumentParser(description="Download/verify the Geo-FNO Elasticity dataset (3 files).")
ap.add_argument("--out", default="data", help="output directory")
ap.add_argument("--report-only", action="store_true", help="skip download, just verify + print shapes")
ap.add_argument("--force", action="store_true", help="re-download even if files exist")
args = ap.parse_args()
if not args.report_only:
download(args.out, force=args.force)
return report(args.out)
if __name__ == "__main__":
sys.exit(main())