{"id":26270725,"url":"https://github.com/emapco/chem-mrl","last_synced_at":"2025-05-01T13:20:22.038Z","repository":{"id":272897728,"uuid":"918081383","full_name":"emapco/chem-mrl","owner":"emapco","description":"Chem-MRL: SMILES Matryoshka Representation Learning Embedding Model","archived":false,"fork":false,"pushed_at":"2025-03-29T09:44:04.000Z","size":32967,"stargazers_count":0,"open_issues_count":0,"forks_count":0,"subscribers_count":1,"default_branch":"main","last_synced_at":"2025-05-01T13:19:44.495Z","etag":null,"topics":["chemoinformatics","embedding-models","latent-space","matryoshka-representation-learning","smiles"],"latest_commit_sha":null,"homepage":"","language":"Python","has_issues":true,"has_wiki":null,"has_pages":null,"mirror_url":null,"source_name":null,"license":"apache-2.0","status":null,"scm":"git","pull_requests_enabled":true,"icon_url":"https://github.com/emapco.png","metadata":{"files":{"readme":"README.md","changelog":null,"contributing":null,"funding":null,"license":"LICENSE","code_of_conduct":null,"threat_model":null,"audit":null,"citation":null,"codeowners":null,"security":null,"support":null,"governance":null,"roadmap":null,"authors":null,"dei":null,"publiccode":null,"codemeta":null}},"created_at":"2025-01-17T08:05:03.000Z","updated_at":"2025-03-29T09:44:08.000Z","dependencies_parsed_at":"2025-02-02T00:15:27.619Z","dependency_job_id":"5a3f3b06-81cc-496b-ad43-c8040e5a5281","html_url":"https://github.com/emapco/chem-mrl","commit_stats":null,"previous_names":["emapco/chem-mrl-public","emapco/chem-mrl"],"tags_count":0,"template":false,"template_full_name":null,"repository_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/emapco%2Fchem-mrl","tags_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/emapco%2Fchem-mrl/tags","releases_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/emapco%2Fchem-mrl/releases","manifests_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories/emapco%2Fchem-mrl/manifests","owner_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners/emapco","download_url":"https://codeload.github.com/emapco/chem-mrl/tar.gz/refs/heads/main","host":{"name":"GitHub","url":"https://github.com","kind":"github","repositories_count":251879306,"owners_count":21658730,"icon_url":"https://github.com/github.png","version":null,"created_at":"2022-05-30T11:31:42.601Z","updated_at":"2022-07-04T15:15:14.044Z","host_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub","repositories_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repositories","repository_names_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/repository_names","owners_url":"https://repos.ecosyste.ms/api/v1/hosts/GitHub/owners"}},"keywords":["chemoinformatics","embedding-models","latent-space","matryoshka-representation-learning","smiles"],"created_at":"2025-03-14T06:17:00.278Z","updated_at":"2025-05-01T13:20:22.005Z","avatar_url":"https://github.com/emapco.png","language":"Python","funding_links":[],"categories":[],"sub_categories":[],"readme":"# CHEM-MRL\n\nChem-MRL is a SMILES embedding transformer model that leverages Matryoshka Representation Learning (MRL) to generate efficient, truncatable embeddings for downstream tasks such as classification, clustering, and database querying.\n\nThe model employs [SentenceTransformers' (SBERT)](https://sbert.net/) [2D Matryoshka Sentence Embeddings](https://sbert.net/examples/training/matryoshka/README.html) (`Matryoshka2dLoss`) to enable truncatable embeddings with minimal accuracy loss, improving query performance and flexibility in downstream applications.\n\nDatasets should consists of SMILES pairs and their corresponding [Morgan fingerprint](https://www.rdkit.org/docs/GettingStartedInPython.html#morgan-fingerprints-circular-fingerprints) Tanimoto similarity scores. Currently, datasets must be in Parquet format.\n\nHyperparameter tuning indicates that a custom Tanimoto similarity loss function, `TanimotoSentLoss`, based on [CoSENTLoss](https://kexue.fm/archives/8847), outperforms [Tanimoto similarity](https://jcheminf.biomedcentral.com/articles/10.1186/s13321-015-0069-3/tables/2), CoSENTLoss, [AnglELoss](https://arxiv.org/pdf/2309.12871), and cosine similarity.\n\n## Installation\n\n**Install with pip**\n\n```bash\npip install chem-mrl\n```\n\n**Install from source code**\n\n```bash\npip install -e .\n```\n\n## Usage\n\n### Hydra \u0026 Training Scripts\n\nHydra configuration files are in `chem_mrl/conf`. The base config defines shared arguments, while model-specific configs are located in `chem_mrl/conf/model`. Use `chem_mrl_config.yaml` or `classifier_config.yaml` to run specific models.\n\nThe `scripts` directory provides training scripts with Hydra for parameter management:\n\n- **Train Chem-MRL model:**\n  ```bash\n  python scripts/train_chem_mrl.py train_dataset_path=/path/to/training.parquet val_dataset_path=/path/to/val.parquet\n  ```\n- **Train a linear classifier:**\n  ```bash\n  python scripts/train_classifier.py train_dataset_path=/path/to/training.parquet val_dataset_path=/path/to/val.parquet\n  ```\n\n### Basic Training Workflow\n\nTo train a model, initialize the configuration with dataset paths and model parameters, then pass it to `ChemMRLTrainer` for training.\n\n```python\nfrom chem_mrl.schemas import BaseConfig, ChemMRLConfig\nfrom chem_mrl.constants import BASE_MODEL_NAME\nfrom chem_mrl.trainers import ChemMRLTrainer\n\n# Define training configuration\nconfig = BaseConfig(\n    model=ChemMRLConfig(\n        model_name=BASE_MODEL_NAME,  # Predefined model name - Can be any transformer model name or path that is compatible with sentence-transformers\n        n_dims_per_step=3,  # Model-specific hyperparameter\n        use_2d_matryoshka=True,  # Enable 2d MRL\n        # Additional parameters specific to 2D MRL models\n        n_layers_per_step=2,\n        kl_div_weight=0.7,  # Weight for KL divergence regularization\n        kl_temperature=0.5,  # Temperature parameter for KL loss\n    ),\n    train_dataset_path=\"train.parquet\",  # Path to training data\n    val_dataset_path=\"val.parquet\",  # Path to validation data\n    test_dataset_path=\"test.parquet\",  # Optional test dataset\n    smiles_a_column_name=\"smiles_a\",  # Column with first molecule SMILES representation\n    smiles_b_column_name=\"smiles_b\",  # Column with second molecule SMILES representation\n    label_column_name=\"similarity\",  # Similarity score between molecules\n)\n\n# Initialize trainer and start training\ntrainer = ChemMRLTrainer(config)\ntest_eval_metric = (\n    trainer.train()\n)  # Returns the test evaluation metric if a test dataset is provided.\n# Otherwise returns the final validation eval metric\n```\n\n### Experimental\n\n#### Train a Query Model\n\nTo train a querying model, configure the model to utilize the specialized query tokenizer.\n\nThe query tokenizer supports the following query types:\n\n- similar: Computes SMILES similarity between two molecular structures. For retrieving similar SMILES.\n- substructure: Determines the presence of a substructure within the second SMILES string.\n\nSupported query formats for `smiles_a` column:\n\n- `similar {smiles}`\n- `substructure {smiles}`\n\n```python\nfrom chem_mrl.schemas import BaseConfig, ChemMRLConfig\nfrom chem_mrl.constants import BASE_MODEL_NAME\nfrom chem_mrl.trainers import ChemMRLTrainer\n\nconfig = BaseConfig(\n    model=ChemMRLConfig(\n        model_name=BASE_MODEL_NAME,\n        use_query_tokenizer=True,  # Train a query model\n    ),\n    train_dataset_path=\"train.parquet\",\n    val_dataset_path=\"val.parquet\",\n    smiles_a_column_name=\"query\",\n    smiles_b_column_name=\"target_smiles\",\n    label_column_name=\"similarity\",\n)\ntrainer = ChemMRLTrainer(config)\n```\n\n#### Latent Attention Layer\n\nThe Latent Attention Layer model is an experimental component designed to enhance the representation learning of transformer-based models by introducing a trainable latent dictionary. This mechanism applies cross-attention between token embeddings and a set of learnable latent vectors before pooling. The output of this layer contributes to both **1D Matryoshka loss** (as the final layer output) and **2D Matryoshka loss** (by integrating into all-layer outputs). Note: initial tests suggests that when using default configuration, the latent attention layer leads to overfitting.\n\n```python\nfrom chem_mrl.models import LatentAttentionLayer\nfrom chem_mrl.schemas import BaseConfig, ChemMRLConfig, LatentAttentionConfig\nfrom chem_mrl.constants import BASE_MODEL_NAME\nfrom chem_mrl.trainers import ChemMRLTrainer\n\nconfig = BaseConfig(\n    model=ChemMRLConfig(\n        model_name=BASE_MODEL_NAME,\n        latent_attention_config=LatentAttentionConfig(\n            hidden_dim=768,  # Transformer hidden size\n            num_latents=512,  # Number of learnable latents\n            num_cross_heads=8,  # Number of attention heads\n            cross_head_dim=32,  # Dimensionality of each head\n            output_normalize=True,  # Apply L2 normalization to outputs\n        ),\n        use_2d_matryoshka=True,\n    ),\n    train_dataset_path=\"train.parquet\",\n    val_dataset_path=\"val.parquet\",\n)\n\n# Train a model with latent attention\ntrainer = ChemMRLTrainer(config)\n```\n\n### Custom Evaluation Callbacks\n\nYou can provide a callback function that is executed every `evaluation_steps` steps, allowing custom logic such as logging, early stopping, or model checkpointing.\n\n```python\nfrom chem_mrl.schemas import BaseConfig, ChemMRLConfig\nfrom chem_mrl.constants import BASE_MODEL_NAME\nfrom chem_mrl.trainers import ChemMRLTrainer\n\n\n# Define a callback function for logging evaluation metrics\ndef eval_callback(score: float, epoch: int, steps: int):\n    print(f\"Step {steps}, Epoch {epoch}: Evaluation Score = {score}\")\n\n\nconfig = BaseConfig(\n    model=ChemMRLConfig(\n        model_name=BASE_MODEL_NAME,\n    ),\n    train_dataset_path=\"train.parquet\",\n    val_dataset_path=\"val.parquet\",\n    smiles_a_column_name=\"smiles_a\",\n    smiles_b_column_name=\"smiles_b\",\n    label_column_name=\"similarity\",\n)\n\n# Train with callback\ntrainer = ChemMRLTrainer(config)\nval_eval_metric = trainer.train(\n    eval_callback=eval_callback\n)  # Callback executed every `evaluation_steps`\n```\n\n### W\u0026B Integration\n\nThis library includes a `WandBTrainerExecutor` class for seamless Weights \u0026 Biases (W\u0026B) integration. It handles authentication, initialization, and logging at the frequency specified by `evaluation_steps`.\n\n```python\nfrom chem_mrl.schemas import BaseConfig, WandbConfig, ChemMRLConfig\nfrom chem_mrl.constants import BASE_MODEL_NAME\nfrom chem_mrl.trainers import ChemMRLTrainer, WandBTrainerExecutor\nfrom chem_mrl.schemas.Enums import WatchLogOption\n\n# Define W\u0026B configuration for experiment tracking\nwandb_config = WandbConfig(\n    project_name=\"chem_mrl_test\",  # W\u0026B project name\n    run_name=\"test\",  # Name for the experiment run\n    use_watch=True,  # Enables model watching for tracking gradients\n    watch_log=WatchLogOption.all,  # Logs all model parameters and gradients\n    watch_log_freq=1000,  # Logging frequency\n    watch_log_graph=True,  # Logs model computation graph\n)\n\n# Configure training with W\u0026B integration\nconfig = BaseConfig(\n    model=ChemMRLConfig(\n        model_name=BASE_MODEL_NAME,\n    ),\n    train_dataset_path=\"train.parquet\",\n    val_dataset_path=\"val.parquet\",\n    smiles_a_column_name=\"smiles_a\",\n    smiles_b_column_name=\"smiles_b\",\n    label_column_name=\"similarity\",\n    evaluation_steps=1000,\n    wandb=wandb_config,\n)\n\n# Initialize trainer and W\u0026B executor\ntrainer = ChemMRLTrainer(config)\nexecutor = WandBTrainerExecutor(trainer)\nexecutor.execute()  # Handles training and W\u0026B logging\n```\n\n## Classifier\n\nThis repository includes code for training a linear classifier with optional dropout regularization. The classifier categorizes substances based on SMILES and category features.\n\nHyperparameter tuning shows that cross-entropy loss (`softmax` option) outperforms self-adjusting dice loss in terms of accuracy, making it the preferred choice for molecular property classification.\n\n### Usage\n\n#### Basic Classification Training\n\nTo train a classifier, configure the model with dataset paths and column names, then initialize `ClassifierTrainer` to start training.\n\n```python\nfrom chem_mrl.schemas import BaseConfig, ClassifierConfig\nfrom chem_mrl.trainers import ClassifierTrainer\n\n# Define classification training configuration\nconfig = BaseConfig(\n    model=ClassifierConfig(\n        model_name=\"path/to/trained_mrl_model\",  # Pretrained MRL model path\n    ),\n    train_dataset_path=\"train_classification.parquet\",  # Path to training dataset\n    val_dataset_path=\"val_classification.parquet\",  # Path to validation dataset\n    smiles_a_column_name=\"smiles\",  # Column containing SMILES representations of molecules\n    label_column_name=\"label\",  # Column containing classification labels\n)\n\n# Initialize and train the classifier\ntrainer = ClassifierTrainer(config)\ntrainer.train()\n```\n\n#### Training with Dice Loss\n\nFor imbalanced classification tasks, **Dice Loss** can improve performance by focusing on hard-to-classify samples. Below is a configuration using `DiceLossClassifierConfig`, which introduces additional hyperparameters.\n\n```python\nfrom chem_mrl.schemas import BaseConfig, ClassifierConfig\nfrom chem_mrl.trainers import ClassifierTrainer\nfrom chem_mrl.schemas.Enums import ClassifierLossFctOption, DiceReductionOption\n\n# Define classification training configuration with Dice Loss\nconfig = BaseConfig(\n    model=ClassifierConfig(\n        model_name=\"path/to/trained_mrl_model\",\n        loss_func=ClassifierLossFctOption.selfadjdice,\n        dice_reduction=DiceReductionOption.sum,  # Reduction method for Dice Loss (e.g., 'mean' or 'sum')\n        dice_gamma=1.0,  # Smoothing factor hyperparameter\n    ),\n    train_dataset_path=\"train_classification.parquet\",  # Path to training dataset\n    val_dataset_path=\"val_classification.parquet\",  # Path to validation dataset\n    smiles_a_column_name=\"smiles\",\n    label_column_name=\"label\",\n)\n\n# Initialize and train the classifier with Dice Loss\ntrainer = ClassifierTrainer(config)\ntrainer.train()\n```\n\n## References:\n\n- Chithrananda, Seyone, et al. \"ChemBERTa: Large-Scale Self-Supervised Pretraining for Molecular Property Prediction.\" _arXiv [Cs.LG]_, 2020. [Link](http://arxiv.org/abs/2010.09885).\n- Ahmad, Walid, et al. \"ChemBERTa-2: Towards Chemical Foundation Models.\" _arXiv [Cs.LG]_, 2022. [Link](http://arxiv.org/abs/2209.01712).\n- Kusupati, Aditya, et al. \"Matryoshka Representation Learning.\" _arXiv [Cs.LG]_, 2022. [Link](https://arxiv.org/abs/2205.13147).\n- Li, Xianming, et al. \"2D Matryoshka Sentence Embeddings.\" _arXiv [Cs.CL]_, 2024. [Link](http://arxiv.org/abs/2402.14776).\n- Bajusz, Dávid, et al. \"Why is the Tanimoto Index an Appropriate Choice for Fingerprint-Based Similarity Calculations?\" _J Cheminform_, 7, 20 (2015). [Link](https://doi.org/10.1186/s13321-015-0069-3).\n- Li, Xiaoya, et al. \"Dice Loss for Data-imbalanced NLP Tasks.\" _arXiv [Cs.CL]_, 2020. [Link](https://arxiv.org/abs/1911.02855)\n- Reimers, Nils, and Gurevych, Iryna. \"Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks.\" _Proceedings of the 2019 Conference on Empirical Methods in Natural Language Processing_, 2019. [Link](https://arxiv.org/abs/1908.10084).\n- Lee, Chankyu, et al. \"NV-Embed: Improved Techniques for Training LLMs as Generalist Embedding Models.\" _arXiv [Cs.CL]_, 2025. [Link](https://arxiv.org/abs/2405.17428).\n","project_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Femapco%2Fchem-mrl","html_url":"https://awesome.ecosyste.ms/projects/github.com%2Femapco%2Fchem-mrl","lists_url":"https://awesome.ecosyste.ms/api/v1/projects/github.com%2Femapco%2Fchem-mrl/lists"}