| 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 |
| else: |
| raise ValueError("pool must be 'mean', 'max', or None") |
| |
| return latent |
|
|