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