schrodinger.application.matsci.flywheel.ldmol.bert_decoder module¶
PyTorch BERT model for SMILES encoding and decoding.
Based on ALBEF (Salesforce) BERT implementation.
- schrodinger.application.matsci.flywheel.ldmol.bert_decoder.gelu(tensor)¶
GELU activation (Gaussian Error Linear Unit).
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertConfig(vocab_size=30522, hidden_size=768, num_hidden_layers=12, num_attention_heads=12, intermediate_size=3072, hidden_act='gelu', hidden_dropout_prob=0.1, attention_probs_dropout_prob=0.1, max_position_embeddings=512, type_vocab_size=2, layer_norm_eps=1e-12, pad_token_id=0, fusion_layer=0, encoder_width=768, autoregressive=0, **kwargs)¶
Bases:
objectConfiguration loaded from JSON model info files.
Replaces transformers.BertConfig to eliminate the external dependency.
- __init__(vocab_size=30522, hidden_size=768, num_hidden_layers=12, num_attention_heads=12, intermediate_size=3072, hidden_act='gelu', hidden_dropout_prob=0.1, attention_probs_dropout_prob=0.1, max_position_embeddings=512, type_vocab_size=2, layer_norm_eps=1e-12, pad_token_id=0, fusion_layer=0, encoder_width=768, autoregressive=0, **kwargs)¶
- Parameters:
vocab_size (int) – Vocabulary size
hidden_size (int) – Hidden layer dimension
num_hidden_layers (int) – Number of transformer layers
num_attention_heads (int) – Number of attention heads
intermediate_size (int) – FFN inner dimension
hidden_act (str) – Activation function name
hidden_dropout_prob (float) – Hidden layer dropout
attention_probs_dropout_prob (float) – Attention dropout
max_position_embeddings (int) – Max sequence length
type_vocab_size (int) – Token type vocabulary size
layer_norm_eps (float) – LayerNorm epsilon
pad_token_id (int) – Padding token ID
fusion_layer (int) – Layer index where cross-attention begins
encoder_width (int) – Cross-attention key/value dimension
autoregressive (int) – Non-zero enables causal masking
kwargs – Extra keys from JSON (e.g. ‘architectures’, ‘model_type’, ‘add_cross_attention’) are silently ignored so that fromJsonFile can pass the full config_dict without filtering.
- classmethod fromJsonFile(json_file)¶
Load config from a JSON file.
- Parameters:
json_file (str) – Path to JSON config
- Return type:
- Returns:
Config instance with JSON values
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertEmbeddings(config)¶
Bases:
ModuleWord, position, and token type embeddings.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(input_ids=None, token_type_ids=None, position_ids=None, inputs_embeds=None, past_key_values_length=0)¶
Compute combined embeddings for a token sequence.
- Parameters:
input_ids (torch.Tensor) – Token IDs (batch, seq). It is mutually exclusive with inputs_embeds
token_type_ids (torch.Tensor) – Segment IDs (batch, seq); defaults to zeros
position_ids (torch.Tensor) – Position indices (batch, seq). The values are inferred from sequence length if None
inputs_embeds (torch.Tensor) – Pre-computed embeddings (batch, seq, hidden_size)
past_key_values_length (int) – Cached key/value length for incremental decoding
- Return type:
torch.Tensor
- Returns:
Embeddings (batch, seq, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertSelfAttention(config, is_cross_attention)¶
Bases:
ModuleMulti-head self/cross-attention.
Scaled dot-product attention flow:
hidden ─► Q ─┐ ─► K ─┼─► Q@K^T/√d ─► +mask ─► softmax ─► @V ─► output ─► V ─┘- __init__(config, is_cross_attention)¶
- Parameters:
config (BertConfig) – Model hyperparameters
is_cross_attention (bool) – If True, K and V project from encoder_width instead of hidden_size
- transposeForScores(tensor)¶
Reshape flat hidden dim into per-head slices for attention scoring.
- Parameters:
tensor (torch.Tensor) – Shape (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Shape (batch, num_heads, seq, head_dim)
- forward(hidden_states, attention_mask=None, head_mask=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_value=None, output_attentions=False)¶
Compute attended context vectors.
- Parameters:
hidden_states (torch.Tensor) – Query input (batch, seq, hidden_size)
attention_mask (torch.Tensor) – Extended mask (batch, 1, 1, seq), values 0 or -10000
head_mask (torch.Tensor) – Per-head mask or None
encoder_hidden_states (torch.Tensor) – Cross-attention keys/values source (batch, enc_seq, encoder_width)
encoder_attention_mask (torch.Tensor) – Extended encoder mask (batch, 1, 1, enc_seq)
past_key_value (tuple) – Cached (key, value) tensors for incremental decoding
output_attentions (bool) – Include attention weights in output tuple
- Return type:
tuple
- Returns:
(context, [attention_probs], present_key_value)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertSelfOutput(config)¶
Bases:
ModuleResidual connection after self-attention.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states, input_tensor)¶
Apply dense projection, dropout, and residual LayerNorm.
- Parameters:
hidden_states (torch.Tensor) – Attention output (batch, seq, hidden_size)
input_tensor (torch.Tensor) – Residual input (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Normalized output (batch, seq, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertAttention(config, is_cross_attention=False)¶
Bases:
ModuleSelf-attention + residual output.
- __init__(config, is_cross_attention=False)¶
- Parameters:
config (BertConfig) – Model hyperparameters
is_cross_attention (bool) – Enable cross-attention mode
- forward(hidden_states, attention_mask=None, head_mask=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_value=None, output_attentions=False)¶
Run self-attention and apply residual output projection.
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
attention_mask (torch.Tensor) – Extended mask or None
head_mask (torch.Tensor) – Per-head mask or None
encoder_hidden_states (torch.Tensor) – Cross-attention source or None
encoder_attention_mask (torch.Tensor) – Encoder mask or None
past_key_value (tuple) – Cached keys/values or None
output_attentions (bool) – Include attention weights
- Return type:
tuple
- Returns:
(attention_output, [attn_probs], present_kv)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertIntermediate(config)¶
Bases:
ModuleFeed-forward expansion: hidden_size -> intermediate_size.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states)¶
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Expanded output (batch, seq, intermediate_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertOutput(config)¶
Bases:
ModuleFeed-forward projection with residual connection.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states, input_tensor)¶
Project back to hidden_size and apply residual LayerNorm.
- Parameters:
hidden_states (torch.Tensor) – FFN expanded output (batch, seq, intermediate_size)
input_tensor (torch.Tensor) – Residual attention output (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Layer output (batch, seq, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertLayer(config, layer_num)¶
Bases:
ModuleSingle transformer layer with optional cross-attention.
- __init__(config, layer_num)¶
- Parameters:
config (BertConfig) – Model hyperparameters
layer_num (int) – Zero-based layer index; layers at or above fusion_layer receive a cross-attention sublayer
- forward(hidden_states, attention_mask=None, head_mask=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_value=None, output_attentions=False)¶
Run self-attention, optional cross-attention, and FFN.
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
attention_mask (torch.Tensor) – Extended self-attention mask or None
head_mask (torch.Tensor) – Per-head mask or None
encoder_hidden_states – Cross-attention source; list of tensors for multi-scale or single tensor
encoder_attention_mask – Encoder mask; list or tensor
past_key_value (tuple) – Cached keys/values or None
output_attentions (bool) – Include attention weights
- Return type:
tuple
- Returns:
(layer_output, [attn_probs], present_kv)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertEncoder(config)¶
Bases:
ModuleStack of BertLayer modules.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states, attention_mask=None, head_mask=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_values=None, output_attentions=False, mode='multi_modal')¶
Run a slice of transformer layers determined by mode.
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
attention_mask (torch.Tensor) – Extended mask or None
head_mask (torch.Tensor) – Per-layer head mask or None
encoder_hidden_states – Cross-attention source or None
encoder_attention_mask – Encoder mask or None
past_key_values (list) – Cached keys/values per layer
output_attentions (bool) – Include attention weights
mode (str) – MODE_TEXT (layers 0..fusion_layer), MODE_FUSION (fusion_layer..end), or MODE_MULTI_MODAL (all)
- Return type:
torch.Tensor
- Returns:
Final hidden states (batch, seq, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertPooler(config)¶
Bases:
ModulePool [CLS] token through dense + tanh.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states)¶
- Parameters:
hidden_states (torch.Tensor) – Encoder output (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Pooled [CLS] representation (batch, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertPreTrainedModel(config)¶
Bases:
ModuleBase class for BERT models (inference only).
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- init_weights()¶
No-op for inference.
- getExtendedAttentionMask(attention_mask, input_shape, device, is_decoder)¶
Expand a 2-D or 3-D mask to 4-D and convert to additive form.
Additive masks add 0.0 to attended positions and -10000.0 to masked positions so that after softmax the masked positions become ~0.
- Parameters:
attention_mask (torch.Tensor) – Boolean/float mask shaped (batch, seq) or (batch, tgt_seq, src_seq)
input_shape (tuple) – (batch_size, seq_length) of the input
device (torch.device) – Device for new tensors
is_decoder (bool) – If True, combine with causal mask so each position can only attend to earlier positions
- Return type:
torch.Tensor
- Returns:
Additive mask (batch, 1, tgt_seq, src_seq)
- invertAttentionMask(encoder_attention_mask)¶
Convert an encoder padding mask to additive form for cross-attention.
Expands a 2-D or 3-D boolean/float mask to 4-D and flips the convention: 1 (valid) becomes 0.0 and 0 (padded) becomes -10000.0.
- Parameters:
encoder_attention_mask (torch.Tensor) – Mask shaped (batch, enc_seq) or (batch, 1, tgt_seq, enc_seq)
- Return type:
torch.Tensor
- Returns:
Additive mask (batch, 1, 1, enc_seq) or (batch, 1, tgt_seq, enc_seq)
- getHeadMask(head_mask, num_hidden_layers)¶
Expand a flat or per-layer head mask to 5-D for BertLayer.
- Parameters:
head_mask (torch.Tensor or None) – 1-D tensor (num_heads,) to apply the same mask to all layers, or 2-D tensor (num_layers, num_heads) for per-layer masks. None means all heads are kept.
num_hidden_layers (int) – Number of transformer layers
- Return type:
list
- Returns:
List of length num_hidden_layers; each element is a (1, 1, num_heads, 1, 1) tensor or None
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertModel(config, add_pooling_layer=True)¶
Bases:
BertPreTrainedModelBERT model for encoding or decoding with cross-attention.
When configured as a decoder, encoder_hidden_states are expected as input to the forward pass.
- __init__(config, add_pooling_layer=True)¶
- Parameters:
config (BertConfig) – Model hyperparameters
add_pooling_layer (bool) – Attach BertPooler to output
- forward(input_ids=None, attention_mask=None, token_type_ids=None, position_ids=None, head_mask=None, inputs_embeds=None, encoder_embeds=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_values=None, is_decoder=False, mode='multi_modal')¶
Run the full BERT forward pass.
Exactly one of input_ids, inputs_embeds, or encoder_embeds must be provided.
- Parameters:
input_ids (torch.Tensor) – Token IDs (batch, seq)
attention_mask (torch.Tensor) – 1/0 mask (batch, seq); defaults to all-ones
token_type_ids (torch.Tensor) – Segment IDs (batch, seq)
position_ids (torch.Tensor) – Position indices (batch, seq)
head_mask (torch.Tensor) – Per-layer head mask or None
inputs_embeds (torch.Tensor) – Pre-computed token embeddings (batch, seq, hidden_size)
encoder_embeds (torch.Tensor) – Pre-computed embeddings that bypass the embedding layer entirely
encoder_hidden_states – Cross-attention source; tensor or list
encoder_attention_mask – Encoder mask; tensor or list
past_key_values (list) – Cached key/value pairs per layer
is_decoder (bool) – Enable causal attention masking
mode (str) – Layer range (MODE_TEXT, MODE_FUSION, MODE_MULTI_MODAL)
- Return type:
tuple
- Returns:
(sequence_output, pooled_output)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertPredictionHeadTransform(config)¶
Bases:
ModuleDense + activation + LayerNorm before LM projection.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states)¶
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Transformed output (batch, seq, hidden_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertLMPredictionHead(config)¶
Bases:
ModuleTransform then project to vocab logits.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(hidden_states)¶
- Parameters:
hidden_states (torch.Tensor) – Input (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Vocab logits (batch, seq, vocab_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertOnlyMLMHead(config)¶
Bases:
ModuleMLM head: predictions only (no NSP).
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(sequence_output)¶
- Parameters:
sequence_output (torch.Tensor) – Encoder output (batch, seq, hidden_size)
- Return type:
torch.Tensor
- Returns:
Vocab logits (batch, seq, vocab_size)
- class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertForMaskedLM(config)¶
Bases:
BertPreTrainedModelBERT with masked language modeling head.
Used as the SMILES decoder in the VAE autoencoder.
- __init__(config)¶
- Parameters:
config (BertConfig) – Model hyperparameters
- forward(input_ids=None, attention_mask=None, token_type_ids=None, position_ids=None, head_mask=None, inputs_embeds=None, encoder_embeds=None, encoder_hidden_states=None, encoder_attention_mask=None, labels=None, is_decoder=False, mode='multi_modal', return_logits=False)¶
Run masked language model forward pass.
- Parameters:
input_ids (torch.Tensor) – Token IDs (batch, seq)
attention_mask (torch.Tensor) – 1/0 mask (batch, seq)
token_type_ids (torch.Tensor) – Segment IDs (batch, seq)
position_ids (torch.Tensor) – Position indices (batch, seq)
head_mask (torch.Tensor) – Per-layer head mask or None
inputs_embeds (torch.Tensor) – Pre-computed embeddings (batch, seq, hidden_size)
encoder_embeds (torch.Tensor) – Bypass embedding layer
encoder_hidden_states – Cross-attention source or None
encoder_attention_mask – Encoder mask or None
labels (torch.Tensor) – Target token IDs for MLM loss. If None, skips loss computation
is_decoder (bool) – Enable causal masking
mode (str) – Layer range (‘text’, ‘fusion’, ‘multi_modal’)
return_logits (bool) – If True, return raw logits tensor instead of the standard output tuple
- Return type:
torch.Tensor or tuple
- Returns:
Logits tensor if return_logits, else (prediction_scores, …) or (loss, prediction_scores, …)