-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
53 lines (43 loc) · 1.8 KB
/
Copy pathtrain.py
File metadata and controls
53 lines (43 loc) · 1.8 KB
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
import torch
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from vae import ChemVAE
from dataset import ChemDataset
import os
def vae_loss(recon_x, x, mu, logvar):
# ignore padding index in bce
BCE = F.cross_entropy(recon_x.view(-1, recon_x.size(-1)), x.view(-1), ignore_index=0, reduction='sum')
KLD = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
return BCE + KLD
def train():
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
smiles_data = ["CCO", "CCN", "CCC", "C=C", "CC=O", "c1ccccc1", "CC(=O)O", "CN1CCC1"] * 100
chars = set("".join(smiles_data))
char_to_idx = {c: i+1 for i, c in enumerate(chars)}
char_to_idx['<PAD>'] = 0
vocab_size = len(char_to_idx)
max_len = 20
dataset = ChemDataset(smiles_data, char_to_idx, max_len)
loader = DataLoader(dataset, batch_size=32, shuffle=True)
model = ChemVAE(vocab_size=vocab_size, embed_size=128, hidden_size=256, latent_size=64).to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
best_loss = float('inf')
for epoch in range(1, 21):
model.train()
train_loss = 0
for batch in loader:
batch = batch.to(device)
optimizer.zero_grad()
recon_batch, mu, logvar = model(batch)
loss = vae_loss(recon_batch, batch, mu, logvar)
loss.backward()
train_loss += loss.item()
optimizer.step()
avg_loss = train_loss / len(loader.dataset)
print(f'Epoch: {epoch} Average loss: {avg_loss:.4f}')
if avg_loss < best_loss:
best_loss = avg_loss
torch.save(model.state_dict(), 'vae_final_weights.pt')
if __name__ == '__main__':
train()