Download src/remedi/io.py from ChatterjeeLab/ReMEDi: direct link, hf CLI and curl.
- Browser
- Download file 4 kB
-
https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/src/remedi/io.py
- Command line
-
hf download hf://ChatterjeeLab/ReMEDi/src/remedi/io.py
-
curl -L -o io.py https://huggingface.co/ChatterjeeLab/ReMEDi/resolve/main/src/remedi/io.py
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() | |
| 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) | |
| 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() | |