import torch import torch.nn as nn from transformers import PreTrainedModel from .configuration_diffusion import DiffusionConfig class DiffusionTransformer(nn.Module): def __init__(self, vocab_size, d_model=512, n_heads=16, n_layers=6, dropout=0.1, max_len=2250, num_steps=1000, pad_token_id=0): super().__init__() self.vocab_size = vocab_size self.max_len = max_len self.num_steps = num_steps self.pad_token_id = pad_token_id self.token_embed = nn.Embedding(vocab_size, d_model) self.pos_embed = nn.Embedding(max_len, d_model) self.time_embed = nn.Embedding(num_steps, d_model) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=n_heads, dim_feedforward=4 * d_model, dropout=dropout, batch_first=True ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_layers) self.output_head = nn.Linear(d_model, vocab_size) def forward(self, input_ids, t=None): B, T = input_ids.shape device = input_ids.device positions = torch.arange(T, device=device).unsqueeze(0).expand(B, T) tok_emb = self.token_embed(input_ids) pos_emb = self.pos_embed(positions) if t is not None: t = t.to(device) t_ids = (t * (self.num_steps - 1)).long() t_emb = self.time_embed(t_ids).unsqueeze(1) else: t_emb = torch.zeros_like(tok_emb[:, :1]) x = tok_emb + pos_emb + t_emb pad_mask = input_ids == self.pad_token_id x = self.transformer(x, src_key_padding_mask=pad_mask) logits = self.output_head(x) return logits class HuggingFaceDiffusionModel(PreTrainedModel): config_class = DiffusionConfig def __init__(self, config): super().__init__(config) self.model = DiffusionTransformer( vocab_size=config.vocab_size, d_model=config.d_model, n_heads=config.n_heads, n_layers=config.n_layers, dropout=config.dropout, max_len=config.max_len, num_steps=config.num_steps, pad_token_id=config.pad_token ) def forward(self, input_ids, t=None): return self.model(input_ids, t) def get_latent(self, input_ids, t=None, pool="mean"): """ Returns embeddings for a SMILES sequence. input_ids: torch.LongTensor (B, T) pool: "mean", "max", or None """ with torch.no_grad(): x = self.model.token_embed(input_ids) + \ self.model.pos_embed( torch.arange(input_ids.size(1), device=input_ids.device) .unsqueeze(0) .expand(input_ids.size(0), -1) ) if t is not None: t_ids = (t * (self.model.num_steps - 1)).long().to(input_ids.device) t_emb = self.model.time_embed(t_ids).unsqueeze(1) else: t_emb = torch.zeros_like(x[:, :1]) x = x + t_emb x = self.model.transformer(x) if pool == "mean": latent = x.mean(dim=1) elif pool == "max": latent, _ = x.max(dim=1) elif pool is None: latent = x # return full sequence embeddings else: raise ValueError("pool must be 'mean', 'max', or None") return latent