Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
79 changes: 54 additions & 25 deletions src/molgen3D/data_processing/data_preprocessing.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,12 @@
import argparse
import ast
import glob
import json
import os
import os.path as osp
import random
from collections import defaultdict
from multiprocessing import Pool
from typing import Any, Dict, Optional, Set, Tuple, List
from typing import Any, Dict, List, Optional, Set, Tuple
import numpy as np
from loguru import logger as log
from rdkit import Chem, RDLogger
Expand All @@ -16,25 +15,32 @@
from molgen3D.data_processing.utils import (
JsonlSplitWriter,
filter_mols,
get_embedding_func_and_config,
parse_coordinate_ranges,
save_processed_pickle,
)
from molgen3D.data_processing.smiles_encoder_decoder import (
encode_cartesian_v2,
BinConfig,
encode_cartesian_binned,
encode_cartesian_binned_v2,
encode_cartesian_with_config,
)
from molgen3D.utils.utils import load_pkl

RDLogger.DisableLog("rdApp.*")


def read_mol(
args: Tuple[str, int, int, Any, float, List[Tuple[float, float]], bool, str, str]
args: Tuple,
) -> Optional[Tuple[List[str], Dict[str, Any]]]:
mol_path, max_confs, precision, embedding_func, bin_size, ranges, do_filter, pickle_dir, _geom_root = args
mol_path, max_confs, precision, embedding_func, bin_size, ranges, do_filter, pickle_dir, _geom_root = args[:9]
bin_config = args[9] if len(args) > 9 else None
use_isomeric_smiles = args[10] if len(args) > 10 else False
try:
return _read_mol_impl(
mol_path, max_confs, precision, embedding_func, bin_size, ranges, do_filter, pickle_dir
mol_path, max_confs, precision, embedding_func, bin_size, ranges, do_filter, pickle_dir,
bin_config=bin_config,
use_isomeric_smiles=use_isomeric_smiles,
)
except Exception as exc:
log.error("Unhandled exception in read_mol | path={} | error={}", mol_path, exc)
Expand All @@ -50,6 +56,8 @@ def _read_mol_impl(
ranges: List[Tuple[float, float]],
do_filter: bool,
pickle_dir: str,
bin_config: Optional[BinConfig] = None,
use_isomeric_smiles: bool = False,
) -> Tuple[List[str], Dict[str, Any]]:
mol_object = load_pkl(mol_path)
geom_smiles = mol_object["smiles"]
Expand All @@ -75,7 +83,9 @@ def _read_mol_impl(
continue

try:
if embedding_func in (encode_cartesian_binned, encode_cartesian_binned_v2):
if embedding_func is encode_cartesian_with_config:
embedded_smile, iso_smile = embedding_func(mol, bin_config)
elif embedding_func in (encode_cartesian_binned, encode_cartesian_binned_v2):
embedded_smile, iso_smile = embedding_func(mol, bin_size=bin_size, ranges=ranges)
else:
embedded_smile, iso_smile = embedding_func(mol, precision=precision)
Expand All @@ -85,6 +95,7 @@ def _read_mol_impl(
continue

# Compute nonisomeric SMILES only for conformers that encoded successfully
noniso = None
try:
noniso = Chem.MolToSmiles(Chem.RemoveHs(mol, sanitize=False), canonical=True, isomericSmiles=False)
nonisomeric_smiles.add(noniso)
Expand All @@ -93,10 +104,11 @@ def _read_mol_impl(
except Exception:
pass

canonical_smiles = iso_smile if use_isomeric_smiles else (noniso or iso_smile)
samples.append(
json.dumps(
{
"canonical_smiles": iso_smile,
"canonical_smiles": canonical_smiles,
"embedded_smiles": embedded_smile,
},
separators=(",", ":"),
Expand Down Expand Up @@ -158,19 +170,16 @@ def preprocess(
bin_size: float = 0.104,
ranges: str = "[-13.0, 13.0], [-13.0, 13.0], [-13.0, 13.0]",
filter_ranges: str = None,
bin_config_path: Optional[str] = None,
use_isomeric_smiles: bool = False,
) -> None:
if dest_path is None:
raise ValueError("dest_path must be provided for preprocessing output")

embedding_registry = {
"cartesian_v2": encode_cartesian_v2,
"cartesian": encode_cartesian_v2,
"cartesian_binned": encode_cartesian_binned,
"cartesian_binned_v2": encode_cartesian_binned_v2,
}
if embedding_type not in embedding_registry:
raise ValueError(f"Unsupported embedding_type '{embedding_type}'. Options: {sorted(embedding_registry)}")
embedding_func = embedding_registry[embedding_type]
embedding_func, bin_config = get_embedding_func_and_config(
embedding_type=embedding_type,
bin_config_path=bin_config_path,
)

overall_total_input_mols = overall_total_confs = overall_total_mols = 0
overall_multi_distinct_graphs = overall_mol_with_dotted_smiles = overall_total_dotted_smiles = 0
Expand Down Expand Up @@ -199,13 +208,7 @@ def preprocess(
if pickle_paths.size == 0:
raise FileNotFoundError(f"No pickle files found under pattern {pickle_glob}")

# Parse ranges string once
try:
parsed_ranges = ast.literal_eval(f"[{ranges}]")
parsed_ranges = [tuple(r) for r in parsed_ranges]
except Exception as e:
log.error(f"Failed to parse ranges: {ranges}. Error: {e}")
parsed_ranges = [(-13.0, 13.0), (-13.0, 13.0), (-13.0, 13.0)]
parsed_ranges = parse_coordinate_ranges(ranges)

do_filter = False
if filter_ranges is not None:
Expand Down Expand Up @@ -250,6 +253,8 @@ def preprocess(
do_filter,
split_pickle_dirs[split_name],
geom_raw_path,
bin_config,
use_isomeric_smiles,
)
for path in mol_paths
]
Expand Down Expand Up @@ -379,7 +384,8 @@ def preprocess(
"--embedding_type",
"-et",
type=str,
choices=["cartesian", "cartesian_v2", "cartesian_binned", "cartesian_binned_v2"],
choices=["cartesian", "cartesian_v2", "cartesian_binned", "cartesian_binned_v2",
"uniform_binned", "quantile_binned"],
default="cartesian_v2",
help="Embedding type to use for enrichment.",
)
Expand Down Expand Up @@ -446,6 +452,25 @@ def preprocess(
default=None,
help="Filter ranges for binned embedding.",
)
parser.add_argument(
"--bin_config_path",
type=str,
default=None,
help="Path to BinConfig JSON (required for uniform_binned / quantile_binned).",
)

parser.add_argument(
"--isomeric",
action="store_true",
help="Alias for --use_isomeric_smiles.",
)
parser.add_argument(
"--sort_by",
type=str,
choices=["energy", "weight", "none"],
default="energy",
help="Sort conformers by energy, weight, or keep original order.",
)
args = parser.parse_args()

dest_path = osp.join(args.dest, args.run_name)
Expand All @@ -458,6 +483,8 @@ def preprocess(
diagnose=False,
)

use_isomeric = args.isomeric

preprocess(
geom_raw_path=args.geom_raw_path,
indices_path=args.indices_path,
Expand All @@ -471,5 +498,7 @@ def preprocess(
bin_size=args.bin_size,
ranges=args.ranges,
filter_ranges=args.filter_ranges,
bin_config_path=args.bin_config_path,
use_isomeric_smiles=use_isomeric,
)

Loading