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: object

Configuration 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:

BertConfig

Returns:

Config instance with JSON values

class schrodinger.application.matsci.flywheel.ldmol.bert_decoder.BertEmbeddings(config)

Bases: Module

Word, 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: Module

Multi-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: Module

Residual 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: Module

Self-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: Module

Feed-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: Module

Feed-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: Module

Single 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: Module

Stack 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: Module

Pool [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: Module

Base 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: BertPreTrainedModel

BERT 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: Module

Dense + 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: Module

Transform 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: Module

MLM 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: BertPreTrainedModel

BERT 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, …)