poincare-hyper / src /convert_well_active_matter.py
DHDRL's picture
Rename convert_well_active_matter.py to src/convert_well_active_matter.py
eb3c6a9 verified
Raw
History Blame Contribute Delete
3.32 kB
#!/usr/bin/env python3
"""
Convert The Well active_matter nested HDF5 into flat (T, C, H, W) files
for LocalWellHDF5 / run_full.
Default C=11 (full spatiotemporal state, float32, lossless gzip):
concentration, vx, vy, D00, D01, D10, D11, E00, E01, E10, E11
Does not reduce velocity to speed, does not drop tensors, does not
downcast precision. Truncated/corrupt source files are skipped.
Example:
python scripts/convert_well_active_matter.py \\
--src data/well/datasets/active_matter/data/train \\
--out data/well_simple \\
--max-trajs 16
"""
from __future__ import annotations
import argparse
from pathlib import Path
import h5py
import numpy as np
CHANNEL_NAMES = [
"concentration", "vx", "vy",
"D00", "D01", "D10", "D11",
"E00", "E01", "E10", "E11",
]
def convert_one(fp: Path, out_dir: Path, n_written: int, max_trajs: int) -> int:
try:
with h5py.File(fp, "r") as f:
conc = f["t0_fields/concentration"][:]
vel = f["t1_fields/velocity"][:]
D = f["t2_fields/D"][:]
E = f["t2_fields/E"][:]
meta = {}
for k in ("L", "alpha", "zeta"):
if f"scalars/{k}" in f:
meta[k] = float(f[f"scalars/{k}"][()])
except OSError as e:
print(f"SKIP truncated/corrupt: {fp.name} | {e}")
return n_written
N, T, H, W = conc.shape
for i in range(N):
if n_written >= max_trajs:
break
chans = [
conc[i],
vel[i, ..., 0], vel[i, ..., 1],
D[i, ..., 0, 0], D[i, ..., 0, 1], D[i, ..., 1, 0], D[i, ..., 1, 1],
E[i, ..., 0, 0], E[i, ..., 0, 1], E[i, ..., 1, 0], E[i, ..., 1, 1],
]
fields = np.stack(chans, axis=1).astype(np.float32)
assert fields.shape == (T, 11, H, W), fields.shape
out = out_dir / f"traj_{n_written:03d}.hdf5"
with h5py.File(out, "w") as g:
g.create_dataset("fields", data=fields, compression="gzip")
g.attrs["channel_names"] = np.array(CHANNEL_NAMES, dtype="S")
g.attrs["source_file"] = fp.name
for k, v in meta.items():
g.attrs[k] = v
print(f"wrote {out.name} {fields.shape}")
n_written += 1
return n_written
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--src", type=Path, required=True,
help="Directory of active_matter train *.hdf5 files")
ap.add_argument("--out", type=Path, required=True,
help="Output directory for flat traj_*.hdf5")
ap.add_argument("--max-trajs", type=int, default=16)
ap.add_argument("--clear-out", action="store_true",
help="Delete existing traj_*.hdf5 in --out first")
args = ap.parse_args()
args.out.mkdir(parents=True, exist_ok=True)
if args.clear_out:
for old in args.out.glob("traj_*.hdf5"):
old.unlink()
n = 0
for fp in sorted(args.src.glob("*.hdf5")):
n = convert_one(fp, args.out, n, args.max_trajs)
if n >= args.max_trajs:
break
print(f"total trajectories: {n}")
print(f"out_dir: {args.out}")
print("Train with: python -m src.run_full --data-root <out> --channels 11")
if __name__ == "__main__":
main()