315 lines
9.6 KiB
Python
315 lines
9.6 KiB
Python
"""Pinned RDKit 2D descriptor computation for the descriptor-batch workload.
|
|
|
|
The scientific contract of ``descriptor-batch`` is deliberately small and
|
|
fully pinned:
|
|
|
|
- exactly one output row per valid input molecule, in input order;
|
|
- RDKit canonical SMILES recomputed with ``MolToSmiles(..., canonical=True)``;
|
|
- the descriptor set is an explicit, versioned tuple of RDKit
|
|
``Descriptors.descList`` names (2D only), not a scan of installed names;
|
|
- float values are serialized with fixed ``%.6f`` formatting so that the
|
|
output is byte-identical for identical inputs and a pinned environment;
|
|
- invalid SMILES rows are either skipped (counted) or fail the run, selected
|
|
by the explicit ``skip_invalid`` parameter.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import csv
|
|
from dataclasses import dataclass
|
|
from functools import lru_cache
|
|
from pathlib import Path
|
|
from typing import Any, Iterator, Mapping, Sequence
|
|
|
|
from rdkit import Chem
|
|
from rdkit.ML.Descriptors.MoleculeDescriptors import MolecularDescriptorCalculator
|
|
|
|
from scimesh.chemistry.dataset import iter_rows
|
|
|
|
# Explicit pinned list. Names must exist in the installed RDKit ``descList``;
|
|
# the list itself is the reproducibility contract and must change version
|
|
# together with the workload (descriptor-batch@1.0.0).
|
|
DESCRIPTOR_NAMES: tuple[str, ...] = (
|
|
"ExactMolWt",
|
|
"MolWt",
|
|
"HeavyAtomMolWt",
|
|
"HeavyAtomCount",
|
|
"NumHDonors",
|
|
"NumHAcceptors",
|
|
"NumRotatableBonds",
|
|
"NumHeteroatoms",
|
|
"NumRadicalElectrons",
|
|
"NumValenceElectrons",
|
|
"FractionCSP3",
|
|
"RingCount",
|
|
"NumAromaticRings",
|
|
"NumSaturatedRings",
|
|
"NumAliphaticRings",
|
|
"NumAromaticHeterocycles",
|
|
"NumSaturatedHeterocycles",
|
|
"NumAliphaticHeterocycles",
|
|
"NumAromaticCarbocycles",
|
|
"NumSaturatedCarbocycles",
|
|
"NumAliphaticCarbocycles",
|
|
"TPSA",
|
|
"LabuteASA",
|
|
"MolLogP",
|
|
"MolMR",
|
|
"BalabanJ",
|
|
"BertzCT",
|
|
"HallKierAlpha",
|
|
"Kappa1",
|
|
"Kappa2",
|
|
"Kappa3",
|
|
"Chi0",
|
|
"Chi1",
|
|
"Chi0n",
|
|
"Chi1n",
|
|
"Chi2n",
|
|
"Chi3n",
|
|
"Chi4n",
|
|
"Chi0v",
|
|
"Chi1v",
|
|
"Chi2v",
|
|
"Chi3v",
|
|
"Chi4v",
|
|
"PEOE_VSA1",
|
|
"PEOE_VSA2",
|
|
"PEOE_VSA3",
|
|
"PEOE_VSA4",
|
|
"PEOE_VSA5",
|
|
"PEOE_VSA6",
|
|
"PEOE_VSA7",
|
|
"PEOE_VSA8",
|
|
"PEOE_VSA9",
|
|
"PEOE_VSA10",
|
|
"PEOE_VSA11",
|
|
"PEOE_VSA12",
|
|
"PEOE_VSA13",
|
|
"PEOE_VSA14",
|
|
"SMR_VSA1",
|
|
"SMR_VSA2",
|
|
"SMR_VSA3",
|
|
"SMR_VSA4",
|
|
"SMR_VSA5",
|
|
"SMR_VSA6",
|
|
"SMR_VSA7",
|
|
"SMR_VSA8",
|
|
"SMR_VSA9",
|
|
"SMR_VSA10",
|
|
"SlogP_VSA1",
|
|
"SlogP_VSA2",
|
|
"SlogP_VSA3",
|
|
"SlogP_VSA4",
|
|
"SlogP_VSA5",
|
|
"SlogP_VSA6",
|
|
"SlogP_VSA7",
|
|
"SlogP_VSA8",
|
|
"SlogP_VSA9",
|
|
"SlogP_VSA10",
|
|
"SlogP_VSA11",
|
|
"SlogP_VSA12",
|
|
"NHOHCount",
|
|
"NOCount",
|
|
)
|
|
|
|
DESCRIPTOR_COLUMNS: tuple[str, ...] = (
|
|
"chembl_id",
|
|
"canonical_smiles",
|
|
) + DESCRIPTOR_NAMES
|
|
|
|
|
|
@lru_cache(maxsize=1)
|
|
def descriptor_calculator() -> MolecularDescriptorCalculator:
|
|
"""Build the pinned calculator once per process."""
|
|
return MolecularDescriptorCalculator(DESCRIPTOR_NAMES)
|
|
|
|
|
|
def validate_descriptor_names() -> None:
|
|
"""Fail fast when the pinned list is unavailable in the installed RDKit."""
|
|
from rdkit.Chem import Descriptors
|
|
|
|
available = {name for name, _ in Descriptors.descList}
|
|
missing = [name for name in DESCRIPTOR_NAMES if name not in available]
|
|
if missing:
|
|
raise ValueError(
|
|
"pinned descriptor-batch descriptors are missing from RDKit: "
|
|
+ ", ".join(missing)
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class DescriptorRow:
|
|
"""One canonical descriptor row for a valid input molecule."""
|
|
|
|
molecule_id: str
|
|
canonical_smiles: str
|
|
values: tuple[float, ...]
|
|
|
|
|
|
class DescriptorStats:
|
|
"""Row counters collected while computing a descriptor batch."""
|
|
|
|
def __init__(self) -> None:
|
|
self.scanned = 0
|
|
self.invalid = 0
|
|
self.emitted = 0
|
|
|
|
def as_metrics(self) -> dict[str, int]:
|
|
return {
|
|
"rows_scanned": self.scanned,
|
|
"invalid_rows": self.invalid,
|
|
"rows_emitted": self.emitted,
|
|
}
|
|
|
|
|
|
def iter_descriptor_rows(
|
|
input_path: Path,
|
|
*,
|
|
skip_invalid: bool = True,
|
|
) -> tuple[Iterator[DescriptorRow], DescriptorStats]:
|
|
"""Yield canonical descriptor rows in input order with streaming stats."""
|
|
calculator = descriptor_calculator()
|
|
stats = DescriptorStats()
|
|
|
|
def generate() -> Iterator[DescriptorRow]:
|
|
for row in iter_rows(input_path):
|
|
stats.scanned += 1
|
|
smiles = row.get("canonical_smiles", "")
|
|
molecule = Chem.MolFromSmiles(smiles)
|
|
if molecule is None:
|
|
stats.invalid += 1
|
|
if not skip_invalid:
|
|
raise ValueError(
|
|
f"row {stats.scanned} has an invalid canonical_smiles"
|
|
)
|
|
continue
|
|
canonical = Chem.MolToSmiles(molecule, canonical=True)
|
|
values = tuple(
|
|
float(value) for value in calculator.CalcDescriptors(molecule)
|
|
)
|
|
stats.emitted += 1
|
|
yield DescriptorRow(row.get("chembl_id", ""), canonical, values)
|
|
|
|
return generate(), stats
|
|
|
|
|
|
def write_descriptor_rows(
|
|
output_path: Path,
|
|
rows: Sequence[DescriptorRow],
|
|
) -> None:
|
|
"""Write a canonical one-row-per-input descriptor CSV with one header."""
|
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
|
writer = csv.DictWriter(
|
|
destination, fieldnames=list(DESCRIPTOR_COLUMNS), lineterminator="\n"
|
|
)
|
|
writer.writeheader()
|
|
for row in rows:
|
|
writer.writerow(
|
|
{
|
|
"chembl_id": row.molecule_id,
|
|
"canonical_smiles": row.canonical_smiles,
|
|
**{
|
|
name: f"{value:.6f}"
|
|
for name, value in zip(DESCRIPTOR_NAMES, row.values)
|
|
},
|
|
}
|
|
)
|
|
|
|
|
|
def compute_descriptor_batch(
|
|
input_path: Path,
|
|
output_path: Path,
|
|
*,
|
|
skip_invalid: bool = True,
|
|
) -> dict[str, int]:
|
|
"""Single-process reference: read the whole input and write the CSV."""
|
|
rows, stats = iter_descriptor_rows(input_path, skip_invalid=skip_invalid)
|
|
materialized = list(rows)
|
|
write_descriptor_rows(output_path, materialized)
|
|
return stats.as_metrics()
|
|
|
|
|
|
def write_descriptor_shards(
|
|
input_path: Path,
|
|
workspace: Path,
|
|
shard_rows: int,
|
|
) -> list[Path]:
|
|
"""Split the input TSV into deterministic row-bounded shards with headers."""
|
|
if (
|
|
isinstance(shard_rows, bool)
|
|
or not isinstance(shard_rows, int)
|
|
or shard_rows < 1
|
|
):
|
|
raise ValueError("shard_rows must be a positive integer")
|
|
paths: list[Path] = []
|
|
current: Path | None = None
|
|
destination = None
|
|
writer = None
|
|
rows_in_shard = 0
|
|
try:
|
|
with input_path.open("r", encoding="utf-8", newline="") as source:
|
|
reader = csv.DictReader(source, delimiter="\t")
|
|
fieldnames = tuple(reader.fieldnames or ())
|
|
if not {"chembl_id", "canonical_smiles"}.issubset(set(fieldnames)):
|
|
raise ValueError(
|
|
"dataset is missing required columns: chembl_id, canonical_smiles"
|
|
)
|
|
for row in reader:
|
|
if destination is None or rows_in_shard == shard_rows:
|
|
if destination is not None:
|
|
destination.close()
|
|
current = workspace / f"shard-{len(paths)}.tsv"
|
|
destination = current.open("w", encoding="utf-8", newline="")
|
|
writer = csv.DictWriter(
|
|
destination,
|
|
fieldnames=list(fieldnames),
|
|
delimiter="\t",
|
|
lineterminator="\n",
|
|
)
|
|
writer.writeheader()
|
|
paths.append(current)
|
|
rows_in_shard = 0
|
|
assert writer is not None
|
|
writer.writerow(row)
|
|
rows_in_shard += 1
|
|
finally:
|
|
if destination is not None:
|
|
destination.close()
|
|
if not paths:
|
|
raise ValueError("dataset has no data rows")
|
|
return paths
|
|
|
|
|
|
def concatenate_descriptor_shards(
|
|
partial_paths: Sequence[Path],
|
|
output_path: Path,
|
|
) -> dict[str, int]:
|
|
"""Merge shard partial CSVs by shard index with exactly one header.
|
|
|
|
Every partial is a full CSV with the same header. The first partial is
|
|
copied verbatim; each later partial contributes only its data rows, so the
|
|
merged file is byte-identical to the single-process reference for the same
|
|
input rows.
|
|
"""
|
|
if not partial_paths:
|
|
raise ValueError("descriptor reducer requires at least one partial")
|
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
rows_emitted = 0
|
|
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
|
for index, partial in enumerate(partial_paths):
|
|
with partial.open("r", encoding="utf-8", newline="") as source:
|
|
for line_index, line in enumerate(source):
|
|
if line_index == 0:
|
|
if index > 0:
|
|
continue
|
|
if line.rstrip("\r\n") != ",".join(DESCRIPTOR_COLUMNS):
|
|
raise ValueError(
|
|
"partial descriptor CSV has an invalid header"
|
|
)
|
|
destination.write(line)
|
|
if line_index > 0:
|
|
rows_emitted += 1
|
|
return {"partial_count": len(partial_paths), "rows_emitted": rows_emitted}
|