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