From afc1652ac19c43330715844e6cb5134702ed927c Mon Sep 17 00:00:00 2001 From: T4ras123 Date: Fri, 20 Mar 2026 09:07:23 +0000 Subject: [PATCH] Enhance SMILES decoding and inference capabilities Made-with: Cursor --- .../data_processing/smiles_encoder_decoder.py | 122 ++++++++++++++++++ src/molgen3D/evaluation/inference.py | 84 ++++++++++-- .../evaluation/inference_multiconf.py | 80 +++++++++--- 3 files changed, 262 insertions(+), 24 deletions(-) diff --git a/src/molgen3D/data_processing/smiles_encoder_decoder.py b/src/molgen3D/data_processing/smiles_encoder_decoder.py index bc853ae..8f1f4ca 100644 --- a/src/molgen3D/data_processing/smiles_encoder_decoder.py +++ b/src/molgen3D/data_processing/smiles_encoder_decoder.py @@ -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 @@ -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) @@ -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): """ @@ -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}") diff --git a/src/molgen3D/evaluation/inference.py b/src/molgen3D/evaluation/inference.py index 25b3e6d..85eb396 100644 --- a/src/molgen3D/evaluation/inference.py +++ b/src/molgen3D/evaluation/inference.py @@ -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, @@ -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: @@ -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=}") @@ -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") @@ -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])) @@ -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): @@ -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(): @@ -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 @@ -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: @@ -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") @@ -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 diff --git a/src/molgen3D/evaluation/inference_multiconf.py b/src/molgen3D/evaluation/inference_multiconf.py index d3297c6..9334f9c 100644 --- a/src/molgen3D/evaluation/inference_multiconf.py +++ b/src/molgen3D/evaluation/inference_multiconf.py @@ -252,6 +252,9 @@ def generate_multiple_conformers( stats: Counter, geom_smiles: str, current_output: str = None, + serialization_tag: str = "cartesian", + uniform_bin_config_path: str = None, + quantile_bin_config_path: str = None, ) -> tuple[List, str]: """ Generate multiple conformers for a single SMILES by forcing conformer continuation. @@ -276,10 +279,9 @@ def generate_multiple_conformers( ConformerCountStoppingCriteria, ) from molgen3D.data_processing.smiles_encoder_decoder import ( - decode_cartesian_v2, + decode_conformer_by_serialization, strip_smiles, - decode_cartesian_binned_v2, - get_bins_for_coords + get_bins_for_coords, ) from molgen3D.evaluation.utils import same_molecular_graph, log_mfu @@ -416,10 +418,13 @@ def generate_multiple_conformers( else: # Try to decode the conformer 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, + ) mol_objects.append(mol_obj) except Exception as e: if stats["mol_parse_fail"] < 20: @@ -439,6 +444,9 @@ def generate_multiple_conformers_batched( gen_config, binned: bool, stats: Counter, + serialization_tag: str = "cartesian", + uniform_bin_config_path: str = None, + quantile_bin_config_path: str = None, ) -> List[List]: """ Generate multiple conformers for a BATCH of SMILES in parallel. @@ -463,10 +471,9 @@ def generate_multiple_conformers_batched( ConformerCountStoppingCriteria, ) from molgen3D.data_processing.smiles_encoder_decoder import ( - decode_cartesian_v2, + decode_conformer_by_serialization, strip_smiles, - decode_cartesian_binned_v2, - get_bins_for_coords + get_bins_for_coords, ) from molgen3D.evaluation.utils import same_molecular_graph, log_mfu @@ -599,10 +606,13 @@ def generate_multiple_conformers_batched( 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, + ) mol_objects.append(mol_obj) except Exception as e: stats["mol_parse_fail"] += 1 @@ -705,12 +715,21 @@ def run_multiconf_inference(inference_config: dict): # Get configuration parameters conformers_per_batch = inference_config.get("conformers_per_batch", 8) binned = inference_config.get("binned", False) + from molgen3D.data_processing.smiles_encoder_decoder import normalize_serialization_tag + + serialization_tag = normalize_serialization_tag( + inference_config.get("serialization_tag", "cartesian") + ) logger.info(f"Conformers per batch: {conformers_per_batch}") if not binned and "binned" in str(inference_config["model_path"]): logger.info("Auto-detecting binned=True based on model path") binned = True + if binned and serialization_tag == "cartesian": + serialization_tag = "cartesian_binned" + + logger.info(f"Using serialization_tag={serialization_tag}") # Initialize statistics and results stats = Counter({ @@ -781,6 +800,9 @@ def run_multiconf_inference(inference_config: dict): gen_config=inference_config["gen_config"], binned=binned, stats=stats, + 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"), ) # Accumulate results and update remaining counts @@ -833,6 +855,9 @@ def launch_multiconf_inference_from_cli( conformer_multiplier: int = 2, limit: Optional[int] = None, binned: bool = False, + serialization_tag: str = "cartesian", + uniform_bin_config_path: str = None, + quantile_bin_config_path: str = None, parallel_jobs: int = 1, ) -> None: """Launch multi-conformer inference from CLI arguments. @@ -922,7 +947,10 @@ def launch_multiconf_inference_from_cli( "conformers_per_batch": conformers_per_batch, "conformer_multiplier": conformer_multiplier, "limit": limit, - "binned": True, # Auto-enable for binned models + "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: @@ -1027,6 +1055,25 @@ def launch_multiconf_inference_from_cli( help="Limit number of unique SMILES to process (default: 10 for testing)") parser.add_argument("--binned", action="store_true", default=False, help="Use binned decoding") + 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("--parallel_jobs", type=int, default=1, help="Number of parallel inference jobs for local execution") @@ -1045,5 +1092,8 @@ def launch_multiconf_inference_from_cli( conformer_multiplier=args.conformer_multiplier, 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, parallel_jobs=args.parallel_jobs, )