File size: 3,524 Bytes
65f1711 3200eec e4b08c4 3200eec f628a4e 3200eec 5da209a 3200eec fca4acc 8b94352 3200eec 0fbf8be d7f042d 3200eec f061e29 06b7798 f061e29 06b7798 f061e29 06b7798 f061e29 06b7798 f061e29 06b7798 f061e29 06b7798 f061e29 06b7798 f061e29 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 | 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
|