ReMEDi / src /remedi /io.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw History Blame Contribute Delete
4 kB
"""Validated condition-level artifacts and reproducible metadata."""
from dataclasses import dataclass
from pathlib import Path
import hashlib
import json
import numpy as np
import pandas as pd
REQUIRED = ["molecule_id", "smiles", "dose_um", "context", "block", "control_id", "split"]
def write_json(path, value):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(value, indent=2, allow_nan=False, default=str) + "\n")
def file_hash(path):
h = hashlib.sha256()
with open(path, "rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
h.update(chunk)
return h.hexdigest()
def fingerprint_config(value):
return hashlib.sha256(json.dumps(value, sort_keys=True).encode()).hexdigest()
@dataclass
class ResponseData:
obs: pd.DataFrame
response: np.ndarray
control: np.ndarray
treated_var: np.ndarray
control_var: np.ndarray
metadata: dict
def validate(self):
missing = set(REQUIRED) - set(self.obs.columns)
if missing:
raise ValueError(f"Missing condition fields: {sorted(missing)}")
n = len(self.obs)
if n == 0:
raise ValueError("No eligible treated conditions")
if self.response.ndim != 2 or self.response.shape[0] != n:
raise ValueError("Response matrix and condition table are misaligned")
for name in ["response", "control", "treated_var", "control_var"]:
a = getattr(self, name)
if a.shape != self.response.shape or not np.isfinite(a).all():
raise ValueError(f"Invalid {name} matrix")
if (self.treated_var < 0).any() or (self.control_var < 0).any():
raise ValueError("Sampling variances must be nonnegative")
if not np.isfinite(self.obs.dose_um).all() or (self.obs.dose_um <= 0).any():
raise ValueError("Treated concentrations must be positive, finite micromolar values")
if self.obs[REQUIRED].isna().any().any():
raise ValueError("Condition identifiers and splits cannot be missing")
if self.obs.groupby("molecule_id").split.nunique().max() > 1:
raise ValueError("A molecule appears in multiple splits")
if self.obs.groupby("smiles").split.nunique().max() > 1:
raise ValueError("An identical structure appears in multiple splits")
for cid, group in self.obs.groupby("control_id", sort=False):
idx = group.index.to_numpy()
if not np.allclose(self.control[idx], self.control[idx[0]]):
raise ValueError(f"Shared control {cid} has inconsistent summaries")
return self
def subset(self, mask):
idx = np.flatnonzero(np.asarray(mask))
return ResponseData(self.obs.iloc[idx].reset_index(drop=True), self.response[idx],
self.control[idx], self.treated_var[idx], self.control_var[idx],
dict(self.metadata)).validate()
def save(self, directory):
self.validate()
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
self.obs.to_csv(directory / "conditions.csv", index=False)
np.savez_compressed(directory / "responses.npz", response=self.response,
control=self.control, treated_var=self.treated_var,
control_var=self.control_var)
write_json(directory / "metadata.json", self.metadata)
@classmethod
def load(cls, directory):
directory = Path(directory)
obs = pd.read_csv(directory / "conditions.csv", dtype={k: str for k in REQUIRED if k != "dose_um"})
with np.load(directory / "responses.npz", allow_pickle=False) as a:
values = {k: np.array(a[k], dtype=float) for k in ["response", "control", "treated_var", "control_var"]}
return cls(obs, **values, metadata=json.loads((directory / "metadata.json").read_text())).validate()