Safetensors
esm

CreiLOV RLXF-aligned ESM-2 650M

This repository contains an RLXF-aligned ESM-2 masked language model developed for the design of fluorescent CreiLOV variants.

The model was initialized from facebook/esm2_t33_650M_UR50D and aligned toward experimentally informed CreiLOV fluorescence preferences.

Model details

  • Architecture: ESM-2 masked language model
  • Base model: facebook/esm2_t33_650M_UR50D
  • Protein target: CreiLOV
  • Framework: PyTorch and Hugging Face Transformers
  • Weight format: Safetensors

Intended use

The model is intended for research involving:

  • Generating candidate CreiLOV variants for experimental testing
  • Examining aligned amino-acid preferences
  • Ranking amino-acid substitutions

The model was aligned specifically for CreiLOV and should not be assumed to provide improved predictions or designs for unrelated protein families.

Loading the model

from transformers import AutoModelForMaskedLM, AutoTokenizer

model_id = "RomeroLab-Duke/creilov-rlxf-esm2-650m"

tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForMaskedLM.from_pretrained(model_id)
model.eval()

Sampling CreiLOV variants

The sampling procedure used in this work (https://www.biorxiv.org/content/10.1101/2025.05.02.651993v2.article-metrics) generates variants relative to the wild-type CreiLOV sequence in two stages:

  1. Each position in the wild-type sequence is masked independently to calculate its conditional amino-acid distribution.
  2. Non-wild-type substitutions with probability greater than the high-confidence threshold are introduced deterministically.
  3. Conditional amino-acid probabilities are recalculated using this high-confidence sequence as context.
  4. Positions whose cumulative non-wild-type probability exceeds the candidate-position threshold are eligible for additional mutation.
  5. Candidate positions are sampled in proportion to their cumulative non-wild-type probability. The amino acid at each selected position is then sampled from the model's conditional distribution over the 20 canonical amino acids.
  6. Sampling continues until the requested number of mutations, num_mutations are introduced.

The experiments associated with this model used a high-confidence threshold of 0.90 and a candidate-position threshold of 0.25.

The following self-contained implementation samples variants:

import random
import torch
from transformers import AutoModelForMaskedLM, AutoTokenizer

MODEL_ID = "RomeroLab-Duke/creilov-rlxf-esm2-650m"
WT = "MAGLRHTFVVADATLPDCPLVYASEGFYAMTGYGPDEVLGHNARFLQGEGTDPKEVQKIRDAIKKGEACSVRLLNYRKDGTPFWNLLTVTPIKTPDGRVSKFVGVQVDVTSKTEGKALA"
AMINO_ACIDS = list("ACDEFGHIKLMNPQRSTVWY")

def masked_probabilities(sequence, model, tokenizer, device):
    """Calculate the conditional distribution at every sequence position."""

    position_probabilities = []

    with torch.inference_mode():
        for position in range(len(sequence)):
            masked_sequence = (
                sequence[:position]
                + tokenizer.mask_token
                + sequence[position + 1:]
            )

            inputs = tokenizer(
                masked_sequence,
                return_tensors="pt",
            ).to(device)

            logits = model(**inputs).logits

            # Token position zero is the beginning-of-sequence token.
            mask_logits = logits[0, position + 1]
            position_probabilities.append(
                torch.softmax(mask_logits, dim=-1).cpu()
            )

    return torch.stack(position_probabilities)


def mutation_labels(wild_type, sequence):
    """Return mutations using one-based protein numbering."""

    return [
        f"{wt_residue}{position}{new_residue}"
        for position, (wt_residue, new_residue) in enumerate(
            zip(wild_type, sequence),
            start=1,
        )
        if wt_residue != new_residue
    ]


def sample_creilov_variants(
    model,
    tokenizer,
    wild_type=WT,
    num_designs=10,
    num_mutations=6,
    high_confidence_threshold=0.90,
    candidate_position_threshold=0.25,
    seed=7028,
):
    """Sample variants at an exact Hamming distance from wild-type."""

    random.seed(seed)
    torch.manual_seed(seed)

    device = next(model.parameters()).device
    amino_acid_ids = {
        amino_acid: tokenizer.convert_tokens_to_ids(amino_acid)
        for amino_acid in AMINO_ACIDS
    }

    # First pass: evaluate every single substitution from wild-type.
    wt_probabilities = masked_probabilities(
        wild_type,
        model,
        tokenizer,
        device,
    )

    # Introduce the most probable non-WT residue wherever its individual
    # probability exceeds the high-confidence threshold.
    scaffold = list(wild_type)

    for position, wt_residue in enumerate(wild_type):
        alternatives = [
            (amino_acid, wt_probabilities[position, token_id].item())
            for amino_acid, token_id in amino_acid_ids.items()
            if amino_acid != wt_residue
        ]

        best_amino_acid, best_probability = max(
            alternatives,
            key=lambda item: item[1],
        )

        if best_probability > high_confidence_threshold:
            scaffold[position] = best_amino_acid

    scaffold = "".join(scaffold)
    scaffold_distance = sum(
        residue != wt_residue
        for residue, wt_residue in zip(scaffold, wild_type)
    )

    if scaffold_distance > num_mutations:
        raise ValueError(
            f"The high-confidence scaffold already contains "
            f"{scaffold_distance} mutations, which exceeds the requested "
            f"target of {num_mutations}. Increase num_mutations or increase "
            "high_confidence_threshold."
        )

    # Recalculate position probabilities using the scaffold as context.
    scaffold_probabilities = masked_probabilities(
        scaffold,
        model,
        tokenizer,
        device,
    )

    candidate_weights = {}

    for position, wt_residue in enumerate(wild_type):
        # Preserve mutations already introduced into the scaffold.
        if scaffold[position] != wt_residue:
            continue

        non_wt_probability = sum(
            scaffold_probabilities[position, token_id].item()
            for amino_acid, token_id in amino_acid_ids.items()
            if amino_acid != wt_residue
        )

        if non_wt_probability > candidate_position_threshold:
            candidate_weights[position] = non_wt_probability

    mutations_needed = num_mutations - scaffold_distance

    if len(candidate_weights) < mutations_needed:
        raise ValueError(
            "There are not enough eligible positions to reach the requested "
            "number of mutations. Lower candidate_position_threshold."
        )

    sampled_sequences = []

    for _ in range(num_designs):
        sequence = list(scaffold)

        while sum(
            residue != wt_residue
            for residue, wt_residue in zip(sequence, wild_type)
        ) < num_mutations:
            available_positions = [
                position
                for position in candidate_weights
                if sequence[position] == wild_type[position]
            ]

            position = random.choices(
                available_positions,
                weights=[
                    candidate_weights[p] for p in available_positions
                ],
                k=1,
            )[0]

            masked_sequence = sequence.copy()
            masked_sequence[position] = tokenizer.mask_token

            inputs = tokenizer(
                "".join(masked_sequence),
                return_tensors="pt",
            ).to(device)

            with torch.inference_mode():
                logits = model(**inputs).logits[0, position + 1]

            # Sample only from non-WT canonical amino acids so that every
            # iteration adds exactly one mutation.
            allowed_amino_acids = [
                amino_acid
                for amino_acid in AMINO_ACIDS
                if amino_acid != wild_type[position]
            ]
            allowed_ids = torch.tensor(
                [amino_acid_ids[aa] for aa in allowed_amino_acids],
                device=logits.device,
            )

            probabilities = torch.softmax(
                logits[allowed_ids],
                dim=-1,
            )

            sampled_index = torch.multinomial(
                probabilities,
                num_samples=1,
            ).item()

            sequence[position] = allowed_amino_acids[sampled_index]

        sampled_sequences.append("".join(sequence))

    return sampled_sequences


tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForMaskedLM.from_pretrained(MODEL_ID)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)
model.eval()

designs = sample_creilov_variants(
    model,
    tokenizer,
    num_designs=10,
    num_mutations=6,
    high_confidence_threshold=0.90,
    candidate_position_threshold=0.25,
    seed=7028,
)

for index, sequence in enumerate(designs, start=1):
    mutations = ", ".join(mutation_labels(WT, sequence))
    print(f"Design {index}: {mutations}")
    print(sequence)

Base-model citation

@article{lin2023evolutionary,
  title={Evolutionary-scale prediction of atomic-level protein structure with a language model},
  author={Lin, Zeming and others},
  journal={Science},
  volume={379},
  number={6637},
  pages={1123--1130},
  year={2023}
}

Authors

Romero Lab at Duke University. Contact nlb51@duke.edu or philip.romero@duke.edu with any questions.

Downloads last month
26
Safetensors
Model size
0.7B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support