File size: 3,996 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
"""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()