From 404cce9718a388834156fd4f3b5543ad88b0cfe8 Mon Sep 17 00:00:00 2001 From: m21hm9 Date: Wed, 9 Sep 2026 21:59:35 +0800 Subject: [PATCH 1/5] Add a Hugging Face pretrained molecular generator with NovoMolGen, GP-MoLFormer, MolGen, and Molexar support. --- README.md | 10 + docs/source/install.rst | 34 +- pyproject.toml | 3 + tests/generator/hfpretrained.py | 383 ++++++++++++ tests/generator/test_causal_lm.py | 122 ++++ tests/generator/test_compat.py | 64 ++ tests/generator/test_finetune.py | 197 ++++++ tests/generator/test_molexar.py | 38 ++ tests/generator/test_seq2seq.py | 67 +++ torch_molecule/__init__.py | 2 + .../generator/pretrained/__init__.py | 3 + .../generator/pretrained/checkpoint.py | 46 ++ torch_molecule/generator/pretrained/compat.py | 115 ++++ .../generator/pretrained/families/__init__.py | 10 + .../pretrained/families/causal_lm.py | 97 +++ .../generator/pretrained/families/molexar.py | 129 ++++ .../generator/pretrained/families/seq2seq.py | 77 +++ .../generator/pretrained/finetune.py | 288 +++++++++ .../pretrained/modeling_pretrained.py | 560 +++++++++++++++++- .../generator/pretrained/registry.py | 43 ++ torch_molecule/generator/pretrained/utils.py | 155 +++++ 21 files changed, 2433 insertions(+), 10 deletions(-) create mode 100644 tests/generator/hfpretrained.py create mode 100644 tests/generator/test_causal_lm.py create mode 100644 tests/generator/test_compat.py create mode 100644 tests/generator/test_finetune.py create mode 100644 tests/generator/test_molexar.py create mode 100644 tests/generator/test_seq2seq.py create mode 100644 torch_molecule/generator/pretrained/__init__.py create mode 100644 torch_molecule/generator/pretrained/checkpoint.py create mode 100644 torch_molecule/generator/pretrained/compat.py create mode 100644 torch_molecule/generator/pretrained/families/__init__.py create mode 100644 torch_molecule/generator/pretrained/families/causal_lm.py create mode 100644 torch_molecule/generator/pretrained/families/molexar.py create mode 100644 torch_molecule/generator/pretrained/families/seq2seq.py create mode 100644 torch_molecule/generator/pretrained/finetune.py create mode 100644 torch_molecule/generator/pretrained/registry.py create mode 100644 torch_molecule/generator/pretrained/utils.py diff --git a/README.md b/README.md index a73941b..09e5780 100644 --- a/README.md +++ b/README.md @@ -56,6 +56,10 @@ See the [List of Supported Models](#list-of-supported-models) section for all av | Model | Required Packages | |-------|-------------------| | HFPretrainedMolecularEncoder | transformers | +| HFPretrainedMolecularGenerator | transformers | +| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed) | +| HFPretrainedMolecularGenerator (GP-MoLFormer) | transformers<=4.56.2 | +| HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | | BFGNNMolecularPredictor | torch-scatter | | GRINMolecularPredictor | torch-scatter | | GRINMolecularPredictor (if enable `repetition_augmentation=True`) | CombineMols | @@ -66,6 +70,12 @@ See the [List of Supported Models](#list-of-supported-models) section for all av **For models that require `transformers`:** `pip install transformers` +**For MolGen (`selfies`):** `pip install "selfies>=2.1"` (tested on 2.x; 3.x is not guaranteed). + +**For GP-MoLFormer:** `pip install "transformers>=4.40,<=4.56.2"`. Do not use this with Molexar in the same environment (Molexar needs `transformers>=5.8`). + +**For Molexar:** `pip install fragment-selfies loguru` and `pip install git+https://github.com/fairydance/Molexar.git`. + ## Usage > More examples can be found in the `examples` and `tests` folders. diff --git a/docs/source/install.rst b/docs/source/install.rst index e37d2f2..d465557 100644 --- a/docs/source/install.rst +++ b/docs/source/install.rst @@ -66,12 +66,28 @@ Additional Packages Some models require extra libraries. Install these packages if you use the corresponding model: -+------------------------------+-------------------+ -| Model | Required Package | -+==============================+===================+ -| HFPretrainedMolecularEncoder | transformers | -+------------------------------+-------------------+ -| BFGNNMolecularPredictor | torch-scatter | -+------------------------------+-------------------+ -| GRINMolecularPredictor | torch-scatter | -+------------------------------+-------------------+ ++----------------------------------------------+----------------------------------------------+ +| Model | Required Package | ++==============================================+==============================================+ +| HFPretrainedMolecularEncoder | transformers | ++----------------------------------------------+----------------------------------------------+ +| HFPretrainedMolecularGenerator | transformers | ++----------------------------------------------+----------------------------------------------+ +| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed)| ++----------------------------------------------+----------------------------------------------+ +| HFPretrainedMolecularGenerator (GP-MoLFormer)| transformers<=4.56.2 | ++----------------------------------------------+----------------------------------------------+ +| HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | ++----------------------------------------------+----------------------------------------------+ +| BFGNNMolecularPredictor | torch-scatter | ++----------------------------------------------+----------------------------------------------+ +| GRINMolecularPredictor | torch-scatter | ++----------------------------------------------+----------------------------------------------+ + +**For models that require** ``transformers``: ``pip install transformers`` + +**For MolGen** (``selfies``): ``pip install "selfies>=2.1"`` (tested on 2.x; 3.x is not guaranteed). + +**For GP-MoLFormer:** ``pip install "transformers>=4.40,<=4.56.2"``. Do not use this with Molexar in the same environment (Molexar needs ``transformers>=5.8``). + +**For Molexar:** ``pip install fragment-selfies loguru`` and ``pip install git+https://github.com/fairydance/Molexar.git``. diff --git a/pyproject.toml b/pyproject.toml index af1055a..b1599a7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,6 +57,9 @@ include-package-data = true [tool.pytest.ini_options] addopts = "--verbose" testpaths = ["tests"] +markers = [ + "integration: tests that download models or require network access", +] [project.optional-dependencies] dev = [ diff --git a/tests/generator/hfpretrained.py b/tests/generator/hfpretrained.py new file mode 100644 index 0000000..25c9ed6 --- /dev/null +++ b/tests/generator/hfpretrained.py @@ -0,0 +1,383 @@ +import pytest + +from torch_molecule.generator.pretrained.registry import resolve_family + + +@pytest.mark.parametrize( + "repo_id,expected", + [ + ("chandar-lab/NovoMolGen_32M_SMILES_BPE", "novomolgen"), + ("ibm-research/GP-MoLFormer-Uniq", "gp_molformer"), + ("zjunlp/MolGen-large", "molgen"), + ("zjunlp/MolGen-large-opt", "molgen"), + ("fairydance/molexar-10m-base", "molexar"), + ("fairydance/molexar-10m-omni", "molexar"), + ("some-user/custom-causal-lm", "causal_lm"), + ], +) +def test_resolve_family(repo_id, expected): + assert resolve_family(repo_id) == expected + + +def test_hf_pretrained_generator_requires_transformers(): + pytest.importorskip("transformers") + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + ) + assert model.repo_id == "chandar-lab/NovoMolGen_32M_SMILES_BPE" + assert model.is_fitted_ is False + + +@pytest.mark.integration +def test_hf_pretrained_generator_novomolgen_smoke(): + pytest.importorskip("transformers") + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + generate_max_length=64, + ) + model.fit() + assert model.is_fitted_ is True + + smiles_list = model.generate(n_samples=2, temperature=1.0) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 2 + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + + +@pytest.mark.integration +def test_hf_pretrained_generator_finetune_smoke(tmp_path): + pytest.importorskip("transformers") + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + batch_size=2, + epochs=1, + generate_max_length=32, + ) + train_smiles = ["CCO", "CC(=O)O", "c1ccccc1"] + model.fit(train_smiles) + + assert model.is_fitted_ is True + assert len(model.fitting_loss) == 1 + assert model.fitting_epoch == 0 + + save_dir = tmp_path / "novomolgen-finetuned" + model.save_to_local(str(save_dir)) + + reloaded = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + generate_max_length=32, + ) + reloaded.load_from_local(str(save_dir)) + assert reloaded.is_fitted_ is True + assert reloaded.repo_id == model.repo_id + assert reloaded._family == "novomolgen" + + smiles_list = reloaded.generate(n_samples=1, temperature=1.0, do_sample=True) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 1 + assert isinstance(smiles_list[0], str) + + +@pytest.mark.integration +def test_hf_pretrained_generator_finetune_warns_on_y(): + pytest.importorskip("transformers") + import numpy as np + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + batch_size=2, + epochs=1, + ) + with pytest.warns(UserWarning, match="Conditional fine-tuning"): + model.fit(["CCO", "CC(=O)O"], y=np.array([0.0, 1.0])) + assert model.is_fitted_ is True + + +def test_smiles_selfies_roundtrip(): + pytest.importorskip("selfies") + + from torch_molecule.generator.pretrained.utils import selfies_to_smiles, smiles_to_selfies + + smiles = ["CCO", "c1ccccc1", "CC(=O)O"] + recovered = selfies_to_smiles(smiles_to_selfies(smiles)) + assert recovered == smiles + + +def test_smiles_to_selfies_invalid_smiles(): + pytest.importorskip("selfies") + + from torch_molecule.generator.pretrained.utils import smiles_to_selfies + + with pytest.raises(ValueError, match="Invalid SMILES"): + smiles_to_selfies(["not-a-smiles"]) + + +def test_smiles_to_selfies_encoder_error_has_index(monkeypatch): + pytest.importorskip("selfies") + import selfies as sf + + from torch_molecule.generator.pretrained.utils import smiles_to_selfies + + def _raise_encoder(_smiles): + raise sf.EncoderError("synthetic encoder failure") + + monkeypatch.setattr(sf, "encoder", _raise_encoder) + with pytest.raises(ValueError, match="index 0.*not SELFIES-encodable"): + smiles_to_selfies(["CCO"]) + + +def test_selfies_to_smiles_drops_invalid_entries(): + pytest.importorskip("selfies") + + from torch_molecule.generator.pretrained.utils import selfies_to_smiles, smiles_to_selfies + + valid = smiles_to_selfies(["CCO"])[0] + with pytest.warns(UserWarning, match="dropped 2 invalid SELFIES"): + recovered = selfies_to_smiles([valid, "not-valid-selfies-[[[", ""]) + assert recovered == ["CCO"] + assert "" not in recovered + + +def test_decode_outputs_molgen_drops_empty_and_warns(): + pytest.importorskip("transformers") + pytest.importorskip("selfies") + + from torch_molecule import HFPretrainedMolecularGenerator + from torch_molecule.generator.pretrained.utils import smiles_to_selfies + + model = HFPretrainedMolecularGenerator(repo_id="zjunlp/MolGen-large") + model._family = "molgen" + valid = smiles_to_selfies(["c1ccccc1"])[0] + with pytest.warns(UserWarning, match="got 1/2 valid SMILES"): + out = model._decode_outputs([valid, "not-a-selfies"]) + assert len(out) == 1 + assert "" not in out + + +def test_known_family_prefix_does_not_warn_unknown_repo(): + pytest.importorskip("transformers") + import warnings + + from torch_molecule import HFPretrainedMolecularGenerator + + with warnings.catch_warnings(record=True) as recorded: + warnings.simplefilter("always") + HFPretrainedMolecularGenerator(repo_id="chandar-lab/NovoMolGen_157M") + assert not any("Unknown repo_id" in str(item.message) for item in recorded) + + +def test_unknown_repo_fallback_warns(): + pytest.importorskip("transformers") + + from torch_molecule import HFPretrainedMolecularGenerator + + with pytest.warns(UserWarning, match="Unknown repo_id"): + HFPretrainedMolecularGenerator(repo_id="some-user/custom-causal-lm") + + +def _transformers_supports_gp_molformer() -> bool: + pytest.importorskip("transformers") + import transformers + + major, minor, _ = map(int, transformers.__version__.split(".")[:3]) + return major < 5 and (major < 4 or minor < 57) + + +@pytest.mark.integration +def test_hf_pretrained_generator_gp_molformer_denovo(): + if not _transformers_supports_gp_molformer(): + pytest.skip("GP-MoLFormer requires transformers<=4.56.2") + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="ibm-research/GP-MoLFormer-Uniq", + generate_max_length=128, + ) + model.fit() + assert model.is_fitted_ is True + + smiles_list = model.generate(n_samples=2, temperature=1.0) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 2 + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + + +@pytest.mark.integration +def test_hf_pretrained_generator_gp_molformer_scaffold(): + if not _transformers_supports_gp_molformer(): + pytest.skip("GP-MoLFormer requires transformers<=4.56.2") + + from torch_molecule import HFPretrainedMolecularGenerator + + # IBM's official conditional prompt is a *partial* SMILES, not a closed ring. + scaffold = "c1cccc" + model = HFPretrainedMolecularGenerator( + repo_id="ibm-research/GP-MoLFormer-Uniq", + generate_max_length=128, + ) + model.fit() + + smiles_list = model.generate(n_samples=2, scaffold=scaffold, temperature=1.0) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 2 + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + assert all(smiles.startswith(scaffold) for smiles in smiles_list) + + +def test_gp_molformer_transformers_version_guard(): + pytest.importorskip("transformers") + import transformers + + from torch_molecule import HFPretrainedMolecularGenerator + + major, minor, _ = map(int, transformers.__version__.split(".")[:3]) + model = HFPretrainedMolecularGenerator( + repo_id="ibm-research/GP-MoLFormer-Uniq", + ) + + if major >= 5 or (major == 4 and minor >= 57): + with pytest.raises(ImportError, match="transformers<=4.56.2"): + model.fit() + else: + pytest.skip("GP-MoLFormer version guard only applies to transformers>=4.57") + + +def _molexar_available() -> bool: + try: + import molexar # noqa: F401 + import fragment_selfies # noqa: F401 + return True + except ImportError: + return False + + +@pytest.mark.integration +@pytest.mark.parametrize("repo_id", ["fairydance/molexar-10m-base"]) +def test_hf_pretrained_generator_molexar_denovo(repo_id): + if not _molexar_available(): + pytest.skip("Molexar requires fragment-selfies and molexar") + + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator(repo_id=repo_id) + model.fit() + assert model.is_fitted_ is True + + smiles_list = model.generate(n_samples=2, temperature=0.8) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 2 + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) + + +@pytest.mark.integration +def test_hf_pretrained_generator_molexar_fragment_constraint(): + if not _molexar_available(): + pytest.skip("Molexar requires fragment-selfies and molexar") + + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + + start_smiles = "[*]C1(CC#N)CN(S(=O)(=O)CC)C1" + model = HFPretrainedMolecularGenerator(repo_id="fairydance/molexar-10m-base") + model.fit() + + smiles_list = model.generate( + n_samples=2, + start_smiles=start_smiles, + generation_task="motif_extension", + temperature=0.8, + ) + assert len(smiles_list) == 2 + assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) + + +def test_smiles_fragment_selfies_roundtrip(): + if not _molexar_available(): + pytest.skip("Molexar requires fragment-selfies and molexar") + + from torch_molecule.generator.pretrained.utils import ( + fragment_selfies_to_smiles, + smiles_to_fragment_selfies, + ) + + smiles = ["CCO", "c1ccccc1"] + recovered = fragment_selfies_to_smiles(smiles_to_fragment_selfies(smiles)) + assert len(recovered) == 2 + assert all(smiles_string for smiles_string in recovered) + + +@pytest.mark.parametrize( + "repo_id", + [ + "zjunlp/MolGen-large", + "zjunlp/MolGen-large-opt", + ], +) +@pytest.mark.integration +def test_hf_pretrained_generator_molgen_smoke(repo_id): + pytest.importorskip("transformers") + pytest.importorskip("selfies") + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id=repo_id, + generate_max_length=20, + ) + model.fit() + assert model.is_fitted_ is True + + smiles_list = model.generate(n_samples=2, num_beams=5) + assert isinstance(smiles_list, list) + assert len(smiles_list) == 2 + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) + + +@pytest.mark.integration +def test_hf_pretrained_generator_molgen_scaffold_prefix(): + pytest.importorskip("transformers") + pytest.importorskip("selfies") + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + from torch_molecule.generator.pretrained.utils import smiles_to_selfies + + scaffold = "c1ccccc1" + benzene = Chem.MolFromSmiles(scaffold) + prefix_selfies = smiles_to_selfies([scaffold])[0] + + model = HFPretrainedMolecularGenerator( + repo_id="zjunlp/MolGen-large", + generate_max_length=20, + ) + model.fit() + + smiles_from_scaffold = model.generate(n_samples=2, scaffold=scaffold, num_beams=5) + smiles_from_prefix = model.generate(n_samples=2, prefix_selfies=prefix_selfies, num_beams=5) + + assert len(smiles_from_scaffold) == 2 + assert len(smiles_from_prefix) == 2 + assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_from_scaffold) + assert all( + Chem.MolFromSmiles(smiles) is not None + and Chem.MolFromSmiles(smiles).HasSubstructMatch(benzene) + for smiles in smiles_from_scaffold + ) diff --git a/tests/generator/test_causal_lm.py b/tests/generator/test_causal_lm.py new file mode 100644 index 0000000..4dc5cbb --- /dev/null +++ b/tests/generator/test_causal_lm.py @@ -0,0 +1,122 @@ +from unittest.mock import MagicMock + +import pytest +import torch + +from torch_molecule.generator.pretrained.families.causal_lm import generate_causal_lm + + +class _FakeTokenizer: + bos_token_id = 1 + pad_token_id = 0 + eos_token_id = 2 + + def __call__(self, text, return_tensors="pt", add_special_tokens=True): + token_ids = [10 + len(text), 11 + len(text)] + if add_special_tokens: + token_ids = [self.bos_token_id] + token_ids + [self.eos_token_id] + return {"input_ids": torch.tensor([token_ids])} + + def batch_decode(self, outputs, skip_special_tokens=True): + return [f"SMILES_{idx}" for idx in range(outputs.shape[0])] + + +class _FakeModel: + def generate(self, **kwargs): + batch_size = kwargs["input_ids"].shape[0] if "input_ids" in kwargs else kwargs["num_return_sequences"] + seq_len = kwargs.get("max_length", 8) + return torch.zeros(batch_size, seq_len, dtype=torch.long) + + +def test_generate_causal_lm_bos_path(): + outputs = generate_causal_lm( + _FakeModel(), + _FakeTokenizer(), + torch.device("cpu"), + n_samples=3, + family="novomolgen", + max_length=12, + do_sample=False, + ) + assert outputs == ["SMILES_0", "SMILES_1", "SMILES_2"] + + +def test_generate_causal_lm_scaffold_path(): + model = _FakeModel() + model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) + + outputs = generate_causal_lm( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=2, + family="gp_molformer", + scaffold="c1ccccc1", + max_length=12, + do_sample=False, + ) + assert outputs == ["SMILES_0", "SMILES_1"] + kwargs = model.generate.call_args.kwargs + assert kwargs["use_cache"] is False + assert kwargs["top_k"] is None + # Default tokenize is [BOS, ..., EOS]; IBM drops the trailing special token. + assert kwargs["input_ids"].tolist() == [[1, 18, 19], [1, 18, 19]] + + +def test_generate_causal_lm_novomolgen_scaffold_keeps_all_tokens(): + model = _FakeModel() + model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) + + generate_causal_lm( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=2, + family="novomolgen", + scaffold="c1ccccc1", + max_length=12, + do_sample=False, + ) + + assert model.generate.call_args.kwargs["input_ids"].tolist() == [ + [18, 19], + [18, 19], + ] + assert "top_k" not in model.generate.call_args.kwargs + + +def test_generate_causal_lm_gp_molformer_denovo_path(): + model = _FakeModel() + model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) + + outputs = generate_causal_lm( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=2, + family="gp_molformer", + max_length=12, + do_sample=True, + ) + + assert outputs == ["SMILES_0", "SMILES_1"] + assert model.generate.call_args.kwargs["num_return_sequences"] == 2 + assert model.generate.call_args.kwargs["use_cache"] is False + assert "input_ids" not in model.generate.call_args.kwargs + + +def test_generate_causal_lm_novomolgen_does_not_force_use_cache_false(): + model = _FakeModel() + model.generate = MagicMock(return_value=torch.zeros(3, 8, dtype=torch.long)) + + generate_causal_lm( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=3, + family="novomolgen", + max_length=12, + do_sample=False, + ) + + assert "use_cache" not in model.generate.call_args.kwargs diff --git a/tests/generator/test_compat.py b/tests/generator/test_compat.py new file mode 100644 index 0000000..98f2e60 --- /dev/null +++ b/tests/generator/test_compat.py @@ -0,0 +1,64 @@ +from types import SimpleNamespace + +import torch + +from torch_molecule.generator.pretrained.compat import ( + patch_gp_molformer_generation_cache, + to_legacy_past_key_values, +) + + +class _EmptyCache: + def to_legacy_cache(self): + return ((None, None), (None, None)) + + +class _PopulatedCache: + def to_legacy_cache(self): + key = torch.zeros(1, 2, 3, 4) + value = torch.zeros(1, 2, 3, 4) + return ((key, value),) + + +def test_to_legacy_past_key_values_empty_cache_becomes_none(): + assert to_legacy_past_key_values(None) is None + assert to_legacy_past_key_values(_EmptyCache()) is None + assert to_legacy_past_key_values(((None, None),)) is None + assert to_legacy_past_key_values(SimpleNamespace()) is None + + +def test_to_legacy_past_key_values_keeps_legacy_tensors(): + key = torch.zeros(1, 2, 3, 4) + value = torch.zeros(1, 2, 3, 4) + legacy = ((key, value),) + assert to_legacy_past_key_values(legacy) == legacy + converted = to_legacy_past_key_values(_PopulatedCache()) + assert converted[0][0].shape == (1, 2, 3, 4) + + +def test_patch_gp_molformer_generation_cache_converts_empty_cache(): + captured = {} + + class _Model: + def __init__(self): + self.config = SimpleNamespace(use_cache=True) + self.generation_config = SimpleNamespace(use_cache=True) + + def prepare_inputs_for_generation( + self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs + ): + captured["past_key_values"] = past_key_values + return {"input_ids": input_ids, "past_key_values": past_key_values} + + model = _Model() + patch_gp_molformer_generation_cache(model) + patch_gp_molformer_generation_cache(model) + + result = model.prepare_inputs_for_generation( + torch.tensor([[1]]), + past_key_values=_EmptyCache(), + ) + assert captured["past_key_values"] is None + assert result["past_key_values"] is None + assert model.config.use_cache is False + assert model.generation_config.use_cache is False diff --git a/tests/generator/test_finetune.py b/tests/generator/test_finetune.py new file mode 100644 index 0000000..ad13570 --- /dev/null +++ b/tests/generator/test_finetune.py @@ -0,0 +1,197 @@ +import json +import os + +import pytest +import torch + +from torch_molecule.generator.pretrained.checkpoint import ( + METADATA_FILENAME, + build_metadata, + load_metadata, + save_metadata, +) +from torch_molecule.generator.pretrained.finetune import ( + corrupt_token_ids, + finetune_causal_lm, + finetune_generator, + finetune_seq2seq, +) + + +class _FakeOutput: + def __init__(self, loss: torch.Tensor): + self.loss = loss + + +class _FakeLM(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(1.0)) + + def forward(self, **kwargs): + return _FakeOutput(self.weight * 0.0 + 1.0) + + +def test_finetune_causal_lm_runs_one_epoch(): + pytest.importorskip("transformers") + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained("gpt2") + if tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.eos_token + + losses, last_epoch = finetune_causal_lm( + _FakeLM(), + tokenizer, + ["CCO", "CC(=O)O"], + torch.device("cpu"), + max_length=16, + batch_size=2, + epochs=1, + learning_rate=1e-3, + weight_decay=0.0, + grad_norm_clip=1.0, + verbose="none", + ) + assert last_epoch == 0 + assert len(losses) == 1 + assert losses[0] == pytest.approx(1.0) + + +class _FakeTokenizer: + mask_token_id = 99 + pad_token_id = 0 + all_special_ids = [0, 1, 2] + + +def _seq2seq_tokenizer(): + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained("gpt2") + if tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.eos_token + if tokenizer.mask_token is None: + tokenizer.add_special_tokens({"mask_token": ""}) + return tokenizer + + +def test_corrupt_token_ids_masks_non_special_tokens(): + input_ids = torch.tensor([1, 10, 11, 12, 2]) + generator = torch.Generator().manual_seed(0) + corrupted = corrupt_token_ids( + input_ids, + _FakeTokenizer(), + mask_prob=1.0, + generator=generator, + ) + assert corrupted[0].item() == 1 + assert corrupted[-1].item() == 2 + assert torch.equal(corrupted[1:-1], torch.tensor([99, 99, 99])) + + +def test_corrupt_token_ids_requires_mask_token(): + class _NoMask: + mask_token_id = None + + with pytest.raises(ValueError, match="mask_token_id"): + corrupt_token_ids(torch.tensor([1, 2, 3]), _NoMask()) + + +def test_finetune_seq2seq_runs_one_epoch(): + pytest.importorskip("transformers") + + tokenizer = _seq2seq_tokenizer() + losses, last_epoch = finetune_seq2seq( + _FakeLM(), + tokenizer, + ["[C][C][O]", "[C][C][Branch1][C][O]"], + torch.device("cpu"), + max_length=16, + batch_size=1, + epochs=1, + learning_rate=1e-3, + weight_decay=0.0, + grad_norm_clip=None, + verbose="none", + mask_prob=1.0, + ) + assert last_epoch == 0 + assert len(losses) == 1 + + +def test_finetune_seq2seq_labels_stay_clean(): + pytest.importorskip("transformers") + + tokenizer = _seq2seq_tokenizer() + captured = {} + + class _CaptureLM(_FakeLM): + def forward(self, **kwargs): + captured["input_ids"] = kwargs["input_ids"].detach().clone() + captured["labels"] = kwargs["labels"].detach().clone() + return super().forward(**kwargs) + + finetune_seq2seq( + _CaptureLM(), + tokenizer, + ["[C][C][O][C][C][O]"], + torch.device("cpu"), + max_length=16, + batch_size=1, + epochs=1, + learning_rate=1e-3, + weight_decay=0.0, + grad_norm_clip=None, + verbose="none", + mask_prob=1.0, + ) + + labels = captured["labels"] + input_ids = captured["input_ids"] + ignore_index = -100 + content = labels[0] != ignore_index + assert not torch.equal(input_ids[0][content], labels[0][content]) + assert tokenizer.mask_token_id in input_ids[0].tolist() + + +def test_finetune_generator_dispatches_seq2seq(): + pytest.importorskip("transformers") + + tokenizer = _seq2seq_tokenizer() + + losses, last_epoch = finetune_generator( + "molgen", + _FakeLM(), + tokenizer, + ["[C][C][O]"], + torch.device("cpu"), + max_length=16, + batch_size=1, + epochs=1, + learning_rate=1e-3, + weight_decay=0.0, + grad_norm_clip=1.0, + verbose="none", + ) + assert last_epoch == 0 + assert len(losses) == 1 + + +def test_checkpoint_metadata_roundtrip(tmp_path): + metadata = build_metadata( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", + family="novomolgen", + max_length=128, + revision="hf-checkpoint", + trust_remote_code=False, + tokenizer_repo_id=None, + generate_max_length=64, + model_name="HFPretrainedMolecularGenerator", + ) + save_metadata(str(tmp_path), metadata) + loaded = load_metadata(str(tmp_path)) + assert loaded == metadata + assert os.path.exists(os.path.join(tmp_path, METADATA_FILENAME)) + with open(os.path.join(tmp_path, METADATA_FILENAME), encoding="utf-8") as handle: + on_disk = json.load(handle) + assert on_disk["family"] == "novomolgen" diff --git a/tests/generator/test_molexar.py b/tests/generator/test_molexar.py new file mode 100644 index 0000000..1d3da9e --- /dev/null +++ b/tests/generator/test_molexar.py @@ -0,0 +1,38 @@ +import importlib.util + +import pytest + +from torch_molecule.generator.pretrained.families.molexar import ( + extract_conditions, + resolve_start_string, +) + + +def test_resolve_start_string_none_for_de_novo(): + assert resolve_start_string() is None + + +def test_resolve_start_string_accepts_literal_prefix(): + prefix = "[Frag][C][C][Attach:0]" + assert resolve_start_string(start_string=prefix) == prefix + + +@pytest.mark.skipif( + importlib.util.find_spec("molexar") is None, + reason="molexar not installed", +) +def test_resolve_start_string_from_smiles_fragment(): + start_smiles = "[*]C1(CC#N)CN(S(=O)(=O)CC)C1" + resolved = resolve_start_string( + start_smiles=start_smiles, + generation_task="motif_extension", + ) + assert resolved is not None + assert "[Attach:0]" in resolved + + +def test_extract_conditions_from_kwargs(): + kwargs = {"temperature": 0.8, "mol_qed": 0.9, "conditions": {"mol_logp": 2.5}} + conditions = extract_conditions(kwargs) + assert conditions == {"mol_logp": 2.5, "mol_qed": 0.9} + assert kwargs == {"temperature": 0.8} diff --git a/tests/generator/test_seq2seq.py b/tests/generator/test_seq2seq.py new file mode 100644 index 0000000..a653e6d --- /dev/null +++ b/tests/generator/test_seq2seq.py @@ -0,0 +1,67 @@ +from unittest.mock import MagicMock + +import torch + +from torch_molecule.generator.pretrained.families.seq2seq import ( + DEFAULT_MOLGEN_PREFIX_SELFIES, + generate_seq2seq, +) + + +class _FakeTokenizer: + def __call__(self, text, return_tensors="pt"): + length = len(text) + return { + "input_ids": torch.tensor([[1, length, 3]]), + "attention_mask": torch.tensor([[1, 1, 1]]), + } + + def decode(self, sequence, skip_special_tokens=True, clean_up_tokenization_spaces=True): + return f"[SELFIES_{int(sequence[0].item())}]" + + +class _FakeModel: + def generate(self, **kwargs): + num_sequences = kwargs["num_return_sequences"] + seq_len = kwargs.get("max_length", 8) + return torch.arange(num_sequences * seq_len, dtype=torch.long).reshape(num_sequences, seq_len) + + +def test_generate_seq2seq_default_prefix(): + model = _FakeModel() + model.generate = MagicMock(side_effect=_FakeModel().generate) + + outputs = generate_seq2seq( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=3, + max_length=12, + min_length=4, + num_beams=5, + ) + + assert len(outputs) == 3 + call_kwargs = model.generate.call_args.kwargs + assert call_kwargs["num_return_sequences"] == 3 + assert call_kwargs["num_beams"] == 5 + assert call_kwargs["max_length"] == 12 + + +def test_generate_seq2seq_expands_beams_for_sample_count(): + model = MagicMock() + model.generate.return_value = torch.zeros(7, 5, dtype=torch.long) + + generate_seq2seq( + model, + _FakeTokenizer(), + torch.device("cpu"), + n_samples=7, + num_beams=3, + ) + + assert model.generate.call_args.kwargs["num_beams"] == 7 + + +def test_default_molgen_prefix_is_benzene_selfies(): + assert "Ring1" in DEFAULT_MOLGEN_PREFIX_SELFIES diff --git a/torch_molecule/__init__.py b/torch_molecule/__init__.py index 5b6ebff..11a1e4a 100644 --- a/torch_molecule/__init__.py +++ b/torch_molecule/__init__.py @@ -36,6 +36,7 @@ from .generator.lstm import LSTMMolecularGenerator from .generator.molgpt import MolGPTMolecularGenerator from .generator.defog import DeFoGMolecularGenerator +from .generator.pretrained import HFPretrainedMolecularGenerator __all__ = [ # 'BaseMolecularPredictor', @@ -69,4 +70,5 @@ 'MolGPTMolecularGenerator', 'LSTMMolecularGenerator', 'DeFoGMolecularGenerator', + 'HFPretrainedMolecularGenerator', ] \ No newline at end of file diff --git a/torch_molecule/generator/pretrained/__init__.py b/torch_molecule/generator/pretrained/__init__.py new file mode 100644 index 0000000..212000f --- /dev/null +++ b/torch_molecule/generator/pretrained/__init__.py @@ -0,0 +1,3 @@ +from .modeling_pretrained import HFPretrainedMolecularGenerator + +__all__ = ["HFPretrainedMolecularGenerator"] diff --git a/torch_molecule/generator/pretrained/checkpoint.py b/torch_molecule/generator/pretrained/checkpoint.py new file mode 100644 index 0000000..1d46604 --- /dev/null +++ b/torch_molecule/generator/pretrained/checkpoint.py @@ -0,0 +1,46 @@ +"""Local save/load helpers for HF pretrained generators.""" + +import json +import os +from typing import Any, Dict, Optional + +METADATA_FILENAME = "hf_generator_metadata.json" + + +def build_metadata( + repo_id: str, + family: str, + max_length: int, + revision: Optional[str], + trust_remote_code: bool, + tokenizer_repo_id: Optional[str], + generate_max_length: int, + model_name: str, +) -> Dict[str, Any]: + return { + "repo_id": repo_id, + "family": family, + "max_length": max_length, + "revision": revision, + "trust_remote_code": trust_remote_code, + "tokenizer_repo_id": tokenizer_repo_id, + "generate_max_length": generate_max_length, + "model_name": model_name, + } + + +def save_metadata(path: str, metadata: Dict[str, Any]) -> None: + os.makedirs(path, exist_ok=True) + with open(os.path.join(path, METADATA_FILENAME), "w", encoding="utf-8") as handle: + json.dump(metadata, handle, indent=2) + + +def load_metadata(path: str) -> Dict[str, Any]: + metadata_path = os.path.join(path, METADATA_FILENAME) + if not os.path.exists(metadata_path): + raise FileNotFoundError( + f"Missing {METADATA_FILENAME} in '{path}'. Expected a directory saved by " + "HFPretrainedMolecularGenerator.save_to_local()." + ) + with open(metadata_path, "r", encoding="utf-8") as handle: + return json.load(handle) diff --git a/torch_molecule/generator/pretrained/compat.py b/torch_molecule/generator/pretrained/compat.py new file mode 100644 index 0000000..5614a99 --- /dev/null +++ b/torch_molecule/generator/pretrained/compat.py @@ -0,0 +1,115 @@ +"""Compatibility helpers for Hugging Face pretrained generators.""" + +import sys +import types +from typing import Any, Optional, Tuple + + +def _parse_transformers_version(version: str) -> Tuple[int, int, int]: + parts = version.split(".") + return tuple(int(part) for part in parts[:3]) + + +def ensure_gp_molformer_transformers_compat() -> None: + """Validate that the installed ``transformers`` version can load GP-MoLFormer.""" + import transformers + + major, minor, _ = _parse_transformers_version(transformers.__version__) + if major >= 5 or (major == 4 and minor >= 57): + raise ImportError( + "GP-MoLFormer remote code is not compatible with transformers " + f"{transformers.__version__}. Install transformers<=4.56.2, for example:\n" + " pip install 'transformers>=4.40,<=4.56.2'" + ) + + +def ensure_transformers_onnx_compat() -> None: + """Provide a stub ``transformers.onnx`` module for legacy remote code. + + IBM MoLFormer remote configs import ``OnnxConfig`` from ``transformers.onnx``, + which was removed in recent ``transformers`` releases. The ONNX export class + is not needed for generation, so a lightweight stub is sufficient. + """ + if "transformers.onnx" in sys.modules: + return + + try: + from transformers.onnx import OnnxConfig # noqa: F401 + return + except ModuleNotFoundError: + pass + + onnx_module = types.ModuleType("transformers.onnx") + + class OnnxConfig: + """Minimal stub for legacy remote configuration modules.""" + + onnx_module.OnnxConfig = OnnxConfig + sys.modules["transformers.onnx"] = onnx_module + + +def to_legacy_past_key_values(past_key_values: Any) -> Optional[tuple]: + """Convert HF Cache objects to the tuple layout IBM MoLFormer expects. + + ``transformers`` 4.47+ injects an empty ``DynamicCache`` into + ``prepare_inputs_for_generation``. IBM remote code then does + ``past_key_values[0][0].shape``, which raises because the empty cache + stores ``None`` instead of tensors. + """ + if past_key_values is None: + return None + + if not isinstance(past_key_values, (tuple, list)): + converter = getattr(past_key_values, "to_legacy_cache", None) + if converter is None: + return None + past_key_values = converter() + + if not past_key_values: + return None + first_layer = past_key_values[0] + if not first_layer: + return None + if first_layer[0] is None: + return None + return tuple(past_key_values) + + +def patch_gp_molformer_generation_cache(model: Any) -> None: + """Make GP-MoLFormer generation tolerate DynamicCache from transformers 4.56. + + IBM ``MolformerForCausalLM.prepare_inputs_for_generation`` only understands + the legacy tuple KV cache. Convert Cache objects (or empty caches) before + that method runs, and default ``use_cache`` off so later steps do not + reintroduce an incompatible cache layout. + """ + if getattr(model, "_gp_molformer_cache_patched", False): + return + + original = model.prepare_inputs_for_generation + + def prepare_inputs_for_generation( + input_ids, + past_key_values=None, + attention_mask=None, + inputs_embeds=None, + **kwargs, + ): + past_key_values = to_legacy_past_key_values(past_key_values) + return original( + input_ids, + past_key_values=past_key_values, + attention_mask=attention_mask, + inputs_embeds=inputs_embeds, + **kwargs, + ) + + model.prepare_inputs_for_generation = prepare_inputs_for_generation + model._gp_molformer_cache_patched = True + + generation_config = getattr(model, "generation_config", None) + if generation_config is not None: + generation_config.use_cache = False + config = getattr(model, "config", None) + if config is not None: + config.use_cache = False diff --git a/torch_molecule/generator/pretrained/families/__init__.py b/torch_molecule/generator/pretrained/families/__init__.py new file mode 100644 index 0000000..476efaa --- /dev/null +++ b/torch_molecule/generator/pretrained/families/__init__.py @@ -0,0 +1,10 @@ +from .causal_lm import generate_causal_lm +from .molexar import generate_molexar +from .seq2seq import DEFAULT_MOLGEN_PREFIX_SELFIES, generate_seq2seq + +__all__ = [ + "generate_causal_lm", + "generate_molexar", + "generate_seq2seq", + "DEFAULT_MOLGEN_PREFIX_SELFIES", +] diff --git a/torch_molecule/generator/pretrained/families/causal_lm.py b/torch_molecule/generator/pretrained/families/causal_lm.py new file mode 100644 index 0000000..1bb798d --- /dev/null +++ b/torch_molecule/generator/pretrained/families/causal_lm.py @@ -0,0 +1,97 @@ +"""Causal language model generation for SMILES-based HF generators.""" + +from typing import Any, List, Optional + +import torch + + +def generate_causal_lm( + model: torch.nn.Module, + tokenizer: Any, + device: torch.device, + n_samples: int, + *, + family: Optional[str] = None, + max_length: int = 64, + temperature: float = 1.0, + do_sample: bool = True, + scaffold: Optional[str] = None, + **kwargs: Any, +) -> List[str]: + """Generate SMILES strings with a causal language model. + + Parameters + ---------- + model : torch.nn.Module + A Hugging Face causal LM. + tokenizer : transformers.PreTrainedTokenizer + Tokenizer paired with the model. + device : torch.device + Device used for generation. + n_samples : int + Number of molecules to generate. + family : Optional[str], default=None + Generator family name. GP-MoLFormer uses a model-specific de novo path. + max_length : int, default=64 + Maximum generated sequence length passed to ``model.generate``. + temperature : float, default=1.0 + Sampling temperature. + do_sample : bool, default=True + Whether to use sampling during generation. + scaffold : Optional[str], default=None + Optional SMILES prefix for scaffold completion. For GP-MoLFormer this + should be a *partial* SMILES string (IBM's example is ``c1cccc``); + the official tokenizer appends a trailing special token that is then + dropped so generation continues the prefix. + + Returns + ------- + List[str] + Raw decoded strings from the tokenizer (may contain spaces). + """ + pad_token_id = tokenizer.pad_token_id + if pad_token_id is None: + pad_token_id = tokenizer.eos_token_id + + generate_kwargs = { + "max_length": max_length, + "do_sample": do_sample, + "pad_token_id": pad_token_id, + } + if do_sample: + generate_kwargs["temperature"] = temperature + if family == "gp_molformer": + # IBM remote code indexes tuple KV caches; transformers 4.56 injects + # an empty DynamicCache that crashes prepare_inputs_for_generation. + generate_kwargs["use_cache"] = False + generate_kwargs["top_k"] = None + + generate_kwargs.update(kwargs) + + if scaffold: + if family == "gp_molformer": + # Match IBM/gp-molformer scripts/conditional_generation.py: + # tokenize with special tokens, then drop the trailing SEP/EOS. + input_ids = tokenizer(scaffold, return_tensors="pt")["input_ids"] + if input_ids.shape[1] > 1: + input_ids = input_ids[:, :-1] + else: + encoded = tokenizer(scaffold, return_tensors="pt", add_special_tokens=False) + input_ids = encoded["input_ids"] + input_ids = input_ids.to(device).expand(n_samples, -1).contiguous() + generate_kwargs["input_ids"] = input_ids + elif family == "gp_molformer": + generate_kwargs["num_return_sequences"] = n_samples + else: + if tokenizer.bos_token_id is None: + raise ValueError( + "Tokenizer has no BOS token. Provide `scaffold=` or use a model " + "with a defined bos_token_id." + ) + input_ids = torch.tensor([[tokenizer.bos_token_id]], device=device) + generate_kwargs["input_ids"] = input_ids.expand(n_samples, -1).contiguous() + + with torch.no_grad(): + outputs = model.generate(**generate_kwargs) + + return tokenizer.batch_decode(outputs, skip_special_tokens=True) diff --git a/torch_molecule/generator/pretrained/families/molexar.py b/torch_molecule/generator/pretrained/families/molexar.py new file mode 100644 index 0000000..7d3fdc8 --- /dev/null +++ b/torch_molecule/generator/pretrained/families/molexar.py @@ -0,0 +1,129 @@ +"""Molexar Fragment-SELFIES generation.""" + +from typing import Any, Dict, List, Optional + +SINGLE_FRAGMENT_TASKS = frozenset({"motif_extension", "scaffold_decoration"}) +TWO_FRAGMENT_TASKS = frozenset({"linker_design", "scaffold_morphing"}) +FRAGMENT_CONSTRAINED_TASKS = SINGLE_FRAGMENT_TASKS | TWO_FRAGMENT_TASKS | frozenset({"superstructure"}) + +PROPERTY_KEYS = ( + "mol_hac", + "mol_hbdc", + "mol_hbac", + "mol_rotbc", + "mol_wt", + "mol_logp", + "mol_tpsa", + "mol_qed", + "mol_sas", +) + + +def _require_molexar(): + try: + from molexar.inference import MolexarInference # noqa: F401 + except ImportError as exc: + raise ImportError( + "The 'molexar' package is required for Molexar generation. " + "Install it with `pip install git+https://github.com/fairydance/Molexar.git`." + ) from exc + + +def resolve_start_string( + *, + start_string: Optional[str] = None, + start_smiles: Optional[str] = None, + start_fragment_selfies: Optional[str] = None, + generation_task: Optional[str] = None, +) -> Optional[str]: + """Resolve the Fragment-SELFIES prefix placed after ````.""" + if start_string is not None: + return start_string + if start_fragment_selfies is not None: + return start_fragment_selfies + if start_smiles is None: + return None + + task = generation_task or "motif_extension" + if task == "de_novo": + raise ValueError("de_novo generation does not accept start_smiles") + + from molexar.data.converter import smiles_fragment_to_fragment_selfies + + if task in SINGLE_FRAGMENT_TASKS | {"superstructure"}: + encoded = smiles_fragment_to_fragment_selfies(start_smiles, randomized=True) + return f"{encoded}[Attach:0]" + + if task in TWO_FRAGMENT_TASKS: + fragments = [fragment.strip() for fragment in start_smiles.split(".") if fragment.strip()] + if len(fragments) != 2: + raise ValueError( + "linker_design and scaffold_morphing require exactly two " + "dot-separated SMILES fragments" + ) + encoded_fragments = [ + smiles_fragment_to_fragment_selfies(fragment, randomized=True) for fragment in fragments + ] + return "".join(encoded_fragments) + "[Attach:0]" + + raise ValueError( + f"Unsupported generation_task '{task}'. Supported tasks: de_novo, " + "motif_extension, scaffold_decoration, linker_design, scaffold_morphing, superstructure" + ) + + +def extract_conditions(kwargs: Dict[str, Any]) -> Dict[str, Any]: + """Extract Molexar condition kwargs into a conditions dictionary.""" + conditions = dict(kwargs.pop("conditions", {}) or {}) + for key in PROPERTY_KEYS: + if key in kwargs: + conditions[key] = kwargs.pop(key) + for key in ("mol_pharma_fp", "prot_seq_esm_emb", "prot_poc_gvp_emb"): + if key in kwargs: + conditions[key] = kwargs.pop(key) + return conditions + + +def generate_molexar( + engine: Any, + n_samples: int, + *, + start_string: Optional[str] = None, + start_smiles: Optional[str] = None, + start_fragment_selfies: Optional[str] = None, + generation_task: Optional[str] = None, + conditions: Optional[Dict[str, Any]] = None, + max_new_tokens: Optional[int] = None, + temperature: float = 0.8, + top_p: float = 0.95, + top_k: int = 50, + do_sample: bool = True, + repetition_penalty: float = 1.0, + batch_size: int = 100, + **kwargs: Any, +) -> List[str]: + """Generate Fragment-SELFIES strings with a Molexar inference engine.""" + _require_molexar() + + merged_conditions = dict(conditions or {}) + + resolved_start = resolve_start_string( + start_string=start_string, + start_smiles=start_smiles, + start_fragment_selfies=start_fragment_selfies, + generation_task=generation_task, + ) + + return engine.generate( + conditions=merged_conditions, + start_string=resolved_start, + max_new_tokens=max_new_tokens, + num_samples=n_samples, + temperature=temperature, + top_p=top_p, + top_k=top_k, + do_sample=do_sample, + repetition_penalty=repetition_penalty, + batch_size=batch_size, + **kwargs, + ) diff --git a/torch_molecule/generator/pretrained/families/seq2seq.py b/torch_molecule/generator/pretrained/families/seq2seq.py new file mode 100644 index 0000000..99d11f1 --- /dev/null +++ b/torch_molecule/generator/pretrained/families/seq2seq.py @@ -0,0 +1,77 @@ +"""Seq2Seq generation for SELFIES-based HF generators such as MolGen.""" + +from typing import Any, List, Optional + +import torch + +DEFAULT_MOLGEN_PREFIX_SELFIES = "[C][=C][C][=C][C][=C][Ring1][=Branch1]" + + +def generate_seq2seq( + model: torch.nn.Module, + tokenizer: Any, + device: torch.device, + n_samples: int, + *, + prefix_selfies: Optional[str] = None, + max_length: int = 15, + min_length: int = 5, + num_beams: int = 5, + **kwargs: Any, +) -> List[str]: + """Generate SELFIES strings with a seq2seq language model. + + MolGen uses a corrupted SELFIES prefix as input and generates a completed + SELFIES sequence via beam search. + + Parameters + ---------- + model : torch.nn.Module + A Hugging Face seq2seq model. + tokenizer : transformers.PreTrainedTokenizer + Tokenizer paired with the model. + device : torch.device + Device used for generation. + n_samples : int + Number of molecules to generate. + prefix_selfies : Optional[str], default=None + SELFIES prefix used as model input. Defaults to a benzene ring fragment. + max_length : int, default=15 + Maximum generated sequence length. + min_length : int, default=5 + Minimum generated sequence length. + num_beams : int, default=5 + Beam width for beam search. + + Returns + ------- + List[str] + Raw decoded SELFIES strings from the tokenizer. + """ + prefix = prefix_selfies or DEFAULT_MOLGEN_PREFIX_SELFIES + encoded = tokenizer(prefix, return_tensors="pt") + input_ids = encoded["input_ids"].to(device) + attention_mask = encoded["attention_mask"].to(device) + + beam_width = max(num_beams, n_samples) + generate_kwargs = { + "input_ids": input_ids, + "attention_mask": attention_mask, + "max_length": max_length, + "min_length": min_length, + "num_return_sequences": n_samples, + "num_beams": beam_width, + } + generate_kwargs.update(kwargs) + + with torch.no_grad(): + outputs = model.generate(**generate_kwargs) + + return [ + tokenizer.decode( + sequence, + skip_special_tokens=True, + clean_up_tokenization_spaces=True, + ) + for sequence in outputs + ] diff --git a/torch_molecule/generator/pretrained/finetune.py b/torch_molecule/generator/pretrained/finetune.py new file mode 100644 index 0000000..e9199fd --- /dev/null +++ b/torch_molecule/generator/pretrained/finetune.py @@ -0,0 +1,288 @@ +"""Fine-tuning utilities for Hugging Face pretrained generators.""" + +from typing import Any, Dict, List, Optional, Tuple + +import numpy as np +import torch +from torch.utils.data import DataLoader, Dataset +from tqdm import tqdm + +from .registry import MOLEXAR_FAMILIES, SEQ2SEQ_FAMILIES + + +def corrupt_token_ids( + input_ids: torch.Tensor, + tokenizer: Any, + *, + mask_prob: float = 0.15, + generator: Optional[torch.Generator] = None, +) -> torch.Tensor: + """Replace a random subset of non-special tokens with the tokenizer mask id. + + MolGen is trained as a denoising seq2seq model: corrupted SELFIES in, + clean SELFIES as labels. Special tokens (BOS/EOS/PAD/mask) are left intact. + """ + if tokenizer.mask_token_id is None: + raise ValueError( + "MolGen denoising fine-tuning requires tokenizer.mask_token_id. " + "MolGen tokenizers provide a token." + ) + + corrupted = input_ids.clone() + special = torch.zeros_like(corrupted, dtype=torch.bool) + for special_id in tokenizer.all_special_ids: + special |= corrupted == special_id + if tokenizer.pad_token_id is not None: + special |= corrupted == tokenizer.pad_token_id + + probs = torch.rand(corrupted.shape, generator=generator, device=corrupted.device) + to_mask = (probs < mask_prob) & ~special + corrupted[to_mask] = tokenizer.mask_token_id + return corrupted + + +class _TokenizedDataset(Dataset): + def __init__(self, input_ids: List[List[int]], attention_mask: List[List[int]]): + self.input_ids = input_ids + self.attention_mask = attention_mask + + def __len__(self) -> int: + return len(self.input_ids) + + def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: + return { + "input_ids": torch.tensor(self.input_ids[idx]), + "attention_mask": torch.tensor(self.attention_mask[idx]), + } + + +def _batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -> Dict[str, torch.Tensor]: + return {key: value.to(device) for key, value in batch.items()} + + +def _run_training_loop( + model: torch.nn.Module, + train_loader: DataLoader, + optimizer: torch.optim.Optimizer, + device: torch.device, + epochs: int, + grad_norm_clip: Optional[float], + verbose: str, +) -> Tuple[List[float], int]: + model.train() + epoch_losses: List[float] = [] + last_epoch = 0 + + for epoch in range(epochs): + last_epoch = epoch + batch_losses: List[float] = [] + iterator = train_loader + if verbose in {"progress_bar", "print_statement"}: + iterator = tqdm(train_loader, desc=f"Fine-tuning epoch {epoch + 1}/{epochs}") + + for batch in iterator: + batch = _batch_to_device(batch, device) + optimizer.zero_grad() + outputs = model(**batch) + loss = outputs.loss + loss.backward() + + if grad_norm_clip is not None: + torch.nn.utils.clip_grad_norm_(model.parameters(), grad_norm_clip) + + optimizer.step() + batch_losses.append(float(loss.detach().cpu())) + + epoch_losses.append(float(np.mean(batch_losses)) if batch_losses else 0.0) + if verbose == "print_statement": + print(f"Epoch {epoch + 1}/{epochs} loss: {epoch_losses[-1]:.4f}") + + model.eval() + return epoch_losses, last_epoch + + +def finetune_causal_lm( + model: torch.nn.Module, + tokenizer: Any, + texts: List[str], + device: torch.device, + *, + max_length: int, + batch_size: int, + epochs: int, + learning_rate: float, + weight_decay: float, + grad_norm_clip: Optional[float], + verbose: str, +) -> Tuple[List[float], int]: + """Fine-tune a causal language model with next-token prediction.""" + import transformers + + tokenized = tokenizer( + texts, + truncation=True, + max_length=max_length, + padding=False, + ) + dataset = _TokenizedDataset(tokenized["input_ids"], tokenized["attention_mask"]) + collator = transformers.DataCollatorForLanguageModeling(tokenizer, mlm=False) + train_loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=collator) + + optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay) + return _run_training_loop(model, train_loader, optimizer, device, epochs, grad_norm_clip, verbose) + + +def finetune_seq2seq( + model: torch.nn.Module, + tokenizer: Any, + texts: List[str], + device: torch.device, + *, + max_length: int, + batch_size: int, + epochs: int, + learning_rate: float, + weight_decay: float, + grad_norm_clip: Optional[float], + verbose: str, + mask_prob: float = 0.15, +) -> Tuple[List[float], int]: + """Fine-tune a seq2seq model with MolGen-style denoising. + + Encoder inputs are token-masked SELFIES; labels remain the original + clean sequence. Causal LM and Molexar fine-tuning are unchanged. + """ + import transformers + + class _Seq2SeqDataset(Dataset): + def __init__(self, items: List[str]): + self.items = items + + def __len__(self): + return len(self.items) + + def __getitem__(self, idx): + encoded = tokenizer( + self.items[idx], + truncation=True, + max_length=max_length, + padding=False, + ) + item = {key: torch.tensor(value) for key, value in encoded.items()} + item["labels"] = item["input_ids"].clone() + item["input_ids"] = corrupt_token_ids( + item["input_ids"], + tokenizer, + mask_prob=mask_prob, + ) + return item + + collator = transformers.DataCollatorForSeq2Seq(tokenizer, model=model, padding=True) + train_loader = DataLoader( + _Seq2SeqDataset(texts), + batch_size=batch_size, + shuffle=True, + collate_fn=collator, + ) + optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay) + return _run_training_loop(model, train_loader, optimizer, device, epochs, grad_norm_clip, verbose) + + +def finetune_molexar( + model: torch.nn.Module, + tokenizer: Any, + config: Any, + texts: List[str], + device: torch.device, + *, + max_length: int, + batch_size: int, + epochs: int, + learning_rate: float, + weight_decay: float, + grad_norm_clip: Optional[float], + verbose: str, +) -> Tuple[List[float], int]: + """Fine-tune a Molexar model on Fragment-SELFIES training templates.""" + from molexar.templates import build_condition_template, build_training_text + + condition_block, _ = build_condition_template(config) + training_texts = [build_training_text(config, condition_block, text) for text in texts] + return finetune_causal_lm( + model, + tokenizer, + training_texts, + device, + max_length=max_length, + batch_size=batch_size, + epochs=epochs, + learning_rate=learning_rate, + weight_decay=weight_decay, + grad_norm_clip=grad_norm_clip, + verbose=verbose, + ) + + +def finetune_generator( + family: str, + model: torch.nn.Module, + tokenizer: Any, + texts: List[str], + device: torch.device, + *, + config: Optional[Any] = None, + max_length: int, + batch_size: int, + epochs: int, + learning_rate: float, + weight_decay: float, + grad_norm_clip: Optional[float], + verbose: str, +) -> Tuple[List[float], int]: + """Dispatch fine-tuning to the family-specific routine.""" + if family in SEQ2SEQ_FAMILIES: + return finetune_seq2seq( + model, + tokenizer, + texts, + device, + max_length=max_length, + batch_size=batch_size, + epochs=epochs, + learning_rate=learning_rate, + weight_decay=weight_decay, + grad_norm_clip=grad_norm_clip, + verbose=verbose, + ) + + if family in MOLEXAR_FAMILIES: + if config is None: + raise ValueError("Molexar fine-tuning requires a model config.") + return finetune_molexar( + model, + tokenizer, + config, + texts, + device, + max_length=max_length, + batch_size=batch_size, + epochs=epochs, + learning_rate=learning_rate, + weight_decay=weight_decay, + grad_norm_clip=grad_norm_clip, + verbose=verbose, + ) + + return finetune_causal_lm( + model, + tokenizer, + texts, + device, + max_length=max_length, + batch_size=batch_size, + epochs=epochs, + learning_rate=learning_rate, + weight_decay=weight_decay, + grad_norm_clip=grad_norm_clip, + verbose=verbose, + ) diff --git a/torch_molecule/generator/pretrained/modeling_pretrained.py b/torch_molecule/generator/pretrained/modeling_pretrained.py index f87f5c1..2311902 100644 --- a/torch_molecule/generator/pretrained/modeling_pretrained.py +++ b/torch_molecule/generator/pretrained/modeling_pretrained.py @@ -1 +1,559 @@ -# TODO \ No newline at end of file +import warnings +import os +from typing import Any, Dict, List, Optional, Tuple, Union + +import numpy as np +import torch + +from ...base import BaseMolecularGenerator +from .checkpoint import build_metadata, load_metadata, save_metadata +from .families.causal_lm import generate_causal_lm +from .families.molexar import extract_conditions, generate_molexar +from .families.seq2seq import generate_seq2seq +from .finetune import finetune_generator +from .registry import ( + CAUSAL_LM_FAMILIES, + MOLEXAR_FAMILIES, + SEQ2SEQ_FAMILIES, + resolve_family, +) + + +class HFPretrainedMolecularGenerator(BaseMolecularGenerator): + """Hugging Face pretrained models as molecular generators. + + This class loads pretrained generative models from Hugging Face and exposes + a sklearn-style ``fit`` / ``generate`` interface consistent with other + generators in torch-molecule. + + Supported generation modes depend on the model family: + + - NovoMolGen: de novo SMILES generation from BOS. + - GP-MoLFormer: de novo generation and scaffold completion via ``scaffold=``. + - MolGen-large / MolGen-large-opt: SELFIES seq2seq generation via ``prefix_selfies=`` + or ``scaffold=`` (SMILES converted internally). + - Molexar: Fragment-SELFIES de novo and fragment-constrained generation via + ``start_smiles`` / ``start_string`` / ``conditions`` (omni). + + Other registered families can be loaded but may raise ``NotImplementedError`` + until later phases are implemented. + + Tested models include: + + - NovoMolGen: Causal LM pretrained on ZINC-22 for de novo SMILES generation. + + repo_id: ``"chandar-lab/NovoMolGen_32M_SMILES_BPE"`` + (https://huggingface.co/chandar-lab/NovoMolGen_32M_SMILES_BPE) + + - GP-MoLFormer: Causal LM for de novo generation and scaffold decoration. + + repo_id: ``"ibm-research/GP-MoLFormer-Uniq"`` + (https://huggingface.co/ibm-research/GP-MoLFormer-Uniq) + + - MolGen-large: Seq2Seq SELFIES generator with high chemical validity. + + repo_id: ``"zjunlp/MolGen-large"`` + (https://huggingface.co/zjunlp/MolGen-large) + + - MolGen-large-opt: MolGen-large fine-tuned for QED / p-logP optimization. + + repo_id: ``"zjunlp/MolGen-large-opt"`` + (https://huggingface.co/zjunlp/MolGen-large-opt) + + - Molexar-10M-base: Fragment-SELFIES de novo and fragment-constrained generation. + + repo_id: ``"fairydance/molexar-10m-base"`` + (https://huggingface.co/fairydance/molexar-10m-base) + + - Molexar-10M-omni: Multi-condition Molexar model for property-guided generation. + + repo_id: ``"fairydance/molexar-10m-omni"`` + (https://huggingface.co/fairydance/molexar-10m-omni) + + Parameters + ---------- + repo_id : str + Hugging Face repository id of the pretrained generator. + max_length : int, default=128 + Maximum sequence length used when loading the tokenizer. + revision : Optional[str], default=None + Model revision on the Hugging Face Hub. NovoMolGen defaults to + ``"hf-checkpoint"`` so standard ``model.generate`` works out of the box. + trust_remote_code : bool, default=False + Whether to trust remote code when loading from Hugging Face. + Automatically enabled for GP-MoLFormer and Molexar. + tokenizer_repo_id : Optional[str], default=None + Optional Hugging Face repo for the tokenizer. GP-MoLFormer defaults to + ``"ibm-research/MoLFormer-XL-both-10pct"``. + generate_max_length : int, default=64 + Default ``max_length`` passed to ``generate()``. + batch_size : int, default=8 + Batch size used when fine-tuning on SMILES data. + epochs : int, default=1 + Number of fine-tuning epochs when ``fit(X)`` is called. + learning_rate : float, default=5e-5 + Learning rate for fine-tuning. + weight_decay : float, default=0.01 + Weight decay for fine-tuning. + grad_norm_clip : Optional[float], default=1.0 + Maximum gradient norm during fine-tuning. Set to ``None`` to disable clipping. + device : Optional[Union[torch.device, str]], default=None + Device to run the model on. + model_name : str, default="HFPretrainedMolecularGenerator" + Name identifier for the model instance. + verbose : str, default="none" + Progress display mode: ``"none"``, ``"progress_bar"``, or + ``"print_statement"``. + """ + + def __init__( + self, + repo_id: str, + max_length: int = 128, + revision: Optional[str] = None, + trust_remote_code: bool = False, + tokenizer_repo_id: Optional[str] = None, + generate_max_length: int = 64, + batch_size: int = 8, + epochs: int = 1, + learning_rate: float = 5e-5, + weight_decay: float = 0.01, + grad_norm_clip: Optional[float] = 1.0, + *, + device: Optional[Union[torch.device, str]] = None, + model_name: str = "HFPretrainedMolecularGenerator", + verbose: str = "none", + ): + super().__init__(device=device, model_name=model_name, verbose=verbose) + + self.repo_id = repo_id + self.max_length = max_length + self.revision = revision + self.trust_remote_code = trust_remote_code + self.tokenizer_repo_id = tokenizer_repo_id + self.generate_max_length = generate_max_length + self.batch_size = batch_size + self.epochs = epochs + self.learning_rate = learning_rate + self.weight_decay = weight_decay + self.grad_norm_clip = grad_norm_clip + self.fitting_loss: List[float] = [] + + self._family: Optional[str] = None + self.tokenizer = None + self._molexar_engine = None + self._model_local_path: Optional[str] = None + self.fitting_epoch = -1 + + self._require_transformers() + + if resolve_family(self.repo_id) == "causal_lm": + warnings.warn( + f"Unknown repo_id: {self.repo_id}. The class will try to load the " + "model from Hugging Face as a causal LM, but generation may fail " + "if the architecture is not supported.", + stacklevel=2, + ) + + @staticmethod + def _get_param_names() -> List[str]: + return [ + "repo_id", + "max_length", + "revision", + "trust_remote_code", + "tokenizer_repo_id", + "generate_max_length", + "batch_size", + "epochs", + "learning_rate", + "weight_decay", + "grad_norm_clip", + "model_name", + ] + + def _get_model_params(self) -> Dict[str, Any]: + return { + "repo_id": self.repo_id, + "max_length": self.max_length, + "generate_max_length": self.generate_max_length, + } + + def _setup_optimizers(self) -> Tuple[torch.optim.Optimizer, Optional[Any]]: + optimizer = torch.optim.AdamW( + self.model.parameters(), + lr=self.learning_rate, + weight_decay=self.weight_decay, + ) + return optimizer, None + + def _train_epoch(self, train_loader, optimizer) -> Dict[str, float]: + raise NotImplementedError( + "Use fit(X) for fine-tuning HFPretrainedMolecularGenerator." + ) + + def save_to_local(self, path: str) -> None: + """Save the model and tokenizer to a local directory.""" + self._check_is_fitted() + os.makedirs(path, exist_ok=True) + + if self._family in MOLEXAR_FAMILIES: + self.model.save_pretrained(path) + self.tokenizer.save_pretrained(path) + else: + self.model.save_pretrained(path) + self.tokenizer.save_pretrained(path) + + save_metadata( + path, + build_metadata( + repo_id=self.repo_id, + family=self._family, + max_length=self.max_length, + revision=self.revision, + trust_remote_code=self.trust_remote_code, + tokenizer_repo_id=self.tokenizer_repo_id, + generate_max_length=self.generate_max_length, + model_name=self.model_name, + ), + ) + self._model_local_path = path + + def load_from_local(self, path: str) -> None: + """Load a model and tokenizer saved by :meth:`save_to_local`.""" + metadata = load_metadata(path) + self.repo_id = metadata["repo_id"] + self._family = metadata["family"] + self.max_length = metadata.get("max_length", self.max_length) + self.revision = metadata.get("revision") + self.trust_remote_code = metadata.get("trust_remote_code", self.trust_remote_code) + self.tokenizer_repo_id = metadata.get("tokenizer_repo_id") + self.generate_max_length = metadata.get("generate_max_length", self.generate_max_length) + self.model_name = metadata.get("model_name", self.model_name) + self._model_local_path = path + + if self._family in MOLEXAR_FAMILIES: + self._load_molexar_pretrained(local_path=path) + else: + self._load_pretrained(local_path=path) + + self.is_fitted_ = True + + def save_to_hf(self, repo_id: str, **kwargs) -> None: + raise NotImplementedError( + "HFPretrainedMolecularGenerator does not support saving to Hugging Face." + ) + + def load_from_hf(self, repo_id: Optional[str] = None, **kwargs) -> None: + """Load the pretrained model from Hugging Face (same as ``fit()``).""" + if repo_id is not None: + self.repo_id = repo_id + self.fit() + + def load(self, path: Optional[str] = None, repo_id: Optional[str] = None, **kwargs) -> None: + """Load the model from a local directory or Hugging Face.""" + if path is not None: + self.load_from_local(path) + return + if repo_id is not None: + self.repo_id = repo_id + self.fit() + + def fit( + self, + X: Optional[List[str]] = None, + y: Optional[np.ndarray] = None, + ) -> "HFPretrainedMolecularGenerator": + """Load the pretrained model from Hugging Face. + + Parameters + ---------- + X : Optional[List[str]], default=None + Optional SMILES strings for fine-tuning. When provided, the pretrained + weights are adapted on the encoded family-specific representation. + y : Optional[np.ndarray], default=None + Reserved for future conditional fine-tuning. Currently ignored with a warning. + + Returns + ------- + HFPretrainedMolecularGenerator + The fitted generator instance. + """ + assert self.repo_id is not None, "repo_id is not set" + self._require_transformers() + + self._family = resolve_family(self.repo_id) + self._load_pretrained() + + if X is not None: + if y is not None: + warnings.warn( + "Conditional fine-tuning with y is not implemented yet; continuing with " + "unconditional language-model fine-tuning.", + stacklevel=2, + ) + y = None + X, y = self._validate_inputs(X, y, return_rdkit_mol=False) + X = self._encode_inputs(X) + self._finetune(X, y) + + self.is_fitted_ = True + return self + + def generate(self, n_samples: int = 10, **kwargs) -> List[str]: + """Generate molecules as SMILES strings. + + Parameters + ---------- + n_samples : int, default=10 + Number of molecules to generate. + **kwargs + Additional arguments forwarded to the family-specific generator. + For causal LMs, common options include ``max_length``, ``temperature``, + ``do_sample``, and ``scaffold``. For MolGen, use ``prefix_selfies`` or + ``scaffold`` plus optional ``num_beams``, ``min_length``, and + ``max_length``. For Molexar, use ``start_smiles``, ``start_string``, + ``generation_task``, or ``conditions`` for omni models. + + Returns + ------- + List[str] + Generated SMILES strings. For MolGen and Molexar, invalid decodes + are dropped, so the list may be shorter than ``n_samples``. + """ + self._check_is_fitted() + + if self._family in CAUSAL_LM_FAMILIES: + raw = generate_causal_lm( + self.model, + self.tokenizer, + self.device, + n_samples, + family=self._family, + max_length=kwargs.pop("max_length", self.generate_max_length), + temperature=kwargs.pop("temperature", 1.0), + do_sample=kwargs.pop("do_sample", True), + scaffold=kwargs.pop("scaffold", None), + **kwargs, + ) + return self._decode_outputs(raw) + + if self._family in SEQ2SEQ_FAMILIES: + raw = generate_seq2seq( + self.model, + self.tokenizer, + self.device, + n_samples, + prefix_selfies=self._resolve_prefix_selfies(kwargs), + max_length=kwargs.pop("max_length", self.generate_max_length), + min_length=kwargs.pop("min_length", 5), + num_beams=kwargs.pop("num_beams", 5), + **kwargs, + ) + return self._decode_outputs(raw) + + if self._family in MOLEXAR_FAMILIES: + conditions = extract_conditions(kwargs) + raw = generate_molexar( + self._molexar_engine, + n_samples, + conditions=conditions or None, + max_new_tokens=kwargs.pop("max_new_tokens", None), + temperature=kwargs.pop("temperature", 0.8), + top_p=kwargs.pop("top_p", 0.95), + top_k=kwargs.pop("top_k", 50), + do_sample=kwargs.pop("do_sample", True), + repetition_penalty=kwargs.pop("repetition_penalty", 1.0), + batch_size=kwargs.pop("batch_size", 100), + start_string=kwargs.pop("start_string", None), + start_smiles=kwargs.pop("start_smiles", None), + start_fragment_selfies=kwargs.pop("start_fragment_selfies", None), + generation_task=kwargs.pop("generation_task", None), + **kwargs, + ) + return self._decode_outputs(raw) + + raise NotImplementedError(f"Generation is not implemented for family '{self._family}'.") + + def _encode_inputs(self, smiles: List[str]) -> List[str]: + """Convert SMILES inputs to the representation expected by the model family.""" + if self._family in SEQ2SEQ_FAMILIES: + from .utils import smiles_to_selfies + + return smiles_to_selfies(smiles) + if self._family in MOLEXAR_FAMILIES: + from .utils import smiles_to_fragment_selfies + + return smiles_to_fragment_selfies(smiles) + return smiles + + def _resolve_prefix_selfies(self, kwargs: Dict[str, Any]) -> Optional[str]: + """Resolve a MolGen SELFIES prefix from kwargs.""" + prefix_selfies = kwargs.pop("prefix_selfies", None) + scaffold = kwargs.pop("scaffold", None) + + if prefix_selfies is not None: + return prefix_selfies + if scaffold is not None: + from .utils import smiles_to_selfies + + return smiles_to_selfies([scaffold])[0] + return None + + def _decode_outputs(self, outputs: List[str]) -> List[str]: + """Normalize raw model strings to SMILES. + + MolGen and Molexar drop strings that cannot be decoded. The returned + list length is the number of successful SMILES, which may be smaller + than ``n_samples``. + """ + n_attempted = len(outputs) + + if self._family in SEQ2SEQ_FAMILIES: + from .utils import selfies_to_smiles + + cleaned = [output.replace(" ", "") for output in outputs] + smiles = selfies_to_smiles(cleaned) + elif self._family in MOLEXAR_FAMILIES: + from .utils import fragment_selfies_to_smiles + + smiles = fragment_selfies_to_smiles(outputs) + else: + return [output.replace(" ", "") for output in outputs] + + if len(smiles) < n_attempted: + warnings.warn( + f"got {len(smiles)}/{n_attempted} valid SMILES", + stacklevel=2, + ) + return smiles + + def _load_pretrained(self, local_path: Optional[str] = None) -> None: + import transformers + + from .compat import ( + ensure_gp_molformer_transformers_compat, + ensure_transformers_onnx_compat, + patch_gp_molformer_generation_cache, + ) + from .registry import DEFAULT_GP_MOLFORMER_TOKENIZER + + if self._family in MOLEXAR_FAMILIES: + self._load_molexar_pretrained(local_path=local_path) + return + + if self._family == "gp_molformer": + ensure_gp_molformer_transformers_compat() + ensure_transformers_onnx_compat() + + load_kwargs = self._get_load_kwargs() + model_cls = self._get_model_class() + model_source = local_path or self.repo_id + tokenizer_repo = local_path or self.repo_id + if self._family == "gp_molformer" and local_path is None: + tokenizer_repo = self.tokenizer_repo_id or DEFAULT_GP_MOLFORMER_TOKENIZER + + self.tokenizer = transformers.AutoTokenizer.from_pretrained( + tokenizer_repo, + model_max_length=self.max_length, + **load_kwargs, + ) + self.model = model_cls.from_pretrained(model_source, **load_kwargs) + if self._family == "gp_molformer": + patch_gp_molformer_generation_cache(self.model) + self._setup_tokenizer() + self.model.to(self.device) + self.model.eval() + + def _load_molexar_pretrained(self, local_path: Optional[str] = None) -> None: + from huggingface_hub import snapshot_download + + from .families.molexar import _require_molexar + + _require_molexar() + from molexar.inference import MolexarInference + + if local_path is None: + self._model_local_path = snapshot_download(self.repo_id) + else: + self._model_local_path = local_path + + device = str(self.device) + self._molexar_engine = MolexarInference( + self._model_local_path, + device=device, + tokenizer_path=self.tokenizer_repo_id, + ) + self.model = self._molexar_engine.model + self.tokenizer = self._molexar_engine.tokenizer + + def _get_molexar_config(self): + if self._molexar_engine is not None: + return self._molexar_engine.config + return getattr(self.model, "config", None) + + def _finetune(self, X: List[str], y: Optional[np.ndarray]) -> None: + if len(X) == 0: + raise ValueError("Fine-tuning requires at least one training example.") + + config = self._get_molexar_config() if self._family in MOLEXAR_FAMILIES else None + losses, last_epoch = finetune_generator( + self._family, + self.model, + self.tokenizer, + X, + self.device, + config=config, + max_length=self.max_length, + batch_size=self.batch_size, + epochs=self.epochs, + learning_rate=self.learning_rate, + weight_decay=self.weight_decay, + grad_norm_clip=self.grad_norm_clip, + verbose=self.verbose, + ) + self.fitting_loss = losses + self.fitting_epoch = last_epoch + self.model.eval() + + def _get_model_class(self): + import transformers + + if self._family in SEQ2SEQ_FAMILIES: + return transformers.AutoModelForSeq2SeqLM + return transformers.AutoModelForCausalLM + + def _get_load_kwargs(self) -> Dict[str, Any]: + load_kwargs: Dict[str, Any] = {} + + if self._family == "novomolgen" and self.revision is None: + load_kwargs["revision"] = "hf-checkpoint" + elif self.revision is not None: + load_kwargs["revision"] = self.revision + + if ( + self._family in MOLEXAR_FAMILIES + or self._family == "gp_molformer" + or self.trust_remote_code + ): + load_kwargs["trust_remote_code"] = True + + return load_kwargs + + def _setup_tokenizer(self) -> None: + if self.tokenizer.pad_token is None: + if self.tokenizer.eos_token is not None: + self.tokenizer.pad_token = self.tokenizer.eos_token + else: + self.tokenizer.add_special_tokens({"pad_token": ""}) + self.model.resize_token_embeddings(len(self.tokenizer)) + + @staticmethod + def _require_transformers() -> None: + try: + import transformers # noqa: F401 + except ImportError as exc: + raise ImportError( + "The 'transformers' package is required for HFPretrainedMolecularGenerator. " + "Please install it using `pip install transformers`." + ) from exc diff --git a/torch_molecule/generator/pretrained/registry.py b/torch_molecule/generator/pretrained/registry.py new file mode 100644 index 0000000..9ed36fe --- /dev/null +++ b/torch_molecule/generator/pretrained/registry.py @@ -0,0 +1,43 @@ +"""Family registry for Hugging Face pretrained molecular generators.""" + +from typing import Dict + +KNOWN_REPOS: Dict[str, str] = { + "chandar-lab/NovoMolGen_32M_SMILES_BPE": "novomolgen", + "ibm-research/GP-MoLFormer-Uniq": "gp_molformer", + "zjunlp/MolGen-large": "molgen", + "zjunlp/MolGen-large-opt": "molgen", + "fairydance/molexar-10m-base": "molexar", + "fairydance/molexar-10m-omni": "molexar", +} + +FAMILY_PREFIXES: Dict[str, str] = { + "chandar-lab/NovoMolGen": "novomolgen", + "ibm-research/GP-MoLFormer": "gp_molformer", + "zjunlp/MolGen": "molgen", + "fairydance/molexar": "molexar", +} + +DEFAULT_GP_MOLFORMER_TOKENIZER = "ibm-research/MoLFormer-XL-both-10pct" + +MOLEXAR_OMNI_REPOS = frozenset( + { + "fairydance/molexar-10m-omni", + } +) + +CAUSAL_LM_FAMILIES = frozenset({"novomolgen", "gp_molformer", "causal_lm"}) +SEQ2SEQ_FAMILIES = frozenset({"molgen"}) +MOLEXAR_FAMILIES = frozenset({"molexar"}) + + +def resolve_family(repo_id: str) -> str: + """Map a Hugging Face repo id to a generator family.""" + if repo_id in KNOWN_REPOS: + return KNOWN_REPOS[repo_id] + + for prefix, family in FAMILY_PREFIXES.items(): + if repo_id.startswith(prefix): + return family + + return "causal_lm" diff --git a/torch_molecule/generator/pretrained/utils.py b/torch_molecule/generator/pretrained/utils.py new file mode 100644 index 0000000..8d3880b --- /dev/null +++ b/torch_molecule/generator/pretrained/utils.py @@ -0,0 +1,155 @@ +"""Representation conversion utilities for HF pretrained generators.""" + +import warnings +from typing import List + +from rdkit import Chem + + +def _require_selfies(): + try: + import selfies # noqa: F401 + except ImportError as exc: + raise ImportError( + "The 'selfies' package is required for SELFIES conversion. " + "Install it with `pip install selfies`." + ) from exc + + +def smiles_to_selfies(smiles: List[str]) -> List[str]: + """Convert SMILES strings to SELFIES representations. + + Parameters + ---------- + smiles : List[str] + Input SMILES strings. + + Returns + ------- + List[str] + SELFIES strings in the same order as the input. + + Raises + ------ + ValueError + If a SMILES string is invalid or cannot be encoded as SELFIES. + """ + _require_selfies() + import selfies as sf + + selfies_list: List[str] = [] + for idx, smiles_string in enumerate(smiles): + mol = Chem.MolFromSmiles(smiles_string) + if mol is None: + raise ValueError(f"Invalid SMILES at index {idx}: {smiles_string}") + canonical = Chem.MolToSmiles(mol) + try: + selfies_list.append(sf.encoder(canonical)) + except sf.EncoderError as exc: + raise ValueError( + f"SMILES at index {idx} is RDKit-valid but not SELFIES-encodable: {smiles_string}" + ) from exc + return selfies_list + + +def selfies_to_smiles(selfies_list: List[str]) -> List[str]: + """Convert SELFIES strings to canonical SMILES. + + Entries that cannot be decoded are dropped and a warning is emitted. + Unexpected errors such as ``ImportError`` are not swallowed. + + Parameters + ---------- + selfies_list : List[str] + Input SELFIES strings. + + Returns + ------- + List[str] + Canonical SMILES strings for entries that decoded successfully. + """ + _require_selfies() + import selfies as sf + + smiles_list: List[str] = [] + n_dropped = 0 + for selfies_string in selfies_list: + if not selfies_string or not str(selfies_string).strip(): + n_dropped += 1 + continue + try: + decoded = sf.decoder(selfies_string) + except sf.DecoderError: + n_dropped += 1 + continue + mol = Chem.MolFromSmiles(decoded) if decoded else None + if mol is None: + n_dropped += 1 + continue + smiles_list.append(Chem.MolToSmiles(mol)) + + if n_dropped: + warnings.warn(f"dropped {n_dropped} invalid SELFIES", stacklevel=2) + return smiles_list + + +def _require_fragment_selfies(): + try: + from fragment_selfies import FragmentSelfiesCodec # noqa: F401 + except ImportError as exc: + raise ImportError( + "The 'fragment-selfies' package is required for Fragment-SELFIES conversion. " + "Install it with `pip install fragment-selfies`." + ) from exc + + +def smiles_to_fragment_selfies(smiles: List[str]) -> List[str]: + """Convert SMILES strings to Fragment-SELFIES representations.""" + _require_fragment_selfies() + + try: + from molexar.data.converter import smiles_to_fragment_selfies as encode_one + except ImportError as exc: + raise ImportError( + "Molexar conversion utilities require the 'molexar' package. " + "Install it with `pip install git+https://github.com/fairydance/Molexar.git`." + ) from exc + + return [encode_one(smiles_string, canonical=True) for smiles_string in smiles] + + +def fragment_selfies_to_smiles(fragment_selfies_list: List[str]) -> List[str]: + """Convert Fragment-SELFIES strings to canonical SMILES. + + Entries that cannot be decoded are dropped and a warning is emitted. + Unexpected errors such as ``ImportError`` are not swallowed. + """ + _require_fragment_selfies() + + try: + from molexar.data.converter import fragment_selfies_to_smiles as decode_one + except ImportError as exc: + raise ImportError( + "Molexar conversion utilities require the 'molexar' package. " + "Install it with `pip install git+https://github.com/fairydance/Molexar.git`." + ) from exc + + smiles_list: List[str] = [] + n_dropped = 0 + for fragment_selfies in fragment_selfies_list: + if not fragment_selfies or not str(fragment_selfies).strip(): + n_dropped += 1 + continue + try: + decoded = decode_one(fragment_selfies, canonical=True, ignore_errors=False) + except (ValueError, TypeError): + n_dropped += 1 + continue + if not decoded: + n_dropped += 1 + continue + smiles_list.append(decoded) + + if n_dropped: + warnings.warn(f"dropped {n_dropped} invalid Fragment-SELFIES", stacklevel=2) + return smiles_list From b2d264405d2d2bf7a57a7841eab100f61c78dd88 Mon Sep 17 00:00:00 2001 From: m21hm9 Date: Fri, 11 Sep 2026 13:35:04 +0800 Subject: [PATCH 2/5] Add a document hub models in the README and replace GP-MoLFormer with SAFE-GPT due to the transformers incompatibility. --- README.md | 25 ++- docs/source/install.rst | 12 +- tests/generator/hfpretrained.py | 204 ++++++++++++------ tests/generator/test_causal_lm.py | 47 +--- tests/generator/test_compat.py | 67 +----- torch_molecule/generator/pretrained/compat.py | 126 ++--------- .../pretrained/families/causal_lm.py | 30 +-- .../pretrained/modeling_pretrained.py | 116 ++++++---- .../generator/pretrained/registry.py | 9 +- torch_molecule/generator/pretrained/utils.py | 113 ++++++++++ 10 files changed, 390 insertions(+), 359 deletions(-) diff --git a/README.md b/README.md index 09e5780..ca795d9 100644 --- a/README.md +++ b/README.md @@ -57,9 +57,9 @@ See the [List of Supported Models](#list-of-supported-models) section for all av |-------|-------------------| | HFPretrainedMolecularEncoder | transformers | | HFPretrainedMolecularGenerator | transformers | -| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed) | -| HFPretrainedMolecularGenerator (GP-MoLFormer) | transformers<=4.56.2 | -| HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | +| HFPretrainedMolecularGenerator (MolGen) | transformers, [selfies](https://github.com/aspuru-guzik-group/selfies) | +| HFPretrainedMolecularGenerator (Molexar) | transformers, [fragment-selfies](https://github.com/fairydance/Fragment-SELFIES), [molexar](https://github.com/fairydance/Molexar) | +| HFPretrainedMolecularGenerator (SAFE-GPT) | transformers, [safe-mol](https://github.com/datamol-io/safe) | | BFGNNMolecularPredictor | torch-scatter | | GRINMolecularPredictor | torch-scatter | | GRINMolecularPredictor (if enable `repetition_augmentation=True`) | CombineMols | @@ -70,11 +70,23 @@ See the [List of Supported Models](#list-of-supported-models) section for all av **For models that require `transformers`:** `pip install transformers` -**For MolGen (`selfies`):** `pip install "selfies>=2.1"` (tested on 2.x; 3.x is not guaranteed). +**For MolGen (`selfies`):** `pip install "selfies>=2.1"`. Source: [aspuru-guzik-group/selfies](https://github.com/aspuru-guzik-group/selfies). -**For GP-MoLFormer:** `pip install "transformers>=4.40,<=4.56.2"`. Do not use this with Molexar in the same environment (Molexar needs `transformers>=5.8`). +**For Molexar:** `pip install fragment-selfies loguru` ([Fragment-SELFIES](https://github.com/fairydance/Fragment-SELFIES)) and `pip install git+https://github.com/fairydance/Molexar.git` ([Molexar](https://github.com/fairydance/Molexar)). Molexar itself requires `transformers>=5.8`. -**For Molexar:** `pip install fragment-selfies loguru` and `pip install git+https://github.com/fairydance/Molexar.git`. +**For SAFE-GPT:** `pip install safe-mol` ([SAFE](https://github.com/datamol-io/safe)). + +```python +from torch_molecule import HFPretrainedMolecularGenerator + +model = HFPretrainedMolecularGenerator( + repo_id="datamol-io/safe-gpt", + generate_max_length=128, +) +model.fit() +print(model.generate(n_samples=5)) +print(model.generate(n_samples=5, scaffold="c1ccccc1")) +``` ## Usage @@ -206,6 +218,7 @@ new_model.load_from_local("qm9_grea.pt") | JTVAE | [Junction Tree Variational Autoencoder for Molecular Graph Generation. ICML 2018.](https://proceedings.mlr.press/v80/jin18a) | | GraphGA | [A Graph-Based Genetic Algorithm and Its Application to the Multiobjective Evolution of Median Molecules. Journal of Chemical Information and Computer Sciences 2004](https://pubs.acs.org/doi/10.1021/ci034290p) | | LSTM (SMILES) | [Long short-term memory (Neural Computation 1997)](https://ieeexplore.ieee.org/abstract/document/6795963) based on SMILES strings | +| Pretrained | [NovoMolGen](https://huggingface.co/chandar-lab/NovoMolGen_32M_SMILES_BPE): Causal LM pretrained on ZINC-22 for de novo SMILES generation.
[MolGen-large](https://huggingface.co/zjunlp/MolGen-large): Seq2Seq SELFIES generator with high chemical validity.
[MolGen-large-opt](https://huggingface.co/zjunlp/MolGen-large-opt): MolGen-large fine-tuned for QED / p-logP optimization.
[Molexar-10M-base](https://huggingface.co/fairydance/molexar-10m-base): Fragment-SELFIES de novo and fragment-constrained generation.
[Molexar-10M-omni](https://huggingface.co/fairydance/molexar-10m-omni): Multi-condition Molexar model for property-guided generation.
[SAFE-GPT](https://huggingface.co/datamol-io/safe-gpt): GPT-2 causal LM pretrained on SAFE strings for de novo generation and scaffold-prefix completion. | ### Representation Models diff --git a/docs/source/install.rst b/docs/source/install.rst index d465557..b20ab38 100644 --- a/docs/source/install.rst +++ b/docs/source/install.rst @@ -73,12 +73,12 @@ Some models require extra libraries. Install these packages if you use the corre +----------------------------------------------+----------------------------------------------+ | HFPretrainedMolecularGenerator | transformers | +----------------------------------------------+----------------------------------------------+ -| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed)| -+----------------------------------------------+----------------------------------------------+ -| HFPretrainedMolecularGenerator (GP-MoLFormer)| transformers<=4.56.2 | +| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies | +----------------------------------------------+----------------------------------------------+ | HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | +----------------------------------------------+----------------------------------------------+ +| HFPretrainedMolecularGenerator (SAFE-GPT) | transformers, safe-mol | ++----------------------------------------------+----------------------------------------------+ | BFGNNMolecularPredictor | torch-scatter | +----------------------------------------------+----------------------------------------------+ | GRINMolecularPredictor | torch-scatter | @@ -86,8 +86,8 @@ Some models require extra libraries. Install these packages if you use the corre **For models that require** ``transformers``: ``pip install transformers`` -**For MolGen** (``selfies``): ``pip install "selfies>=2.1"`` (tested on 2.x; 3.x is not guaranteed). +**For MolGen** (``selfies``): ``pip install "selfies>=2.1"``. Source: `aspuru-guzik-group/selfies `_. -**For GP-MoLFormer:** ``pip install "transformers>=4.40,<=4.56.2"``. Do not use this with Molexar in the same environment (Molexar needs ``transformers>=5.8``). +**For Molexar:** ``pip install fragment-selfies loguru`` (`Fragment-SELFIES `_) and ``pip install git+https://github.com/fairydance/Molexar.git`` (`Molexar `_). Molexar itself requires ``transformers>=5.8``. -**For Molexar:** ``pip install fragment-selfies loguru`` and ``pip install git+https://github.com/fairydance/Molexar.git``. +**For SAFE-GPT:** ``pip install safe-mol`` (`SAFE `_). diff --git a/tests/generator/hfpretrained.py b/tests/generator/hfpretrained.py index 25c9ed6..08ff2e5 100644 --- a/tests/generator/hfpretrained.py +++ b/tests/generator/hfpretrained.py @@ -7,11 +7,11 @@ "repo_id,expected", [ ("chandar-lab/NovoMolGen_32M_SMILES_BPE", "novomolgen"), - ("ibm-research/GP-MoLFormer-Uniq", "gp_molformer"), ("zjunlp/MolGen-large", "molgen"), ("zjunlp/MolGen-large-opt", "molgen"), ("fairydance/molexar-10m-base", "molexar"), ("fairydance/molexar-10m-omni", "molexar"), + ("datamol-io/safe-gpt", "safe_gpt"), ("some-user/custom-causal-lm", "causal_lm"), ], ) @@ -186,74 +186,6 @@ def test_unknown_repo_fallback_warns(): HFPretrainedMolecularGenerator(repo_id="some-user/custom-causal-lm") -def _transformers_supports_gp_molformer() -> bool: - pytest.importorskip("transformers") - import transformers - - major, minor, _ = map(int, transformers.__version__.split(".")[:3]) - return major < 5 and (major < 4 or minor < 57) - - -@pytest.mark.integration -def test_hf_pretrained_generator_gp_molformer_denovo(): - if not _transformers_supports_gp_molformer(): - pytest.skip("GP-MoLFormer requires transformers<=4.56.2") - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="ibm-research/GP-MoLFormer-Uniq", - generate_max_length=128, - ) - model.fit() - assert model.is_fitted_ is True - - smiles_list = model.generate(n_samples=2, temperature=1.0) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 2 - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - - -@pytest.mark.integration -def test_hf_pretrained_generator_gp_molformer_scaffold(): - if not _transformers_supports_gp_molformer(): - pytest.skip("GP-MoLFormer requires transformers<=4.56.2") - - from torch_molecule import HFPretrainedMolecularGenerator - - # IBM's official conditional prompt is a *partial* SMILES, not a closed ring. - scaffold = "c1cccc" - model = HFPretrainedMolecularGenerator( - repo_id="ibm-research/GP-MoLFormer-Uniq", - generate_max_length=128, - ) - model.fit() - - smiles_list = model.generate(n_samples=2, scaffold=scaffold, temperature=1.0) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 2 - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - assert all(smiles.startswith(scaffold) for smiles in smiles_list) - - -def test_gp_molformer_transformers_version_guard(): - pytest.importorskip("transformers") - import transformers - - from torch_molecule import HFPretrainedMolecularGenerator - - major, minor, _ = map(int, transformers.__version__.split(".")[:3]) - model = HFPretrainedMolecularGenerator( - repo_id="ibm-research/GP-MoLFormer-Uniq", - ) - - if major >= 5 or (major == 4 and minor >= 57): - with pytest.raises(ImportError, match="transformers<=4.56.2"): - model.fit() - else: - pytest.skip("GP-MoLFormer version guard only applies to transformers>=4.57") - - def _molexar_available() -> bool: try: import molexar # noqa: F401 @@ -381,3 +313,137 @@ def test_hf_pretrained_generator_molgen_scaffold_prefix(): and Chem.MolFromSmiles(smiles).HasSubstructMatch(benzene) for smiles in smiles_from_scaffold ) + + +def _safe_mol_available() -> bool: + try: + from torch_molecule.generator.pretrained.compat import ensure_safe_transformers_compat + + ensure_safe_transformers_compat() + from safe.converter import encode # noqa: F401 + + return True + except ImportError: + return False + + +def test_smiles_safe_roundtrip(): + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + + from torch_molecule.generator.pretrained.utils import safe_to_smiles, smiles_to_safe + + smiles = ["CCO", "c1ccccc1", "CC(=O)O"] + recovered = safe_to_smiles(smiles_to_safe(smiles)) + assert recovered == smiles + + +def test_smiles_to_safe_invalid_smiles(): + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + + from torch_molecule.generator.pretrained.utils import smiles_to_safe + + with pytest.raises(ValueError, match="Invalid SMILES"): + smiles_to_safe(["not-a-smiles"]) + + +def test_safe_to_smiles_drops_invalid_entries(): + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + + from torch_molecule.generator.pretrained.utils import safe_to_smiles, smiles_to_safe + + valid = smiles_to_safe(["CCO"])[0] + with pytest.warns(UserWarning, match="dropped 2 invalid SAFE"): + recovered = safe_to_smiles([valid, "not-valid-safe-[[[", ""]) + assert recovered == ["CCO"] + assert "" not in recovered + + +def test_decode_outputs_safe_gpt_drops_empty_and_warns(): + pytest.importorskip("transformers") + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + + from torch_molecule import HFPretrainedMolecularGenerator + from torch_molecule.generator.pretrained.utils import smiles_to_safe + + model = HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") + model._family = "safe_gpt" + valid = smiles_to_safe(["CCO"])[0] + with pytest.warns(UserWarning, match="got 1/2 valid SMILES"): + out = model._decode_outputs([valid, "not-a-safe"]) + assert len(out) == 1 + assert "" not in out + + +def test_safe_gpt_uses_gpt2_lm_head(): + pytest.importorskip("transformers") + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") + model._family = "safe_gpt" + assert model._get_model_class().__name__ == "GPT2LMHeadModel" + + +def test_safe_gpt_known_repo_does_not_warn_unknown(): + pytest.importorskip("transformers") + import warnings + + from torch_molecule import HFPretrainedMolecularGenerator + + with warnings.catch_warnings(record=True) as recorded: + warnings.simplefilter("always") + HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") + assert not any("Unknown repo_id" in str(item.message) for item in recorded) + + +@pytest.mark.integration +def test_hf_pretrained_generator_safe_gpt_denovo(): + pytest.importorskip("transformers") + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + + model = HFPretrainedMolecularGenerator( + repo_id="datamol-io/safe-gpt", + generate_max_length=128, + ) + model.fit() + assert model.is_fitted_ is True + + smiles_list = model.generate(n_samples=2, temperature=1.0) + assert isinstance(smiles_list, list) + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) + + +@pytest.mark.integration +def test_hf_pretrained_generator_safe_gpt_scaffold(): + pytest.importorskip("transformers") + if not _safe_mol_available(): + pytest.skip("safe-mol not installed") + from rdkit import Chem + + from torch_molecule import HFPretrainedMolecularGenerator + + scaffold = "c1ccccc1" + benzene = Chem.MolFromSmiles(scaffold) + model = HFPretrainedMolecularGenerator( + repo_id="datamol-io/safe-gpt", + generate_max_length=128, + ) + model.fit() + + smiles_list = model.generate(n_samples=2, scaffold=scaffold, temperature=1.0) + assert isinstance(smiles_list, list) + assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) + assert all( + Chem.MolFromSmiles(smiles) is not None + and Chem.MolFromSmiles(smiles).HasSubstructMatch(benzene) + for smiles in smiles_list + ) diff --git a/tests/generator/test_causal_lm.py b/tests/generator/test_causal_lm.py index 4dc5cbb..940af15 100644 --- a/tests/generator/test_causal_lm.py +++ b/tests/generator/test_causal_lm.py @@ -1,6 +1,5 @@ from unittest.mock import MagicMock -import pytest import torch from torch_molecule.generator.pretrained.families.causal_lm import generate_causal_lm @@ -41,29 +40,7 @@ def test_generate_causal_lm_bos_path(): assert outputs == ["SMILES_0", "SMILES_1", "SMILES_2"] -def test_generate_causal_lm_scaffold_path(): - model = _FakeModel() - model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) - - outputs = generate_causal_lm( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=2, - family="gp_molformer", - scaffold="c1ccccc1", - max_length=12, - do_sample=False, - ) - assert outputs == ["SMILES_0", "SMILES_1"] - kwargs = model.generate.call_args.kwargs - assert kwargs["use_cache"] is False - assert kwargs["top_k"] is None - # Default tokenize is [BOS, ..., EOS]; IBM drops the trailing special token. - assert kwargs["input_ids"].tolist() == [[1, 18, 19], [1, 18, 19]] - - -def test_generate_causal_lm_novomolgen_scaffold_keeps_all_tokens(): +def test_generate_causal_lm_scaffold_keeps_all_tokens(): model = _FakeModel() model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) @@ -85,27 +62,7 @@ def test_generate_causal_lm_novomolgen_scaffold_keeps_all_tokens(): assert "top_k" not in model.generate.call_args.kwargs -def test_generate_causal_lm_gp_molformer_denovo_path(): - model = _FakeModel() - model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) - - outputs = generate_causal_lm( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=2, - family="gp_molformer", - max_length=12, - do_sample=True, - ) - - assert outputs == ["SMILES_0", "SMILES_1"] - assert model.generate.call_args.kwargs["num_return_sequences"] == 2 - assert model.generate.call_args.kwargs["use_cache"] is False - assert "input_ids" not in model.generate.call_args.kwargs - - -def test_generate_causal_lm_novomolgen_does_not_force_use_cache_false(): +def test_generate_causal_lm_does_not_force_use_cache_false(): model = _FakeModel() model.generate = MagicMock(return_value=torch.zeros(3, 8, dtype=torch.long)) diff --git a/tests/generator/test_compat.py b/tests/generator/test_compat.py index 98f2e60..e430aa0 100644 --- a/tests/generator/test_compat.py +++ b/tests/generator/test_compat.py @@ -1,64 +1,13 @@ -from types import SimpleNamespace +import pytest -import torch +from torch_molecule.generator.pretrained.compat import ensure_safe_transformers_compat -from torch_molecule.generator.pretrained.compat import ( - patch_gp_molformer_generation_cache, - to_legacy_past_key_values, -) +def test_ensure_safe_transformers_compat_provides_constraints(): + pytest.importorskip("transformers") -class _EmptyCache: - def to_legacy_cache(self): - return ((None, None), (None, None)) + ensure_safe_transformers_compat() + import transformers.generation as generation - -class _PopulatedCache: - def to_legacy_cache(self): - key = torch.zeros(1, 2, 3, 4) - value = torch.zeros(1, 2, 3, 4) - return ((key, value),) - - -def test_to_legacy_past_key_values_empty_cache_becomes_none(): - assert to_legacy_past_key_values(None) is None - assert to_legacy_past_key_values(_EmptyCache()) is None - assert to_legacy_past_key_values(((None, None),)) is None - assert to_legacy_past_key_values(SimpleNamespace()) is None - - -def test_to_legacy_past_key_values_keeps_legacy_tensors(): - key = torch.zeros(1, 2, 3, 4) - value = torch.zeros(1, 2, 3, 4) - legacy = ((key, value),) - assert to_legacy_past_key_values(legacy) == legacy - converted = to_legacy_past_key_values(_PopulatedCache()) - assert converted[0][0].shape == (1, 2, 3, 4) - - -def test_patch_gp_molformer_generation_cache_converts_empty_cache(): - captured = {} - - class _Model: - def __init__(self): - self.config = SimpleNamespace(use_cache=True) - self.generation_config = SimpleNamespace(use_cache=True) - - def prepare_inputs_for_generation( - self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs - ): - captured["past_key_values"] = past_key_values - return {"input_ids": input_ids, "past_key_values": past_key_values} - - model = _Model() - patch_gp_molformer_generation_cache(model) - patch_gp_molformer_generation_cache(model) - - result = model.prepare_inputs_for_generation( - torch.tensor([[1]]), - past_key_values=_EmptyCache(), - ) - assert captured["past_key_values"] is None - assert result["past_key_values"] is None - assert model.config.use_cache is False - assert model.generation_config.use_cache is False + assert hasattr(generation, "DisjunctiveConstraint") + assert hasattr(generation, "PhrasalConstraint") diff --git a/torch_molecule/generator/pretrained/compat.py b/torch_molecule/generator/pretrained/compat.py index 5614a99..172572d 100644 --- a/torch_molecule/generator/pretrained/compat.py +++ b/torch_molecule/generator/pretrained/compat.py @@ -1,115 +1,33 @@ """Compatibility helpers for Hugging Face pretrained generators.""" -import sys -import types -from typing import Any, Optional, Tuple +from typing import Any -def _parse_transformers_version(version: str) -> Tuple[int, int, int]: - parts = version.split(".") - return tuple(int(part) for part in parts[:3]) +def ensure_safe_transformers_compat() -> None: + """Restore generation constraint symbols removed in transformers 5. - -def ensure_gp_molformer_transformers_compat() -> None: - """Validate that the installed ``transformers`` version can load GP-MoLFormer.""" - import transformers - - major, minor, _ = _parse_transformers_version(transformers.__version__) - if major >= 5 or (major == 4 and minor >= 57): - raise ImportError( - "GP-MoLFormer remote code is not compatible with transformers " - f"{transformers.__version__}. Install transformers<=4.56.2, for example:\n" - " pip install 'transformers>=4.40,<=4.56.2'" - ) - - -def ensure_transformers_onnx_compat() -> None: - """Provide a stub ``transformers.onnx`` module for legacy remote code. - - IBM MoLFormer remote configs import ``OnnxConfig`` from ``transformers.onnx``, - which was removed in recent ``transformers`` releases. The ONNX export class - is not needed for generation, so a lightweight stub is sufficient. + PyPI ``safe-mol`` still imports ``DisjunctiveConstraint`` and + ``PhrasalConstraint`` from ``transformers.generation`` when the package + is imported. Those classes were removed in transformers 5. Dummy + stand-ins are enough to import ``safe.converter`` and run standard GPT-2 + ``generate()``. Official ``SAFEDesign`` constrained beam search is not used. """ - if "transformers.onnx" in sys.modules: - return - try: - from transformers.onnx import OnnxConfig # noqa: F401 + import transformers.generation as generation + except ImportError: return - except ModuleNotFoundError: - pass - - onnx_module = types.ModuleType("transformers.onnx") - - class OnnxConfig: - """Minimal stub for legacy remote configuration modules.""" - - onnx_module.OnnxConfig = OnnxConfig - sys.modules["transformers.onnx"] = onnx_module - -def to_legacy_past_key_values(past_key_values: Any) -> Optional[tuple]: - """Convert HF Cache objects to the tuple layout IBM MoLFormer expects. - - ``transformers`` 4.47+ injects an empty ``DynamicCache`` into - ``prepare_inputs_for_generation``. IBM remote code then does - ``past_key_values[0][0].shape``, which raises because the empty cache - stores ``None`` instead of tensors. - """ - if past_key_values is None: - return None - - if not isinstance(past_key_values, (tuple, list)): - converter = getattr(past_key_values, "to_legacy_cache", None) - if converter is None: - return None - past_key_values = converter() - - if not past_key_values: - return None - first_layer = past_key_values[0] - if not first_layer: - return None - if first_layer[0] is None: - return None - return tuple(past_key_values) - - -def patch_gp_molformer_generation_cache(model: Any) -> None: - """Make GP-MoLFormer generation tolerate DynamicCache from transformers 4.56. - - IBM ``MolformerForCausalLM.prepare_inputs_for_generation`` only understands - the legacy tuple KV cache. Convert Cache objects (or empty caches) before - that method runs, and default ``use_cache`` off so later steps do not - reintroduce an incompatible cache layout. - """ - if getattr(model, "_gp_molformer_cache_patched", False): + if hasattr(generation, "DisjunctiveConstraint") and hasattr(generation, "PhrasalConstraint"): return - original = model.prepare_inputs_for_generation - - def prepare_inputs_for_generation( - input_ids, - past_key_values=None, - attention_mask=None, - inputs_embeds=None, - **kwargs, - ): - past_key_values = to_legacy_past_key_values(past_key_values) - return original( - input_ids, - past_key_values=past_key_values, - attention_mask=attention_mask, - inputs_embeds=inputs_embeds, - **kwargs, - ) - - model.prepare_inputs_for_generation = prepare_inputs_for_generation - model._gp_molformer_cache_patched = True - - generation_config = getattr(model, "generation_config", None) - if generation_config is not None: - generation_config.use_cache = False - config = getattr(model, "config", None) - if config is not None: - config.use_cache = False + class _RemovedConstraint: + def __init__(self, *args: Any, **kwargs: Any) -> None: + raise NotImplementedError( + "Constrained beam search was removed in transformers 5. " + "SAFE-GPT in torch-molecule uses standard GPT-2 generate()." + ) + + if not hasattr(generation, "DisjunctiveConstraint"): + generation.DisjunctiveConstraint = _RemovedConstraint + if not hasattr(generation, "PhrasalConstraint"): + generation.PhrasalConstraint = _RemovedConstraint diff --git a/torch_molecule/generator/pretrained/families/causal_lm.py b/torch_molecule/generator/pretrained/families/causal_lm.py index 1bb798d..3410b7b 100644 --- a/torch_molecule/generator/pretrained/families/causal_lm.py +++ b/torch_molecule/generator/pretrained/families/causal_lm.py @@ -31,7 +31,8 @@ def generate_causal_lm( n_samples : int Number of molecules to generate. family : Optional[str], default=None - Generator family name. GP-MoLFormer uses a model-specific de novo path. + Generator family name. Unused by the standard ``generate()`` path; + kept for call-site compatibility. max_length : int, default=64 Maximum generated sequence length passed to ``model.generate``. temperature : float, default=1.0 @@ -39,16 +40,15 @@ def generate_causal_lm( do_sample : bool, default=True Whether to use sampling during generation. scaffold : Optional[str], default=None - Optional SMILES prefix for scaffold completion. For GP-MoLFormer this - should be a *partial* SMILES string (IBM's example is ``c1cccc``); - the official tokenizer appends a trailing special token that is then - dropped so generation continues the prefix. + Optional tokenized prefix. Callers that need a SMILES-to-SAFE + conversion (SAFE-GPT) should pass the already-encoded prefix. Returns ------- List[str] Raw decoded strings from the tokenizer (may contain spaces). """ + del family # dispatch is done by HFPretrainedMolecularGenerator pad_token_id = tokenizer.pad_token_id if pad_token_id is None: pad_token_id = tokenizer.eos_token_id @@ -60,28 +60,14 @@ def generate_causal_lm( } if do_sample: generate_kwargs["temperature"] = temperature - if family == "gp_molformer": - # IBM remote code indexes tuple KV caches; transformers 4.56 injects - # an empty DynamicCache that crashes prepare_inputs_for_generation. - generate_kwargs["use_cache"] = False - generate_kwargs["top_k"] = None - generate_kwargs.update(kwargs) if scaffold: - if family == "gp_molformer": - # Match IBM/gp-molformer scripts/conditional_generation.py: - # tokenize with special tokens, then drop the trailing SEP/EOS. - input_ids = tokenizer(scaffold, return_tensors="pt")["input_ids"] - if input_ids.shape[1] > 1: - input_ids = input_ids[:, :-1] - else: - encoded = tokenizer(scaffold, return_tensors="pt", add_special_tokens=False) - input_ids = encoded["input_ids"] + encoded = tokenizer(scaffold, return_tensors="pt", add_special_tokens=False) + input_ids = encoded["input_ids"] input_ids = input_ids.to(device).expand(n_samples, -1).contiguous() generate_kwargs["input_ids"] = input_ids - elif family == "gp_molformer": - generate_kwargs["num_return_sequences"] = n_samples + generate_kwargs["attention_mask"] = torch.ones_like(input_ids) else: if tokenizer.bos_token_id is None: raise ValueError( diff --git a/torch_molecule/generator/pretrained/modeling_pretrained.py b/torch_molecule/generator/pretrained/modeling_pretrained.py index 2311902..1b8f303 100644 --- a/torch_molecule/generator/pretrained/modeling_pretrained.py +++ b/torch_molecule/generator/pretrained/modeling_pretrained.py @@ -14,6 +14,7 @@ from .registry import ( CAUSAL_LM_FAMILIES, MOLEXAR_FAMILIES, + SAFE_GPT_FAMILIES, SEQ2SEQ_FAMILIES, resolve_family, ) @@ -29,11 +30,11 @@ class HFPretrainedMolecularGenerator(BaseMolecularGenerator): Supported generation modes depend on the model family: - NovoMolGen: de novo SMILES generation from BOS. - - GP-MoLFormer: de novo generation and scaffold completion via ``scaffold=``. - MolGen-large / MolGen-large-opt: SELFIES seq2seq generation via ``prefix_selfies=`` or ``scaffold=`` (SMILES converted internally). - Molexar: Fragment-SELFIES de novo and fragment-constrained generation via ``start_smiles`` / ``start_string`` / ``conditions`` (omni). + - SAFE-GPT: GPT-2 causal LM on SAFE strings; de novo and ``scaffold=`` prefix. Other registered families can be loaded but may raise ``NotImplementedError`` until later phases are implemented. @@ -45,11 +46,6 @@ class HFPretrainedMolecularGenerator(BaseMolecularGenerator): repo_id: ``"chandar-lab/NovoMolGen_32M_SMILES_BPE"`` (https://huggingface.co/chandar-lab/NovoMolGen_32M_SMILES_BPE) - - GP-MoLFormer: Causal LM for de novo generation and scaffold decoration. - - repo_id: ``"ibm-research/GP-MoLFormer-Uniq"`` - (https://huggingface.co/ibm-research/GP-MoLFormer-Uniq) - - MolGen-large: Seq2Seq SELFIES generator with high chemical validity. repo_id: ``"zjunlp/MolGen-large"`` @@ -70,6 +66,12 @@ class HFPretrainedMolecularGenerator(BaseMolecularGenerator): repo_id: ``"fairydance/molexar-10m-omni"`` (https://huggingface.co/fairydance/molexar-10m-omni) + - SAFE-GPT: GPT-2 causal LM pretrained on 1.1B SAFE strings for de novo + generation and scaffold-prefix completion. + + repo_id: ``"datamol-io/safe-gpt"`` + (https://huggingface.co/datamol-io/safe-gpt) + Parameters ---------- repo_id : str @@ -81,10 +83,9 @@ class HFPretrainedMolecularGenerator(BaseMolecularGenerator): ``"hf-checkpoint"`` so standard ``model.generate`` works out of the box. trust_remote_code : bool, default=False Whether to trust remote code when loading from Hugging Face. - Automatically enabled for GP-MoLFormer and Molexar. + Automatically enabled for Molexar. tokenizer_repo_id : Optional[str], default=None - Optional Hugging Face repo for the tokenizer. GP-MoLFormer defaults to - ``"ibm-research/MoLFormer-XL-both-10pct"``. + Optional Hugging Face repo for the tokenizer. generate_max_length : int, default=64 Default ``max_length`` passed to ``generate()``. batch_size : int, default=8 @@ -313,17 +314,23 @@ def generate(self, n_samples: int = 10, **kwargs) -> List[str]: ``do_sample``, and ``scaffold``. For MolGen, use ``prefix_selfies`` or ``scaffold`` plus optional ``num_beams``, ``min_length``, and ``max_length``. For Molexar, use ``start_smiles``, ``start_string``, - ``generation_task``, or ``conditions`` for omni models. + ``generation_task``, or ``conditions`` for omni models. For SAFE-GPT, + use ``scaffold=`` with a SMILES prefix (converted to SAFE internally). Returns ------- List[str] - Generated SMILES strings. For MolGen and Molexar, invalid decodes - are dropped, so the list may be shorter than ``n_samples``. + Generated SMILES strings. For MolGen, Molexar, and SAFE-GPT, invalid + decodes are dropped, so the list may be shorter than ``n_samples``. """ self._check_is_fitted() if self._family in CAUSAL_LM_FAMILIES: + scaffold = kwargs.pop("scaffold", None) + if scaffold is not None and self._family in SAFE_GPT_FAMILIES: + from .utils import smiles_to_safe + + scaffold = smiles_to_safe([scaffold])[0] raw = generate_causal_lm( self.model, self.tokenizer, @@ -333,7 +340,7 @@ def generate(self, n_samples: int = 10, **kwargs) -> List[str]: max_length=kwargs.pop("max_length", self.generate_max_length), temperature=kwargs.pop("temperature", 1.0), do_sample=kwargs.pop("do_sample", True), - scaffold=kwargs.pop("scaffold", None), + scaffold=scaffold, **kwargs, ) return self._decode_outputs(raw) @@ -385,6 +392,10 @@ def _encode_inputs(self, smiles: List[str]) -> List[str]: from .utils import smiles_to_fragment_selfies return smiles_to_fragment_selfies(smiles) + if self._family in SAFE_GPT_FAMILIES: + from .utils import smiles_to_safe + + return smiles_to_safe(smiles) return smiles def _resolve_prefix_selfies(self, kwargs: Dict[str, Any]) -> Optional[str]: @@ -403,9 +414,9 @@ def _resolve_prefix_selfies(self, kwargs: Dict[str, Any]) -> Optional[str]: def _decode_outputs(self, outputs: List[str]) -> List[str]: """Normalize raw model strings to SMILES. - MolGen and Molexar drop strings that cannot be decoded. The returned - list length is the number of successful SMILES, which may be smaller - than ``n_samples``. + MolGen, Molexar, and SAFE-GPT drop strings that cannot be decoded. The + returned list length is the number of successful SMILES, which may be + smaller than ``n_samples``. """ n_attempted = len(outputs) @@ -418,6 +429,10 @@ def _decode_outputs(self, outputs: List[str]) -> List[str]: from .utils import fragment_selfies_to_smiles smiles = fragment_selfies_to_smiles(outputs) + elif self._family in SAFE_GPT_FAMILIES: + from .utils import safe_to_smiles + + smiles = safe_to_smiles(outputs) else: return [output.replace(" ", "") for output in outputs] @@ -431,40 +446,43 @@ def _decode_outputs(self, outputs: List[str]) -> List[str]: def _load_pretrained(self, local_path: Optional[str] = None) -> None: import transformers - from .compat import ( - ensure_gp_molformer_transformers_compat, - ensure_transformers_onnx_compat, - patch_gp_molformer_generation_cache, - ) - from .registry import DEFAULT_GP_MOLFORMER_TOKENIZER - if self._family in MOLEXAR_FAMILIES: self._load_molexar_pretrained(local_path=local_path) return - if self._family == "gp_molformer": - ensure_gp_molformer_transformers_compat() - ensure_transformers_onnx_compat() - load_kwargs = self._get_load_kwargs() model_cls = self._get_model_class() model_source = local_path or self.repo_id - tokenizer_repo = local_path or self.repo_id - if self._family == "gp_molformer" and local_path is None: - tokenizer_repo = self.tokenizer_repo_id or DEFAULT_GP_MOLFORMER_TOKENIZER - - self.tokenizer = transformers.AutoTokenizer.from_pretrained( - tokenizer_repo, - model_max_length=self.max_length, - **load_kwargs, - ) + tokenizer_repo = local_path or self.tokenizer_repo_id or self.repo_id + + if self._family in SAFE_GPT_FAMILIES: + self.tokenizer = self._load_safe_gpt_tokenizer(tokenizer_repo) + else: + self.tokenizer = transformers.AutoTokenizer.from_pretrained( + tokenizer_repo, + model_max_length=self.max_length, + **load_kwargs, + ) self.model = model_cls.from_pretrained(model_source, **load_kwargs) - if self._family == "gp_molformer": - patch_gp_molformer_generation_cache(self.model) self._setup_tokenizer() self.model.to(self.device) self.model.eval() + def _load_safe_gpt_tokenizer(self, tokenizer_repo: str): + """Load the custom SAFE tokenizer as a Hugging Face fast tokenizer.""" + from .utils import _require_safe + + _require_safe() + from safe.tokenizer import SAFETokenizer + + tokenizer_kwargs = {} + if self.revision is not None: + tokenizer_kwargs["revision"] = self.revision + safe_tokenizer = SAFETokenizer.from_pretrained(tokenizer_repo, **tokenizer_kwargs) + tokenizer = safe_tokenizer.get_pretrained() + tokenizer.model_max_length = self.max_length + return tokenizer + def _load_molexar_pretrained(self, local_path: Optional[str] = None) -> None: from huggingface_hub import snapshot_download @@ -521,6 +539,9 @@ def _get_model_class(self): if self._family in SEQ2SEQ_FAMILIES: return transformers.AutoModelForSeq2SeqLM + if self._family in SAFE_GPT_FAMILIES: + # Hub config lists SAFEDoubleHeadsModel; the LM head is standard GPT-2. + return transformers.GPT2LMHeadModel return transformers.AutoModelForCausalLM def _get_load_kwargs(self) -> Dict[str, Any]: @@ -531,13 +552,13 @@ def _get_load_kwargs(self) -> Dict[str, Any]: elif self.revision is not None: load_kwargs["revision"] = self.revision - if ( - self._family in MOLEXAR_FAMILIES - or self._family == "gp_molformer" - or self.trust_remote_code - ): + if self._family in MOLEXAR_FAMILIES or self.trust_remote_code: load_kwargs["trust_remote_code"] = True + if self._family in SAFE_GPT_FAMILIES: + # Extra property-prediction head in the checkpoint is unused. + load_kwargs["ignore_mismatched_sizes"] = True + return load_kwargs def _setup_tokenizer(self) -> None: @@ -548,6 +569,15 @@ def _setup_tokenizer(self) -> None: self.tokenizer.add_special_tokens({"pad_token": ""}) self.model.resize_token_embeddings(len(self.tokenizer)) + if self._family in SAFE_GPT_FAMILIES: + config = getattr(self.model, "config", None) + if self.tokenizer.bos_token_id is None and getattr(config, "bos_token_id", None) is not None: + self.tokenizer.bos_token_id = config.bos_token_id + if self.tokenizer.eos_token_id is None and getattr(config, "eos_token_id", None) is not None: + self.tokenizer.eos_token_id = config.eos_token_id + if self.tokenizer.pad_token_id is None and getattr(config, "pad_token_id", None) is not None: + self.tokenizer.pad_token_id = config.pad_token_id + @staticmethod def _require_transformers() -> None: try: diff --git a/torch_molecule/generator/pretrained/registry.py b/torch_molecule/generator/pretrained/registry.py index 9ed36fe..c66919f 100644 --- a/torch_molecule/generator/pretrained/registry.py +++ b/torch_molecule/generator/pretrained/registry.py @@ -4,31 +4,30 @@ KNOWN_REPOS: Dict[str, str] = { "chandar-lab/NovoMolGen_32M_SMILES_BPE": "novomolgen", - "ibm-research/GP-MoLFormer-Uniq": "gp_molformer", "zjunlp/MolGen-large": "molgen", "zjunlp/MolGen-large-opt": "molgen", "fairydance/molexar-10m-base": "molexar", "fairydance/molexar-10m-omni": "molexar", + "datamol-io/safe-gpt": "safe_gpt", } FAMILY_PREFIXES: Dict[str, str] = { "chandar-lab/NovoMolGen": "novomolgen", - "ibm-research/GP-MoLFormer": "gp_molformer", "zjunlp/MolGen": "molgen", "fairydance/molexar": "molexar", + "datamol-io/safe": "safe_gpt", } -DEFAULT_GP_MOLFORMER_TOKENIZER = "ibm-research/MoLFormer-XL-both-10pct" - MOLEXAR_OMNI_REPOS = frozenset( { "fairydance/molexar-10m-omni", } ) -CAUSAL_LM_FAMILIES = frozenset({"novomolgen", "gp_molformer", "causal_lm"}) +CAUSAL_LM_FAMILIES = frozenset({"novomolgen", "safe_gpt", "causal_lm"}) SEQ2SEQ_FAMILIES = frozenset({"molgen"}) MOLEXAR_FAMILIES = frozenset({"molexar"}) +SAFE_GPT_FAMILIES = frozenset({"safe_gpt"}) def resolve_family(repo_id: str) -> str: diff --git a/torch_molecule/generator/pretrained/utils.py b/torch_molecule/generator/pretrained/utils.py index 8d3880b..c480dc4 100644 --- a/torch_molecule/generator/pretrained/utils.py +++ b/torch_molecule/generator/pretrained/utils.py @@ -153,3 +153,116 @@ def fragment_selfies_to_smiles(fragment_selfies_list: List[str]) -> List[str]: if n_dropped: warnings.warn(f"dropped {n_dropped} invalid Fragment-SELFIES", stacklevel=2) return smiles_list + + +def _require_safe(): + from .compat import ensure_safe_transformers_compat + + ensure_safe_transformers_compat() + try: + from safe.converter import encode # noqa: F401 + except ImportError as exc: + raise ImportError( + "The 'safe-mol' package is required for SAFE-GPT conversion. " + "Install it with `pip install safe-mol`." + ) from exc + + +def smiles_to_safe(smiles: List[str]) -> List[str]: + """Convert SMILES strings to SAFE representations. + + Parameters + ---------- + smiles : List[str] + Input SMILES strings. + + Returns + ------- + List[str] + SAFE strings in the same order as the input. + + Raises + ------ + ValueError + If a SMILES string is invalid or cannot be encoded as SAFE. + """ + _require_safe() + try: + from safe._exception import SAFEEncodeError, SAFEFragmentationError + except ImportError: + from safe.converter import SAFEEncodeError, SAFEFragmentationError # type: ignore + from safe.converter import encode + + safe_list: List[str] = [] + for idx, smiles_string in enumerate(smiles): + mol = Chem.MolFromSmiles(smiles_string) + if mol is None: + raise ValueError(f"Invalid SMILES at index {idx}: {smiles_string}") + canonical = Chem.MolToSmiles(mol) + try: + try: + encoded = encode(canonical, canonical=True, allow_empty=True) + except TypeError: + encoded = encode(canonical, canonical=True) + except (SAFEEncodeError, SAFEFragmentationError) as exc: + raise ValueError( + f"SMILES at index {idx} is RDKit-valid but not SAFE-encodable: {smiles_string}" + ) from exc + if not encoded: + raise ValueError( + f"SMILES at index {idx} is RDKit-valid but not SAFE-encodable: {smiles_string}" + ) + safe_list.append(encoded) + return safe_list + + +def safe_to_smiles(safe_list: List[str]) -> List[str]: + """Convert SAFE strings to canonical SMILES. + + Entries that cannot be decoded are dropped and a warning is emitted. + Unexpected errors such as ``ImportError`` are not swallowed. + + Parameters + ---------- + safe_list : List[str] + Input SAFE strings. + + Returns + ------- + List[str] + Canonical SMILES strings for entries that decoded successfully. + """ + _require_safe() + try: + from safe._exception import SAFEDecodeError + except ImportError: + from safe.converter import SAFEDecodeError # type: ignore + from safe.converter import decode + + smiles_list: List[str] = [] + n_dropped = 0 + for safe_string in safe_list: + if not safe_string or not str(safe_string).strip(): + n_dropped += 1 + continue + try: + try: + decoded = decode( + str(safe_string).replace(" ", ""), + canonical=True, + ignore_errors=False, + ) + except TypeError: + decoded = decode(str(safe_string).replace(" ", "")) + except SAFEDecodeError: + n_dropped += 1 + continue + mol = Chem.MolFromSmiles(decoded) if decoded else None + if mol is None: + n_dropped += 1 + continue + smiles_list.append(Chem.MolToSmiles(mol)) + + if n_dropped: + warnings.warn(f"dropped {n_dropped} invalid SAFE", stacklevel=2) + return smiles_list From abcb68820efa22fc43ac8a8c038f0397125d278f Mon Sep 17 00:00:00 2001 From: m21hm9 Date: Fri, 18 Sep 2026 11:05:05 +0800 Subject: [PATCH 3/5] Add: SMILESDataset train_test_split with Random, Scaffold, Butina, and Size splits --- README.md | 20 +- tests/datasets/test_split.py | 327 +++++++++++++++++++++ torch_molecule/datasets/__init__.py | 5 + torch_molecule/datasets/constant.py | 55 +++- torch_molecule/datasets/split.py | 421 ++++++++++++++++++++++++++++ 5 files changed, 820 insertions(+), 8 deletions(-) create mode 100644 tests/datasets/test_split.py create mode 100644 torch_molecule/datasets/split.py diff --git a/README.md b/README.md index ca795d9..c23568f 100644 --- a/README.md +++ b/README.md @@ -129,15 +129,21 @@ assert molecular_data.target is None ### Fit a Model -After preparing the dataset, we can easily fit a model similar to how we use sklearn (actually, the coding is even simpler than sklearn, as we still need to do feature engineering in sklearn to convert molecule SMILES into vectors): +After preparing the dataset, split it, then fit a model with an sklearn-style API (no extra SMILES featurization is required): ```python +from torch_molecule.datasets import load_qm9 from torch_molecule import GREAMolecularPredictor -split = int(0.8 * len(smiles_list)) +data = load_qm9(local_dir='torchmol_data') +# "random" | "scaffold" | "butina" | "size" +# scaffold: unseen Bemis-Murcko scaffolds; butina: Tanimoto clusters; size: heavy-atom count +# Split the full dataset. subsample() is only for local debugging / CI — do not +# shrink QM9 (or any benchmark) just to make Butina cheaper. +train, val = data.train_test_split(test_size=0.2, method="scaffold", seed=42) grea = GREAMolecularPredictor( - num_task=num_task, + num_task=1, task_type="regression", evaluate_higher_better=False, verbose="progress_bar" #or "print_statement" recommended for jupyter notebooks, or "none" @@ -145,10 +151,10 @@ grea = GREAMolecularPredictor( # Fit with automatic hyperparameter tuning with 10 attempts, or implement .fit() with the default/manual hyperparameters grea.autofit( - X_train=smiles_list[:split], - y_train=property_np_array[:split], - X_val=smiles_list[split:], - y_val=property_np_array[split:], + X_train=train.data, + y_train=train.target, + X_val=val.data, + y_val=val.target, n_trials=10, ) ``` diff --git a/tests/datasets/test_split.py b/tests/datasets/test_split.py new file mode 100644 index 0000000..1616c18 --- /dev/null +++ b/tests/datasets/test_split.py @@ -0,0 +1,327 @@ +import numpy as np +import pytest +from rdkit import Chem +from rdkit.Chem.Scaffolds import MurckoScaffold + +from torch_molecule.datasets import SMILESDataset, subsample, train_test_split +from torch_molecule.datasets.split import _BUTINA_OOM_MESSAGE + + +def _benzene_family(): + return [ + "c1ccccc1", + "c1ccccc1O", + "c1ccccc1N", + "c1ccccc1C", + "c1ccccc1Cl", + "c1ccccc1F", + "Cc1ccccc1O", + "Nc1ccccc1O", + ] + + +def _other_molecules(): + return [ + "CCO", + "CCN", + "CCC", + "C1CCCCC1", + "n1ccccc1", + "C1CCNCC1", + "CC(=O)O", + "CC(C)O", + ] + + +def _labeled_dataset(): + smiles = _benzene_family() + _other_molecules() + y = np.arange(len(smiles), dtype=np.float32).reshape(-1, 1) + return SMILESDataset(data=smiles, target=y) + + +def _scaffold_smiles(smiles: str) -> str: + mol = Chem.MolFromSmiles(smiles) + scaffold = MurckoScaffold.GetScaffoldForMol(mol) + if scaffold is None or scaffold.GetNumAtoms() == 0: + return Chem.MolToSmiles(mol) + return Chem.MolToSmiles(scaffold) + + +def test_random_split_reproducible(): + data = _labeled_dataset() + train_a, test_a = train_test_split(data, test_size=0.25, method="random", seed=42) + train_b, test_b = train_test_split(data, test_size=0.25, method="random", seed=42) + assert train_a.data == train_b.data + assert test_a.data == test_b.data + np.testing.assert_array_equal(train_a.target, train_b.target) + + +def test_random_split_different_seeds(): + data = _labeled_dataset() + train_a, _ = train_test_split(data, test_size=0.25, method="random", seed=0) + train_b, _ = train_test_split(data, test_size=0.25, method="random", seed=1) + assert train_a.data != train_b.data + + +def test_random_split_ratio_and_coverage(): + data = _labeled_dataset() + train, holdout = train_test_split(data, test_size=0.25, method="random", seed=42) + n = len(data.data) + assert len(train.data) + len(holdout.data) == n + assert abs(len(holdout.data) / n - 0.25) < 1e-9 + assert set(train.data).isdisjoint(holdout.data) + assert set(train.data) | set(holdout.data) == set(data.data) + + +def test_random_preserves_multitask_target(): + smiles = _benzene_family() + _other_molecules() + y = np.column_stack( + [ + np.arange(len(smiles), dtype=np.float32), + np.arange(len(smiles), dtype=np.float32) * 2, + ] + ) + data = SMILESDataset(data=smiles, target=y) + train, holdout = train_test_split(data, test_size=0.2, method="random", seed=7) + assert train.target.shape[1] == 2 + assert holdout.target.shape[1] == 2 + assert train.target.shape[0] == len(train.data) + + +def test_split_unlabeled_dataset(): + data = SMILESDataset(data=_benzene_family() + _other_molecules(), target=None) + train, holdout = train_test_split(data, test_size=0.2, method="random", seed=1) + assert train.target is None + assert holdout.target is None + assert len(train.data) + len(holdout.data) == len(data.data) + + +def test_scaffold_no_leakage(): + data = _labeled_dataset() + train, holdout = train_test_split(data, test_size=0.3, method="scaffold", seed=42) + train_scaffolds = {_scaffold_smiles(s) for s in train.data} + holdout_scaffolds = {_scaffold_smiles(s) for s in holdout.data} + assert train_scaffolds.isdisjoint(holdout_scaffolds) + assert len(train.data) + len(holdout.data) == len(data.data) + assert len(set(train.data) & set(holdout.data)) == 0 + + +def test_scaffold_keeps_benzene_family_together(): + data = _labeled_dataset() + train, holdout = train_test_split(data, test_size=0.3, method="scaffold") + benzene_scaffold = _scaffold_smiles("c1ccccc1") + train_has = any(_scaffold_smiles(s) == benzene_scaffold for s in train.data) + holdout_has = any(_scaffold_smiles(s) == benzene_scaffold for s in holdout.data) + assert train_has ^ holdout_has + + +def test_scaffold_acyclic_molecules_do_not_crash(): + data = SMILESDataset( + data=["CCO", "CCN", "CCC", "CC", "C"], + target=np.arange(5).reshape(-1, 1), + ) + train, holdout = train_test_split(data, test_size=0.4, method="scaffold") + assert len(train.data) >= 1 + assert len(holdout.data) >= 1 + + +def test_scaffold_invalid_smiles_raises(): + data = SMILESDataset(data=["CCO", "not_a_smiles"], target=None) + with pytest.raises(ValueError, match="Invalid SMILES"): + train_test_split(data, method="scaffold") + + +def test_unknown_method_raises(): + data = _labeled_dataset() + with pytest.raises(ValueError, match="Unknown split method"): + train_test_split(data, method="kmeans") + + +def test_subsample_reproducible_and_size(): + data = _labeled_dataset() + a = subsample(data, n=5, seed=0) + b = subsample(data, n=5, seed=0) + c = subsample(data, n=5, seed=1) + assert len(a.data) == 5 + assert a.data == b.data + assert a.data != c.data + assert a.target.shape == (5, 1) + + +def test_subsample_and_split_methods_on_dataset(): + data = _labeled_dataset() + small = data.subsample(n=10, seed=3) + assert len(small.data) == 10 + train, holdout = small.train_test_split(test_size=0.3, method="random", seed=4) + assert len(train.data) + len(holdout.data) == 10 + butina_train, butina_hold = small.train_test_split( + test_size=0.3, method="butina", similarity_cutoff=0.4 + ) + assert len(butina_train.data) + len(butina_hold.data) == 10 + size_train, size_hold = small.train_test_split(test_size=0.3, method="size") + assert len(size_train.data) + len(size_hold.data) == 10 + + +def test_subsample_rejects_too_large_n(): + data = _labeled_dataset() + with pytest.raises(ValueError, match="larger than the dataset size"): + data.subsample(n=len(data.data) + 1) + + +def test_target_row_mismatch_raises(): + data = SMILESDataset(data=["CCO", "CCC"], target=np.array([[1.0]])) + with pytest.raises(ValueError, match="target has"): + train_test_split(data, method="random") + + +def _rdkit_butina_clusters(smiles, similarity_cutoff=0.65): + from rdkit.ML.Cluster import Butina + from rdkit.Chem import rdFingerprintGenerator + from rdkit import DataStructs as RDS + + gen = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=2048) + fps = [gen.GetFingerprint(Chem.MolFromSmiles(s)) for s in smiles] + dists = [] + for i in range(1, len(fps)): + sims = RDS.BulkTanimotoSimilarity(fps[i], fps[:i]) + dists.extend([1.0 - x for x in sims]) + return Butina.ClusterData( + dists, len(fps), 1.0 - similarity_cutoff, isDistData=True + ) + + +def _cluster_membership(clusters): + return {frozenset(cluster) for cluster in clusters} + + +def test_butina_matches_rdkit_clusterdata_membership(): + from torch_molecule.datasets.split import _butina_clusters + + smiles = _benzene_family() + _other_molecules() + for cutoff in (0.3, 0.4, 0.65): + ours = _butina_clusters(smiles, similarity_cutoff=cutoff) + rdkit = _rdkit_butina_clusters(smiles, similarity_cutoff=cutoff) + assert _cluster_membership(ours) == _cluster_membership(rdkit) + for cluster, ref in zip(ours, rdkit): + assert cluster[0] == ref[0] + + +def test_butina_no_cluster_leakage(): + from torch_molecule.datasets.split import _butina_clusters + + data = _labeled_dataset() + cutoff = 0.4 + train, holdout = train_test_split( + data, test_size=0.3, method="butina", similarity_cutoff=cutoff + ) + smiles_to_idx = {s: i for i, s in enumerate(data.data)} + clusters = _butina_clusters(data.data, similarity_cutoff=cutoff) + train_idx = {smiles_to_idx[s] for s in train.data} + holdout_idx = {smiles_to_idx[s] for s in holdout.data} + assert train_idx.isdisjoint(holdout_idx) + assert train_idx | holdout_idx == set(range(len(data.data))) + for cluster in clusters: + members = set(cluster) + assert members <= train_idx or members <= holdout_idx + + +def test_butina_ignores_seed_and_is_reproducible(): + data = _labeled_dataset() + a_train, a_hold = train_test_split( + data, test_size=0.3, method="butina", seed=0, similarity_cutoff=0.4 + ) + b_train, b_hold = train_test_split( + data, test_size=0.3, method="butina", seed=1, similarity_cutoff=0.4 + ) + assert a_train.data == b_train.data + assert a_hold.data == b_hold.data + + +def test_butina_invalid_smiles_raises(): + data = SMILESDataset(data=["CCO", "not_a_smiles"], target=None) + with pytest.raises(ValueError, match="Invalid SMILES"): + train_test_split(data, method="butina") + + +def test_butina_rejects_bad_cutoff(): + data = _labeled_dataset() + with pytest.raises(ValueError, match="similarity_cutoff"): + train_test_split(data, method="butina", similarity_cutoff=0.0) + with pytest.raises(ValueError, match="similarity_cutoff"): + train_test_split(data, method="butina", similarity_cutoff=1.5) + + +def _heavy_atoms(smiles: str) -> int: + return Chem.MolFromSmiles(smiles).GetNumHeavyAtoms() + + +def test_size_split_holds_out_larger_molecules(): + smiles = ["C", "CC", "CCC", "CCCC", "c1ccccc1", "c1ccccc1c1ccccc1"] + y = np.arange(len(smiles), dtype=np.float32).reshape(-1, 1) + data = SMILESDataset(data=smiles, target=y) + train, holdout = train_test_split( + data, test_size=1 / 3, method="size", direction="small_to_large" + ) + assert len(holdout.data) == 2 + assert max(_heavy_atoms(s) for s in train.data) <= min( + _heavy_atoms(s) for s in holdout.data + ) + assert set(train.data) | set(holdout.data) == set(smiles) + + +def test_size_split_large_to_small_holds_out_smaller_molecules(): + smiles = ["C", "CC", "CCC", "CCCC", "c1ccccc1", "c1ccccc1c1ccccc1"] + data = SMILESDataset(data=smiles, target=None) + train, holdout = train_test_split( + data, test_size=1 / 3, method="size", direction="large_to_small" + ) + assert max(_heavy_atoms(s) for s in holdout.data) <= min( + _heavy_atoms(s) for s in train.data + ) + + +def test_size_split_sizeshiftreg_protocol(): + smiles = [ + "C", + "CC", + "CCC", + "CCCC", + "CCCCC", + "CCCCCC", + "CCCCCCC", + "CCCCCCCC", + "CCCCCCCCC", + "c1ccccc1", + ] + data = SMILESDataset(data=smiles, target=None) + train, holdout = train_test_split( + data, test_size=0.2, method="size", mode="sizeshiftreg" + ) + n = len(smiles) + assert len(train.data) == int(round(0.5 * n)) + assert len(holdout.data) == int(round(0.1 * n)) + assert len(train.data) + len(holdout.data) < n + assert max(_heavy_atoms(s) for s in train.data) <= min( + _heavy_atoms(s) for s in holdout.data + ) + + +def test_size_invalid_smiles_raises(): + data = SMILESDataset(data=["CCO", "not_a_smiles"], target=None) + with pytest.raises(ValueError, match="Invalid SMILES"): + train_test_split(data, method="size") + + +def test_butina_oom_does_not_suggest_subsample(monkeypatch): + data = _labeled_dataset() + + def _boom(*args, **kwargs): + raise MemoryError("Unable to allocate array") + + monkeypatch.setattr( + "torch_molecule.datasets.split._butina_groups", _boom + ) + with pytest.raises(MemoryError, match="Do not subsample") as excinfo: + train_test_split(data, method="butina") + assert "chemfp" in str(excinfo.value) + assert "Do not subsample" in _BUTINA_OOM_MESSAGE diff --git a/torch_molecule/datasets/__init__.py b/torch_molecule/datasets/__init__.py index 407ce15..9fcb111 100644 --- a/torch_molecule/datasets/__init__.py +++ b/torch_molecule/datasets/__init__.py @@ -1,7 +1,12 @@ +from .constant import SMILESDataset from .load_hf_dataset import load_qm9, load_chembl2k, load_broad6k, load_toxcast, load_admet, load_zinc250k from .load_local_csv import load_gasperm +from .split import subsample, train_test_split __all__ = [ + "SMILESDataset", + "subsample", + "train_test_split", "load_qm9", "load_chembl2k", "load_broad6k", diff --git a/torch_molecule/datasets/constant.py b/torch_molecule/datasets/constant.py index 859b02c..f717fee 100644 --- a/torch_molecule/datasets/constant.py +++ b/torch_molecule/datasets/constant.py @@ -1,5 +1,5 @@ from dataclasses import dataclass -from typing import List +from typing import List, Tuple import numpy as np @dataclass @@ -14,6 +14,59 @@ class SMILESDataset: data: List[str] target: np.ndarray | None + def subsample(self, n: int, seed: int = 0) -> "SMILESDataset": + """Draw a random subset without replacement. + + Intended for local debugging and CI, not as a way to make structure- + aware splits cheaper. Do not subsample a benchmark (for example QM9) + because Butina is slow or memory-heavy; that changes what the split + measures. + + Parameters + ---------- + n : int + Number of molecules to keep. + seed : int, default=0 + Random seed. + """ + from .split import subsample + + return subsample(self, n=n, seed=seed) + + def train_test_split( + self, + test_size: float = 0.2, + method: str = "random", + seed: int = 42, + **kwargs, + ) -> Tuple["SMILESDataset", "SMILESDataset"]: + """Split into train and holdout ``SMILESDataset`` objects. + + Parameters + ---------- + test_size : float, default=0.2 + Requested holdout fraction. + method : {"random", "scaffold", "butina", "size"}, default="random" + Split protocol. ``random`` is an i.i.d. baseline; ``scaffold`` + holds out unseen Bemis-Murcko scaffolds; ``butina`` holds out + unseen Taylor-Butina clusters; ``size`` splits by heavy-atom count. + seed : int, default=42 + Random seed (used by ``random``). + **kwargs + Extra options forwarded to the splitter (``use_csk`` for scaffold, + ``similarity_cutoff`` for butina, ``direction`` / ``mode`` for size). + + Returns + ------- + train, holdout : SMILESDataset + The second dataset is intended as validation data for ``fit``. + """ + from .split import train_test_split + + return train_test_split( + self, test_size=test_size, method=method, seed=seed, **kwargs + ) + TOXCAST_TASKS = [ 'ACEA_T47D_80hr_Negative', 'ACEA_T47D_80hr_Positive', diff --git a/torch_molecule/datasets/split.py b/torch_molecule/datasets/split.py new file mode 100644 index 0000000..f62c105 --- /dev/null +++ b/torch_molecule/datasets/split.py @@ -0,0 +1,421 @@ +"""Train/test splitting utilities for molecular SMILES datasets. + +Random splitting is an i.i.d. baseline. Scaffold splitting groups molecules by +Bemis-Murcko frameworks so that the same scaffold does not appear in both +splits. Butina splitting groups by Taylor-Butina clusters on Morgan fingerprints +(sparse Tanimoto neighbor graph). Size splitting holds out larger (or smaller) +molecules by heavy-atom count. + +The second split returned by ``train_test_split`` is a holdout set. Pass it to +``fit(..., X_val, y_val)`` as validation data. A disjoint final test set requires +a later three-way split API. +""" + +from __future__ import annotations + +from collections import defaultdict +from typing import Dict, List, Optional, Sequence, Tuple + +import numpy as np +from rdkit import Chem, DataStructs +from rdkit.Chem import rdFingerprintGenerator +from rdkit.Chem.Scaffolds import MurckoScaffold + +from .constant import SMILESDataset + +_SUPPORTED_METHODS = ("random", "scaffold", "butina", "size") +_SIZE_DIRECTIONS = ("small_to_large", "large_to_small") +_SIZE_MODES = ("standard", "sizeshiftreg") +_BUTINA_RADIUS = 2 +_BUTINA_FP_SIZE = 2048 +_BUTINA_DEFAULT_CUTOFF = 0.65 +_SIZESHIFTREG_TRAIN_FRACTION = 0.5 +_SIZESHIFTREG_TEST_FRACTION = 0.1 +_BUTINA_OOM_MESSAGE = ( + "Butina split ran out of memory on the full dataset. " + "Do not subsample to work around a split failure; that changes what " + "the split measures. QM9-scale data is expected to fit; for much " + "larger libraries an optional chemfp backend may be added later." +) + +_MORGAN_FP_GEN = rdFingerprintGenerator.GetMorganGenerator( + radius=_BUTINA_RADIUS, fpSize=_BUTINA_FP_SIZE +) + + +def subsample( + dataset: SMILESDataset, + n: int, + seed: int = 0, +) -> SMILESDataset: + """Draw a random subset without replacement. + + Intended for local debugging and CI, not as a way to make structure-aware + splits cheaper. Do not subsample a benchmark because Butina is slow or + memory-heavy; that changes what the split measures. + + Parameters + ---------- + dataset : SMILESDataset + Input dataset. + n : int + Number of molecules to keep. Must be at least 1 and at most the dataset + size. + seed : int, default=0 + Random seed. + + Returns + ------- + SMILESDataset + Subsampled dataset. If ``n`` equals the dataset size, a copy with the + original order is returned. + """ + _check_dataset(dataset) + n_total = len(dataset.data) + if n < 1: + raise ValueError(f"n must be >= 1, got {n}.") + if n > n_total: + raise ValueError( + f"n={n} is larger than the dataset size ({n_total})." + ) + if n == n_total: + return _subset(dataset, list(range(n_total))) + + rng = np.random.RandomState(seed) + indices = rng.choice(n_total, size=n, replace=False) + return _subset(dataset, indices.tolist()) + + +def train_test_split( + dataset: SMILESDataset, + test_size: float = 0.2, + method: str = "random", + seed: int = 42, + *, + use_csk: bool = False, + similarity_cutoff: float = _BUTINA_DEFAULT_CUTOFF, + direction: str = "small_to_large", + mode: str = "standard", +) -> Tuple[SMILESDataset, SMILESDataset]: + """Split a SMILES dataset into train and holdout subsets. + + Parameters + ---------- + dataset : SMILESDataset + Input dataset. + test_size : float, default=0.2 + Fraction of molecules requested for the holdout set. For scaffold and + Butina splits the realized fraction can differ because whole groups are + assigned together. Ignored when ``method="size"`` and + ``mode="sizeshiftreg"``. + method : {"random", "scaffold", "butina", "size"}, default="random" + ``"random"`` is an i.i.d. baseline (often optimistic for molecules). + ``"scaffold"`` holds out unseen Bemis-Murcko scaffolds. + ``"butina"`` holds out unseen Taylor-Butina clusters (Morgan / Tanimoto). + ``"size"`` holds out molecules by heavy-atom count. + seed : int, default=42 + Random seed. Used by ``random``. Scaffold, Butina clustering, and size + assignment are deterministic; ``seed`` is accepted for API stability + and ignored. + use_csk : bool, default=False + Scaffold only. If True, generic cyclic skeletons are used (all atoms + as carbon). If False, atom types are kept (RDKit default). + similarity_cutoff : float, default=0.65 + Butina only. Tanimoto **similarity** threshold in ``(0, 1]``. Molecules + with similarity at least this value are neighbors. This is not + DeepChem's distance cutoff. + direction : {"small_to_large", "large_to_small"}, default="small_to_large" + Size only (``mode="standard"``). ``small_to_large`` puts smaller + molecules in train and larger ones in holdout. + mode : {"standard", "sizeshiftreg"}, default="standard" + Size only. ``sizeshiftreg`` uses the SizeShiftReg protocol: smallest + 50% train, largest 10% holdout; the middle 40% is unused. + + Returns + ------- + train, holdout : SMILESDataset + The second dataset is a holdout split intended as validation data for + ``fit`` / ``autofit``. + """ + _check_dataset(dataset) + if method not in _SUPPORTED_METHODS: + raise ValueError( + f"Unknown split method {method!r}. " + f"Supported methods: {list(_SUPPORTED_METHODS)}." + ) + if not 0.0 < test_size < 1.0: + raise ValueError(f"test_size must be in (0, 1), got {test_size}.") + + n = len(dataset.data) + if n < 2: + raise ValueError("Need at least 2 molecules to split a dataset.") + + if method == "random": + idx_train, idx_test = _random_split(n, test_size, seed) + elif method == "scaffold": + groups = _scaffold_groups(dataset.data, use_csk=use_csk) + idx_train, idx_test = _group_split(groups, test_size) + elif method == "butina": + if not 0.0 < similarity_cutoff <= 1.0: + raise ValueError( + f"similarity_cutoff must be in (0, 1], got {similarity_cutoff}." + ) + groups = _butina_groups_or_oom( + dataset.data, similarity_cutoff=similarity_cutoff + ) + idx_train, idx_test = _group_split(groups, test_size) + else: + idx_train, idx_test = _size_split( + dataset.data, + test_size=test_size, + direction=direction, + mode=mode, + ) + + return _subset(dataset, idx_train), _subset(dataset, idx_test) + + +def _check_dataset(dataset: SMILESDataset) -> None: + if not isinstance(dataset, SMILESDataset): + raise TypeError( + f"dataset must be a SMILESDataset, got {type(dataset).__name__}." + ) + if not isinstance(dataset.data, list): + raise TypeError("dataset.data must be a list of SMILES strings.") + if dataset.target is not None: + target = np.asarray(dataset.target) + if target.shape[0] != len(dataset.data): + raise ValueError( + f"target has {target.shape[0]} rows but data has " + f"{len(dataset.data)} molecules." + ) + + +def _random_split( + n: int, test_size: float, seed: int +) -> Tuple[List[int], List[int]]: + rng = np.random.RandomState(seed) + perm = rng.permutation(n) + n_test = int(round(n * test_size)) + n_test = min(max(n_test, 1), n - 1) + idx_test = np.sort(perm[:n_test]).tolist() + idx_train = np.sort(perm[n_test:]).tolist() + return idx_train, idx_test + + +def _mols_from_smiles(smiles_list: Sequence[str]) -> List[Chem.Mol]: + from ..utils.checker import MolecularInputChecker + + invalid = [] + mols: List[Optional[Chem.Mol]] = [None] * len(smiles_list) + for i, smiles in enumerate(smiles_list): + if not isinstance(smiles, str): + invalid.append(f"Non-string SMILES at index {i}: {smiles!r}") + continue + is_valid, error_msg, mol = MolecularInputChecker.validate_smiles( + smiles, i + ) + if not is_valid: + invalid.append(error_msg) + continue + mols[i] = mol + + if invalid: + raise ValueError("Invalid SMILES found:\n" + "\n".join(invalid)) + return mols # type: ignore[return-value] + + +def _scaffold_groups( + smiles_list: Sequence[str], use_csk: bool = False +) -> Dict[str, List[int]]: + groups: Dict[str, List[int]] = defaultdict(list) + for i, mol in enumerate(_mols_from_smiles(smiles_list)): + groups[_scaffold_key(mol, use_csk=use_csk)].append(i) + return dict(groups) + + +def _scaffold_key(mol: Chem.Mol, use_csk: bool = False) -> str: + scaffold = MurckoScaffold.GetScaffoldForMol(mol) + if scaffold is None or scaffold.GetNumAtoms() == 0: + return Chem.MolToSmiles(mol) + + if use_csk: + scaffold = MurckoScaffold.MakeScaffoldGeneric(scaffold) + if scaffold is None or scaffold.GetNumAtoms() == 0: + return Chem.MolToSmiles(mol) + return Chem.MolToSmiles(scaffold) + + +def _butina_groups( + smiles_list: Sequence[str], + similarity_cutoff: float = _BUTINA_DEFAULT_CUTOFF, +) -> Dict[str, List[int]]: + clusters = _butina_clusters(smiles_list, similarity_cutoff=similarity_cutoff) + return {f"cluster_{i}": members for i, members in enumerate(clusters)} + + +def _butina_groups_or_oom( + smiles_list: Sequence[str], + similarity_cutoff: float = _BUTINA_DEFAULT_CUTOFF, +) -> Dict[str, List[int]]: + try: + return _butina_groups( + smiles_list, similarity_cutoff=similarity_cutoff + ) + except MemoryError as exc: + raise MemoryError(_BUTINA_OOM_MESSAGE) from exc + + +def _butina_clusters( + smiles_list: Sequence[str], + similarity_cutoff: float = _BUTINA_DEFAULT_CUTOFF, +) -> List[List[int]]: + """Exact Taylor-Butina clusters via a sparse Tanimoto neighbor graph. + + Matches RDKit ``Butina.ClusterData`` membership (``reordering=False``) + without storing the condensed distance matrix. Neighbor edges are pairs + with Tanimoto similarity >= ``similarity_cutoff``. Degree includes self, + matching RDKit's zero self-distance. Ties break like RDKit: higher index + first among equal degrees. + """ + mols = _mols_from_smiles(smiles_list) + fps = [_MORGAN_FP_GEN.GetFingerprint(mol) for mol in mols] + neighbor_lists = _butina_neighbor_lists(fps, similarity_cutoff) + return _butina_exclusion_spheres(neighbor_lists) + + +def _butina_neighbor_lists( + fps: Sequence[DataStructs.ExplicitBitVect], + similarity_cutoff: float, +) -> List[List[int]]: + n = len(fps) + neighbors: List[List[int]] = [[] for _ in range(n)] + for i in range(n): + neighbors[i].append(i) + if i == 0: + continue + sims = DataStructs.BulkTanimotoSimilarity(fps[i], fps[:i]) + if sims: + hits = np.flatnonzero(np.asarray(sims, dtype=np.float64) >= similarity_cutoff) + for j in hits.tolist(): + neighbors[i].append(int(j)) + neighbors[j].append(i) + for i in range(n): + neighbors[i].sort() + return neighbors + + +def _butina_exclusion_spheres( + neighbor_lists: Sequence[Sequence[int]], +) -> List[List[int]]: + n = len(neighbor_lists) + sorted_indices = [ + (len(nbrs), idx) for idx, nbrs in enumerate(neighbor_lists) + ] + sorted_indices.sort(reverse=True) + + clusters: List[List[int]] = [] + seen = np.zeros(n, dtype=bool) + + while sorted_indices and sorted_indices[0][0] > 1: + _, idx = sorted_indices.pop(0) + if seen[idx]: + continue + cluster = [idx] + seen[idx] = True + for neighbor in neighbor_lists[idx]: + if not seen[neighbor]: + cluster.append(neighbor) + seen[neighbor] = True + clusters.append(cluster) + + while sorted_indices: + _, idx = sorted_indices.pop(0) + if seen[idx]: + continue + clusters.append([idx]) + return clusters + + +def _size_split( + smiles_list: Sequence[str], + test_size: float, + direction: str = "small_to_large", + mode: str = "standard", +) -> Tuple[List[int], List[int]]: + if direction not in _SIZE_DIRECTIONS: + raise ValueError( + f"Unknown size direction {direction!r}. " + f"Supported directions: {list(_SIZE_DIRECTIONS)}." + ) + if mode not in _SIZE_MODES: + raise ValueError( + f"Unknown size mode {mode!r}. " + f"Supported modes: {list(_SIZE_MODES)}." + ) + + mols = _mols_from_smiles(smiles_list) + n_atoms = np.array([mol.GetNumHeavyAtoms() for mol in mols], dtype=np.int64) + order = np.argsort(n_atoms, kind="stable") + n = len(order) + + if mode == "sizeshiftreg": + n_train = int(round(_SIZESHIFTREG_TRAIN_FRACTION * n)) + n_test = int(round(_SIZESHIFTREG_TEST_FRACTION * n)) + n_train = min(max(n_train, 1), n - 1) + n_test = min(max(n_test, 1), n - n_train) + idx_train = order[:n_train].tolist() + idx_test = order[-n_test:].tolist() + return idx_train, idx_test + + n_test = int(round(n * test_size)) + n_test = min(max(n_test, 1), n - 1) + if direction == "small_to_large": + idx_train = order[:-n_test].tolist() + idx_test = order[-n_test:].tolist() + else: + idx_train = order[n_test:].tolist() + idx_test = order[:n_test].tolist() + return idx_train, idx_test + + +def _group_split( + groups: Dict[str, List[int]], test_size: float +) -> Tuple[List[int], List[int]]: + """Assign whole groups with DeepChem-style greedy filling. + + Groups are sorted by decreasing size (group id as a tie-break). A group + goes to train if it still fits under the train cutoff; otherwise it goes to + the holdout set. + """ + n = sum(len(idx) for idx in groups.values()) + train_cutoff = (1.0 - test_size) * n + ordered = sorted(groups.items(), key=lambda item: (-len(item[1]), item[0])) + + idx_train: List[int] = [] + idx_test: List[int] = [] + for _, members in ordered: + members_sorted = sorted(members) + if len(idx_train) + len(members_sorted) > train_cutoff: + idx_test.extend(members_sorted) + else: + idx_train.extend(members_sorted) + + if not idx_train or not idx_test: + raise ValueError( + "Group split produced an empty train or holdout set. " + "Try a different test_size or a dataset with more diverse groups." + ) + return idx_train, idx_test + + +def _subset(dataset: SMILESDataset, indices: Sequence[int]) -> SMILESDataset: + indices = list(indices) + data = [dataset.data[i] for i in indices] + if dataset.target is None: + target: Optional[np.ndarray] = None + else: + target = np.asarray(dataset.target)[indices] + if target.ndim == 1: + target = target.reshape(-1, 1) + return SMILESDataset(data=data, target=target) From f23497d570effb74081a536125ce1bf49df6714e Mon Sep 17 00:00:00 2001 From: m21hm9 Date: Fri, 18 Sep 2026 21:39:50 +0800 Subject: [PATCH 4/5] Fix: Causal LM scaffolds using max_new_tokens and add pretrained_{model} generator tests --- README.md | 1 - molecule_notes/butina_optimization.md | 283 +++++++++++ molecule_notes/generation_model.md | 473 ++++++++++++++++++ molecule_notes/train_test_split.md | 385 ++++++++++++++ tests/generator/hfpretrained.py | 449 ----------------- tests/generator/pretrained_molexar.py | 66 +++ tests/generator/pretrained_molgen.py | 74 +++ tests/generator/pretrained_novomolgen.py | 86 ++++ tests/generator/pretrained_safe_gpt.py | 58 +++ tests/generator/test_causal_lm.py | 79 --- tests/generator/test_compat.py | 13 - tests/generator/test_finetune.py | 197 -------- tests/generator/test_molexar.py | 38 -- tests/generator/test_seq2seq.py | 67 --- .../pretrained/families/causal_lm.py | 21 +- .../pretrained/modeling_pretrained.py | 82 ++- 16 files changed, 1508 insertions(+), 864 deletions(-) create mode 100644 molecule_notes/butina_optimization.md create mode 100644 molecule_notes/generation_model.md create mode 100644 molecule_notes/train_test_split.md delete mode 100644 tests/generator/hfpretrained.py create mode 100644 tests/generator/pretrained_molexar.py create mode 100644 tests/generator/pretrained_molgen.py create mode 100644 tests/generator/pretrained_novomolgen.py create mode 100644 tests/generator/pretrained_safe_gpt.py delete mode 100644 tests/generator/test_causal_lm.py delete mode 100644 tests/generator/test_compat.py delete mode 100644 tests/generator/test_finetune.py delete mode 100644 tests/generator/test_molexar.py delete mode 100644 tests/generator/test_seq2seq.py diff --git a/README.md b/README.md index c23568f..1c7cf15 100644 --- a/README.md +++ b/README.md @@ -81,7 +81,6 @@ from torch_molecule import HFPretrainedMolecularGenerator model = HFPretrainedMolecularGenerator( repo_id="datamol-io/safe-gpt", - generate_max_length=128, ) model.fit() print(model.generate(n_samples=5)) diff --git a/molecule_notes/butina_optimization.md b/molecule_notes/butina_optimization.md new file mode 100644 index 0000000..f4ba583 --- /dev/null +++ b/molecule_notes/butina_optimization.md @@ -0,0 +1,283 @@ +# Butina 聚类 / Split:原文 vs 现有实现与优化 + +本文记录 **Taylor–Butina** 在分子 ML 划分中的算法原意,以及 2024–2026 年各实现如何加速。 +核心结论:**精确 Butina 不需要全距离矩阵**;当前工业/新开源优化都走「阈值邻居 + exclusion sphere」。DeepChem/RDKit `ClusterData` 是正确但未针对大库优化的参考实现。 + +与本仓库关系:`method="butina"` 尚未实现。若做全集(如 QM9 ~13 万)且要求 **accuracy**,应对齐 Chalcedon / chemfp 的稀疏路径,而不是 DeepChem 的稠密 `dists` 列表。 + +--- + +## 1. 原文算法(必须对齐的「准确」定义) + +**文献:** Darko Butina, *Unsupervised Data Base Clustering Based on Daylight’s Fingerprint and Tanimoto Similarity*, J. Chem. Inf. Comput. Sci. **1999**, 39, 747–750. +https://doi.org/10.1021/ci9803381 + +**输入:** 分子指纹、Tanimoto 阈值 $t$(例如 0.65 表示「够像才算邻居」)。 + +**Tanimoto(Jaccard):** + +$$ +T(A,B)=\frac{|A\cap B|}{|A\cup B|}=\frac{c}{a+b-c} +$$ + +范围 $[0,1]$。距离常写 $d=1-T$。 + +**三步:** + +1. 生成指纹(原文 Daylight;现代几乎一律 Morgan/ECFP)。 +2. 对每个分子数邻居:$T \ge t$ 的个数;按邻居数 **降序** 排序(潜在簇中心)。 +3. **Exclusion sphere:** 取下一个未标记分子当中心,所有 $T \ge t$ 且未标记的邻居进该簇并标记;已标记者不再当中心、不进别簇。 + +**精确 Butina 真正需要的信息只有布尔邻居关系:** + +$$ +N(i,j)=\mathbf{1}[T(i,j)\ge t] +$$ + +不需要任何 $T < t$ 的数值。因此: + +| 做法 | 是否精确 Butina | +|------|-----------------| +| 全距离矩阵 + `Butina.ClusterData` | 是(小 $n$) | +| 阈值邻接表 + **同一套** 排序和收球 | **同样是**(可上全集) | +| LSH / 漏邻居 / 先抽子集再指派 | 否 | +| BitBIRCH 等别的聚类 | 否(即使质量「差不多」) | + +--- + +## 2. 对比总表 + +| | 原文 1999 | RDKit `ClusterData` | DeepChem / Datamol | Chalcedon (Rowan, 2026) | chemfp `butina` | FPSim2 (ChEMBL) | BitBIRCH | +|--|-----------|---------------------|--------------------|-------------------------|-----------------|-----------------|----------| +| **是不是 Butina** | 定义 | 是 | 是 | 是(对照过 RDKit 簇) | 是 | 否(只做精确相似度搜索) | **否** | +| **指纹** | Daylight | 调用方提供 | Morgan r=2, **1024** bit | Morgan r=2, **2048** bit | 调用方 / fps | RDKit 指纹库 | 二进制指纹 | +| **阈值语义** | 相似度 $T\ge t$ | **距离** `distThresh` | `cutoff` 传给 ClusterData = **距离**(默认 0.6) | `cutoff` = **距离**(博客默认 0.65 ⇒ $T\ge 0.35$) | `--threshold` = **相似度** | `threshold` = **相似度** | 另一套层次阈值 | +| **存什么** | 概念上的邻居 | 全部 condensed 距离 $O(n^2)$ | 全部 `dists` 列表 | 分块;峰值 $O(n)$ + 一块 workspace | **稀疏** NxN(只存 $T\ge t$) | 稀疏 CSR | 树 / iSIM 统计 | +| **加速** | 无 | Bulk 仍填满矩阵 | `BulkTanimotoSimilarity` | NumPy/BLAS、float32、上三角分块、两遍扫描 | POPCNT、稀疏矩阵可存 npz | POPCNT + Swamidass bound、多核/GPU | $O(n)$ 近似层次 | +| **大库** | 当时为替代难调的 Jarvis–Patrick | $n\sim 10^5$ **OOM** | 文档写明 $O(n^2)$,中小集 | **10 万 ≈ 21s / 2.4GB** | 工业级库检索 + 聚类 | 适合建邻居图 | 百万级,但 **不是同簇** | +| **可复现平局** | 按邻居数排序 | 实现相关 | 未强调 | 确定性收球 | 默认 `randomize`;应用 `first`/`last` | N/A | N/A | +| **适合本库** | 语义标准 | 小 $n$ **金标准测试** | API 对标对象,不要抄矩阵 | **全集精确路径的最佳开源参照** | 算法参照;依赖/许可需评估 | 可选搜索后端 | 不要叫 `method="butina"` | + +--- + +## 3. 原文 vs 各实现(分节) + +### 3.1 原文 Butina (1999) + +**问题:** Jarvis–Patrick 要调两个参数,簇要么极大且杂、要么碎;大库要手工调。 + +**方法:** 单阈值、簇中心与每个成员都满足 $T\ge t$、exclusion sphere。 + +**和 split 的关系:** 原文不做 train/test。后来 DeepChem 把 **每个簇当 group**,整组进 train 或 test(与 scaffold 相同贪心)。 + +**局限:** 阈值无唯一正确答案;簇大小极不均匀;指纹种类会改变结果。 + +--- + +### 3.2 RDKit `rdkit.ML.Cluster.Butina` + +**做法:** + +```text +dists = condensed 下三角 (1 - Tanimoto) # 必须全部 pair +clusters = Butina.ClusterData(dists, n, distThresh, isDistData=True) +# 每个簇第一个元素是 centroid +``` + +**优化:** 几乎没有。调用方可用 `DataStructs.BulkTanimotoSimilarity` 加快 **填矩阵**,矩阵本身仍是 $O(n^2)$ 内存。 + +**实测(他人报告):** + +- Macs in Chemistry:15 万分子聚类时内存涨到 **80–267 GB**,进程被杀。 +- Chalcedon 基准:$n=10^5$ **RDKit OOM**;$n=5\times 10^4$ 约 173s / 110GB RSS。 + +**本库用法:** $n \le 2000$ 的单元测试金标准,**禁止** 对 QM9 全量调用。 + +--- + +### 3.3 DeepChem `ButinaSplitter` / Datamol + +**源码要点**(`deepchem/splits/splitters.py`): + +- Morgan radius 2,**1024** bits(不是 2048)。 +- `BulkTanimotoSimilarity(fps[i], fps[:i])`,`dists.extend(1-x)`。 +- `Butina.ClusterData(..., cutoff, isDistData=True)`。 +- 簇按大小降序,再按 scaffold 同一套 cutoff 贪心填 train/val/test。 +- 文档:**$O(n^2)$**,主要为得到 novel chemotypes;默认 cutoff **0.6(距离)**。 + +Datamol 教程同一模式:`BulkTanimotoSimilarity(..., returnDistance=True)` + `ClusterData`。 + +**和原文差别:** + +- 指纹:Morgan ≠ Daylight。 +- 阈值:DeepChem 0.6 **距离** ≈ $T \ge 0.4$,比「相似度 0.65」松得多。 +- 无稀疏化。 + +**本库:** 可对标「split 时整簇分配」;不要对标其矩阵实现。默认阈值不要盲目抄 0.6 距离,除非文档写死语义。 + +--- + +### 3.4 Chalcedon(Rowan, 2026)— 当前最贴近「精确 + 全集」的开源包 + +**链接:** + +- 博客:https://www.rowansci.com/blog/chalcedon (Eli Mann, 2026-05-26) +- 代码:https://github.com/rowansci/chalcedon +- PyPI:`chalcedon` + +**动机:** 按 Walters 建议做 Butina split,但 GEOM ~30 万样本时现有开源实现要 **数 TB 内存**。 + +**仍是精确 Butina:** 分块实现 vs RDKit,在 10k–100k、cutoff=0.65 上他们要求 **簇相同**(bitwise 不完全相同,见浮点)。 + +**算法改写(内存线性的关键):** + +标准形式要持久化 $n\times n$ 距离。Chalcedon 分成三阶段: + +1. **分块**算 pair 相似度,**只持久化每个分子的邻居个数**,丢掉具体相似度。 +2. 按邻居数降序排序。 +3. 再按排序走,对 **仍未分配** 的集合 **分块重算** 该中心的相似度行,收走未分配邻居。 + +峰值 ≈ 一块 batch workspace + $O(n)$ 计数。 + +**其它工程优化:** + +- 全程 float32(sgemm 约 2× 于 dgemm);非二进制描述子他们建议 float64。 +- Cutoff 比较改写成 $|A\cap B| \ge (1-\mathrm{cutoff})\cdot |A\cup B|$,少中间数组、略减 ULP。 +- 预计算每行 $\|A\|^2$(对 0/1 向量即 popcount)。 +- 只走上三角;对角块递归切成矩形 GEMM(类似 SciPy ssyrk 思路)。 +- 预分配 buffer,避免反复 malloc。 + +**Split 部分:** 簇出来后用 Graham **LPT(最长处理时间)** 贪心:每次把最大簇分给「离目标比例最远」的那一份(train/val/test)。与 DeepChem「先填满 train cutoff」不完全相同,但同属整簇分配。 + +**基准(GEOM 子集,cutoff=0.65,Ryzen 9 7950X3D,128GB RAM):** + +Wall time (s): + +| n | Chalcedon chunked f32 | Chalcedon full matrix f32 | RDKit | +|---|----------------------|---------------------------|--------| +| 1,000 | 0.040 | 0.012 | 0.073 | +| 10,000 | 0.357 | 0.281 | 5.16 | +| 50,000 | 5.75 | 5.74 | 172.8 | +| 100,000 | **21.14** | 22.49 | **OOM** | + +Peak RSS (GB): + +| n | Chalcedon chunked f32 | Chalcedon full matrix f32 | RDKit | +|---|----------------------|---------------------------|--------| +| 10,000 | 0.33 | 0.35 | 4.5 | +| 50,000 | 1.23 | 3.16 | 110.4 | +| 100,000 | **2.35** | 11.17 | OOM | + +小 $n$ 全矩阵更快;大 $n$ 分块内存近线性,全矩阵仍 $O(n^2)$。 + +**准确性注意(对「i need accuracy」很重要):** +他们在 cutoff sweep 上发现 float32 重排比较 vs RDKit **不是逐 bit 相同**,部分 cutoff 簇数差 $\le 2.5\%$。原因是浮点,不是故意漏边。 + +要对齐 RDKit / 原文布尔 $T\ge t$:应用 **整数 popcount** 算 $c,a,b$ 再比较,不要用 float32 GEMM 当金标准。 + +**API 示例:** + +```python +splits = chalcedon.butina_split( + smiles, + fractions={"train": 0.8, "val": 0.1, "test": 0.1}, + cutoff=0.65, # 距离 cutoff + dtype="float32", +) +``` + +--- + +### 3.5 chemfp(工业精确稀疏 Butina) + +**文档:** https://chemfp.com/docs/chemfp_butina_command.html + +**做法:** + +1. `threshold_tanimoto_search_symmetric` 生成 **稀疏** 相似度矩阵(只含 $T \ge t_{\mathrm{NxN}}$)。 +2. 可存 `.npz`;用更高的 `--butina-threshold` 调参时 **不必重算 Tanimoto**。 +3. 按行邻居数排序 + exclusion sphere。 +4. `--tiebreaker randomize|first|last`(默认可随机,**可复现必须 first/last + seed**)。 +5. 可选:false singleton 贴最近中心(**原文没有**,属于后处理)。 + +**优化本质:** 与「邻接表精确 Butina」相同,搜索层用 POPCNT 等。 +免费/商业版性能差一截;加依赖前要看许可。 + +--- + +### 3.6 FPSim2(ChEMBL)— 精确搜索后端,不是聚类 + +**文档:** https://chembl.github.io/FPSim2/ + +- CPU POPCNT;高阈值($\ge 0.7$)更合适。 +- **Swamidass & Baldi 2007** bound(doi: [10.1021/ci600358f](https://doi.org/10.1021/ci600358f)): + $T \le \min(a,b)/\max(a,b)$。query 亮位数 $A$、库分子 $B$ 必须落在 $[tA,\ A/t]$,否则 **不可能** 是邻居 → **剪枝不漏真阳性**。 +- 多核、可选 GPU;`symmetric_distance_matrix(threshold=...)` → SciPy CSR。 + +可放在 Butina 流水线的「建邻居图」一步;exclusion 仍要自己写(或交给 chemfp/Chalcedon)。 + +--- + +### 3.7 BitBIRCH — 不要当成优化版 Butina + +Miranda-Quintana 等, *Efficient clustering of large molecular libraries*, bioRxiv 2024. +https://www.biorxiv.org/content/10.1101/2024.08.10.607459v1 + +- 相对 RDKit Taylor–Butina:45 万分子,4TB 仍不够跑矩阵版;BitBIRCH 约 2 分钟。 +- 150 万分子号称 >1000×。 +- 质量指标在部分阈值区间与 Butina「无显著差」或更好。 +- Chalcedon 对比:**同名义阈值下簇数完全不是一回事**($n=10^5$:Chalcedon 9590 vs BitBIRCH-Lean 51232)。 + +这是 **另一种算法**(BIRCH + iSIM)。更快,但不能命名为 `method="butina"`。 + +--- + +## 4. 两种「精确」建邻居方式(和 Chalcedon 的关系) + +```text +A. 全距离矩阵 + ClusterData + 小 n、易与 RDKit/DeepChem 逐簇对齐 + QM9 全量:内存不可行 + +B. 阈值邻接 / 分块邻居 + 同一套排序收球 + 边集完整 ⇒ 与 A 数学等价 + Chalcedon:分块 + 两遍(先计数再收球),线性内存 + chemfp/FPSim2:稀疏阈值搜索(POPCNT ± bound) +``` + +Chalcedon 的巧妙处:第一遍 **甚至可以不存邻接表**,只存度数;第二遍对未分配子集重算。 +代价是部分 pair 会算两次;换来内存 $O(n)$ 而不是 $O(n\bar{k})$ 邻接表。 +chemfp 选择 **存稀疏边**,调阈值更快。都精确,工程权衡不同。 + +--- + +## 5. 对 torch-molecule 的建议 + +1. **小 $n$(测试):** RDKit `BulkTanimotoSimilarity` + `ClusterData`,作为簇相等的金标准。 +2. **大 $n$ / 用户可能丢全集:** Chalcedon 式分块 **或** 稀疏邻接表 + 自写 exclusion;**禁止** 分配 $n(n-1)/2$ 距离数组。 +3. **要比准 RDKit:** 整数 popcount Tanimoto,不要默认 float32 GEMM。 +4. **阈值 API:** 对外只用 `similarity_cutoff`(如 0.65),内部再转距离;文档写清 DeepChem 0.6 是距离。 +5. **指纹默认:** Morgan r=2, 2048 bit(Chalcedon / Walters);若要对齐 DeepChem 再提供 1024。 +6. **平局:** 稳定 `(-度数, index)`,相当于 chemfp `first`,不要默认 randomize。 +7. **不要** 在 `method="butina"` 里偷偷 subsample 或换 BitBIRCH。 +8. 依赖:能零依赖实现稀疏/分块最好;Chalcedon MIT 且专做 split,可作实现参照,不必强绑。 + +--- + +## 6. 参考文献与链接 + +| 主题 | 来源 | +|------|------| +| 原文 | Butina, JCICS 1999, doi:10.1021/ci9803381 | +| 精确剪枝 bound | Swamidass & Baldi, JCIM 2007, doi:10.1021/ci600358f | +| DeepChem splitter | https://github.com/deepchem/deepchem/blob/master/deepchem/splits/splitters.py | +| Chalcedon | https://www.rowansci.com/blog/chalcedon ;https://github.com/rowansci/chalcedon | +| chemfp Butina | https://chemfp.com/docs/chemfp_butina_command.html | +| FPSim2 | https://chembl.github.io/FPSim2/ | +| BitBIRCH | https://www.biorxiv.org/content/10.1101/2024.08.10.607459v1 | +| 大库聚类实践 | https://macinchem.org/2023/03/05/options-for-clustering-large-datasets-of-molecules/ | +| Walters split 评论 | http://practicalcheminformatics.blogspot.com/2024/11/some-thoughts-on-splitting-chemical.html | + +--- + +*记录日期:对照 Chalcedon 2026-05 博客、chemfp 4.x 文档、DeepChem 源码与 FPSim2/BitBIRCH 文献。待 `method="butina"` 实现时以第 5 节为准。* diff --git a/molecule_notes/generation_model.md b/molecule_notes/generation_model.md new file mode 100644 index 0000000..b742987 --- /dev/null +++ b/molecule_notes/generation_model.md @@ -0,0 +1,473 @@ +# Report: HF pretrained molecular generator + +本文记录 `torch-molecule` 接入 Hugging Face 预训练分子生成模型的工作:对外一个类 `HFPretrainedMolecularGenerator`,sklearn 风格 `fit` / `generate`,与 `HFPretrainedMolecularEncoder` 对称。依据仓库内的 `generation_model.md`、`issues_to_be_fixed.md` 以及 `torch_molecule/generator/pretrained/` 的现有代码;不把未实现的接口写成已完成。 + +--- + +## 1. Purpose + +现有 `HFPretrainedMolecularEncoder` 只用 Hugging Face `AutoModel` 做编码,不做生成。仓库里已有 8 个自研生成器(`LSTMMolecularGenerator`、`MolGPTMolecularGenerator`、`DigressMolecularGenerator`、`GDSSMolecularGenerator`、`GraphDITMolecularGenerator`、`DeFoGMolecularGenerator`、`JTVAEMolecularGenerator`、`GraphGAMolecularGenerator`),它们**没有被替换**。 + +这次新增的是第三条路径:从 Hugging Face Hub 加载已预训练的生成权重,用**一个**对外类 `HFPretrainedMolecularGenerator` 提供与 LSTM / MolGPT 相同的 sklearn 风格接口: + +- `fit()` 无数据:下载并加载 Hub 权重 +- `fit(smiles)`:加载预训练权重后再微调 +- `generate()`:返回 `List[str]` SMILES + +六个 Hub 模型的架构不同(因果 LM、seq2seq SELFIES、Fragment-SELFIES + 官方推理引擎),无法共用同一个 `generate()`。家族分流(`novomolgen` / `gp_molformer` / `molgen` / `molexar` / fallback `causal_lm`)只在 `generator/pretrained/` 内部完成。用户不直接实例化 `families/` 里的类;那些模块是内部 dispatch,不是对外 API。 + +数据集层(`load_qm9()`、`load_zinc250k()`、`SMILESDataset`、`MolecularInputChecker`)没有改。用户始终传入 SMILES、始终拿到 SMILES。SELFIES / Fragment-SELFIES 转换只发生在 `generator/pretrained/` 内部。 + +--- + +## 2. What we added (feature work) + +对应 `generation_model.md` 的 Phase 1–5。**Phase 6(文档站点 API 页、README 模型列表、CI 分层)仍未完成**,见第 3 节末尾。 + +### 2.1 对外 API + +```python +from torch_molecule import HFPretrainedMolecularGenerator +``` + +导出链:`torch_molecule/generator/pretrained/__init__.py` → `torch_molecule/__init__.py`(已加入 `__all__`)。 + +与 encoder 的对称关系: + +| | `HFPretrainedMolecularEncoder` | `HFPretrainedMolecularGenerator` | +|---|---|---| +| `fit()` 无参 | 加载 `AutoModel` | 加载 `AutoModelForCausalLM` / `AutoModelForSeq2SeqLM`,或 Molexar 官方 `MolexarInference` | +| `fit(X)` | 不支持 | 可选微调 | +| 主方法 | `encode()` → embedding | `generate()` → SMILES | + +**`fit()` 语义与 LSTM 不同:** + +- `LSTMMolecularGenerator.fit(X_train)`:从零初始化网络并在 SMILES 上训练。 +- `HFPretrainedMolecularGenerator.fit()`:无 `X` 时只从 Hub 拉预训练权重;`fit(smiles)` 是加载权重后再微调,不是从零训练。 + +`generate(n_samples=...)` 返回 `List[str]` SMILES。MolGen / Molexar 解码失败的条目会被丢掉,返回长度可以小于 `n_samples`(见问题 2、7)。 + +条件微调参数 `y` 目前会 warn 后忽略,走无条件语言模型微调(见「仍未做」)。 + +### 2.2 用户侧始终是 SMILES + +| 方向 | 约定 | +|---|---| +| `fit(X)` 输入 | `List[str]` SMILES | +| `generate()` 输出 | `List[str]` SMILES | +| SELFIES / Fragment-SELFIES | 仅内部:`_encode_inputs()` / `_decode_outputs()` | +| 数据集 loader | **未改** | + +内部转换(`modeling_pretrained.py`): + +- MolGen(seq2seq):`smiles_to_selfies` / `selfies_to_smiles` +- Molexar:`smiles_to_fragment_selfies` / `fragment_selfies_to_smiles` +- NovoMolGen / GP-MoLFormer / fallback causal LM:直接用 SMILES(生成后去掉空格) + +### 2.3 六个 Hub 模型与内部家族 + +| 模型 | `repo_id` | 内部 family | 表示 | 加载 / 生成要点 | +|---|---|---|---|---| +| NovoMolGen | `chandar-lab/NovoMolGen_32M_SMILES_BPE` | `novomolgen` | SMILES + BPE | 因果 LM;`revision` 默认 `hf-checkpoint` | +| GP-MoLFormer | `ibm-research/GP-MoLFormer-Uniq` | `gp_molformer` | SMILES | 因果 LM;`scaffold=` 骨架补全;tokenizer 默认 `ibm-research/MoLFormer-XL-both-10pct`;IBM remote code 需要 `transformers<=4.56.2` | +| MolGen-large | `zjunlp/MolGen-large` | `molgen` | SELFIES | seq2seq;默认苯环 prefix | +| MolGen-large-opt | `zjunlp/MolGen-large-opt` | `molgen` | SELFIES | 同上,权重已偏 QED / p-logP | +| Molexar-10M-base | `fairydance/molexar-10m-base` | `molexar` | Fragment-SELFIES | 官方 `MolexarInference` | +| Molexar-10M-omni | `fairydance/molexar-10m-omni` | `molexar` | Fragment-SELFIES | 同引擎;`conditions` / 性质键等 omni kwargs | + +未知 `repo_id`:`resolve_family()` 返回 fallback `causal_lm` 并 `warnings.warn`(问题 6 修过后,前缀能匹配的变体如 `chandar-lab/NovoMolGen_157M` **不再**误报)。 + +`FAMILY_PREFIXES` 允许同一前缀下的变体(例如其它 NovoMolGen 尺寸)映射到对应家族,而不必全部写进 `KNOWN_REPOS`。 + +### 2.4 分阶段(Phase 1–5 已落地) + +**Phase 1 — 骨架 + NovoMolGen(MVP)** + +- `registry.py`、`modeling_pretrained.py`、`families/causal_lm.py` +- `__init__.py` 导出 +- `tests/generator/hfpretrained.py` smoke(含 `@pytest.mark.integration` 的 NovoMolGen `fit()` + `generate()`) +- optional extra:`[hf-gen]` + +**Phase 2 — GP-MoLFormer + SELFIES 工具** + +- `utils.py`:`smiles_to_selfies` / `selfies_to_smiles` +- `families/causal_lm.py`:de novo + `scaffold=` +- `compat.py`:`transformers>=4.57` 时对 GP-MoLFormer raise;为 IBM remote code 提供 `transformers.onnx` stub +- extra:`[gp-molformer]`(`transformers>=4.40,<=4.56.2`) + +**Phase 3 — MolGen-large / MolGen-large-opt** + +- `families/seq2seq.py` +- `generate(prefix_selfies=...)`;未提供时使用默认苯环 SELFIES(见「仍未做」) +- extra:`[molgen]`(`selfies>=2.1.0`,无上界) + +**Phase 4 — Molexar** + +- `utils.py`:Fragment-SELFIES 转换 +- `families/molexar.py`:wrap 官方 `MolexarInference`(de novo、`start_smiles` / fragment 约束、omni `conditions`) +- extra:`[molexar]` + +**Phase 5 — 微调与本地存盘** + +- `finetune.py`:按家族 dispatch + - 因果 LM(NovoMolGen / GP-MoLFormer / fallback):next-token prediction(`finetune_causal_lm`) + - MolGen seq2seq:denoising(`finetune_seq2seq`,问题 9 之后:mask encoder、labels 保持干净) + - Molexar:官方 training template 后再走因果 LM loop(`finetune_molexar`) +- `checkpoint.py` + `save_to_local` / `load_from_local`:HF 目录(`save_pretrained`)+ `hf_generator_metadata.json` +- `save_to_hf()` **仍是** `NotImplementedError` +- `load_from_hf()` / `load()` 无本地 path 时等价于 `fit()` + +**Phase 6 — 仍开放**(见第 3 节)。 + +### 2.5 目录布局 + +``` +torch_molecule/generator/pretrained/ +├── __init__.py # 只导出 HFPretrainedMolecularGenerator +├── modeling_pretrained.py # 唯一对外类 +├── registry.py # repo_id → family +├── utils.py # SMILES ↔ SELFIES / Fragment-SELFIES +├── finetune.py # 家族微调 +├── checkpoint.py # 本地 metadata +├── compat.py # GP-MoLFormer / transformers 兼容 +└── families/ + ├── __init__.py + ├── causal_lm.py # NovoMolGen, GP-MoLFormer, fallback + ├── seq2seq.py # MolGen-large, MolGen-large-opt + └── molexar.py # Molexar base / omni +``` + +### 2.6 Optional extras(`pyproject.toml`) + +| extra | 依赖 | +|---|---| +| `[hf-gen]` | `transformers>=4.40`, `accelerate` | +| `[molgen]` | `selfies>=2.1.0`, `transformers>=4.40`, `accelerate` | +| `[gp-molformer]` | `transformers>=4.40,<=4.56.2`, `accelerate` | +| `[molexar]` | `fragment-selfies>=1.0.0`, `transformers>=4.40`, `accelerate`, `loguru`, `molexar @ git+https://github.com/fairydance/Molexar.git` | + +`selfies` **没有**在 `pyproject.toml` 里 pin `<3`(老师要求;问题 5 用 README / `install.rst` 说明)。 + +### 2.7 存盘 + +- `save_to_local(path)`:`model.save_pretrained` + `tokenizer.save_pretrained` + `hf_generator_metadata.json` +- `load_from_local(path)`:读 metadata,再按家族从该目录加载 +- `save_to_hf(...)`:`NotImplementedError`(「HFPretrainedMolecularGenerator does not support saving to Hugging Face.」) + +注意:`fit()` **每次**都会 `_load_pretrained()` 从 Hub 再拉一遍。即使刚 `load_from_local`,再调用 `fit()` 仍会覆盖为 Hub 权重。老师要求这一轮先不动(见「仍未做」)。 + +### 2.8 微调:MolGen 是 denoising,不是 identity copy + +问题 9 修完后,仅 `finetune_seq2seq` 改变目标: + +- encoder `input_ids`:非 special token 以 `mask_prob=0.15` 换成 `` +- `labels`:仍是干净序列 +- `fit(X)` 签名不变 +- **不影响** NovoMolGen、GP-MoLFormer、Molexar、LSTM + +### 2.9 用法片段(NovoMolGen) + +```python +from torch_molecule import HFPretrainedMolecularGenerator + +model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", +) +model.fit() # 只加载 Hub 预训练权重,不训练 +smiles_list = model.generate(n_samples=10) + +# 微调:仍传 SMILES(与 LSTM 不同:这里是在预训练权重上继续训) +model.fit(["CCO", "CC(=O)O", "c1ccccc1"]) +``` + +--- + +## 3. Issues we found and how we handled them + +依据 `issues_to_be_fixed.md`。原则:进模型(训练数据)坏样本默认 **raise**,失败要带 index;出模型(生成)坏样本可以丢掉,但 `n_samples` 是**尝试条数**,返回长度 = 解码成功条数,并留下计数警告。 + +问题 2(MolGen 用 `""` 撑长度)和问题 7(Molexar 静默变短)**原 bug 不同**;修完后共用同一条返回约定:不要补 `""`,成功几条返回几条。 + +### 3.1 问题 1 — `fit()` 坏 SMILES:keep raise,不 skip + +**原问题:** `fit(X)` 里一条坏样本会让整次微调失败。讨论过是否 skip / 重试。 + +**决策 / 修法:** **不加重试、不加 skip。** 坏样本就 raise,与 `MolecularInputChecker` 一致。 + +- RDKit 过不了:`MolecularInputChecker` 给出带 index 的错误(`Invalid SMILES structure at index {idx}: ...`),`_validate_inputs` 聚合成 `ValueError`。`smiles_to_selfies` 里同样对 `MolFromSmiles is None` raise `ValueError(f"Invalid SMILES at index {idx}: ...")`。 +- RDKit 过了、`selfies.encoder` 挂:见问题 4,同样 raise,不 skip。 + +**状态:solved。** + +**关键文件:** `torch_molecule/utils/checker.py`、`torch_molecule/generator/pretrained/utils.py`、`modeling_pretrained.py`(`fit` → `_validate_inputs`)、`tests/generator/hfpretrained.py`(`test_smiles_to_selfies_invalid_smiles`)。 + +### 3.2 问题 2 — MolGen `generate` 用 `""` 占位,`len == n_samples` 看起来成功 + +**原问题(仅 MolGen 这条路径):** 解码失败时 `selfies_to_smiles` 写入空串。`len(out) == n_samples` 像成功,空串比短列表更坑。 + +这与问题 7 **不是同一个原 bug**:这里是用 `""` **把长度撑满**;问题 7 是 Molexar **直接丢掉且不警告**。 + +**决策 / 修法:** 不要 `""`,也不要重试凑满。成功几条就返回几条(2 条有效 → 返回 2)。`n_samples` = 尝试条数。 + +例:`generate(n_samples=5)`,其中 2 条解不出: + +```text +原来: ["c1ccccc1", "CCO", "", "CC(=O)O", ""] # len=5,两个假成功 +现在: ["c1ccccc1", "CCO", "CC(=O)O"] # len=3,只留真分子 +``` + +实现上 `_decode_outputs` 对 seq2seq 丢掉无效 SMILES,并 warn `got {k}/{n} valid SMILES`。 + +**状态:solved。** + +**关键文件:** `utils.py`(`selfies_to_smiles`)、`modeling_pretrained.py`(`_decode_outputs`)、`tests/generator/hfpretrained.py`(`test_selfies_to_smiles_drops_invalid_entries`、`test_decode_outputs_molgen_drops_empty_and_warns`)。 + +### 3.3 问题 3 — `except Exception` 把 `ImportError` 吞成 `""` + +**原问题:** `selfies_to_smiles` 和 `fragment_selfies_to_smiles` 都是 `except Exception: append("")`。依赖缺失、键盘中断等也会变成空串,且没有计数、没有日志。 + +**决策 / 修法:** **不要** `return_stats` 新 API。删掉 `except Exception`。分两种: + +- 不该发生的错(`ImportError`、其它意外)→ **raise**,不要吞成 `""` +- 分子解不出来(`DecoderError`、RDKit `mol is None`、空串)→ **不要 raise**(否则问题 2 无法返回成功的 2/3 条)。丢掉这条,**warn**(例如 `dropped {n} invalid SELFIES` / `dropped {n} invalid Fragment-SELFIES`) + +**状态:solved。** + +**关键文件:** `torch_molecule/generator/pretrained/utils.py`。 + +### 3.4 问题 4 — RDKit-valid 但 `selfies.encoder` 失败(dummy `*`) + +**原问题:** `smiles_to_selfies` 在 canonical 之后直接 `sf.encoder(canonical)`,没有接 `EncoderError`。`fit` 带着库异常全挂,看不出是第几条。dummy atom(`*CCO`、`[*]CCO`、`[*]c1ccccc1`)是典型触发点。 + +**决策 / 修法:** **不加 skip。** 接住后 **raise** 成带 index 的 `ValueError`: + +```python +raise ValueError( + f"SMILES at index {idx} is RDKit-valid but not SELFIES-encodable: {smiles_string}" +) from exc +``` + +**状态:solved。** + +**关键文件:** `utils.py`(`smiles_to_selfies`)、`tests/generator/hfpretrained.py`(`test_smiles_to_selfies_encoder_error_has_index`)。 + +### 3.5 问题 5 — `selfies>=2.1` 无上界;老师不要求 pyproject pin `<3` + +**原问题:** `pyproject.toml` 是 `selfies>=2.1.0`,没有上界。3.x 字母表若破坏兼容,安装器不会挡。 + +**决策 / 修法:** **不改 pyproject 的版本上界。** 老师意见:放到 README **Additional Packages**(以及 `docs/source/install.rst`),和其他可选依赖一样写明: + +| Model | Required Packages | +|---|---| +| HFPretrainedMolecularEncoder | transformers | +| HFPretrainedMolecularGenerator | transformers | +| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed) | +| HFPretrainedMolecularGenerator (GP-MoLFormer) | transformers<=4.56.2 | +| HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | + +README 另有安装示例:`pip install "selfies>=2.1"`、`pip install torch-molecule[molgen]`、`[gp-molformer]`、`[molexar]`。 + +`[gp-molformer]` extra 在 pyproject 里 **有** `transformers<=4.56.2`;MolGen 的 `selfies` 仍只有 `>=2.1.0`。运行时 GP-MoLFormer 还会在 `compat.ensure_gp_molformer_transformers_compat()` 对 `>=4.57` raise `ImportError`。 + +**状态:solved(文档,不是 pyproject 给 selfies 加上界)。** + +**关键文件:** `README.md`、`docs/source/install.rst`、`pyproject.toml`、`compat.py`。 + +### 3.6 问题 6 — 未知 `repo_id` 警告只查精确 `KNOWN_REPOS` + +**原问题:** `__init__` 只查 `KNOWN_REPOS` 的完整字符串。`chandar-lab/NovoMolGen_157M` 会被 `FAMILY_PREFIXES` 正确分成 `novomolgen`,却仍警告 “family may not be implemented”。 + +**决策 / 修法:** 在 `resolve_family()` 之后,**只对 fallback `causal_lm` 警告**。 + +**状态:solved。** + +**关键文件:** `modeling_pretrained.py`(`__init__`)、`registry.py`(`resolve_family`)、`tests/generator/hfpretrained.py`(`test_known_family_prefix_does_not_warn_unknown_repo`、`test_unknown_repo_fallback_warns`)。 + +### 3.7 问题 7 — Molexar `generate()` 列表变短且无警告 + +**原问题(Molexar,与问题 2 不同):** `_decode_outputs` 对 Molexar `if smiles` 过滤,解码失败直接丢掉。调用方要 10 条可能拿到 6 条,**没有警告**。 + +**决策 / 修法:** 要 10 条、只有 6 条成功,就用这 6 条,**不要补 `""` 凑满**。缺的是警告,例如 `got 6/10 valid SMILES`。 + +修完后与问题 2 **共用返回约定**(成功几条返回几条 + 计数 warn),但原缺陷分别是「假满长度」vs「静默变短」。 + +**状态:solved。** + +**关键文件:** `modeling_pretrained.py`(`_decode_outputs`,seq2seq 与 molexar 共用 warn)、`utils.py`(`fragment_selfies_to_smiles`)。 + +### 3.8 问题 8 — 边角化学:confirmed,不为转换层加特殊 case + +**原问题:** 担心电荷、立体、dummy `*`、叠氮等边角 SMILES 需要单独一套转换逻辑。 + +**结论(已扫 12 大类、约 45 条 SMILES,`selfies` 2.1.1):** + +| 结果 | 条数 | +|---|---| +| OK | 40 | +| RDKit 拒 | 1 | +| `EncoderError` | 4(dummy `*`;`[*]CCO` 在 dummy 和 Molexar 挂点里各计一次) | +| decode 挂 | 0 | + +会 roundtrip 的不必为转换层加特殊处理,包括:电荷、两性离子、四面体/顺反/联烯、萘/桥环/螺环/大环/三元环、吡啶、Kekulé 苯、三键、过氧、高价 S/P、N-oxide、硝基、自由基、卡宾、显式氢、有机叠氮 `CCN=[N+]=[N-]`、重氮、同位素、`.` 断开、Si/Se/膦/硼酸根、`[Cu+2]` / `[Fe]` / 类格氏。 + +会失败的走问题 1 / 4(进模型 raise),不是新分支: + +| SMILES | 卡在哪 | 行为 | +|---|---|---| +| `C[N-]=[N+]=N` | RDKit 不认(价态) | **raise** 问题 1。有机叠氮请用 `CCN=[N+]=[N-]` | +| `*CCO` | RDKit 能吃,`selfies.encoder` 挂 | **raise** 问题 4 | +| `[*]CCO` | 同上 | **raise** 问题 4 | +| `[*]c1ccccc1` | 同上(Molexar 挂点写法) | **raise** 问题 4 | + +只有 **MolGen** 走 SMILES ↔ SELFIES,才会在 `*` 上撞问题 4。NovoMolGen / GP-MoLFormer / LSTM 不走这条转换。 + +**决策:** 转换层不加边角化学 special case。失败路径已由问题 1、4 覆盖。`issues_to_be_fixed.md` 提到可选把 3 个 `*` 和 1 个坏叠氮加进回归测试;当前 `tests/generator/` **没有**这组 SMILES 作为独立回归用例。 + +**状态:confirmed;no extra conversion cases。** + +**关键文件:** `utils.py`(通用 raise / drop,无化学分类表)。 + +### 3.9 问题 9 — MolGen 微调曾是 identity reconstruction + +**原问题:** MolGen 预训练是 **denoising seq2seq**(损坏的 SELFIES → 还原完整 SELFIES)。实现曾把 `labels = input_ids`,等于 identity copy,和论文目标不一致。 + +**决策 / 修法:** **只改** `finetune_seq2seq`: + +- `corrupt_token_ids(..., mask_prob=0.15)`:非 special token 换成 tokenizer 的 `` +- `labels` 仍是干净序列 +- `fit(X)` 签名不变 +- NovoMolGen / GP-MoLFormer / Molexar / LSTM 微调不变 + +**状态:solved(fixed)。** + +**关键文件:** `torch_molecule/generator/pretrained/finetune.py`(`corrupt_token_ids`、`finetune_seq2seq`)、`tests/generator/test_finetune.py`(`test_corrupt_token_ids_masks_non_special_tokens`、`test_finetune_seq2seq_labels_stay_clean`)。 + +### 3.10 仍未做 + +`issues_to_be_fixed.md` 写明这一轮不做,以及 `generation_model.md` Phase 6 仍是未勾选: + +1. **`fit()` 在 `load_from_local` 之后仍会从 Hub 再加载。** `fit()` 无条件调用 `_load_pretrained()`(无 `local_path` 时用 `self.repo_id`)。老师要求先不动。 +2. **MolGen 默认苯环 prefix。** `families/seq2seq.py` 中 `DEFAULT_MOLGEN_PREFIX_SELFIES = "[C][=C][C][=C][C][=C][Ring1][=Branch1]"`;未传 `prefix_selfies` / `scaffold` 时仍用它。这一轮明确不改默认。 +3. **Phase 6 文档与 CI** + - `docs/source/api/generator.rst` **尚未**收录 `HFPretrainedMolecularGenerator`(仍只有 8 个自研生成器)。 + - README「List of Supported Models → Generative Models」仍是 8 个自研模型,没有把 6 个 HF repo 列进去(Additional Packages 表已有 generator extras,那是问题 5,不是 Phase 6 的模型列表)。 + - CI:`pyproject.toml` 已声明 pytest marker `integration`,测试里 Phase 1/2/3/4 的 Hub 下载用例标了 `@pytest.mark.integration`;仓库 `.github/workflows/` 目前只有 `docs.yml`,**没有**「Phase 1 必跑、Phase 3+ 用 optional marker」的测试 CI。 +4. **`save_to_hf()` 对本类未实现**(`NotImplementedError`)。不要把它写成可用。 +5. **带 `y` 的条件微调未实现。** `fit(X, y)` 若 `y is not None` 会 `UserWarning`(「Conditional fine-tuning with y is not implemented yet」),然后 `y = None`,继续无条件 LM 微调。 + +--- + +## 4. Current data flow (SELFIES models) + +用户侧不变:`load_zinc250k().data` 仍是 SMILES;`fit` / `generate` 仍收发 SMILES。 + +``` +用户 SMILES + │ + ▼ +HFPretrainedMolecularGenerator.fit / generate + │ + ├─ MolGen fit + │ SMILES → _validate_inputs (RDKit) + │ → smiles_to_selfies # 坏样本 raise(问题 1 / 4) + │ → tokenize + │ → finetune_seq2seq # denoising:mask encoder,labels 干净 + │ + ├─ MolGen generate + │ prefix SELFIES(默认苯环,或 prefix_selfies= / scaffold=) + │ → seq2seq model.generate + │ → selfies_to_smiles # 解不出:warn + drop,不补 "" + │ → 可能再 warn got k/n valid SMILES + │ + ├─ Molexar + │ fit: SMILES → smiles_to_fragment_selfies → 官方 training template → causal LM loop + │ generate: MolexarInference → fragment_selfies_to_smiles → drop invalids + warn + │ + └─ Causal LM(NovoMolGen / GP-MoLFormer / fallback) + SMILES 直接 tokenize / 生成;不走 smiles_to_selfies / + selfies_to_smiles / smiles_to_fragment_selfies / fragment_selfies_to_smiles + NovoMolGen: BOS → generate;默认 revision hf-checkpoint + GP-MoLFormer: de novo 或 scaffold= 前缀 +``` + +--- + +## 5. Files touched (high level) + +**新增(生成器核心)** + +- `torch_molecule/generator/pretrained/modeling_pretrained.py` +- `torch_molecule/generator/pretrained/registry.py` +- `torch_molecule/generator/pretrained/utils.py` +- `torch_molecule/generator/pretrained/finetune.py` +- `torch_molecule/generator/pretrained/checkpoint.py` +- `torch_molecule/generator/pretrained/compat.py` +- `torch_molecule/generator/pretrained/__init__.py` +- `torch_molecule/generator/pretrained/families/causal_lm.py` +- `torch_molecule/generator/pretrained/families/seq2seq.py` +- `torch_molecule/generator/pretrained/families/molexar.py` +- `torch_molecule/generator/pretrained/families/__init__.py` + +**导出与依赖** + +- `torch_molecule/__init__.py`(加入 `HFPretrainedMolecularGenerator`) +- `pyproject.toml`(`[hf-gen]` / `[molgen]` / `[gp-molformer]` / `[molexar]`,以及 pytest `integration` marker) + +**测试** + +- `tests/generator/hfpretrained.py` +- `tests/generator/test_finetune.py` +- `tests/generator/test_causal_lm.py` +- `tests/generator/test_seq2seq.py` +- `tests/generator/test_molexar.py` + +**文档(问题 5;不是 Phase 6 API 页)** + +- `README.md` Additional Packages +- `docs/source/install.rst` Additional Packages + +**未改(按设计)** + +- `torch_molecule/datasets/*`(loader、CSV) +- `torch_molecule/encoder/pretrained/*` +- 8 个自研生成器实现 + +**Phase 6 仍未改** + +- `docs/source/api/generator.rst` +- README 生成模型一览表(仍 8 个自研) + +--- + +## 6. How to try it + +```bash +pip install -e ".[hf-gen]" +``` + +按模型再装 extras: + +```bash +pip install -e ".[molgen]" # MolGen:selfies 2.x;3.x 不保证 +pip install -e ".[gp-molformer]" # pins transformers<=4.56.2 +pip install -e ".[molexar]" # fragment-selfies + molexar +``` + +NovoMolGen 推理: + +```python +from torch_molecule import HFPretrainedMolecularGenerator + +model = HFPretrainedMolecularGenerator( + repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", +) +model.fit() +print(model.generate(n_samples=5)) +``` + +需要下载 Hub 权重的测试标了 `@pytest.mark.integration`,例如: + +```bash +pytest tests/generator/hfpretrained.py -m "not integration" +pytest tests/generator/test_finetune.py +``` diff --git a/molecule_notes/train_test_split.md b/molecule_notes/train_test_split.md new file mode 100644 index 0000000..bd18a10 --- /dev/null +++ b/molecule_notes/train_test_split.md @@ -0,0 +1,385 @@ +# Train / Test Split 实现方案 + +本文档记录 `torch-molecule` 中分子数据划分(train/test split)的设计与实现计划。 +背景:分子数据不能简单使用 sklearn 的 `train_test_split`,需支持 structure-aware 划分。 + +--- + +## 1. 设计原则 + +1. **Split 是数据工具,不是模型工具** — 不要写进 `BaseMolecularPredictor.fit()`。 +2. **返回值仍是** `SMILESDataset` — 与现有 `load_qm9()` → `fit()` 流程对接。 +3. `method` 支持 `random`、`scaffold`、`butina`、`size`。三路划分见 3.3。 +4. **SizeShiftReg 的 coarsening/CMD 不属于 split** — 那是 SSR 模型的训练正则,split 只按原子数切分。 +5. **默认** `method="random"` — 保证可复现、与旧脚本一致;文档说明 random 分数往往偏乐观。 + +--- + +## 2. 文件结构 + +``` +torch_molecule/datasets/ + constant.py # SMILESDataset 增加 train_test_split / subsample + split.py # 新增:各划分方法实现 + __init__.py # 导出 train_test_split, SMILESDataset + +tests/datasets/ + test_split.py # 新增单元测试 + +README.md # 更新:加载 → 划分 → fit 完整示例 +``` + +**不要修改:** `predictor/*/modeling_*.py`、`base/predictor.py`(训练接口已足够)。 + +**依赖:** RDKit(已在 `pyproject.toml`)只负责 SMILES → 指纹。无需 DeepChem、无需 chemfp。 + +**增强版 Butina:** 不是 RDKit `ClusterData` 的稠密距离矩阵。生产路径是 **稀疏阈值邻居图 + exclusion sphere**(chemfp / Chalcedon 同一精确算法,内存 $O(n+E)$)。原文、RDKit、chemfp、Chalcedon 的对照表见下方 5.3;更长的文献笔记见 `molecule_notes/butina_optimization.md`。 + +--- + +## 3. 用户 API + +### 3.1 挂在数据集上(推荐) + +```python +from torch_molecule.datasets import load_qm9 + +data = load_qm9(local_dir="torchmol_data") +# subsample 仅用于本地调试 / CI,不要为了 Butina 而缩小 QM9 +# data = data.subsample(n=5000, seed=0) + +train, val = data.train_test_split( + test_size=0.2, + method="scaffold", # "random" | "scaffold" | "butina" | "size" + seed=42, +) + +train, val = data.train_test_split( + test_size=0.2, + method="scaffold", # "random" | "scaffold" | "butina" | "size" + seed=42, +) + +predictor.fit(train.data, train.target, val.data, val.target) +``` + +### 3.2 函数式 + +```python +from torch_molecule.datasets import train_test_split + +train, test = train_test_split(data, test_size=0.2, method="random", seed=42) +``` + +### 3.3 三路划分(对标 MoleculeNet 80/10/10) + +```python +train, val, test = data.train_val_test_split( + train_size=0.8, val_size=0.1, test_size=0.1, + method="scaffold", + seed=42, +) +``` + +### 3.4 可选参数 + + +| 参数 | 适用 method | 说明 | +| ------------------- | --------------------------------------- | ----------------------------------------------------------- | +| `test_size` | 全部 | 测试集比例,默认 0.2 | +| `seed` | random;scaffold/butina 的 `_group_split` | 随机种子,默认 42。Butina **聚类本身**用 index 平局,不用 seed | +| `similarity_cutoff` | butina | Tanimoto **相似度**阈值,默认 0.65(chemfp 语义;不是 DeepChem 距离 cutoff) | +| `use_csk` | scaffold | 是否用 cyclic skeleton(全碳),默认 False | +| `direction` | size | `small_to_large`(默认)等 | + + +--- + +## 4. `split.py` 内部结构 + +```python +def train_test_split(dataset, test_size=0.2, method="random", seed=42, **kwargs): + smiles = dataset.data + y = dataset.target + n = len(smiles) + + if method == "random": + idx_train, idx_test = _random_split(n, test_size, seed) + elif method == "scaffold": + groups = _scaffold_groups(smiles, use_csk=kwargs.get("use_csk", False)) + idx_train, idx_test = _group_split(groups, test_size, seed=seed) + elif method == "butina": + groups = _butina_groups_or_oom( + smiles, + similarity_cutoff=kwargs.get("similarity_cutoff", 0.65), + ) + idx_train, idx_test = _group_split(groups, test_size, seed=seed) + elif method == "size": + idx_train, idx_test = _size_split(smiles, test_size, seed=seed, **kwargs) + else: + raise ValueError(f"Unknown split method: {method}") + + return _subset(dataset, idx_train), _subset(dataset, idx_test) +``` + +核心:**按样本切**(random, size) vs **按组切**(scaffold, butina)→ 共用 `_group_split`。 + +--- + +## 5. 各方法实现要点 + +### 5.1 Random split + +- 与 sklearn 等价:`np.random.RandomState(seed).permutation(n)`。 +- 固定 `seed` 保证可复现。 +- 支持 `target is None`(如 ZINC)。 +- **测什么泛化:** 同分布插值;分数往往偏高,仅作对照。 + +**参考:** MoleculeNet (Wu et al., 2018) — baseline split。 + +--- + +### 5.2 Scaffold split + +**步骤:** + +1. SMILES → RDKit Mol(无效则 `ValueError`,与 `MolecularInputChecker` 一致)。 +2. `MurckoScaffold.GetScaffoldForMol(mol)` → scaffold Mol。 +3. `MolToSmiles(scaffold)` 作为 group id;无环小分子用 `"_acyclic_"` 或规范 SMILES。 +4. `dict[scaffold_smiles] -> list[index]`。 +5. **组级别**分配到 train/test(同一 scaffold 不能跨集合)。 + +**组分配策略(对齐 DeepChem / scikit-fingerprints):** + +- 按组大小降序排列 scaffold 组。 +- 贪心:将组填入 train 或 test,使 test 比例接近 `test_size`(或最小组优先进 test,测稀有骨架)。 + +**测试断言:** + +- train 与 test 的 scaffold 集合 **交集为空**。 +- 每个 index 恰好出现一次。 +- 实际比例可能偏离 `test_size`(group split 正常现象,需在文档说明)。 + +**可选:** `use_csk=False`(默认,保留原子类型);`True` 时用 `MakeScaffoldGeneric`(更粗)。 + +**测什么泛化:** 未见过的 Bemis–Murcko 骨架。 + +**参考:** + +- Bemis & Murcko (1996) — 骨架定义。 +- MoleculeNet (2018) — ML 标准协议。 + +**局限:** 差一个原子可能换 scaffold,但分子仍很相似(Walters 2024);RDKit 实现与原文细节略有差异。 + +--- + +### 5.3 Butina split(增强 / 稀疏精确版) + +代码:`torch_molecule/datasets/split.py` 中 `_butina_clusters` / `_butina_neighbor_lists` / `_butina_exclusion_spheres`。 +**精确 Taylor–Butina(Butina 1999)**,不是 BitBIRCH,也不是「先 subsample 再聚类」。 + +#### 为什么叫增强 + +朴素实现(DeepChem / Datamol / 直接 `Butina.ClusterData`)要先填满 condensed 距离: + +$$ +\text{dists 长度} = n(n-1)/2,\quad \text{内存 } O(n^2) +$$ + +QM9($n \approx 133885$)这条路会 OOM(他人报告 15 万分子涨到几十~上百 GB)。 +**精确 Butina 其实只需要布尔邻居** $N(i,j)=\mathbf{1}[T(i,j)\ge t]$。增强版只存 $T \ge t$ 的边,再按原文做度数排序 + exclusion sphere。簇与小集上的 RDKit `ClusterData` **membership 一致**(`tests/datasets/test_split.py`)。 + + +| | 朴素 `ClusterData` | 本库增强版 | +| ------ | ---------------------------------- | -------------------------------------------------------- | +| 算法 | 原文 Butina | **同一套** exclusion sphere | +| 指纹 | 调用方 / DeepChem 1024 bit | Morgan r=2,**2048** bit | +| 阈值 | DeepChem 默认 **距离 0.6**($T\ge 0.4$) | `**similarity_cutoff=0.65` 是相似度**(不是 Chalcedon 的距离 0.65) | +| 存储 | 全部 pair 距离 | 稀疏邻接表 $O(n+E)$ | +| 搜索 | Python 填矩阵 | RDKit `BulkTanimotoSimilarity`(C++/POPCNT)+ numpy 筛边 | +| 平局 | `(degree, index)` 降序(高 index 优先) | **对齐 RDKit**(不是 chemfp 默认 `randomize`) | +| 度数 | 含自身(对角距离 0) | 含自身 | +| QM9 全集 | OOM | 可跑完(Colab 约十几~几十分钟,不是 2 小时) | + + +chemfp 的 `threshold_tanimoto_search_symmetric`、Chalcedon 的分块两遍,是同一精确算法的另两种工程布局。本库 **不 `import chemfp`**;OOM 时明确报错,**禁止 subsample 来「修好」split**。预设簇文件按教授要求放 **Hugging Face**,不进 git。 + +#### 步骤(与代码一致) + +1. `GetMorganGenerator(radius=2, fpSize=2048)`。 +2. 对 $i=1\ldots n-1$:`BulkTanimotoSimilarity(fps[i], fps[:i])`,只把 $T \ge t$ 的 $j$ 写入双方邻居表(另加自身,以对齐 RDKit 零自距)。 +3. 按 `(邻居数, index)` **降序**(`list.sort(reverse=True)`)。 +4. Exclusion sphere:未标记且度数 $>1$ 的点当中心,收走未标记邻居;剩余单点各自成簇。 +5. 每个 cluster → `_group_split`(与 scaffold 同一套 DeepChem 式贪心)。 + +Tanimoto: + +$$ +T(A,B)=\frac{|A\cap B|}{|A\cup B|}=\frac{c}{a+b-c},\quad d=1-T +$$ + +$0.65$ 是 Walters / 常见实践默认,**不是**唯一验证过的最优阈值。对齐 DeepChem 应设 `similarity_cutoff=0.4`。Chalcedon 博客的 `cutoff=0.65` 是**距离**($\Rightarrow T\ge 0.35$),不要和本 API 混用。 + +#### QM9 默认 $t=0.65$ 实测(全集) + +- 133885 分子 → **100653** 簇;单点簇约 **77.6%**;簇大小 min/median/mean/max = 1 / 1 / 1.33 / **23**。 +- 最大簇成员对中心 $T$:min=0.65,mean≈0.71,max=1.0(exclusion sphere 成立)。 +- 阈值偏严:多数分子没有 $T\ge 0.65$ 的邻居,group split 会接近「按分子切」,但多成员簇仍整组进同一边。 + +#### 不要做 + +- 全距离 / `ClusterData` 当生产路径(仅测试、$n\le 2000$)。 +- LSH、先 subsample 再叫 `method="butina"`(教授:不要为 split 缩小 QM9)。 +- BitBIRCH(另一种算法)。 +- chemfp false-singleton 后处理(原文没有)。 +- 把预设 `.json.gz` 提交进 `torch_molecule/datasets/data/`。 + +**测什么泛化:** 指纹空间上不相似的分子(通常比 scaffold 更严;在 QM9+$0.65$ 下因大量 singleton 会略接近 random)。 + +**参考:** Butina, JCICS 1999, doi:10.1021/ci9803381;chemfp `butina`;Chalcedon (Rowan, 2026);DeepChem `ButinaSplitter`(只对标整簇分配);Walters 2024;`molecule_notes/butina_optimization.md`。 + +--- + +### 5.4 Size / atom-count split + +**只做划分,不包含 SizeShiftReg 的训练正则。** + +```python +n_atoms = [mol.GetNumHeavyAtoms() for mol in mols] +order = np.argsort(n_atoms) +n_test = int(round(n * test_size)) +idx_test = order[-n_test:] # 最大分子进 test +idx_train = order[:-n_test] +``` + +- 按 **重原子数**,不是分子量(DeepChem `MolecularWeightSplitter` 不同)。 +- 默认 `direction="small_to_large"`:小 train、大 test(对齐 SizeShiftReg 评估思路)。 +- 可选 `mode="sizeshiftreg"`:50% 最小 train / 10% 最大 test(论文协议)。 + +**测什么泛化:** 尺度外推(小分子 → 大分子),与骨架无关。 + +**参考:** Buffelli et al., SizeShiftReg, NeurIPS 2022。 + +**不要:** 在 `split.py` 中实现 graph coarsening 或 CMD loss(属于 `SSRMolecularPredictor`)。 + +--- + +## 6. `SMILESDataset` 扩展 + +当前定义(`torch_molecule/datasets/constant.py`): + +```python +@dataclass +class SMILESDataset: + data: List[str] + target: np.ndarray | None +``` + +建议新增方法(逻辑委托 `split.py`): + +```python +def subsample(self, n: int, seed: int = 0) -> "SMILESDataset": + """随机抽取 n 条。仅调试 / CI,不要为了 Butina 缩小基准集。""" + +def train_test_split(self, test_size=0.2, method="random", seed=42, **kwargs): + """返回 (train_dataset, test_dataset)。""" + +def train_val_test_split(self, train_size=0.8, val_size=0.1, test_size=0.1, ...): + """三路划分。""" +``` + +`subsample`:`RandomState(seed).choice(n, size=min(n, n_sub), replace=False)`,保持 `data`/`target` 对齐。 + +--- + +## 7. 与现有训练流程对接 + +```python +from torch_molecule.datasets import load_qm9 +from torch_molecule import GREAMolecularPredictor + +data = load_qm9(local_dir="torchmol_data") +train, val = data.train_test_split(test_size=0.2, method="scaffold", seed=42) + +predictor = GREAMolecularPredictor(num_task=1, task_type="regression") +predictor.fit(train.data, train.target, val.data, val.target) +predictions = predictor.predict(val.data) +``` + +**不要在** `fit()` **内加** `split=` **参数** — 用户需在同一划分上比较 GREA vs GNN。 + +--- + +## 8. 测试计划(`tests/datasets/test_split.py`) + + +| 测试项 | 断言 | +| ---------------- | ----------------------------------------------------------------- | +| random 可复现 | 同一 seed → 相同索引 | +| random 比例 | test 约等于 `test_size` | +| scaffold 无泄漏 | train/test scaffold 集合交集为空 | +| scaffold 覆盖 | 所有 index 仅用一次 | +| 无环分子 | 不崩溃 | +| 无效 SMILES | `ValueError` | +| `target is None` | 可划分 | +| 多任务 `y` | `y.shape[1]` 保持 | +| subsample | 长度与 seed 可复现 | +| butina | 同 cluster 不跨集合;小 $n$ 簇与 RDKit `ClusterData` 一致;OOM 文案禁止 subsample | +| size | test 平均原子数 > train | + + +可选慢测:`load_qm9` + subsample(1000) + scaffold,CI 可 skip。 + +--- + +## 9. 文档说明(每种 method 的 docstring) + + +| method | 测什么 | 注意 | +| ---------- | ----------- | ------------------------------------------------ | +| `random` | 同分布插值 | 分数往往偏高,作 baseline | +| `scaffold` | 未见 scaffold | MoleculeNet 推荐用于 HIV/BACE/BBBP | +| `butina` | 结构不相似 | 增强版:稀疏精确 Butina;默认 $T\ge 0.65$;不是 ClusterData 矩阵 | +| `size` | 小→大原子数 | 与 scaffold 正交;SSR 正则另见模型 | + + +Group split 时实际比例可能偏离 `test_size` — 文档中明确说明。 + +--- + +## 10. 明确不做的事 + + +| 不做 | 原因 | +| --------------------------------- | ------------------------- | +| 在 `fit()` 内自动 split | 无法固定划分比较模型 | +| split 内实现 CMD / coarsening | 属于 SSR 训练,非划分 | +| 默认 `method=scaffold` | 破坏旧脚本可复现 | +| 用 `ClusterData` / 全距离矩阵跑全集 Butina | 精确算法不需要矩阵;QM9 会 OOM | +| 为了 Butina 把 QM9 subsample | 教授:QM9 不算大库;缩小数据改变的是划分本身 | +| 把 QM9 簇 `.json.gz` 放进 git 包内 | 预设放 Hugging Face | +| 把 BitBIRCH / LSH 命名为 `butina` | 不是原文算法 | +| 引入 DeepChem 依赖 | RDKit 指纹 + 自写 chemfp 收球即可 | +| 混淆 stratified(QM7 排序切)与 random | 若做 stratified,单独 `method` | + + +--- + +## 11. 参考文献与资源 + + +| 主题 | 文献 / 链接 | +| ------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| Scaffold 定义 | Bemis & Murcko, J. Med. Chem. 1996, doi:10.1021/jm9602928 | +| MoleculeNet 协议 | Wu et al., Chem. Sci. 2018, doi:10.1039/C7SC02664A | +| Butina 聚类 | Butina, J. Chem. Inf. Comput. Sci. 1999, doi:10.1021/ci9803381 | +| chemfp 稀疏 Butina | [https://chemfp.com/docs/chemfp_butina_command.html](https://chemfp.com/docs/chemfp_butina_command.html) | +| 实现对照(Chalcedon 等) | 本仓库 `molecule_notes/butina_optimization.md` | +| Split 实践对比 | Walters, Practical Cheminformatics, 2024 | +| Size 泛化 | Buffelli et al., SizeShiftReg, NeurIPS 2022, arXiv:2206.07096 | +| DeepChem splitters | [https://deepchem.readthedocs.io/en/stable/api_reference/splitters.html](https://deepchem.readthedocs.io/en/stable/api_reference/splitters.html) | +| scikit-fingerprints | [https://scikit-fingerprints.readthedocs.io/stable/examples/06_dataset_splits.html](https://scikit-fingerprints.readthedocs.io/stable/examples/06_dataset_splits.html) | + + +--- + diff --git a/tests/generator/hfpretrained.py b/tests/generator/hfpretrained.py deleted file mode 100644 index 08ff2e5..0000000 --- a/tests/generator/hfpretrained.py +++ /dev/null @@ -1,449 +0,0 @@ -import pytest - -from torch_molecule.generator.pretrained.registry import resolve_family - - -@pytest.mark.parametrize( - "repo_id,expected", - [ - ("chandar-lab/NovoMolGen_32M_SMILES_BPE", "novomolgen"), - ("zjunlp/MolGen-large", "molgen"), - ("zjunlp/MolGen-large-opt", "molgen"), - ("fairydance/molexar-10m-base", "molexar"), - ("fairydance/molexar-10m-omni", "molexar"), - ("datamol-io/safe-gpt", "safe_gpt"), - ("some-user/custom-causal-lm", "causal_lm"), - ], -) -def test_resolve_family(repo_id, expected): - assert resolve_family(repo_id) == expected - - -def test_hf_pretrained_generator_requires_transformers(): - pytest.importorskip("transformers") - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - ) - assert model.repo_id == "chandar-lab/NovoMolGen_32M_SMILES_BPE" - assert model.is_fitted_ is False - - -@pytest.mark.integration -def test_hf_pretrained_generator_novomolgen_smoke(): - pytest.importorskip("transformers") - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - generate_max_length=64, - ) - model.fit() - assert model.is_fitted_ is True - - smiles_list = model.generate(n_samples=2, temperature=1.0) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 2 - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - - -@pytest.mark.integration -def test_hf_pretrained_generator_finetune_smoke(tmp_path): - pytest.importorskip("transformers") - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - batch_size=2, - epochs=1, - generate_max_length=32, - ) - train_smiles = ["CCO", "CC(=O)O", "c1ccccc1"] - model.fit(train_smiles) - - assert model.is_fitted_ is True - assert len(model.fitting_loss) == 1 - assert model.fitting_epoch == 0 - - save_dir = tmp_path / "novomolgen-finetuned" - model.save_to_local(str(save_dir)) - - reloaded = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - generate_max_length=32, - ) - reloaded.load_from_local(str(save_dir)) - assert reloaded.is_fitted_ is True - assert reloaded.repo_id == model.repo_id - assert reloaded._family == "novomolgen" - - smiles_list = reloaded.generate(n_samples=1, temperature=1.0, do_sample=True) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 1 - assert isinstance(smiles_list[0], str) - - -@pytest.mark.integration -def test_hf_pretrained_generator_finetune_warns_on_y(): - pytest.importorskip("transformers") - import numpy as np - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - batch_size=2, - epochs=1, - ) - with pytest.warns(UserWarning, match="Conditional fine-tuning"): - model.fit(["CCO", "CC(=O)O"], y=np.array([0.0, 1.0])) - assert model.is_fitted_ is True - - -def test_smiles_selfies_roundtrip(): - pytest.importorskip("selfies") - - from torch_molecule.generator.pretrained.utils import selfies_to_smiles, smiles_to_selfies - - smiles = ["CCO", "c1ccccc1", "CC(=O)O"] - recovered = selfies_to_smiles(smiles_to_selfies(smiles)) - assert recovered == smiles - - -def test_smiles_to_selfies_invalid_smiles(): - pytest.importorskip("selfies") - - from torch_molecule.generator.pretrained.utils import smiles_to_selfies - - with pytest.raises(ValueError, match="Invalid SMILES"): - smiles_to_selfies(["not-a-smiles"]) - - -def test_smiles_to_selfies_encoder_error_has_index(monkeypatch): - pytest.importorskip("selfies") - import selfies as sf - - from torch_molecule.generator.pretrained.utils import smiles_to_selfies - - def _raise_encoder(_smiles): - raise sf.EncoderError("synthetic encoder failure") - - monkeypatch.setattr(sf, "encoder", _raise_encoder) - with pytest.raises(ValueError, match="index 0.*not SELFIES-encodable"): - smiles_to_selfies(["CCO"]) - - -def test_selfies_to_smiles_drops_invalid_entries(): - pytest.importorskip("selfies") - - from torch_molecule.generator.pretrained.utils import selfies_to_smiles, smiles_to_selfies - - valid = smiles_to_selfies(["CCO"])[0] - with pytest.warns(UserWarning, match="dropped 2 invalid SELFIES"): - recovered = selfies_to_smiles([valid, "not-valid-selfies-[[[", ""]) - assert recovered == ["CCO"] - assert "" not in recovered - - -def test_decode_outputs_molgen_drops_empty_and_warns(): - pytest.importorskip("transformers") - pytest.importorskip("selfies") - - from torch_molecule import HFPretrainedMolecularGenerator - from torch_molecule.generator.pretrained.utils import smiles_to_selfies - - model = HFPretrainedMolecularGenerator(repo_id="zjunlp/MolGen-large") - model._family = "molgen" - valid = smiles_to_selfies(["c1ccccc1"])[0] - with pytest.warns(UserWarning, match="got 1/2 valid SMILES"): - out = model._decode_outputs([valid, "not-a-selfies"]) - assert len(out) == 1 - assert "" not in out - - -def test_known_family_prefix_does_not_warn_unknown_repo(): - pytest.importorskip("transformers") - import warnings - - from torch_molecule import HFPretrainedMolecularGenerator - - with warnings.catch_warnings(record=True) as recorded: - warnings.simplefilter("always") - HFPretrainedMolecularGenerator(repo_id="chandar-lab/NovoMolGen_157M") - assert not any("Unknown repo_id" in str(item.message) for item in recorded) - - -def test_unknown_repo_fallback_warns(): - pytest.importorskip("transformers") - - from torch_molecule import HFPretrainedMolecularGenerator - - with pytest.warns(UserWarning, match="Unknown repo_id"): - HFPretrainedMolecularGenerator(repo_id="some-user/custom-causal-lm") - - -def _molexar_available() -> bool: - try: - import molexar # noqa: F401 - import fragment_selfies # noqa: F401 - return True - except ImportError: - return False - - -@pytest.mark.integration -@pytest.mark.parametrize("repo_id", ["fairydance/molexar-10m-base"]) -def test_hf_pretrained_generator_molexar_denovo(repo_id): - if not _molexar_available(): - pytest.skip("Molexar requires fragment-selfies and molexar") - - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator(repo_id=repo_id) - model.fit() - assert model.is_fitted_ is True - - smiles_list = model.generate(n_samples=2, temperature=0.8) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 2 - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) - - -@pytest.mark.integration -def test_hf_pretrained_generator_molexar_fragment_constraint(): - if not _molexar_available(): - pytest.skip("Molexar requires fragment-selfies and molexar") - - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - - start_smiles = "[*]C1(CC#N)CN(S(=O)(=O)CC)C1" - model = HFPretrainedMolecularGenerator(repo_id="fairydance/molexar-10m-base") - model.fit() - - smiles_list = model.generate( - n_samples=2, - start_smiles=start_smiles, - generation_task="motif_extension", - temperature=0.8, - ) - assert len(smiles_list) == 2 - assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) - - -def test_smiles_fragment_selfies_roundtrip(): - if not _molexar_available(): - pytest.skip("Molexar requires fragment-selfies and molexar") - - from torch_molecule.generator.pretrained.utils import ( - fragment_selfies_to_smiles, - smiles_to_fragment_selfies, - ) - - smiles = ["CCO", "c1ccccc1"] - recovered = fragment_selfies_to_smiles(smiles_to_fragment_selfies(smiles)) - assert len(recovered) == 2 - assert all(smiles_string for smiles_string in recovered) - - -@pytest.mark.parametrize( - "repo_id", - [ - "zjunlp/MolGen-large", - "zjunlp/MolGen-large-opt", - ], -) -@pytest.mark.integration -def test_hf_pretrained_generator_molgen_smoke(repo_id): - pytest.importorskip("transformers") - pytest.importorskip("selfies") - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id=repo_id, - generate_max_length=20, - ) - model.fit() - assert model.is_fitted_ is True - - smiles_list = model.generate(n_samples=2, num_beams=5) - assert isinstance(smiles_list, list) - assert len(smiles_list) == 2 - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) - - -@pytest.mark.integration -def test_hf_pretrained_generator_molgen_scaffold_prefix(): - pytest.importorskip("transformers") - pytest.importorskip("selfies") - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - from torch_molecule.generator.pretrained.utils import smiles_to_selfies - - scaffold = "c1ccccc1" - benzene = Chem.MolFromSmiles(scaffold) - prefix_selfies = smiles_to_selfies([scaffold])[0] - - model = HFPretrainedMolecularGenerator( - repo_id="zjunlp/MolGen-large", - generate_max_length=20, - ) - model.fit() - - smiles_from_scaffold = model.generate(n_samples=2, scaffold=scaffold, num_beams=5) - smiles_from_prefix = model.generate(n_samples=2, prefix_selfies=prefix_selfies, num_beams=5) - - assert len(smiles_from_scaffold) == 2 - assert len(smiles_from_prefix) == 2 - assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_from_scaffold) - assert all( - Chem.MolFromSmiles(smiles) is not None - and Chem.MolFromSmiles(smiles).HasSubstructMatch(benzene) - for smiles in smiles_from_scaffold - ) - - -def _safe_mol_available() -> bool: - try: - from torch_molecule.generator.pretrained.compat import ensure_safe_transformers_compat - - ensure_safe_transformers_compat() - from safe.converter import encode # noqa: F401 - - return True - except ImportError: - return False - - -def test_smiles_safe_roundtrip(): - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - - from torch_molecule.generator.pretrained.utils import safe_to_smiles, smiles_to_safe - - smiles = ["CCO", "c1ccccc1", "CC(=O)O"] - recovered = safe_to_smiles(smiles_to_safe(smiles)) - assert recovered == smiles - - -def test_smiles_to_safe_invalid_smiles(): - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - - from torch_molecule.generator.pretrained.utils import smiles_to_safe - - with pytest.raises(ValueError, match="Invalid SMILES"): - smiles_to_safe(["not-a-smiles"]) - - -def test_safe_to_smiles_drops_invalid_entries(): - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - - from torch_molecule.generator.pretrained.utils import safe_to_smiles, smiles_to_safe - - valid = smiles_to_safe(["CCO"])[0] - with pytest.warns(UserWarning, match="dropped 2 invalid SAFE"): - recovered = safe_to_smiles([valid, "not-valid-safe-[[[", ""]) - assert recovered == ["CCO"] - assert "" not in recovered - - -def test_decode_outputs_safe_gpt_drops_empty_and_warns(): - pytest.importorskip("transformers") - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - - from torch_molecule import HFPretrainedMolecularGenerator - from torch_molecule.generator.pretrained.utils import smiles_to_safe - - model = HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") - model._family = "safe_gpt" - valid = smiles_to_safe(["CCO"])[0] - with pytest.warns(UserWarning, match="got 1/2 valid SMILES"): - out = model._decode_outputs([valid, "not-a-safe"]) - assert len(out) == 1 - assert "" not in out - - -def test_safe_gpt_uses_gpt2_lm_head(): - pytest.importorskip("transformers") - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") - model._family = "safe_gpt" - assert model._get_model_class().__name__ == "GPT2LMHeadModel" - - -def test_safe_gpt_known_repo_does_not_warn_unknown(): - pytest.importorskip("transformers") - import warnings - - from torch_molecule import HFPretrainedMolecularGenerator - - with warnings.catch_warnings(record=True) as recorded: - warnings.simplefilter("always") - HFPretrainedMolecularGenerator(repo_id="datamol-io/safe-gpt") - assert not any("Unknown repo_id" in str(item.message) for item in recorded) - - -@pytest.mark.integration -def test_hf_pretrained_generator_safe_gpt_denovo(): - pytest.importorskip("transformers") - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - - model = HFPretrainedMolecularGenerator( - repo_id="datamol-io/safe-gpt", - generate_max_length=128, - ) - model.fit() - assert model.is_fitted_ is True - - smiles_list = model.generate(n_samples=2, temperature=1.0) - assert isinstance(smiles_list, list) - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - assert all(Chem.MolFromSmiles(smiles) is not None for smiles in smiles_list) - - -@pytest.mark.integration -def test_hf_pretrained_generator_safe_gpt_scaffold(): - pytest.importorskip("transformers") - if not _safe_mol_available(): - pytest.skip("safe-mol not installed") - from rdkit import Chem - - from torch_molecule import HFPretrainedMolecularGenerator - - scaffold = "c1ccccc1" - benzene = Chem.MolFromSmiles(scaffold) - model = HFPretrainedMolecularGenerator( - repo_id="datamol-io/safe-gpt", - generate_max_length=128, - ) - model.fit() - - smiles_list = model.generate(n_samples=2, scaffold=scaffold, temperature=1.0) - assert isinstance(smiles_list, list) - assert all(isinstance(smiles, str) and smiles for smiles in smiles_list) - assert all( - Chem.MolFromSmiles(smiles) is not None - and Chem.MolFromSmiles(smiles).HasSubstructMatch(benzene) - for smiles in smiles_list - ) diff --git a/tests/generator/pretrained_molexar.py b/tests/generator/pretrained_molexar.py new file mode 100644 index 0000000..26ee478 --- /dev/null +++ b/tests/generator/pretrained_molexar.py @@ -0,0 +1,66 @@ +import os +import shutil + +from torch_molecule import HFPretrainedMolecularGenerator + +REPO_ID = "fairydance/molexar-10m-base" +N_SAMPLES = 2 +START_SMILES = "[*]C1(CC#N)CN(S(=O)(=O)CC)C1" + + +def test_molexar_generator(): + print("\n=== Testing Molexar initialization ===") + model = HFPretrainedMolecularGenerator( + repo_id=REPO_ID, + verbose="progress_bar", + ) + print("Molexar initialized successfully") + + print("\n=== Testing Molexar loading from Hugging Face ===") + model.fit() + print("Molexar loaded successfully") + + print("\n=== Testing Molexar de novo generation ===") + generated_smiles = model.generate( + n_samples=N_SAMPLES, + max_new_tokens=64, + temperature=0.8, + ) + print(f"Generated {len(generated_smiles)} molecules") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing Molexar fragment-constrained generation ===") + generated_smiles = model.generate( + n_samples=N_SAMPLES, + start_smiles=START_SMILES, + generation_task="motif_extension", + max_new_tokens=64, + temperature=0.8, + ) + print(f"Generated {len(generated_smiles)} molecules from start_smiles {START_SMILES}") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing Molexar saving and loading ===") + save_path = "pretrained_molexar_test_model" + model.save_to_local(save_path) + print(f"Molexar saved to {save_path}") + + loaded_model = HFPretrainedMolecularGenerator(repo_id=REPO_ID) + loaded_model.load_from_local(save_path) + print("Molexar loaded from local directory") + + generated_smiles = loaded_model.generate( + n_samples=N_SAMPLES, + max_new_tokens=64, + temperature=0.8, + ) + print(f"Generated {len(generated_smiles)} molecules with loaded model") + print("Example generated SMILES:", generated_smiles[:2]) + + if os.path.exists(save_path): + shutil.rmtree(save_path) + print(f"Cleaned up {save_path}") + + +if __name__ == "__main__": + test_molexar_generator() diff --git a/tests/generator/pretrained_molgen.py b/tests/generator/pretrained_molgen.py new file mode 100644 index 0000000..1f9835a --- /dev/null +++ b/tests/generator/pretrained_molgen.py @@ -0,0 +1,74 @@ +import os +import shutil + +from torch_molecule import HFPretrainedMolecularGenerator + +N_SAMPLES = 2 +SCAFFOLD = "c1ccccc1" +PREFIX_SELFIES = "[C][=C][C][=C][C][=C][Ring1][=Branch1]" + + +def test_molgen_generator(): + models_to_test = [ + {"repo_id": "zjunlp/MolGen-large", "model_name": "MolGen-large"}, + {"repo_id": "zjunlp/MolGen-large-opt", "model_name": "MolGen-large-opt"}, + ] + + for model_config in models_to_test: + name = model_config["model_name"] + repo_id = model_config["repo_id"] + + print(f"\n=== Testing {name} initialization ===") + model = HFPretrainedMolecularGenerator( + repo_id=repo_id, + verbose="progress_bar", + ) + print(f"{name} initialized successfully") + + print(f"\n=== Testing {name} loading from Hugging Face ===") + model.fit() + print(f"{name} loaded successfully") + + print(f"\n=== Testing {name} de novo generation ===") + generated_smiles = model.generate(n_samples=N_SAMPLES, num_beams=5) + print(f"Generated {len(generated_smiles)} molecules") + print("Example generated SMILES:", generated_smiles[:2]) + + print(f"\n=== Testing {name} scaffold generation ===") + generated_smiles = model.generate( + n_samples=N_SAMPLES, + scaffold=SCAFFOLD, + num_beams=5, + ) + print(f"Generated {len(generated_smiles)} molecules from scaffold {SCAFFOLD}") + print("Example generated SMILES:", generated_smiles[:2]) + + print(f"\n=== Testing {name} prefix_selfies generation ===") + generated_smiles = model.generate( + n_samples=N_SAMPLES, + prefix_selfies=PREFIX_SELFIES, + num_beams=5, + ) + print(f"Generated {len(generated_smiles)} molecules from prefix_selfies") + print("Example generated SMILES:", generated_smiles[:2]) + + print(f"\n=== Testing {name} saving and loading ===") + save_path = f"pretrained_{name.lower().replace('-', '_')}_test_model" + model.save_to_local(save_path) + print(f"{name} saved to {save_path}") + + loaded_model = HFPretrainedMolecularGenerator(repo_id=repo_id) + loaded_model.load_from_local(save_path) + print(f"{name} loaded from local directory") + + generated_smiles = loaded_model.generate(n_samples=N_SAMPLES, num_beams=5) + print(f"Generated {len(generated_smiles)} molecules with loaded model") + print("Example generated SMILES:", generated_smiles[:2]) + + if os.path.exists(save_path): + shutil.rmtree(save_path) + print(f"Cleaned up {save_path}") + + +if __name__ == "__main__": + test_molgen_generator() diff --git a/tests/generator/pretrained_novomolgen.py b/tests/generator/pretrained_novomolgen.py new file mode 100644 index 0000000..d04cfbe --- /dev/null +++ b/tests/generator/pretrained_novomolgen.py @@ -0,0 +1,86 @@ +import os +import shutil + +from torch_molecule import HFPretrainedMolecularGenerator + +REPO_ID = "chandar-lab/NovoMolGen_32M_SMILES_BPE" +N_SAMPLES = 5 +TRAIN_SMILES = [ + "CC(=O)O", + "CCO", + "CCCC", + "c1ccccc1", + "CCN", +] + + +def test_novomolgen_generator(): + print("\n=== Testing NovoMolGen initialization ===") + model = HFPretrainedMolecularGenerator( + repo_id=REPO_ID, + verbose="progress_bar", + ) + print("NovoMolGen initialized successfully") + + print("\n=== Testing NovoMolGen loading from Hugging Face ===") + model.fit() + print("NovoMolGen loaded successfully") + + print("\n=== Testing NovoMolGen de novo generation ===") + generated_smiles = model.generate(n_samples=N_SAMPLES) + print(f"Generated {len(generated_smiles)} molecules") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing NovoMolGen saving and loading ===") + save_path = "pretrained_novomolgen_test_model" + model.save_to_local(save_path) + print(f"NovoMolGen saved to {save_path}") + + loaded_model = HFPretrainedMolecularGenerator(repo_id=REPO_ID) + loaded_model.load_from_local(save_path) + print("NovoMolGen loaded from local directory") + + generated_smiles = loaded_model.generate(n_samples=2) + print(f"Generated {len(generated_smiles)} molecules with loaded model") + print("Example generated SMILES:", generated_smiles[:2]) + + if os.path.exists(save_path): + shutil.rmtree(save_path) + print(f"Cleaned up {save_path}") + + print("\n=== Testing NovoMolGen fine-tuning ===") + finetune_model = HFPretrainedMolecularGenerator( + repo_id=REPO_ID, + batch_size=2, + epochs=1, + verbose="progress_bar", + ) + finetune_model.fit(TRAIN_SMILES) + print("Fine-tuning completed") + print(f"Fitting epochs: {finetune_model.fitting_epoch + 1}") + print(f"Fitting loss: {finetune_model.fitting_loss}") + + generated_smiles = finetune_model.generate(n_samples=2) + print(f"Generated {len(generated_smiles)} molecules after fine-tuning") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing fine-tuned NovoMolGen saving and loading ===") + save_path = "pretrained_novomolgen_finetune_test_model" + finetune_model.save_to_local(save_path) + print(f"Fine-tuned NovoMolGen saved to {save_path}") + + loaded_finetune = HFPretrainedMolecularGenerator(repo_id=REPO_ID) + loaded_finetune.load_from_local(save_path) + print("Fine-tuned NovoMolGen loaded from local directory") + + generated_smiles = loaded_finetune.generate(n_samples=2) + print(f"Generated {len(generated_smiles)} molecules with loaded fine-tuned model") + print("Example generated SMILES:", generated_smiles[:2]) + + if os.path.exists(save_path): + shutil.rmtree(save_path) + print(f"Cleaned up {save_path}") + + +if __name__ == "__main__": + test_novomolgen_generator() diff --git a/tests/generator/pretrained_safe_gpt.py b/tests/generator/pretrained_safe_gpt.py new file mode 100644 index 0000000..7d179bd --- /dev/null +++ b/tests/generator/pretrained_safe_gpt.py @@ -0,0 +1,58 @@ +import os +import shutil + +from torch_molecule import HFPretrainedMolecularGenerator + +REPO_ID = "datamol-io/safe-gpt" +N_SAMPLES = 5 +SHORT_SCAFFOLD = "c1ccccc1" +LONG_SCAFFOLD = "CC1=CC=C(C=C1)C2=CC(=NN2C3=CC=C(C=C3)S(=O)(=O)N)C(F)(F)F" + + +def test_safe_gpt_generator(): + print("\n=== Testing SAFE-GPT initialization ===") + model = HFPretrainedMolecularGenerator( + repo_id=REPO_ID, + verbose="progress_bar", + ) + print("SAFE-GPT initialized successfully") + + print("\n=== Testing SAFE-GPT loading from Hugging Face ===") + model.fit() + print("SAFE-GPT loaded successfully") + + print("\n=== Testing SAFE-GPT de novo generation ===") + generated_smiles = model.generate(n_samples=N_SAMPLES) + print(f"Generated {len(generated_smiles)} molecules") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing SAFE-GPT short scaffold generation ===") + generated_smiles = model.generate(n_samples=N_SAMPLES, scaffold=SHORT_SCAFFOLD) + print(f"Generated {len(generated_smiles)} molecules from scaffold {SHORT_SCAFFOLD}") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing SAFE-GPT long scaffold generation ===") + generated_smiles = model.generate(n_samples=2, scaffold=LONG_SCAFFOLD) + print(f"Generated {len(generated_smiles)} molecules from long scaffold") + print("Example generated SMILES:", generated_smiles[:2]) + + print("\n=== Testing SAFE-GPT saving and loading ===") + save_path = "pretrained_safe_gpt_test_model" + model.save_to_local(save_path) + print(f"SAFE-GPT saved to {save_path}") + + loaded_model = HFPretrainedMolecularGenerator(repo_id=REPO_ID) + loaded_model.load_from_local(save_path) + print("SAFE-GPT loaded from local directory") + + generated_smiles = loaded_model.generate(n_samples=2, scaffold=SHORT_SCAFFOLD) + print(f"Generated {len(generated_smiles)} molecules with loaded model") + print("Example generated SMILES:", generated_smiles[:2]) + + if os.path.exists(save_path): + shutil.rmtree(save_path) + print(f"Cleaned up {save_path}") + + +if __name__ == "__main__": + test_safe_gpt_generator() diff --git a/tests/generator/test_causal_lm.py b/tests/generator/test_causal_lm.py deleted file mode 100644 index 940af15..0000000 --- a/tests/generator/test_causal_lm.py +++ /dev/null @@ -1,79 +0,0 @@ -from unittest.mock import MagicMock - -import torch - -from torch_molecule.generator.pretrained.families.causal_lm import generate_causal_lm - - -class _FakeTokenizer: - bos_token_id = 1 - pad_token_id = 0 - eos_token_id = 2 - - def __call__(self, text, return_tensors="pt", add_special_tokens=True): - token_ids = [10 + len(text), 11 + len(text)] - if add_special_tokens: - token_ids = [self.bos_token_id] + token_ids + [self.eos_token_id] - return {"input_ids": torch.tensor([token_ids])} - - def batch_decode(self, outputs, skip_special_tokens=True): - return [f"SMILES_{idx}" for idx in range(outputs.shape[0])] - - -class _FakeModel: - def generate(self, **kwargs): - batch_size = kwargs["input_ids"].shape[0] if "input_ids" in kwargs else kwargs["num_return_sequences"] - seq_len = kwargs.get("max_length", 8) - return torch.zeros(batch_size, seq_len, dtype=torch.long) - - -def test_generate_causal_lm_bos_path(): - outputs = generate_causal_lm( - _FakeModel(), - _FakeTokenizer(), - torch.device("cpu"), - n_samples=3, - family="novomolgen", - max_length=12, - do_sample=False, - ) - assert outputs == ["SMILES_0", "SMILES_1", "SMILES_2"] - - -def test_generate_causal_lm_scaffold_keeps_all_tokens(): - model = _FakeModel() - model.generate = MagicMock(return_value=torch.zeros(2, 8, dtype=torch.long)) - - generate_causal_lm( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=2, - family="novomolgen", - scaffold="c1ccccc1", - max_length=12, - do_sample=False, - ) - - assert model.generate.call_args.kwargs["input_ids"].tolist() == [ - [18, 19], - [18, 19], - ] - assert "top_k" not in model.generate.call_args.kwargs - - -def test_generate_causal_lm_does_not_force_use_cache_false(): - model = _FakeModel() - model.generate = MagicMock(return_value=torch.zeros(3, 8, dtype=torch.long)) - - generate_causal_lm( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=3, - family="novomolgen", - max_length=12, - do_sample=False, - ) - - assert "use_cache" not in model.generate.call_args.kwargs diff --git a/tests/generator/test_compat.py b/tests/generator/test_compat.py deleted file mode 100644 index e430aa0..0000000 --- a/tests/generator/test_compat.py +++ /dev/null @@ -1,13 +0,0 @@ -import pytest - -from torch_molecule.generator.pretrained.compat import ensure_safe_transformers_compat - - -def test_ensure_safe_transformers_compat_provides_constraints(): - pytest.importorskip("transformers") - - ensure_safe_transformers_compat() - import transformers.generation as generation - - assert hasattr(generation, "DisjunctiveConstraint") - assert hasattr(generation, "PhrasalConstraint") diff --git a/tests/generator/test_finetune.py b/tests/generator/test_finetune.py deleted file mode 100644 index ad13570..0000000 --- a/tests/generator/test_finetune.py +++ /dev/null @@ -1,197 +0,0 @@ -import json -import os - -import pytest -import torch - -from torch_molecule.generator.pretrained.checkpoint import ( - METADATA_FILENAME, - build_metadata, - load_metadata, - save_metadata, -) -from torch_molecule.generator.pretrained.finetune import ( - corrupt_token_ids, - finetune_causal_lm, - finetune_generator, - finetune_seq2seq, -) - - -class _FakeOutput: - def __init__(self, loss: torch.Tensor): - self.loss = loss - - -class _FakeLM(torch.nn.Module): - def __init__(self): - super().__init__() - self.weight = torch.nn.Parameter(torch.tensor(1.0)) - - def forward(self, **kwargs): - return _FakeOutput(self.weight * 0.0 + 1.0) - - -def test_finetune_causal_lm_runs_one_epoch(): - pytest.importorskip("transformers") - from transformers import AutoTokenizer - - tokenizer = AutoTokenizer.from_pretrained("gpt2") - if tokenizer.pad_token is None: - tokenizer.pad_token = tokenizer.eos_token - - losses, last_epoch = finetune_causal_lm( - _FakeLM(), - tokenizer, - ["CCO", "CC(=O)O"], - torch.device("cpu"), - max_length=16, - batch_size=2, - epochs=1, - learning_rate=1e-3, - weight_decay=0.0, - grad_norm_clip=1.0, - verbose="none", - ) - assert last_epoch == 0 - assert len(losses) == 1 - assert losses[0] == pytest.approx(1.0) - - -class _FakeTokenizer: - mask_token_id = 99 - pad_token_id = 0 - all_special_ids = [0, 1, 2] - - -def _seq2seq_tokenizer(): - from transformers import AutoTokenizer - - tokenizer = AutoTokenizer.from_pretrained("gpt2") - if tokenizer.pad_token is None: - tokenizer.pad_token = tokenizer.eos_token - if tokenizer.mask_token is None: - tokenizer.add_special_tokens({"mask_token": ""}) - return tokenizer - - -def test_corrupt_token_ids_masks_non_special_tokens(): - input_ids = torch.tensor([1, 10, 11, 12, 2]) - generator = torch.Generator().manual_seed(0) - corrupted = corrupt_token_ids( - input_ids, - _FakeTokenizer(), - mask_prob=1.0, - generator=generator, - ) - assert corrupted[0].item() == 1 - assert corrupted[-1].item() == 2 - assert torch.equal(corrupted[1:-1], torch.tensor([99, 99, 99])) - - -def test_corrupt_token_ids_requires_mask_token(): - class _NoMask: - mask_token_id = None - - with pytest.raises(ValueError, match="mask_token_id"): - corrupt_token_ids(torch.tensor([1, 2, 3]), _NoMask()) - - -def test_finetune_seq2seq_runs_one_epoch(): - pytest.importorskip("transformers") - - tokenizer = _seq2seq_tokenizer() - losses, last_epoch = finetune_seq2seq( - _FakeLM(), - tokenizer, - ["[C][C][O]", "[C][C][Branch1][C][O]"], - torch.device("cpu"), - max_length=16, - batch_size=1, - epochs=1, - learning_rate=1e-3, - weight_decay=0.0, - grad_norm_clip=None, - verbose="none", - mask_prob=1.0, - ) - assert last_epoch == 0 - assert len(losses) == 1 - - -def test_finetune_seq2seq_labels_stay_clean(): - pytest.importorskip("transformers") - - tokenizer = _seq2seq_tokenizer() - captured = {} - - class _CaptureLM(_FakeLM): - def forward(self, **kwargs): - captured["input_ids"] = kwargs["input_ids"].detach().clone() - captured["labels"] = kwargs["labels"].detach().clone() - return super().forward(**kwargs) - - finetune_seq2seq( - _CaptureLM(), - tokenizer, - ["[C][C][O][C][C][O]"], - torch.device("cpu"), - max_length=16, - batch_size=1, - epochs=1, - learning_rate=1e-3, - weight_decay=0.0, - grad_norm_clip=None, - verbose="none", - mask_prob=1.0, - ) - - labels = captured["labels"] - input_ids = captured["input_ids"] - ignore_index = -100 - content = labels[0] != ignore_index - assert not torch.equal(input_ids[0][content], labels[0][content]) - assert tokenizer.mask_token_id in input_ids[0].tolist() - - -def test_finetune_generator_dispatches_seq2seq(): - pytest.importorskip("transformers") - - tokenizer = _seq2seq_tokenizer() - - losses, last_epoch = finetune_generator( - "molgen", - _FakeLM(), - tokenizer, - ["[C][C][O]"], - torch.device("cpu"), - max_length=16, - batch_size=1, - epochs=1, - learning_rate=1e-3, - weight_decay=0.0, - grad_norm_clip=1.0, - verbose="none", - ) - assert last_epoch == 0 - assert len(losses) == 1 - - -def test_checkpoint_metadata_roundtrip(tmp_path): - metadata = build_metadata( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", - family="novomolgen", - max_length=128, - revision="hf-checkpoint", - trust_remote_code=False, - tokenizer_repo_id=None, - generate_max_length=64, - model_name="HFPretrainedMolecularGenerator", - ) - save_metadata(str(tmp_path), metadata) - loaded = load_metadata(str(tmp_path)) - assert loaded == metadata - assert os.path.exists(os.path.join(tmp_path, METADATA_FILENAME)) - with open(os.path.join(tmp_path, METADATA_FILENAME), encoding="utf-8") as handle: - on_disk = json.load(handle) - assert on_disk["family"] == "novomolgen" diff --git a/tests/generator/test_molexar.py b/tests/generator/test_molexar.py deleted file mode 100644 index 1d3da9e..0000000 --- a/tests/generator/test_molexar.py +++ /dev/null @@ -1,38 +0,0 @@ -import importlib.util - -import pytest - -from torch_molecule.generator.pretrained.families.molexar import ( - extract_conditions, - resolve_start_string, -) - - -def test_resolve_start_string_none_for_de_novo(): - assert resolve_start_string() is None - - -def test_resolve_start_string_accepts_literal_prefix(): - prefix = "[Frag][C][C][Attach:0]" - assert resolve_start_string(start_string=prefix) == prefix - - -@pytest.mark.skipif( - importlib.util.find_spec("molexar") is None, - reason="molexar not installed", -) -def test_resolve_start_string_from_smiles_fragment(): - start_smiles = "[*]C1(CC#N)CN(S(=O)(=O)CC)C1" - resolved = resolve_start_string( - start_smiles=start_smiles, - generation_task="motif_extension", - ) - assert resolved is not None - assert "[Attach:0]" in resolved - - -def test_extract_conditions_from_kwargs(): - kwargs = {"temperature": 0.8, "mol_qed": 0.9, "conditions": {"mol_logp": 2.5}} - conditions = extract_conditions(kwargs) - assert conditions == {"mol_logp": 2.5, "mol_qed": 0.9} - assert kwargs == {"temperature": 0.8} diff --git a/tests/generator/test_seq2seq.py b/tests/generator/test_seq2seq.py deleted file mode 100644 index a653e6d..0000000 --- a/tests/generator/test_seq2seq.py +++ /dev/null @@ -1,67 +0,0 @@ -from unittest.mock import MagicMock - -import torch - -from torch_molecule.generator.pretrained.families.seq2seq import ( - DEFAULT_MOLGEN_PREFIX_SELFIES, - generate_seq2seq, -) - - -class _FakeTokenizer: - def __call__(self, text, return_tensors="pt"): - length = len(text) - return { - "input_ids": torch.tensor([[1, length, 3]]), - "attention_mask": torch.tensor([[1, 1, 1]]), - } - - def decode(self, sequence, skip_special_tokens=True, clean_up_tokenization_spaces=True): - return f"[SELFIES_{int(sequence[0].item())}]" - - -class _FakeModel: - def generate(self, **kwargs): - num_sequences = kwargs["num_return_sequences"] - seq_len = kwargs.get("max_length", 8) - return torch.arange(num_sequences * seq_len, dtype=torch.long).reshape(num_sequences, seq_len) - - -def test_generate_seq2seq_default_prefix(): - model = _FakeModel() - model.generate = MagicMock(side_effect=_FakeModel().generate) - - outputs = generate_seq2seq( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=3, - max_length=12, - min_length=4, - num_beams=5, - ) - - assert len(outputs) == 3 - call_kwargs = model.generate.call_args.kwargs - assert call_kwargs["num_return_sequences"] == 3 - assert call_kwargs["num_beams"] == 5 - assert call_kwargs["max_length"] == 12 - - -def test_generate_seq2seq_expands_beams_for_sample_count(): - model = MagicMock() - model.generate.return_value = torch.zeros(7, 5, dtype=torch.long) - - generate_seq2seq( - model, - _FakeTokenizer(), - torch.device("cpu"), - n_samples=7, - num_beams=3, - ) - - assert model.generate.call_args.kwargs["num_beams"] == 7 - - -def test_default_molgen_prefix_is_benzene_selfies(): - assert "Ring1" in DEFAULT_MOLGEN_PREFIX_SELFIES diff --git a/torch_molecule/generator/pretrained/families/causal_lm.py b/torch_molecule/generator/pretrained/families/causal_lm.py index 3410b7b..4cddf23 100644 --- a/torch_molecule/generator/pretrained/families/causal_lm.py +++ b/torch_molecule/generator/pretrained/families/causal_lm.py @@ -12,7 +12,8 @@ def generate_causal_lm( n_samples: int, *, family: Optional[str] = None, - max_length: int = 64, + max_new_tokens: Optional[int] = None, + max_length: Optional[int] = None, temperature: float = 1.0, do_sample: bool = True, scaffold: Optional[str] = None, @@ -33,8 +34,13 @@ def generate_causal_lm( family : Optional[str], default=None Generator family name. Unused by the standard ``generate()`` path; kept for call-site compatibility. - max_length : int, default=64 - Maximum generated sequence length passed to ``model.generate``. + max_new_tokens : Optional[int], default=None + Maximum number of newly generated tokens. Independent of prefix length, + so a long ``scaffold=`` does not consume the generation budget. + Used when ``max_length`` is omitted; defaults to 64. + max_length : Optional[int], default=None + Optional Hugging Face total sequence length (prefix + new tokens). + When set without ``max_new_tokens``, this is passed through instead. temperature : float, default=1.0 Sampling temperature. do_sample : bool, default=True @@ -54,10 +60,17 @@ def generate_causal_lm( pad_token_id = tokenizer.eos_token_id generate_kwargs = { - "max_length": max_length, "do_sample": do_sample, "pad_token_id": pad_token_id, } + # Hugging Face rejects passing both; prefer max_new_tokens unless the caller + # explicitly asks for total-length max_length. + if max_length is not None and max_new_tokens is None: + generate_kwargs["max_length"] = max_length + else: + generate_kwargs["max_new_tokens"] = ( + max_new_tokens if max_new_tokens is not None else 64 + ) if do_sample: generate_kwargs["temperature"] = temperature generate_kwargs.update(kwargs) diff --git a/torch_molecule/generator/pretrained/modeling_pretrained.py b/torch_molecule/generator/pretrained/modeling_pretrained.py index 1b8f303..010364d 100644 --- a/torch_molecule/generator/pretrained/modeling_pretrained.py +++ b/torch_molecule/generator/pretrained/modeling_pretrained.py @@ -87,7 +87,9 @@ class HFPretrainedMolecularGenerator(BaseMolecularGenerator): tokenizer_repo_id : Optional[str], default=None Optional Hugging Face repo for the tokenizer. generate_max_length : int, default=64 - Default ``max_length`` passed to ``generate()``. + Default generation length. For causal LMs (NovoMolGen, SAFE-GPT) this is + ``max_new_tokens`` so a ``scaffold=`` prefix does not consume the budget. + For MolGen this is still decoder ``max_length``. batch_size : int, default=8 Batch size used when fine-tuning on SMILES data. epochs : int, default=1 @@ -142,6 +144,7 @@ def __init__( self._family: Optional[str] = None self.tokenizer = None + self._safe_tokenizer = None self._molexar_engine = None self._model_local_path: Optional[str] = None self.fitting_epoch = -1 @@ -198,12 +201,8 @@ def save_to_local(self, path: str) -> None: self._check_is_fitted() os.makedirs(path, exist_ok=True) - if self._family in MOLEXAR_FAMILIES: - self.model.save_pretrained(path) - self.tokenizer.save_pretrained(path) - else: - self.model.save_pretrained(path) - self.tokenizer.save_pretrained(path) + self.model.save_pretrained(path) + self._save_generator_tokenizer(path) save_metadata( path, @@ -310,12 +309,14 @@ def generate(self, n_samples: int = 10, **kwargs) -> List[str]: Number of molecules to generate. **kwargs Additional arguments forwarded to the family-specific generator. - For causal LMs, common options include ``max_length``, ``temperature``, - ``do_sample``, and ``scaffold``. For MolGen, use ``prefix_selfies`` or - ``scaffold`` plus optional ``num_beams``, ``min_length``, and - ``max_length``. For Molexar, use ``start_smiles``, ``start_string``, - ``generation_task``, or ``conditions`` for omni models. For SAFE-GPT, - use ``scaffold=`` with a SMILES prefix (converted to SAFE internally). + For causal LMs, common options include ``max_new_tokens``, + ``temperature``, ``do_sample``, and ``scaffold``. ``max_length`` is + still accepted as a Hugging Face total-length override. For MolGen, + use ``prefix_selfies`` or ``scaffold`` plus optional ``num_beams``, + ``min_length``, and ``max_length``. For Molexar, use ``start_smiles``, + ``start_string``, ``generation_task``, or ``conditions`` for omni + models. For SAFE-GPT, use ``scaffold=`` with a SMILES prefix + (converted to SAFE internally). Returns ------- @@ -331,13 +332,24 @@ def generate(self, n_samples: int = 10, **kwargs) -> List[str]: from .utils import smiles_to_safe scaffold = smiles_to_safe([scaffold])[0] + max_new_tokens = kwargs.pop("max_new_tokens", None) + max_length = kwargs.pop("max_length", None) + if max_new_tokens is not None and max_length is not None: + warnings.warn( + "Both max_new_tokens and max_length were passed; using max_new_tokens.", + stacklevel=2, + ) + max_length = None + elif max_new_tokens is None and max_length is None: + max_new_tokens = self.generate_max_length raw = generate_causal_lm( self.model, self.tokenizer, self.device, n_samples, family=self._family, - max_length=kwargs.pop("max_length", self.generate_max_length), + max_new_tokens=max_new_tokens, + max_length=max_length, temperature=kwargs.pop("temperature", 1.0), do_sample=kwargs.pop("do_sample", True), scaffold=scaffold, @@ -468,6 +480,44 @@ def _load_pretrained(self, local_path: Optional[str] = None) -> None: self.model.to(self.device) self.model.eval() + def _uses_safe_tokenizer(self) -> bool: + if self._safe_tokenizer is not None: + return True + if self._family in SAFE_GPT_FAMILIES: + return True + if self.repo_id is not None and resolve_family(self.repo_id) in SAFE_GPT_FAMILIES: + return True + return False + + def _save_generator_tokenizer(self, path: str) -> None: + """Save the tokenizer. SAFE's custom pre-tokenizer cannot use HF serialization.""" + if self._uses_safe_tokenizer(): + self._save_safe_gpt_tokenizer(path) + return + try: + self.tokenizer.save_pretrained(path) + except Exception as exc: + if "cannot be serialized" not in str(exc): + raise + self._save_safe_gpt_tokenizer(path) + + def _save_safe_gpt_tokenizer(self, path: str) -> None: + """Save the SAFE tokenizer JSON. The HF fast wrapper cannot be serialized.""" + from .utils import _require_safe + + _require_safe() + from safe.tokenizer import SAFETokenizer + + if self._safe_tokenizer is None: + tokenizer_kwargs = {} + if self.revision is not None: + tokenizer_kwargs["revision"] = self.revision + self._safe_tokenizer = SAFETokenizer.from_pretrained( + self.tokenizer_repo_id or self.repo_id, + **tokenizer_kwargs, + ) + self._safe_tokenizer.save_pretrained(path) + def _load_safe_gpt_tokenizer(self, tokenizer_repo: str): """Load the custom SAFE tokenizer as a Hugging Face fast tokenizer.""" from .utils import _require_safe @@ -478,8 +528,8 @@ def _load_safe_gpt_tokenizer(self, tokenizer_repo: str): tokenizer_kwargs = {} if self.revision is not None: tokenizer_kwargs["revision"] = self.revision - safe_tokenizer = SAFETokenizer.from_pretrained(tokenizer_repo, **tokenizer_kwargs) - tokenizer = safe_tokenizer.get_pretrained() + self._safe_tokenizer = SAFETokenizer.from_pretrained(tokenizer_repo, **tokenizer_kwargs) + tokenizer = self._safe_tokenizer.get_pretrained() tokenizer.model_max_length = self.max_length return tokenizer From 63fe00ed6e48ce8451b13687c8e6507d531fd88a Mon Sep 17 00:00:00 2001 From: m21hm9 Date: Mon, 21 Sep 2026 16:15:34 +0800 Subject: [PATCH 5/5] fix: removed personal notes --- molecule_notes/butina_optimization.md | 283 --------------- molecule_notes/generation_model.md | 473 -------------------------- molecule_notes/train_test_split.md | 385 --------------------- 3 files changed, 1141 deletions(-) delete mode 100644 molecule_notes/butina_optimization.md delete mode 100644 molecule_notes/generation_model.md delete mode 100644 molecule_notes/train_test_split.md diff --git a/molecule_notes/butina_optimization.md b/molecule_notes/butina_optimization.md deleted file mode 100644 index f4ba583..0000000 --- a/molecule_notes/butina_optimization.md +++ /dev/null @@ -1,283 +0,0 @@ -# Butina 聚类 / Split:原文 vs 现有实现与优化 - -本文记录 **Taylor–Butina** 在分子 ML 划分中的算法原意,以及 2024–2026 年各实现如何加速。 -核心结论:**精确 Butina 不需要全距离矩阵**;当前工业/新开源优化都走「阈值邻居 + exclusion sphere」。DeepChem/RDKit `ClusterData` 是正确但未针对大库优化的参考实现。 - -与本仓库关系:`method="butina"` 尚未实现。若做全集(如 QM9 ~13 万)且要求 **accuracy**,应对齐 Chalcedon / chemfp 的稀疏路径,而不是 DeepChem 的稠密 `dists` 列表。 - ---- - -## 1. 原文算法(必须对齐的「准确」定义) - -**文献:** Darko Butina, *Unsupervised Data Base Clustering Based on Daylight’s Fingerprint and Tanimoto Similarity*, J. Chem. Inf. Comput. Sci. **1999**, 39, 747–750. -https://doi.org/10.1021/ci9803381 - -**输入:** 分子指纹、Tanimoto 阈值 $t$(例如 0.65 表示「够像才算邻居」)。 - -**Tanimoto(Jaccard):** - -$$ -T(A,B)=\frac{|A\cap B|}{|A\cup B|}=\frac{c}{a+b-c} -$$ - -范围 $[0,1]$。距离常写 $d=1-T$。 - -**三步:** - -1. 生成指纹(原文 Daylight;现代几乎一律 Morgan/ECFP)。 -2. 对每个分子数邻居:$T \ge t$ 的个数;按邻居数 **降序** 排序(潜在簇中心)。 -3. **Exclusion sphere:** 取下一个未标记分子当中心,所有 $T \ge t$ 且未标记的邻居进该簇并标记;已标记者不再当中心、不进别簇。 - -**精确 Butina 真正需要的信息只有布尔邻居关系:** - -$$ -N(i,j)=\mathbf{1}[T(i,j)\ge t] -$$ - -不需要任何 $T < t$ 的数值。因此: - -| 做法 | 是否精确 Butina | -|------|-----------------| -| 全距离矩阵 + `Butina.ClusterData` | 是(小 $n$) | -| 阈值邻接表 + **同一套** 排序和收球 | **同样是**(可上全集) | -| LSH / 漏邻居 / 先抽子集再指派 | 否 | -| BitBIRCH 等别的聚类 | 否(即使质量「差不多」) | - ---- - -## 2. 对比总表 - -| | 原文 1999 | RDKit `ClusterData` | DeepChem / Datamol | Chalcedon (Rowan, 2026) | chemfp `butina` | FPSim2 (ChEMBL) | BitBIRCH | -|--|-----------|---------------------|--------------------|-------------------------|-----------------|-----------------|----------| -| **是不是 Butina** | 定义 | 是 | 是 | 是(对照过 RDKit 簇) | 是 | 否(只做精确相似度搜索) | **否** | -| **指纹** | Daylight | 调用方提供 | Morgan r=2, **1024** bit | Morgan r=2, **2048** bit | 调用方 / fps | RDKit 指纹库 | 二进制指纹 | -| **阈值语义** | 相似度 $T\ge t$ | **距离** `distThresh` | `cutoff` 传给 ClusterData = **距离**(默认 0.6) | `cutoff` = **距离**(博客默认 0.65 ⇒ $T\ge 0.35$) | `--threshold` = **相似度** | `threshold` = **相似度** | 另一套层次阈值 | -| **存什么** | 概念上的邻居 | 全部 condensed 距离 $O(n^2)$ | 全部 `dists` 列表 | 分块;峰值 $O(n)$ + 一块 workspace | **稀疏** NxN(只存 $T\ge t$) | 稀疏 CSR | 树 / iSIM 统计 | -| **加速** | 无 | Bulk 仍填满矩阵 | `BulkTanimotoSimilarity` | NumPy/BLAS、float32、上三角分块、两遍扫描 | POPCNT、稀疏矩阵可存 npz | POPCNT + Swamidass bound、多核/GPU | $O(n)$ 近似层次 | -| **大库** | 当时为替代难调的 Jarvis–Patrick | $n\sim 10^5$ **OOM** | 文档写明 $O(n^2)$,中小集 | **10 万 ≈ 21s / 2.4GB** | 工业级库检索 + 聚类 | 适合建邻居图 | 百万级,但 **不是同簇** | -| **可复现平局** | 按邻居数排序 | 实现相关 | 未强调 | 确定性收球 | 默认 `randomize`;应用 `first`/`last` | N/A | N/A | -| **适合本库** | 语义标准 | 小 $n$ **金标准测试** | API 对标对象,不要抄矩阵 | **全集精确路径的最佳开源参照** | 算法参照;依赖/许可需评估 | 可选搜索后端 | 不要叫 `method="butina"` | - ---- - -## 3. 原文 vs 各实现(分节) - -### 3.1 原文 Butina (1999) - -**问题:** Jarvis–Patrick 要调两个参数,簇要么极大且杂、要么碎;大库要手工调。 - -**方法:** 单阈值、簇中心与每个成员都满足 $T\ge t$、exclusion sphere。 - -**和 split 的关系:** 原文不做 train/test。后来 DeepChem 把 **每个簇当 group**,整组进 train 或 test(与 scaffold 相同贪心)。 - -**局限:** 阈值无唯一正确答案;簇大小极不均匀;指纹种类会改变结果。 - ---- - -### 3.2 RDKit `rdkit.ML.Cluster.Butina` - -**做法:** - -```text -dists = condensed 下三角 (1 - Tanimoto) # 必须全部 pair -clusters = Butina.ClusterData(dists, n, distThresh, isDistData=True) -# 每个簇第一个元素是 centroid -``` - -**优化:** 几乎没有。调用方可用 `DataStructs.BulkTanimotoSimilarity` 加快 **填矩阵**,矩阵本身仍是 $O(n^2)$ 内存。 - -**实测(他人报告):** - -- Macs in Chemistry:15 万分子聚类时内存涨到 **80–267 GB**,进程被杀。 -- Chalcedon 基准:$n=10^5$ **RDKit OOM**;$n=5\times 10^4$ 约 173s / 110GB RSS。 - -**本库用法:** $n \le 2000$ 的单元测试金标准,**禁止** 对 QM9 全量调用。 - ---- - -### 3.3 DeepChem `ButinaSplitter` / Datamol - -**源码要点**(`deepchem/splits/splitters.py`): - -- Morgan radius 2,**1024** bits(不是 2048)。 -- `BulkTanimotoSimilarity(fps[i], fps[:i])`,`dists.extend(1-x)`。 -- `Butina.ClusterData(..., cutoff, isDistData=True)`。 -- 簇按大小降序,再按 scaffold 同一套 cutoff 贪心填 train/val/test。 -- 文档:**$O(n^2)$**,主要为得到 novel chemotypes;默认 cutoff **0.6(距离)**。 - -Datamol 教程同一模式:`BulkTanimotoSimilarity(..., returnDistance=True)` + `ClusterData`。 - -**和原文差别:** - -- 指纹:Morgan ≠ Daylight。 -- 阈值:DeepChem 0.6 **距离** ≈ $T \ge 0.4$,比「相似度 0.65」松得多。 -- 无稀疏化。 - -**本库:** 可对标「split 时整簇分配」;不要对标其矩阵实现。默认阈值不要盲目抄 0.6 距离,除非文档写死语义。 - ---- - -### 3.4 Chalcedon(Rowan, 2026)— 当前最贴近「精确 + 全集」的开源包 - -**链接:** - -- 博客:https://www.rowansci.com/blog/chalcedon (Eli Mann, 2026-05-26) -- 代码:https://github.com/rowansci/chalcedon -- PyPI:`chalcedon` - -**动机:** 按 Walters 建议做 Butina split,但 GEOM ~30 万样本时现有开源实现要 **数 TB 内存**。 - -**仍是精确 Butina:** 分块实现 vs RDKit,在 10k–100k、cutoff=0.65 上他们要求 **簇相同**(bitwise 不完全相同,见浮点)。 - -**算法改写(内存线性的关键):** - -标准形式要持久化 $n\times n$ 距离。Chalcedon 分成三阶段: - -1. **分块**算 pair 相似度,**只持久化每个分子的邻居个数**,丢掉具体相似度。 -2. 按邻居数降序排序。 -3. 再按排序走,对 **仍未分配** 的集合 **分块重算** 该中心的相似度行,收走未分配邻居。 - -峰值 ≈ 一块 batch workspace + $O(n)$ 计数。 - -**其它工程优化:** - -- 全程 float32(sgemm 约 2× 于 dgemm);非二进制描述子他们建议 float64。 -- Cutoff 比较改写成 $|A\cap B| \ge (1-\mathrm{cutoff})\cdot |A\cup B|$,少中间数组、略减 ULP。 -- 预计算每行 $\|A\|^2$(对 0/1 向量即 popcount)。 -- 只走上三角;对角块递归切成矩形 GEMM(类似 SciPy ssyrk 思路)。 -- 预分配 buffer,避免反复 malloc。 - -**Split 部分:** 簇出来后用 Graham **LPT(最长处理时间)** 贪心:每次把最大簇分给「离目标比例最远」的那一份(train/val/test)。与 DeepChem「先填满 train cutoff」不完全相同,但同属整簇分配。 - -**基准(GEOM 子集,cutoff=0.65,Ryzen 9 7950X3D,128GB RAM):** - -Wall time (s): - -| n | Chalcedon chunked f32 | Chalcedon full matrix f32 | RDKit | -|---|----------------------|---------------------------|--------| -| 1,000 | 0.040 | 0.012 | 0.073 | -| 10,000 | 0.357 | 0.281 | 5.16 | -| 50,000 | 5.75 | 5.74 | 172.8 | -| 100,000 | **21.14** | 22.49 | **OOM** | - -Peak RSS (GB): - -| n | Chalcedon chunked f32 | Chalcedon full matrix f32 | RDKit | -|---|----------------------|---------------------------|--------| -| 10,000 | 0.33 | 0.35 | 4.5 | -| 50,000 | 1.23 | 3.16 | 110.4 | -| 100,000 | **2.35** | 11.17 | OOM | - -小 $n$ 全矩阵更快;大 $n$ 分块内存近线性,全矩阵仍 $O(n^2)$。 - -**准确性注意(对「i need accuracy」很重要):** -他们在 cutoff sweep 上发现 float32 重排比较 vs RDKit **不是逐 bit 相同**,部分 cutoff 簇数差 $\le 2.5\%$。原因是浮点,不是故意漏边。 - -要对齐 RDKit / 原文布尔 $T\ge t$:应用 **整数 popcount** 算 $c,a,b$ 再比较,不要用 float32 GEMM 当金标准。 - -**API 示例:** - -```python -splits = chalcedon.butina_split( - smiles, - fractions={"train": 0.8, "val": 0.1, "test": 0.1}, - cutoff=0.65, # 距离 cutoff - dtype="float32", -) -``` - ---- - -### 3.5 chemfp(工业精确稀疏 Butina) - -**文档:** https://chemfp.com/docs/chemfp_butina_command.html - -**做法:** - -1. `threshold_tanimoto_search_symmetric` 生成 **稀疏** 相似度矩阵(只含 $T \ge t_{\mathrm{NxN}}$)。 -2. 可存 `.npz`;用更高的 `--butina-threshold` 调参时 **不必重算 Tanimoto**。 -3. 按行邻居数排序 + exclusion sphere。 -4. `--tiebreaker randomize|first|last`(默认可随机,**可复现必须 first/last + seed**)。 -5. 可选:false singleton 贴最近中心(**原文没有**,属于后处理)。 - -**优化本质:** 与「邻接表精确 Butina」相同,搜索层用 POPCNT 等。 -免费/商业版性能差一截;加依赖前要看许可。 - ---- - -### 3.6 FPSim2(ChEMBL)— 精确搜索后端,不是聚类 - -**文档:** https://chembl.github.io/FPSim2/ - -- CPU POPCNT;高阈值($\ge 0.7$)更合适。 -- **Swamidass & Baldi 2007** bound(doi: [10.1021/ci600358f](https://doi.org/10.1021/ci600358f)): - $T \le \min(a,b)/\max(a,b)$。query 亮位数 $A$、库分子 $B$ 必须落在 $[tA,\ A/t]$,否则 **不可能** 是邻居 → **剪枝不漏真阳性**。 -- 多核、可选 GPU;`symmetric_distance_matrix(threshold=...)` → SciPy CSR。 - -可放在 Butina 流水线的「建邻居图」一步;exclusion 仍要自己写(或交给 chemfp/Chalcedon)。 - ---- - -### 3.7 BitBIRCH — 不要当成优化版 Butina - -Miranda-Quintana 等, *Efficient clustering of large molecular libraries*, bioRxiv 2024. -https://www.biorxiv.org/content/10.1101/2024.08.10.607459v1 - -- 相对 RDKit Taylor–Butina:45 万分子,4TB 仍不够跑矩阵版;BitBIRCH 约 2 分钟。 -- 150 万分子号称 >1000×。 -- 质量指标在部分阈值区间与 Butina「无显著差」或更好。 -- Chalcedon 对比:**同名义阈值下簇数完全不是一回事**($n=10^5$:Chalcedon 9590 vs BitBIRCH-Lean 51232)。 - -这是 **另一种算法**(BIRCH + iSIM)。更快,但不能命名为 `method="butina"`。 - ---- - -## 4. 两种「精确」建邻居方式(和 Chalcedon 的关系) - -```text -A. 全距离矩阵 + ClusterData - 小 n、易与 RDKit/DeepChem 逐簇对齐 - QM9 全量:内存不可行 - -B. 阈值邻接 / 分块邻居 + 同一套排序收球 - 边集完整 ⇒ 与 A 数学等价 - Chalcedon:分块 + 两遍(先计数再收球),线性内存 - chemfp/FPSim2:稀疏阈值搜索(POPCNT ± bound) -``` - -Chalcedon 的巧妙处:第一遍 **甚至可以不存邻接表**,只存度数;第二遍对未分配子集重算。 -代价是部分 pair 会算两次;换来内存 $O(n)$ 而不是 $O(n\bar{k})$ 邻接表。 -chemfp 选择 **存稀疏边**,调阈值更快。都精确,工程权衡不同。 - ---- - -## 5. 对 torch-molecule 的建议 - -1. **小 $n$(测试):** RDKit `BulkTanimotoSimilarity` + `ClusterData`,作为簇相等的金标准。 -2. **大 $n$ / 用户可能丢全集:** Chalcedon 式分块 **或** 稀疏邻接表 + 自写 exclusion;**禁止** 分配 $n(n-1)/2$ 距离数组。 -3. **要比准 RDKit:** 整数 popcount Tanimoto,不要默认 float32 GEMM。 -4. **阈值 API:** 对外只用 `similarity_cutoff`(如 0.65),内部再转距离;文档写清 DeepChem 0.6 是距离。 -5. **指纹默认:** Morgan r=2, 2048 bit(Chalcedon / Walters);若要对齐 DeepChem 再提供 1024。 -6. **平局:** 稳定 `(-度数, index)`,相当于 chemfp `first`,不要默认 randomize。 -7. **不要** 在 `method="butina"` 里偷偷 subsample 或换 BitBIRCH。 -8. 依赖:能零依赖实现稀疏/分块最好;Chalcedon MIT 且专做 split,可作实现参照,不必强绑。 - ---- - -## 6. 参考文献与链接 - -| 主题 | 来源 | -|------|------| -| 原文 | Butina, JCICS 1999, doi:10.1021/ci9803381 | -| 精确剪枝 bound | Swamidass & Baldi, JCIM 2007, doi:10.1021/ci600358f | -| DeepChem splitter | https://github.com/deepchem/deepchem/blob/master/deepchem/splits/splitters.py | -| Chalcedon | https://www.rowansci.com/blog/chalcedon ;https://github.com/rowansci/chalcedon | -| chemfp Butina | https://chemfp.com/docs/chemfp_butina_command.html | -| FPSim2 | https://chembl.github.io/FPSim2/ | -| BitBIRCH | https://www.biorxiv.org/content/10.1101/2024.08.10.607459v1 | -| 大库聚类实践 | https://macinchem.org/2023/03/05/options-for-clustering-large-datasets-of-molecules/ | -| Walters split 评论 | http://practicalcheminformatics.blogspot.com/2024/11/some-thoughts-on-splitting-chemical.html | - ---- - -*记录日期:对照 Chalcedon 2026-05 博客、chemfp 4.x 文档、DeepChem 源码与 FPSim2/BitBIRCH 文献。待 `method="butina"` 实现时以第 5 节为准。* diff --git a/molecule_notes/generation_model.md b/molecule_notes/generation_model.md deleted file mode 100644 index b742987..0000000 --- a/molecule_notes/generation_model.md +++ /dev/null @@ -1,473 +0,0 @@ -# Report: HF pretrained molecular generator - -本文记录 `torch-molecule` 接入 Hugging Face 预训练分子生成模型的工作:对外一个类 `HFPretrainedMolecularGenerator`,sklearn 风格 `fit` / `generate`,与 `HFPretrainedMolecularEncoder` 对称。依据仓库内的 `generation_model.md`、`issues_to_be_fixed.md` 以及 `torch_molecule/generator/pretrained/` 的现有代码;不把未实现的接口写成已完成。 - ---- - -## 1. Purpose - -现有 `HFPretrainedMolecularEncoder` 只用 Hugging Face `AutoModel` 做编码,不做生成。仓库里已有 8 个自研生成器(`LSTMMolecularGenerator`、`MolGPTMolecularGenerator`、`DigressMolecularGenerator`、`GDSSMolecularGenerator`、`GraphDITMolecularGenerator`、`DeFoGMolecularGenerator`、`JTVAEMolecularGenerator`、`GraphGAMolecularGenerator`),它们**没有被替换**。 - -这次新增的是第三条路径:从 Hugging Face Hub 加载已预训练的生成权重,用**一个**对外类 `HFPretrainedMolecularGenerator` 提供与 LSTM / MolGPT 相同的 sklearn 风格接口: - -- `fit()` 无数据:下载并加载 Hub 权重 -- `fit(smiles)`:加载预训练权重后再微调 -- `generate()`:返回 `List[str]` SMILES - -六个 Hub 模型的架构不同(因果 LM、seq2seq SELFIES、Fragment-SELFIES + 官方推理引擎),无法共用同一个 `generate()`。家族分流(`novomolgen` / `gp_molformer` / `molgen` / `molexar` / fallback `causal_lm`)只在 `generator/pretrained/` 内部完成。用户不直接实例化 `families/` 里的类;那些模块是内部 dispatch,不是对外 API。 - -数据集层(`load_qm9()`、`load_zinc250k()`、`SMILESDataset`、`MolecularInputChecker`)没有改。用户始终传入 SMILES、始终拿到 SMILES。SELFIES / Fragment-SELFIES 转换只发生在 `generator/pretrained/` 内部。 - ---- - -## 2. What we added (feature work) - -对应 `generation_model.md` 的 Phase 1–5。**Phase 6(文档站点 API 页、README 模型列表、CI 分层)仍未完成**,见第 3 节末尾。 - -### 2.1 对外 API - -```python -from torch_molecule import HFPretrainedMolecularGenerator -``` - -导出链:`torch_molecule/generator/pretrained/__init__.py` → `torch_molecule/__init__.py`(已加入 `__all__`)。 - -与 encoder 的对称关系: - -| | `HFPretrainedMolecularEncoder` | `HFPretrainedMolecularGenerator` | -|---|---|---| -| `fit()` 无参 | 加载 `AutoModel` | 加载 `AutoModelForCausalLM` / `AutoModelForSeq2SeqLM`,或 Molexar 官方 `MolexarInference` | -| `fit(X)` | 不支持 | 可选微调 | -| 主方法 | `encode()` → embedding | `generate()` → SMILES | - -**`fit()` 语义与 LSTM 不同:** - -- `LSTMMolecularGenerator.fit(X_train)`:从零初始化网络并在 SMILES 上训练。 -- `HFPretrainedMolecularGenerator.fit()`:无 `X` 时只从 Hub 拉预训练权重;`fit(smiles)` 是加载权重后再微调,不是从零训练。 - -`generate(n_samples=...)` 返回 `List[str]` SMILES。MolGen / Molexar 解码失败的条目会被丢掉,返回长度可以小于 `n_samples`(见问题 2、7)。 - -条件微调参数 `y` 目前会 warn 后忽略,走无条件语言模型微调(见「仍未做」)。 - -### 2.2 用户侧始终是 SMILES - -| 方向 | 约定 | -|---|---| -| `fit(X)` 输入 | `List[str]` SMILES | -| `generate()` 输出 | `List[str]` SMILES | -| SELFIES / Fragment-SELFIES | 仅内部:`_encode_inputs()` / `_decode_outputs()` | -| 数据集 loader | **未改** | - -内部转换(`modeling_pretrained.py`): - -- MolGen(seq2seq):`smiles_to_selfies` / `selfies_to_smiles` -- Molexar:`smiles_to_fragment_selfies` / `fragment_selfies_to_smiles` -- NovoMolGen / GP-MoLFormer / fallback causal LM:直接用 SMILES(生成后去掉空格) - -### 2.3 六个 Hub 模型与内部家族 - -| 模型 | `repo_id` | 内部 family | 表示 | 加载 / 生成要点 | -|---|---|---|---|---| -| NovoMolGen | `chandar-lab/NovoMolGen_32M_SMILES_BPE` | `novomolgen` | SMILES + BPE | 因果 LM;`revision` 默认 `hf-checkpoint` | -| GP-MoLFormer | `ibm-research/GP-MoLFormer-Uniq` | `gp_molformer` | SMILES | 因果 LM;`scaffold=` 骨架补全;tokenizer 默认 `ibm-research/MoLFormer-XL-both-10pct`;IBM remote code 需要 `transformers<=4.56.2` | -| MolGen-large | `zjunlp/MolGen-large` | `molgen` | SELFIES | seq2seq;默认苯环 prefix | -| MolGen-large-opt | `zjunlp/MolGen-large-opt` | `molgen` | SELFIES | 同上,权重已偏 QED / p-logP | -| Molexar-10M-base | `fairydance/molexar-10m-base` | `molexar` | Fragment-SELFIES | 官方 `MolexarInference` | -| Molexar-10M-omni | `fairydance/molexar-10m-omni` | `molexar` | Fragment-SELFIES | 同引擎;`conditions` / 性质键等 omni kwargs | - -未知 `repo_id`:`resolve_family()` 返回 fallback `causal_lm` 并 `warnings.warn`(问题 6 修过后,前缀能匹配的变体如 `chandar-lab/NovoMolGen_157M` **不再**误报)。 - -`FAMILY_PREFIXES` 允许同一前缀下的变体(例如其它 NovoMolGen 尺寸)映射到对应家族,而不必全部写进 `KNOWN_REPOS`。 - -### 2.4 分阶段(Phase 1–5 已落地) - -**Phase 1 — 骨架 + NovoMolGen(MVP)** - -- `registry.py`、`modeling_pretrained.py`、`families/causal_lm.py` -- `__init__.py` 导出 -- `tests/generator/hfpretrained.py` smoke(含 `@pytest.mark.integration` 的 NovoMolGen `fit()` + `generate()`) -- optional extra:`[hf-gen]` - -**Phase 2 — GP-MoLFormer + SELFIES 工具** - -- `utils.py`:`smiles_to_selfies` / `selfies_to_smiles` -- `families/causal_lm.py`:de novo + `scaffold=` -- `compat.py`:`transformers>=4.57` 时对 GP-MoLFormer raise;为 IBM remote code 提供 `transformers.onnx` stub -- extra:`[gp-molformer]`(`transformers>=4.40,<=4.56.2`) - -**Phase 3 — MolGen-large / MolGen-large-opt** - -- `families/seq2seq.py` -- `generate(prefix_selfies=...)`;未提供时使用默认苯环 SELFIES(见「仍未做」) -- extra:`[molgen]`(`selfies>=2.1.0`,无上界) - -**Phase 4 — Molexar** - -- `utils.py`:Fragment-SELFIES 转换 -- `families/molexar.py`:wrap 官方 `MolexarInference`(de novo、`start_smiles` / fragment 约束、omni `conditions`) -- extra:`[molexar]` - -**Phase 5 — 微调与本地存盘** - -- `finetune.py`:按家族 dispatch - - 因果 LM(NovoMolGen / GP-MoLFormer / fallback):next-token prediction(`finetune_causal_lm`) - - MolGen seq2seq:denoising(`finetune_seq2seq`,问题 9 之后:mask encoder、labels 保持干净) - - Molexar:官方 training template 后再走因果 LM loop(`finetune_molexar`) -- `checkpoint.py` + `save_to_local` / `load_from_local`:HF 目录(`save_pretrained`)+ `hf_generator_metadata.json` -- `save_to_hf()` **仍是** `NotImplementedError` -- `load_from_hf()` / `load()` 无本地 path 时等价于 `fit()` - -**Phase 6 — 仍开放**(见第 3 节)。 - -### 2.5 目录布局 - -``` -torch_molecule/generator/pretrained/ -├── __init__.py # 只导出 HFPretrainedMolecularGenerator -├── modeling_pretrained.py # 唯一对外类 -├── registry.py # repo_id → family -├── utils.py # SMILES ↔ SELFIES / Fragment-SELFIES -├── finetune.py # 家族微调 -├── checkpoint.py # 本地 metadata -├── compat.py # GP-MoLFormer / transformers 兼容 -└── families/ - ├── __init__.py - ├── causal_lm.py # NovoMolGen, GP-MoLFormer, fallback - ├── seq2seq.py # MolGen-large, MolGen-large-opt - └── molexar.py # Molexar base / omni -``` - -### 2.6 Optional extras(`pyproject.toml`) - -| extra | 依赖 | -|---|---| -| `[hf-gen]` | `transformers>=4.40`, `accelerate` | -| `[molgen]` | `selfies>=2.1.0`, `transformers>=4.40`, `accelerate` | -| `[gp-molformer]` | `transformers>=4.40,<=4.56.2`, `accelerate` | -| `[molexar]` | `fragment-selfies>=1.0.0`, `transformers>=4.40`, `accelerate`, `loguru`, `molexar @ git+https://github.com/fairydance/Molexar.git` | - -`selfies` **没有**在 `pyproject.toml` 里 pin `<3`(老师要求;问题 5 用 README / `install.rst` 说明)。 - -### 2.7 存盘 - -- `save_to_local(path)`:`model.save_pretrained` + `tokenizer.save_pretrained` + `hf_generator_metadata.json` -- `load_from_local(path)`:读 metadata,再按家族从该目录加载 -- `save_to_hf(...)`:`NotImplementedError`(「HFPretrainedMolecularGenerator does not support saving to Hugging Face.」) - -注意:`fit()` **每次**都会 `_load_pretrained()` 从 Hub 再拉一遍。即使刚 `load_from_local`,再调用 `fit()` 仍会覆盖为 Hub 权重。老师要求这一轮先不动(见「仍未做」)。 - -### 2.8 微调:MolGen 是 denoising,不是 identity copy - -问题 9 修完后,仅 `finetune_seq2seq` 改变目标: - -- encoder `input_ids`:非 special token 以 `mask_prob=0.15` 换成 `` -- `labels`:仍是干净序列 -- `fit(X)` 签名不变 -- **不影响** NovoMolGen、GP-MoLFormer、Molexar、LSTM - -### 2.9 用法片段(NovoMolGen) - -```python -from torch_molecule import HFPretrainedMolecularGenerator - -model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", -) -model.fit() # 只加载 Hub 预训练权重,不训练 -smiles_list = model.generate(n_samples=10) - -# 微调:仍传 SMILES(与 LSTM 不同:这里是在预训练权重上继续训) -model.fit(["CCO", "CC(=O)O", "c1ccccc1"]) -``` - ---- - -## 3. Issues we found and how we handled them - -依据 `issues_to_be_fixed.md`。原则:进模型(训练数据)坏样本默认 **raise**,失败要带 index;出模型(生成)坏样本可以丢掉,但 `n_samples` 是**尝试条数**,返回长度 = 解码成功条数,并留下计数警告。 - -问题 2(MolGen 用 `""` 撑长度)和问题 7(Molexar 静默变短)**原 bug 不同**;修完后共用同一条返回约定:不要补 `""`,成功几条返回几条。 - -### 3.1 问题 1 — `fit()` 坏 SMILES:keep raise,不 skip - -**原问题:** `fit(X)` 里一条坏样本会让整次微调失败。讨论过是否 skip / 重试。 - -**决策 / 修法:** **不加重试、不加 skip。** 坏样本就 raise,与 `MolecularInputChecker` 一致。 - -- RDKit 过不了:`MolecularInputChecker` 给出带 index 的错误(`Invalid SMILES structure at index {idx}: ...`),`_validate_inputs` 聚合成 `ValueError`。`smiles_to_selfies` 里同样对 `MolFromSmiles is None` raise `ValueError(f"Invalid SMILES at index {idx}: ...")`。 -- RDKit 过了、`selfies.encoder` 挂:见问题 4,同样 raise,不 skip。 - -**状态:solved。** - -**关键文件:** `torch_molecule/utils/checker.py`、`torch_molecule/generator/pretrained/utils.py`、`modeling_pretrained.py`(`fit` → `_validate_inputs`)、`tests/generator/hfpretrained.py`(`test_smiles_to_selfies_invalid_smiles`)。 - -### 3.2 问题 2 — MolGen `generate` 用 `""` 占位,`len == n_samples` 看起来成功 - -**原问题(仅 MolGen 这条路径):** 解码失败时 `selfies_to_smiles` 写入空串。`len(out) == n_samples` 像成功,空串比短列表更坑。 - -这与问题 7 **不是同一个原 bug**:这里是用 `""` **把长度撑满**;问题 7 是 Molexar **直接丢掉且不警告**。 - -**决策 / 修法:** 不要 `""`,也不要重试凑满。成功几条就返回几条(2 条有效 → 返回 2)。`n_samples` = 尝试条数。 - -例:`generate(n_samples=5)`,其中 2 条解不出: - -```text -原来: ["c1ccccc1", "CCO", "", "CC(=O)O", ""] # len=5,两个假成功 -现在: ["c1ccccc1", "CCO", "CC(=O)O"] # len=3,只留真分子 -``` - -实现上 `_decode_outputs` 对 seq2seq 丢掉无效 SMILES,并 warn `got {k}/{n} valid SMILES`。 - -**状态:solved。** - -**关键文件:** `utils.py`(`selfies_to_smiles`)、`modeling_pretrained.py`(`_decode_outputs`)、`tests/generator/hfpretrained.py`(`test_selfies_to_smiles_drops_invalid_entries`、`test_decode_outputs_molgen_drops_empty_and_warns`)。 - -### 3.3 问题 3 — `except Exception` 把 `ImportError` 吞成 `""` - -**原问题:** `selfies_to_smiles` 和 `fragment_selfies_to_smiles` 都是 `except Exception: append("")`。依赖缺失、键盘中断等也会变成空串,且没有计数、没有日志。 - -**决策 / 修法:** **不要** `return_stats` 新 API。删掉 `except Exception`。分两种: - -- 不该发生的错(`ImportError`、其它意外)→ **raise**,不要吞成 `""` -- 分子解不出来(`DecoderError`、RDKit `mol is None`、空串)→ **不要 raise**(否则问题 2 无法返回成功的 2/3 条)。丢掉这条,**warn**(例如 `dropped {n} invalid SELFIES` / `dropped {n} invalid Fragment-SELFIES`) - -**状态:solved。** - -**关键文件:** `torch_molecule/generator/pretrained/utils.py`。 - -### 3.4 问题 4 — RDKit-valid 但 `selfies.encoder` 失败(dummy `*`) - -**原问题:** `smiles_to_selfies` 在 canonical 之后直接 `sf.encoder(canonical)`,没有接 `EncoderError`。`fit` 带着库异常全挂,看不出是第几条。dummy atom(`*CCO`、`[*]CCO`、`[*]c1ccccc1`)是典型触发点。 - -**决策 / 修法:** **不加 skip。** 接住后 **raise** 成带 index 的 `ValueError`: - -```python -raise ValueError( - f"SMILES at index {idx} is RDKit-valid but not SELFIES-encodable: {smiles_string}" -) from exc -``` - -**状态:solved。** - -**关键文件:** `utils.py`(`smiles_to_selfies`)、`tests/generator/hfpretrained.py`(`test_smiles_to_selfies_encoder_error_has_index`)。 - -### 3.5 问题 5 — `selfies>=2.1` 无上界;老师不要求 pyproject pin `<3` - -**原问题:** `pyproject.toml` 是 `selfies>=2.1.0`,没有上界。3.x 字母表若破坏兼容,安装器不会挡。 - -**决策 / 修法:** **不改 pyproject 的版本上界。** 老师意见:放到 README **Additional Packages**(以及 `docs/source/install.rst`),和其他可选依赖一样写明: - -| Model | Required Packages | -|---|---| -| HFPretrainedMolecularEncoder | transformers | -| HFPretrainedMolecularGenerator | transformers | -| HFPretrainedMolecularGenerator (MolGen) | transformers, selfies 2.x (3.x not guaranteed) | -| HFPretrainedMolecularGenerator (GP-MoLFormer) | transformers<=4.56.2 | -| HFPretrainedMolecularGenerator (Molexar) | transformers, fragment-selfies, molexar | - -README 另有安装示例:`pip install "selfies>=2.1"`、`pip install torch-molecule[molgen]`、`[gp-molformer]`、`[molexar]`。 - -`[gp-molformer]` extra 在 pyproject 里 **有** `transformers<=4.56.2`;MolGen 的 `selfies` 仍只有 `>=2.1.0`。运行时 GP-MoLFormer 还会在 `compat.ensure_gp_molformer_transformers_compat()` 对 `>=4.57` raise `ImportError`。 - -**状态:solved(文档,不是 pyproject 给 selfies 加上界)。** - -**关键文件:** `README.md`、`docs/source/install.rst`、`pyproject.toml`、`compat.py`。 - -### 3.6 问题 6 — 未知 `repo_id` 警告只查精确 `KNOWN_REPOS` - -**原问题:** `__init__` 只查 `KNOWN_REPOS` 的完整字符串。`chandar-lab/NovoMolGen_157M` 会被 `FAMILY_PREFIXES` 正确分成 `novomolgen`,却仍警告 “family may not be implemented”。 - -**决策 / 修法:** 在 `resolve_family()` 之后,**只对 fallback `causal_lm` 警告**。 - -**状态:solved。** - -**关键文件:** `modeling_pretrained.py`(`__init__`)、`registry.py`(`resolve_family`)、`tests/generator/hfpretrained.py`(`test_known_family_prefix_does_not_warn_unknown_repo`、`test_unknown_repo_fallback_warns`)。 - -### 3.7 问题 7 — Molexar `generate()` 列表变短且无警告 - -**原问题(Molexar,与问题 2 不同):** `_decode_outputs` 对 Molexar `if smiles` 过滤,解码失败直接丢掉。调用方要 10 条可能拿到 6 条,**没有警告**。 - -**决策 / 修法:** 要 10 条、只有 6 条成功,就用这 6 条,**不要补 `""` 凑满**。缺的是警告,例如 `got 6/10 valid SMILES`。 - -修完后与问题 2 **共用返回约定**(成功几条返回几条 + 计数 warn),但原缺陷分别是「假满长度」vs「静默变短」。 - -**状态:solved。** - -**关键文件:** `modeling_pretrained.py`(`_decode_outputs`,seq2seq 与 molexar 共用 warn)、`utils.py`(`fragment_selfies_to_smiles`)。 - -### 3.8 问题 8 — 边角化学:confirmed,不为转换层加特殊 case - -**原问题:** 担心电荷、立体、dummy `*`、叠氮等边角 SMILES 需要单独一套转换逻辑。 - -**结论(已扫 12 大类、约 45 条 SMILES,`selfies` 2.1.1):** - -| 结果 | 条数 | -|---|---| -| OK | 40 | -| RDKit 拒 | 1 | -| `EncoderError` | 4(dummy `*`;`[*]CCO` 在 dummy 和 Molexar 挂点里各计一次) | -| decode 挂 | 0 | - -会 roundtrip 的不必为转换层加特殊处理,包括:电荷、两性离子、四面体/顺反/联烯、萘/桥环/螺环/大环/三元环、吡啶、Kekulé 苯、三键、过氧、高价 S/P、N-oxide、硝基、自由基、卡宾、显式氢、有机叠氮 `CCN=[N+]=[N-]`、重氮、同位素、`.` 断开、Si/Se/膦/硼酸根、`[Cu+2]` / `[Fe]` / 类格氏。 - -会失败的走问题 1 / 4(进模型 raise),不是新分支: - -| SMILES | 卡在哪 | 行为 | -|---|---|---| -| `C[N-]=[N+]=N` | RDKit 不认(价态) | **raise** 问题 1。有机叠氮请用 `CCN=[N+]=[N-]` | -| `*CCO` | RDKit 能吃,`selfies.encoder` 挂 | **raise** 问题 4 | -| `[*]CCO` | 同上 | **raise** 问题 4 | -| `[*]c1ccccc1` | 同上(Molexar 挂点写法) | **raise** 问题 4 | - -只有 **MolGen** 走 SMILES ↔ SELFIES,才会在 `*` 上撞问题 4。NovoMolGen / GP-MoLFormer / LSTM 不走这条转换。 - -**决策:** 转换层不加边角化学 special case。失败路径已由问题 1、4 覆盖。`issues_to_be_fixed.md` 提到可选把 3 个 `*` 和 1 个坏叠氮加进回归测试;当前 `tests/generator/` **没有**这组 SMILES 作为独立回归用例。 - -**状态:confirmed;no extra conversion cases。** - -**关键文件:** `utils.py`(通用 raise / drop,无化学分类表)。 - -### 3.9 问题 9 — MolGen 微调曾是 identity reconstruction - -**原问题:** MolGen 预训练是 **denoising seq2seq**(损坏的 SELFIES → 还原完整 SELFIES)。实现曾把 `labels = input_ids`,等于 identity copy,和论文目标不一致。 - -**决策 / 修法:** **只改** `finetune_seq2seq`: - -- `corrupt_token_ids(..., mask_prob=0.15)`:非 special token 换成 tokenizer 的 `` -- `labels` 仍是干净序列 -- `fit(X)` 签名不变 -- NovoMolGen / GP-MoLFormer / Molexar / LSTM 微调不变 - -**状态:solved(fixed)。** - -**关键文件:** `torch_molecule/generator/pretrained/finetune.py`(`corrupt_token_ids`、`finetune_seq2seq`)、`tests/generator/test_finetune.py`(`test_corrupt_token_ids_masks_non_special_tokens`、`test_finetune_seq2seq_labels_stay_clean`)。 - -### 3.10 仍未做 - -`issues_to_be_fixed.md` 写明这一轮不做,以及 `generation_model.md` Phase 6 仍是未勾选: - -1. **`fit()` 在 `load_from_local` 之后仍会从 Hub 再加载。** `fit()` 无条件调用 `_load_pretrained()`(无 `local_path` 时用 `self.repo_id`)。老师要求先不动。 -2. **MolGen 默认苯环 prefix。** `families/seq2seq.py` 中 `DEFAULT_MOLGEN_PREFIX_SELFIES = "[C][=C][C][=C][C][=C][Ring1][=Branch1]"`;未传 `prefix_selfies` / `scaffold` 时仍用它。这一轮明确不改默认。 -3. **Phase 6 文档与 CI** - - `docs/source/api/generator.rst` **尚未**收录 `HFPretrainedMolecularGenerator`(仍只有 8 个自研生成器)。 - - README「List of Supported Models → Generative Models」仍是 8 个自研模型,没有把 6 个 HF repo 列进去(Additional Packages 表已有 generator extras,那是问题 5,不是 Phase 6 的模型列表)。 - - CI:`pyproject.toml` 已声明 pytest marker `integration`,测试里 Phase 1/2/3/4 的 Hub 下载用例标了 `@pytest.mark.integration`;仓库 `.github/workflows/` 目前只有 `docs.yml`,**没有**「Phase 1 必跑、Phase 3+ 用 optional marker」的测试 CI。 -4. **`save_to_hf()` 对本类未实现**(`NotImplementedError`)。不要把它写成可用。 -5. **带 `y` 的条件微调未实现。** `fit(X, y)` 若 `y is not None` 会 `UserWarning`(「Conditional fine-tuning with y is not implemented yet」),然后 `y = None`,继续无条件 LM 微调。 - ---- - -## 4. Current data flow (SELFIES models) - -用户侧不变:`load_zinc250k().data` 仍是 SMILES;`fit` / `generate` 仍收发 SMILES。 - -``` -用户 SMILES - │ - ▼ -HFPretrainedMolecularGenerator.fit / generate - │ - ├─ MolGen fit - │ SMILES → _validate_inputs (RDKit) - │ → smiles_to_selfies # 坏样本 raise(问题 1 / 4) - │ → tokenize - │ → finetune_seq2seq # denoising:mask encoder,labels 干净 - │ - ├─ MolGen generate - │ prefix SELFIES(默认苯环,或 prefix_selfies= / scaffold=) - │ → seq2seq model.generate - │ → selfies_to_smiles # 解不出:warn + drop,不补 "" - │ → 可能再 warn got k/n valid SMILES - │ - ├─ Molexar - │ fit: SMILES → smiles_to_fragment_selfies → 官方 training template → causal LM loop - │ generate: MolexarInference → fragment_selfies_to_smiles → drop invalids + warn - │ - └─ Causal LM(NovoMolGen / GP-MoLFormer / fallback) - SMILES 直接 tokenize / 生成;不走 smiles_to_selfies / - selfies_to_smiles / smiles_to_fragment_selfies / fragment_selfies_to_smiles - NovoMolGen: BOS → generate;默认 revision hf-checkpoint - GP-MoLFormer: de novo 或 scaffold= 前缀 -``` - ---- - -## 5. Files touched (high level) - -**新增(生成器核心)** - -- `torch_molecule/generator/pretrained/modeling_pretrained.py` -- `torch_molecule/generator/pretrained/registry.py` -- `torch_molecule/generator/pretrained/utils.py` -- `torch_molecule/generator/pretrained/finetune.py` -- `torch_molecule/generator/pretrained/checkpoint.py` -- `torch_molecule/generator/pretrained/compat.py` -- `torch_molecule/generator/pretrained/__init__.py` -- `torch_molecule/generator/pretrained/families/causal_lm.py` -- `torch_molecule/generator/pretrained/families/seq2seq.py` -- `torch_molecule/generator/pretrained/families/molexar.py` -- `torch_molecule/generator/pretrained/families/__init__.py` - -**导出与依赖** - -- `torch_molecule/__init__.py`(加入 `HFPretrainedMolecularGenerator`) -- `pyproject.toml`(`[hf-gen]` / `[molgen]` / `[gp-molformer]` / `[molexar]`,以及 pytest `integration` marker) - -**测试** - -- `tests/generator/hfpretrained.py` -- `tests/generator/test_finetune.py` -- `tests/generator/test_causal_lm.py` -- `tests/generator/test_seq2seq.py` -- `tests/generator/test_molexar.py` - -**文档(问题 5;不是 Phase 6 API 页)** - -- `README.md` Additional Packages -- `docs/source/install.rst` Additional Packages - -**未改(按设计)** - -- `torch_molecule/datasets/*`(loader、CSV) -- `torch_molecule/encoder/pretrained/*` -- 8 个自研生成器实现 - -**Phase 6 仍未改** - -- `docs/source/api/generator.rst` -- README 生成模型一览表(仍 8 个自研) - ---- - -## 6. How to try it - -```bash -pip install -e ".[hf-gen]" -``` - -按模型再装 extras: - -```bash -pip install -e ".[molgen]" # MolGen:selfies 2.x;3.x 不保证 -pip install -e ".[gp-molformer]" # pins transformers<=4.56.2 -pip install -e ".[molexar]" # fragment-selfies + molexar -``` - -NovoMolGen 推理: - -```python -from torch_molecule import HFPretrainedMolecularGenerator - -model = HFPretrainedMolecularGenerator( - repo_id="chandar-lab/NovoMolGen_32M_SMILES_BPE", -) -model.fit() -print(model.generate(n_samples=5)) -``` - -需要下载 Hub 权重的测试标了 `@pytest.mark.integration`,例如: - -```bash -pytest tests/generator/hfpretrained.py -m "not integration" -pytest tests/generator/test_finetune.py -``` diff --git a/molecule_notes/train_test_split.md b/molecule_notes/train_test_split.md deleted file mode 100644 index bd18a10..0000000 --- a/molecule_notes/train_test_split.md +++ /dev/null @@ -1,385 +0,0 @@ -# Train / Test Split 实现方案 - -本文档记录 `torch-molecule` 中分子数据划分(train/test split)的设计与实现计划。 -背景:分子数据不能简单使用 sklearn 的 `train_test_split`,需支持 structure-aware 划分。 - ---- - -## 1. 设计原则 - -1. **Split 是数据工具,不是模型工具** — 不要写进 `BaseMolecularPredictor.fit()`。 -2. **返回值仍是** `SMILESDataset` — 与现有 `load_qm9()` → `fit()` 流程对接。 -3. `method` 支持 `random`、`scaffold`、`butina`、`size`。三路划分见 3.3。 -4. **SizeShiftReg 的 coarsening/CMD 不属于 split** — 那是 SSR 模型的训练正则,split 只按原子数切分。 -5. **默认** `method="random"` — 保证可复现、与旧脚本一致;文档说明 random 分数往往偏乐观。 - ---- - -## 2. 文件结构 - -``` -torch_molecule/datasets/ - constant.py # SMILESDataset 增加 train_test_split / subsample - split.py # 新增:各划分方法实现 - __init__.py # 导出 train_test_split, SMILESDataset - -tests/datasets/ - test_split.py # 新增单元测试 - -README.md # 更新:加载 → 划分 → fit 完整示例 -``` - -**不要修改:** `predictor/*/modeling_*.py`、`base/predictor.py`(训练接口已足够)。 - -**依赖:** RDKit(已在 `pyproject.toml`)只负责 SMILES → 指纹。无需 DeepChem、无需 chemfp。 - -**增强版 Butina:** 不是 RDKit `ClusterData` 的稠密距离矩阵。生产路径是 **稀疏阈值邻居图 + exclusion sphere**(chemfp / Chalcedon 同一精确算法,内存 $O(n+E)$)。原文、RDKit、chemfp、Chalcedon 的对照表见下方 5.3;更长的文献笔记见 `molecule_notes/butina_optimization.md`。 - ---- - -## 3. 用户 API - -### 3.1 挂在数据集上(推荐) - -```python -from torch_molecule.datasets import load_qm9 - -data = load_qm9(local_dir="torchmol_data") -# subsample 仅用于本地调试 / CI,不要为了 Butina 而缩小 QM9 -# data = data.subsample(n=5000, seed=0) - -train, val = data.train_test_split( - test_size=0.2, - method="scaffold", # "random" | "scaffold" | "butina" | "size" - seed=42, -) - -train, val = data.train_test_split( - test_size=0.2, - method="scaffold", # "random" | "scaffold" | "butina" | "size" - seed=42, -) - -predictor.fit(train.data, train.target, val.data, val.target) -``` - -### 3.2 函数式 - -```python -from torch_molecule.datasets import train_test_split - -train, test = train_test_split(data, test_size=0.2, method="random", seed=42) -``` - -### 3.3 三路划分(对标 MoleculeNet 80/10/10) - -```python -train, val, test = data.train_val_test_split( - train_size=0.8, val_size=0.1, test_size=0.1, - method="scaffold", - seed=42, -) -``` - -### 3.4 可选参数 - - -| 参数 | 适用 method | 说明 | -| ------------------- | --------------------------------------- | ----------------------------------------------------------- | -| `test_size` | 全部 | 测试集比例,默认 0.2 | -| `seed` | random;scaffold/butina 的 `_group_split` | 随机种子,默认 42。Butina **聚类本身**用 index 平局,不用 seed | -| `similarity_cutoff` | butina | Tanimoto **相似度**阈值,默认 0.65(chemfp 语义;不是 DeepChem 距离 cutoff) | -| `use_csk` | scaffold | 是否用 cyclic skeleton(全碳),默认 False | -| `direction` | size | `small_to_large`(默认)等 | - - ---- - -## 4. `split.py` 内部结构 - -```python -def train_test_split(dataset, test_size=0.2, method="random", seed=42, **kwargs): - smiles = dataset.data - y = dataset.target - n = len(smiles) - - if method == "random": - idx_train, idx_test = _random_split(n, test_size, seed) - elif method == "scaffold": - groups = _scaffold_groups(smiles, use_csk=kwargs.get("use_csk", False)) - idx_train, idx_test = _group_split(groups, test_size, seed=seed) - elif method == "butina": - groups = _butina_groups_or_oom( - smiles, - similarity_cutoff=kwargs.get("similarity_cutoff", 0.65), - ) - idx_train, idx_test = _group_split(groups, test_size, seed=seed) - elif method == "size": - idx_train, idx_test = _size_split(smiles, test_size, seed=seed, **kwargs) - else: - raise ValueError(f"Unknown split method: {method}") - - return _subset(dataset, idx_train), _subset(dataset, idx_test) -``` - -核心:**按样本切**(random, size) vs **按组切**(scaffold, butina)→ 共用 `_group_split`。 - ---- - -## 5. 各方法实现要点 - -### 5.1 Random split - -- 与 sklearn 等价:`np.random.RandomState(seed).permutation(n)`。 -- 固定 `seed` 保证可复现。 -- 支持 `target is None`(如 ZINC)。 -- **测什么泛化:** 同分布插值;分数往往偏高,仅作对照。 - -**参考:** MoleculeNet (Wu et al., 2018) — baseline split。 - ---- - -### 5.2 Scaffold split - -**步骤:** - -1. SMILES → RDKit Mol(无效则 `ValueError`,与 `MolecularInputChecker` 一致)。 -2. `MurckoScaffold.GetScaffoldForMol(mol)` → scaffold Mol。 -3. `MolToSmiles(scaffold)` 作为 group id;无环小分子用 `"_acyclic_"` 或规范 SMILES。 -4. `dict[scaffold_smiles] -> list[index]`。 -5. **组级别**分配到 train/test(同一 scaffold 不能跨集合)。 - -**组分配策略(对齐 DeepChem / scikit-fingerprints):** - -- 按组大小降序排列 scaffold 组。 -- 贪心:将组填入 train 或 test,使 test 比例接近 `test_size`(或最小组优先进 test,测稀有骨架)。 - -**测试断言:** - -- train 与 test 的 scaffold 集合 **交集为空**。 -- 每个 index 恰好出现一次。 -- 实际比例可能偏离 `test_size`(group split 正常现象,需在文档说明)。 - -**可选:** `use_csk=False`(默认,保留原子类型);`True` 时用 `MakeScaffoldGeneric`(更粗)。 - -**测什么泛化:** 未见过的 Bemis–Murcko 骨架。 - -**参考:** - -- Bemis & Murcko (1996) — 骨架定义。 -- MoleculeNet (2018) — ML 标准协议。 - -**局限:** 差一个原子可能换 scaffold,但分子仍很相似(Walters 2024);RDKit 实现与原文细节略有差异。 - ---- - -### 5.3 Butina split(增强 / 稀疏精确版) - -代码:`torch_molecule/datasets/split.py` 中 `_butina_clusters` / `_butina_neighbor_lists` / `_butina_exclusion_spheres`。 -**精确 Taylor–Butina(Butina 1999)**,不是 BitBIRCH,也不是「先 subsample 再聚类」。 - -#### 为什么叫增强 - -朴素实现(DeepChem / Datamol / 直接 `Butina.ClusterData`)要先填满 condensed 距离: - -$$ -\text{dists 长度} = n(n-1)/2,\quad \text{内存 } O(n^2) -$$ - -QM9($n \approx 133885$)这条路会 OOM(他人报告 15 万分子涨到几十~上百 GB)。 -**精确 Butina 其实只需要布尔邻居** $N(i,j)=\mathbf{1}[T(i,j)\ge t]$。增强版只存 $T \ge t$ 的边,再按原文做度数排序 + exclusion sphere。簇与小集上的 RDKit `ClusterData` **membership 一致**(`tests/datasets/test_split.py`)。 - - -| | 朴素 `ClusterData` | 本库增强版 | -| ------ | ---------------------------------- | -------------------------------------------------------- | -| 算法 | 原文 Butina | **同一套** exclusion sphere | -| 指纹 | 调用方 / DeepChem 1024 bit | Morgan r=2,**2048** bit | -| 阈值 | DeepChem 默认 **距离 0.6**($T\ge 0.4$) | `**similarity_cutoff=0.65` 是相似度**(不是 Chalcedon 的距离 0.65) | -| 存储 | 全部 pair 距离 | 稀疏邻接表 $O(n+E)$ | -| 搜索 | Python 填矩阵 | RDKit `BulkTanimotoSimilarity`(C++/POPCNT)+ numpy 筛边 | -| 平局 | `(degree, index)` 降序(高 index 优先) | **对齐 RDKit**(不是 chemfp 默认 `randomize`) | -| 度数 | 含自身(对角距离 0) | 含自身 | -| QM9 全集 | OOM | 可跑完(Colab 约十几~几十分钟,不是 2 小时) | - - -chemfp 的 `threshold_tanimoto_search_symmetric`、Chalcedon 的分块两遍,是同一精确算法的另两种工程布局。本库 **不 `import chemfp`**;OOM 时明确报错,**禁止 subsample 来「修好」split**。预设簇文件按教授要求放 **Hugging Face**,不进 git。 - -#### 步骤(与代码一致) - -1. `GetMorganGenerator(radius=2, fpSize=2048)`。 -2. 对 $i=1\ldots n-1$:`BulkTanimotoSimilarity(fps[i], fps[:i])`,只把 $T \ge t$ 的 $j$ 写入双方邻居表(另加自身,以对齐 RDKit 零自距)。 -3. 按 `(邻居数, index)` **降序**(`list.sort(reverse=True)`)。 -4. Exclusion sphere:未标记且度数 $>1$ 的点当中心,收走未标记邻居;剩余单点各自成簇。 -5. 每个 cluster → `_group_split`(与 scaffold 同一套 DeepChem 式贪心)。 - -Tanimoto: - -$$ -T(A,B)=\frac{|A\cap B|}{|A\cup B|}=\frac{c}{a+b-c},\quad d=1-T -$$ - -$0.65$ 是 Walters / 常见实践默认,**不是**唯一验证过的最优阈值。对齐 DeepChem 应设 `similarity_cutoff=0.4`。Chalcedon 博客的 `cutoff=0.65` 是**距离**($\Rightarrow T\ge 0.35$),不要和本 API 混用。 - -#### QM9 默认 $t=0.65$ 实测(全集) - -- 133885 分子 → **100653** 簇;单点簇约 **77.6%**;簇大小 min/median/mean/max = 1 / 1 / 1.33 / **23**。 -- 最大簇成员对中心 $T$:min=0.65,mean≈0.71,max=1.0(exclusion sphere 成立)。 -- 阈值偏严:多数分子没有 $T\ge 0.65$ 的邻居,group split 会接近「按分子切」,但多成员簇仍整组进同一边。 - -#### 不要做 - -- 全距离 / `ClusterData` 当生产路径(仅测试、$n\le 2000$)。 -- LSH、先 subsample 再叫 `method="butina"`(教授:不要为 split 缩小 QM9)。 -- BitBIRCH(另一种算法)。 -- chemfp false-singleton 后处理(原文没有)。 -- 把预设 `.json.gz` 提交进 `torch_molecule/datasets/data/`。 - -**测什么泛化:** 指纹空间上不相似的分子(通常比 scaffold 更严;在 QM9+$0.65$ 下因大量 singleton 会略接近 random)。 - -**参考:** Butina, JCICS 1999, doi:10.1021/ci9803381;chemfp `butina`;Chalcedon (Rowan, 2026);DeepChem `ButinaSplitter`(只对标整簇分配);Walters 2024;`molecule_notes/butina_optimization.md`。 - ---- - -### 5.4 Size / atom-count split - -**只做划分,不包含 SizeShiftReg 的训练正则。** - -```python -n_atoms = [mol.GetNumHeavyAtoms() for mol in mols] -order = np.argsort(n_atoms) -n_test = int(round(n * test_size)) -idx_test = order[-n_test:] # 最大分子进 test -idx_train = order[:-n_test] -``` - -- 按 **重原子数**,不是分子量(DeepChem `MolecularWeightSplitter` 不同)。 -- 默认 `direction="small_to_large"`:小 train、大 test(对齐 SizeShiftReg 评估思路)。 -- 可选 `mode="sizeshiftreg"`:50% 最小 train / 10% 最大 test(论文协议)。 - -**测什么泛化:** 尺度外推(小分子 → 大分子),与骨架无关。 - -**参考:** Buffelli et al., SizeShiftReg, NeurIPS 2022。 - -**不要:** 在 `split.py` 中实现 graph coarsening 或 CMD loss(属于 `SSRMolecularPredictor`)。 - ---- - -## 6. `SMILESDataset` 扩展 - -当前定义(`torch_molecule/datasets/constant.py`): - -```python -@dataclass -class SMILESDataset: - data: List[str] - target: np.ndarray | None -``` - -建议新增方法(逻辑委托 `split.py`): - -```python -def subsample(self, n: int, seed: int = 0) -> "SMILESDataset": - """随机抽取 n 条。仅调试 / CI,不要为了 Butina 缩小基准集。""" - -def train_test_split(self, test_size=0.2, method="random", seed=42, **kwargs): - """返回 (train_dataset, test_dataset)。""" - -def train_val_test_split(self, train_size=0.8, val_size=0.1, test_size=0.1, ...): - """三路划分。""" -``` - -`subsample`:`RandomState(seed).choice(n, size=min(n, n_sub), replace=False)`,保持 `data`/`target` 对齐。 - ---- - -## 7. 与现有训练流程对接 - -```python -from torch_molecule.datasets import load_qm9 -from torch_molecule import GREAMolecularPredictor - -data = load_qm9(local_dir="torchmol_data") -train, val = data.train_test_split(test_size=0.2, method="scaffold", seed=42) - -predictor = GREAMolecularPredictor(num_task=1, task_type="regression") -predictor.fit(train.data, train.target, val.data, val.target) -predictions = predictor.predict(val.data) -``` - -**不要在** `fit()` **内加** `split=` **参数** — 用户需在同一划分上比较 GREA vs GNN。 - ---- - -## 8. 测试计划(`tests/datasets/test_split.py`) - - -| 测试项 | 断言 | -| ---------------- | ----------------------------------------------------------------- | -| random 可复现 | 同一 seed → 相同索引 | -| random 比例 | test 约等于 `test_size` | -| scaffold 无泄漏 | train/test scaffold 集合交集为空 | -| scaffold 覆盖 | 所有 index 仅用一次 | -| 无环分子 | 不崩溃 | -| 无效 SMILES | `ValueError` | -| `target is None` | 可划分 | -| 多任务 `y` | `y.shape[1]` 保持 | -| subsample | 长度与 seed 可复现 | -| butina | 同 cluster 不跨集合;小 $n$ 簇与 RDKit `ClusterData` 一致;OOM 文案禁止 subsample | -| size | test 平均原子数 > train | - - -可选慢测:`load_qm9` + subsample(1000) + scaffold,CI 可 skip。 - ---- - -## 9. 文档说明(每种 method 的 docstring) - - -| method | 测什么 | 注意 | -| ---------- | ----------- | ------------------------------------------------ | -| `random` | 同分布插值 | 分数往往偏高,作 baseline | -| `scaffold` | 未见 scaffold | MoleculeNet 推荐用于 HIV/BACE/BBBP | -| `butina` | 结构不相似 | 增强版:稀疏精确 Butina;默认 $T\ge 0.65$;不是 ClusterData 矩阵 | -| `size` | 小→大原子数 | 与 scaffold 正交;SSR 正则另见模型 | - - -Group split 时实际比例可能偏离 `test_size` — 文档中明确说明。 - ---- - -## 10. 明确不做的事 - - -| 不做 | 原因 | -| --------------------------------- | ------------------------- | -| 在 `fit()` 内自动 split | 无法固定划分比较模型 | -| split 内实现 CMD / coarsening | 属于 SSR 训练,非划分 | -| 默认 `method=scaffold` | 破坏旧脚本可复现 | -| 用 `ClusterData` / 全距离矩阵跑全集 Butina | 精确算法不需要矩阵;QM9 会 OOM | -| 为了 Butina 把 QM9 subsample | 教授:QM9 不算大库;缩小数据改变的是划分本身 | -| 把 QM9 簇 `.json.gz` 放进 git 包内 | 预设放 Hugging Face | -| 把 BitBIRCH / LSH 命名为 `butina` | 不是原文算法 | -| 引入 DeepChem 依赖 | RDKit 指纹 + 自写 chemfp 收球即可 | -| 混淆 stratified(QM7 排序切)与 random | 若做 stratified,单独 `method` | - - ---- - -## 11. 参考文献与资源 - - -| 主题 | 文献 / 链接 | -| ------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| Scaffold 定义 | Bemis & Murcko, J. Med. Chem. 1996, doi:10.1021/jm9602928 | -| MoleculeNet 协议 | Wu et al., Chem. Sci. 2018, doi:10.1039/C7SC02664A | -| Butina 聚类 | Butina, J. Chem. Inf. Comput. Sci. 1999, doi:10.1021/ci9803381 | -| chemfp 稀疏 Butina | [https://chemfp.com/docs/chemfp_butina_command.html](https://chemfp.com/docs/chemfp_butina_command.html) | -| 实现对照(Chalcedon 等) | 本仓库 `molecule_notes/butina_optimization.md` | -| Split 实践对比 | Walters, Practical Cheminformatics, 2024 | -| Size 泛化 | Buffelli et al., SizeShiftReg, NeurIPS 2022, arXiv:2206.07096 | -| DeepChem splitters | [https://deepchem.readthedocs.io/en/stable/api_reference/splitters.html](https://deepchem.readthedocs.io/en/stable/api_reference/splitters.html) | -| scikit-fingerprints | [https://scikit-fingerprints.readthedocs.io/stable/examples/06_dataset_splits.html](https://scikit-fingerprints.readthedocs.io/stable/examples/06_dataset_splits.html) | - - ---- -