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
122 changes: 122 additions & 0 deletions src/molgen3D/data_processing/smiles_encoder_decoder.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,10 @@
import ast
import json
import math
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
from rdkit import Chem
from rdkit.Chem import AllChem
from rdkit.Chem.rdchem import ChiralType
Expand Down Expand Up @@ -85,6 +89,7 @@ def strip_smiles(s: str) -> str:
# Remove any tags that might be present
s = s.replace("[CONFORMER]", "").replace("[/CONFORMER]", "")
s = s.replace("[SMILES]", "").replace("[/SMILES]", "")
s = re.sub(r"\[SERIALIZATION\].*?\[/SERIALIZATION\]", "", s, flags=re.IGNORECASE | re.DOTALL)
s = s.replace(";", "")

s = _WHITESPACE_RE.sub('', s)
Expand Down Expand Up @@ -616,6 +621,95 @@ def bins_to_coords(bin_indices, bins, use_bin_center=False):
return np.array(coords)


@dataclass
class BinConfig:
mode: str
L: float
H: float
n_bins: int
edges: np.ndarray
digit_width: int = 3

def __post_init__(self):
self.digit_width = max(3, len(str(self.n_bins + 1)))

@classmethod
def load(cls, path: str) -> "BinConfig":
with open(path) as f:
obj = json.load(f)
return cls(
mode=obj["mode"],
L=obj["L"],
H=obj["H"],
n_bins=obj["n_bins"],
edges=np.array(obj["edges"], dtype=np.float64),
)


def _decode_scalar(idx: int, config: BinConfig) -> float:
if idx <= 0:
return config.L
if idx > config.n_bins:
return config.H
return float((config.edges[idx - 1] + config.edges[idx]) / 2.0)


def decode_cartesian_with_config(enriched_string: str, config: BinConfig):
normalized = enriched_string.replace(";", "")
tokens = tokenize_enriched_v2(normalized, config.digit_width)

smiles_parts = []
coords = []
for token in tokens:
if token["type"] == "atom_with_coords":
desc = token["atom_desc"]
desc_inner = desc[1:-1]
if desc_inner in _ORGANIC_SUBSET:
smiles_parts.append(desc_inner)
else:
smiles_parts.append(desc)

ix, iy, iz = (int(round(v)) for v in token["coords"])
x = _decode_scalar(ix, config)
y = _decode_scalar(iy, config)
z = _decode_scalar(iz, config)
coords.append((x, y, z))
else:
smiles_parts.append(token["text"])

smiles = "".join(smiles_parts)
mol = Chem.MolFromSmiles(smiles, sanitize=False)
if mol is None:
raise ValueError(f"Failed to parse rebuilt SMILES: {smiles}")
if mol.GetNumAtoms() != len(coords):
raise ValueError(
f"Atom count mismatch: mol has {mol.GetNumAtoms()} atoms, "
f"coords list has {len(coords)} entries."
)

Chem.SanitizeMol(mol)

conformer = Chem.Conformer(mol.GetNumAtoms())
for idx, (x, y, z) in enumerate(coords):
conformer.SetAtomPosition(idx, Point3D(x, y, z))
mol.AddConformer(conformer, assignId=True)
return mol


def normalize_serialization_tag(serialization_tag: str) -> str:
mode = str(serialization_tag).strip().lower()
aliases = {
"binned": "cartesian_binned",
"cartesian_binned_v2": "cartesian_binned",
"uniform_binned": "uniform",
"quantile_binned": "quantile",
}
normalized = aliases.get(mode, mode)
valid_modes = {"cartesian", "cartesian_binned", "uniform", "quantile"}
if normalized not in valid_modes:
raise ValueError(f"Unsupported serialization mode: {serialization_tag}")
return normalized


def encode_cartesian_binned(mol, bin_size, ranges=None):
"""
Expand Down Expand Up @@ -855,3 +949,31 @@ def decode_cartesian_binned_v2(enriched_string, bins, use_bin_center=True):
conformer.SetAtomPosition(idx, Point3D(x, y, z))
mol.AddConformer(conformer, assignId=True)
return mol


def decode_conformer_by_serialization(
enriched_string: str,
serialization_tag: str,
*,
bins=None,
uniform_config_path: Optional[str] = None,
quantile_config_path: Optional[str] = None,
):
mode = str(serialization_tag)
if mode == "cartesian":
return decode_cartesian_v2(enriched_string)
if mode == "cartesian_binned":
if bins is None:
raise ValueError("`bins` must be provided for cartesian_binned decoding.")
return decode_cartesian_binned_v2(enriched_string, bins)
if mode == "uniform":
cfg_path = uniform_config_path or str(
Path(__file__).resolve().parents[1] / "config" / "bin_configs" / "uniform_bins.json"
)
return decode_cartesian_with_config(enriched_string, BinConfig.load(cfg_path))
if mode == "quantile":
cfg_path = quantile_config_path or str(
Path(__file__).resolve().parents[1] / "config" / "bin_configs" / "quantile_bins.json"
)
return decode_cartesian_with_config(enriched_string, BinConfig.load(cfg_path))
raise ValueError(f"Unsupported serialization mode: {serialization_tag}")
84 changes: 75 additions & 9 deletions src/molgen3D/evaluation/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,12 @@

# from utils import parse_molecule_with_coordinates
from molgen3D.data_processing.utils import decode_cartesian_raw
from molgen3D.data_processing.smiles_encoder_decoder import decode_cartesian_v2, strip_smiles, decode_cartesian_binned_v2, get_bins_for_coords
from molgen3D.data_processing.smiles_encoder_decoder import (
decode_conformer_by_serialization,
normalize_serialization_tag,
strip_smiles,
get_bins_for_coords,
)
from molgen3D.evaluation.utils import (
extract_between,
same_molecular_graph,
Expand Down Expand Up @@ -98,7 +103,17 @@ def save_results(results_path, generations, stats):
with open(os.path.join(results_path, "generation_results.txt"), 'w') as results_file_txt:
results_file_txt.write(f"{stats=}")

def process_batch(model, tokenizer, batch: list[list], gen_config, eos_token_id, binned: bool):
def process_batch(
model,
tokenizer,
batch: list[list],
gen_config,
eos_token_id,
binned: bool,
serialization_tag: str = "cartesian",
uniform_bin_config_path: str = None,
quantile_bin_config_path: str = None,
):
# Create bins for binned decoding (must match encoding bins)
bins = None
if binned:
Expand Down Expand Up @@ -155,10 +170,13 @@ def process_batch(model, tokenizer, batch: list[list], gen_config, eos_token_id,
stats["smiles_mismatch"] += 1
else:
try:
if binned:
mol_obj = decode_cartesian_binned_v2(generated_conformer, bins)
else:
mol_obj = decode_cartesian_v2(generated_conformer)
mol_obj = decode_conformer_by_serialization(
generated_conformer,
serialization_tag,
bins=bins,
uniform_config_path=uniform_bin_config_path,
quantile_config_path=quantile_bin_config_path,
)
generations[geom_smiles].append(mol_obj)
except Exception as e:
logger.info(f"smiles fails parsing: \n{canonical_smiles=}\n{generated_smiles=}\n{generated_conformer=}")
Expand Down Expand Up @@ -231,9 +249,7 @@ def run_inference(inference_config: dict):
# Validation set format: {smiles: [mol_obj1, mol_obj2, ...]}
for geom_smiles, mol_list in test_data.items():
if isinstance(mol_list, list):
# Each entry has a list of ground truth conformers
num_ground_truths = len(mol_list)
# Generate 2x the ground truth count
mols_list.extend([(geom_smiles, f"[SMILES]{geom_smiles}[/SMILES]")] * num_ground_truths * 2)
else:
logger.warning(f"Unexpected data format for {geom_smiles}, skipping")
Expand All @@ -243,6 +259,11 @@ def run_inference(inference_config: dict):
icl_prompt = data.get('icl_prompt')
if icl_prompt:
mols_list.extend([(geom_smiles, icl_prompt)] * data.get("num_confs", 1) * 2)
elif test_set == "revisited":
logger.info("Processing as revisited dataset")
for geom_smiles, data in test_data.items():
for sub_smiles, count in data["sub_smiles_counts"].items():
mols_list.extend([(geom_smiles, f"[SMILES]{sub_smiles}[/SMILES]")] * count * 2)
logger.info(f"mols_list length: {len(mols_list)}, mols_list_distinct: {len(set(mols_list))}, mols_list: {mols_list[:10]}")

mols_list.sort(key=lambda x: len(x[0]))
Expand All @@ -258,6 +279,15 @@ def run_inference(inference_config: dict):
logger.info("Auto-detecting binned=True based on model path")
binned = True

serialization_tag = normalize_serialization_tag(
inference_config.get("serialization_tag", "cartesian")
)
if binned and serialization_tag == "cartesian":
serialization_tag = "cartesian_binned"
if serialization_tag == "cartesian_binned":
binned = True
logger.info(f"Using serialization_tag={serialization_tag}")

for start in tqdm(range(0, len(mols_list), batch_size), desc="generating"):
batch = mols_list[start:start + batch_size]
for sub_batch in split_batch_on_geom_size(batch, max_geom_len=80):
Expand All @@ -268,6 +298,9 @@ def run_inference(inference_config: dict):
gen_config=inference_config["gen_config"],
eos_token_id=eos_token_id,
binned=binned,
serialization_tag=serialization_tag,
uniform_bin_config_path=inference_config.get("uniform_bin_config_path"),
quantile_bin_config_path=inference_config.get("quantile_bin_config_path"),
)
stats.update(stats_)
for k, v in outputs.items():
Expand All @@ -287,6 +320,9 @@ def launch_inference_from_cli(
valid: bool = False,
limit: int = None,
binned: bool = False,
serialization_tag: str = "cartesian",
uniform_bin_config_path: str = None,
quantile_bin_config_path: str = None,
icl: bool = False,
icl_n: int = 5,
parallel_jobs: int = 1
Expand Down Expand Up @@ -372,6 +408,9 @@ def launch_inference_from_cli(
"run_name": "qwen_pre_binned",
"limit": limit,
"binned": binned,
"serialization_tag": serialization_tag,
"uniform_bin_config_path": uniform_bin_config_path,
"quantile_bin_config_path": quantile_bin_config_path,
}

if grid_run_inference:
Expand Down Expand Up @@ -469,8 +508,32 @@ def launch_inference_from_cli(
parser = argparse.ArgumentParser()
parser.add_argument("--device", type=str, choices=["local", "a100", "h100", "all"], default="local")
parser.add_argument("--grid_run_inference", action="store_true")
parser.add_argument("--test_set", type=str, choices=["clean", "distinct", "corrected"], default=None)
parser.add_argument(
"--test_set",
type=str,
choices=["clean", "distinct", "corrected", "xl", "qm9", "valid", "revisited"],
default=None,
)
parser.add_argument("--binned", action="store_true", default=False)
parser.add_argument(
"--serialization_tag",
type=str,
choices=["cartesian", "cartesian_binned", "uniform", "quantile", "binned"],
default="cartesian",
help="Select decoding scheme.",
)
parser.add_argument(
"--uniform_bin_config_path",
type=str,
default=None,
help="Optional BinConfig path for uniform decoding.",
)
parser.add_argument(
"--quantile_bin_config_path",
type=str,
default=None,
help="Optional BinConfig path for quantile decoding.",
)
parser.add_argument("--xl", action="store_true")
parser.add_argument("--qm9", action="store_true")
parser.add_argument("--valid", action="store_true", help="Run inference on validation set")
Expand All @@ -488,6 +551,9 @@ def launch_inference_from_cli(
valid=args.valid,
limit=args.limit,
binned=args.binned,
serialization_tag=args.serialization_tag,
uniform_bin_config_path=args.uniform_bin_config_path,
quantile_bin_config_path=args.quantile_bin_config_path,
icl=args.icl,
icl_n=args.icl_n,
parallel_jobs=args.parallel_jobs
Expand Down
Loading