Skip to content

Commit 497932f

Browse files
wmaynerclaude
andcommitted
Add save_npz, save_dataframe, and read_metadata provenance writers
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01PEAxNzhDCaTrntX3o1JqMV
1 parent 2fdd639 commit 497932f

2 files changed

Lines changed: 183 additions & 0 deletions

File tree

pyphi/provenance.py

Lines changed: 130 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,11 +22,15 @@
2222
from datetime import UTC
2323
from datetime import datetime
2424
from pathlib import Path
25+
from typing import TYPE_CHECKING
2526
from typing import Any
2627

2728
import numpy as np
2829
import scipy
2930

31+
if TYPE_CHECKING:
32+
import pandas as pd
33+
3034
_PACKAGE_ROOT = Path(__file__).resolve().parent
3135

3236

@@ -269,6 +273,129 @@ def save_json(
269273
return path
270274

271275

276+
def _metadata_json(
277+
prov: Provenance, params: Mapping[str, Any] | None
278+
) -> tuple[str, str]:
279+
"""Serialize the provenance record and params as JSON strings."""
280+
return (
281+
json.dumps(asdict(prov), default=_json_default),
282+
json.dumps(dict(params or {}), default=_json_default),
283+
)
284+
285+
286+
def save_npz(
287+
arrays: Mapping[str, np.ndarray],
288+
directory: Path | str,
289+
name: str,
290+
*,
291+
params: Mapping[str, Any] | None = None,
292+
run_label: str | None = None,
293+
seed: int | None = None,
294+
note: str | None = None,
295+
) -> Path:
296+
"""Write ``arrays`` to a self-describing, non-clobbering ``.npz`` file.
297+
298+
The arrays are stored with :func:`numpy.savez_compressed`, plus two
299+
reserved entries ``_provenance`` and ``_params`` holding JSON strings
300+
(read back with :func:`read_metadata`). Array names beginning with
301+
``_`` raise :class:`ValueError`. Filename, versioning, and seed
302+
resolution follow :func:`save_json`.
303+
304+
Returns the written path.
305+
"""
306+
reserved = [key for key in arrays if key.startswith("_")]
307+
if reserved:
308+
raise ValueError(f"array names beginning with '_' are reserved: {reserved}")
309+
prov = _capture_metadata(params, seed, note)
310+
prov_json, params_json = _metadata_json(prov, params)
311+
path = unique_path(directory, format_stem(name, params, run_label), ".npz")
312+
entries = {
313+
**arrays,
314+
"_provenance": np.array(prov_json),
315+
"_params": np.array(params_json),
316+
}
317+
np.savez_compressed(path, allow_pickle=False, **entries)
318+
return path
319+
320+
321+
def save_dataframe(
322+
df: pd.DataFrame,
323+
directory: Path | str,
324+
name: str,
325+
*,
326+
params: Mapping[str, Any] | None = None,
327+
run_label: str | None = None,
328+
seed: int | None = None,
329+
note: str | None = None,
330+
) -> Path:
331+
"""Write ``df`` to a self-describing, non-clobbering parquet file.
332+
333+
The frame is written with its index preserved; ``pyphi_provenance``
334+
and ``pyphi_params`` entries (JSON strings) are merged into the
335+
parquet schema metadata, so :func:`pandas.read_parquet` reads the
336+
data normally and :func:`read_metadata` recovers the metadata.
337+
Filename, versioning, and seed resolution follow :func:`save_json`.
338+
DataFrame fidelity follows parquet semantics.
339+
340+
Returns the written path.
341+
"""
342+
import pyarrow as pa
343+
import pyarrow.parquet as pq
344+
345+
prov = _capture_metadata(params, seed, note)
346+
prov_json, params_json = _metadata_json(prov, params)
347+
path = unique_path(directory, format_stem(name, params, run_label), ".parquet")
348+
table = pa.Table.from_pandas(df, preserve_index=True)
349+
metadata = dict(table.schema.metadata or {})
350+
metadata[b"pyphi_provenance"] = prov_json.encode()
351+
metadata[b"pyphi_params"] = params_json.encode()
352+
pq.write_table(table.replace_schema_metadata(metadata), path)
353+
return path
354+
355+
356+
def read_metadata(path: Path | str) -> dict[str, Any]:
357+
"""Read the provenance and params embedded in a writer's output file.
358+
359+
Dispatches on the file suffix (``.json``, ``.npz``, or ``.parquet``)
360+
and returns ``{"provenance": dict, "params": dict}``. A file without
361+
the expected metadata (not produced by :func:`save_json`,
362+
:func:`save_npz`, or :func:`save_dataframe`) raises
363+
:class:`ValueError`, as does an unrecognized suffix.
364+
"""
365+
path = Path(path)
366+
missing = ValueError(f"no pyphi provenance metadata in {path}")
367+
if path.suffix == ".json":
368+
document = json.loads(path.read_text())
369+
try:
370+
return {
371+
"provenance": document["provenance"],
372+
"params": document["params"],
373+
}
374+
except (KeyError, TypeError):
375+
raise missing from None
376+
if path.suffix == ".npz":
377+
with np.load(path) as npz:
378+
try:
379+
return {
380+
"provenance": json.loads(str(npz["_provenance"][()])),
381+
"params": json.loads(str(npz["_params"][()])),
382+
}
383+
except KeyError:
384+
raise missing from None
385+
if path.suffix == ".parquet":
386+
import pyarrow.parquet as pq
387+
388+
metadata = pq.read_schema(path).metadata or {}
389+
try:
390+
return {
391+
"provenance": json.loads(metadata[b"pyphi_provenance"]),
392+
"params": json.loads(metadata[b"pyphi_params"]),
393+
}
394+
except KeyError:
395+
raise missing from None
396+
raise ValueError(f"unrecognized suffix {path.suffix!r} for {path}")
397+
398+
272399
def _set_provenance(result: Any, prov: Provenance) -> None:
273400
"""Assign ``prov`` to ``result.provenance``, working around frozen results."""
274401
try:
@@ -319,7 +446,10 @@ def with_provenance(self, **fields: Any) -> HasProvenance:
319446
"HasProvenance",
320447
"Provenance",
321448
"format_stem",
449+
"read_metadata",
450+
"save_dataframe",
322451
"save_json",
452+
"save_npz",
323453
"stamp_wall_time",
324454
"unique_path",
325455
]

test/test_provenance_writers.py

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import json
44

55
import numpy as np
6+
import pandas as pd
67
import pytest
78

89
from pyphi import provenance
@@ -84,3 +85,55 @@ def test_no_clobber(self, tmp_path):
8485
second = provenance.save_json({"x": 2}, tmp_path, "study", params={"seed": 1})
8586
assert second == tmp_path / "study_seed1_v2.json"
8687
assert json.loads(first.read_text())["data"] == {"x": 1}
88+
89+
90+
class TestSaveNpz:
91+
def test_arrays_round_trip_with_metadata(self, tmp_path):
92+
arrays = {"phis": np.linspace(0, 1, 5), "states": np.eye(2)}
93+
path = provenance.save_npz(arrays, tmp_path, "study", params={"seed": 3})
94+
assert path == tmp_path / "study_seed3.npz"
95+
with np.load(path) as npz:
96+
np.testing.assert_array_equal(npz["phis"], arrays["phis"])
97+
np.testing.assert_array_equal(npz["states"], arrays["states"])
98+
metadata = provenance.read_metadata(path)
99+
assert metadata["params"] == {"seed": 3}
100+
assert metadata["provenance"]["seed"] == 3
101+
102+
def test_reserved_names_rejected(self, tmp_path):
103+
with pytest.raises(ValueError, match="reserved"):
104+
provenance.save_npz({"_provenance": np.zeros(1)}, tmp_path, "study")
105+
106+
107+
class TestSaveDataframe:
108+
def test_round_trip_and_metadata(self, tmp_path):
109+
df = pd.DataFrame(
110+
{"phi": [0.1, 0.2], "n": [3, 4]},
111+
index=pd.Index(["a", "b"], name="system"),
112+
)
113+
path = provenance.save_dataframe(df, tmp_path, "study", params={"seed": 9})
114+
assert path == tmp_path / "study_seed9.parquet"
115+
pd.testing.assert_frame_equal(pd.read_parquet(path), df)
116+
metadata = provenance.read_metadata(path)
117+
assert metadata["provenance"]["seed"] == 9
118+
assert metadata["params"] == {"seed": 9}
119+
120+
121+
class TestReadMetadata:
122+
def test_json_metadata(self, tmp_path):
123+
path = provenance.save_json({"x": 1}, tmp_path, "study", params={"seed": 5})
124+
metadata = provenance.read_metadata(path)
125+
assert metadata["params"] == {"seed": 5}
126+
assert metadata["provenance"]["seed"] == 5
127+
assert "pyphi_version" in metadata["provenance"]
128+
129+
def test_plain_json_rejected(self, tmp_path):
130+
path = tmp_path / "plain.json"
131+
path.write_text('{"x": 1}')
132+
with pytest.raises(ValueError, match="no pyphi provenance"):
133+
provenance.read_metadata(path)
134+
135+
def test_unknown_suffix_rejected(self, tmp_path):
136+
path = tmp_path / "file.csv"
137+
path.touch()
138+
with pytest.raises(ValueError, match="unrecognized suffix"):
139+
provenance.read_metadata(path)

0 commit comments

Comments
 (0)