Chemfuser / pipe.py
kelu01's picture
Update pipe.py
cf604be verified
Raw History Blame Contribute Delete
5.71 kB
from transformers import Pipeline
import torch
import torch.nn.functional as F
import numpy as np
import torch
import torch.nn.functional as F
import re
class SmilesDiffusionPipe(Pipeline):
def __init__(self, model, voc, **kwargs):
super().__init__(model=model, **kwargs)
self.voc = voc
def restore_halogen(self, smiles):
smiles = re.sub(r'(?<!\[)L(?!\])', 'Cl', smiles)
smiles = re.sub(r'(?<!\[)R(?!\])', 'Br', smiles)
return smiles
def replace_halogen(self, smiles):
return smiles.replace('Cl', 'L').replace('Br', 'R')
def tokenize_smiles(self, smiles):
regex = r'(\[[^\[\]]{1,6}\])'
smiles = self.replace_halogen(smiles)
char_list = re.split(regex, smiles)
tokenized = ['[GO]']
for char in char_list:
if char.startswith('['):
tokenized.append(char)
else:
chars = [unit for unit in char]
[tokenized.append(unit) for unit in chars]
tokenized.append("[EOS]")
return tokenized
def encode(self, char_list):
smiles_matrix = np.zeros(len(char_list), dtype=np.float32)
unk_id = self.voc.vocab.get("[UNK]")
if unk_id is None:
raise ValueError("[UNK] token not found in vocabulary.")
for i, char in enumerate(char_list):
smiles_matrix[i] = self.voc.vocab.get(char, unk_id)
return smiles_matrix
def decode(self, matrix):
""""Takes an array of indices (can be list of tensors or ints) and returns the corresponding SMILES"""
matrix = [int(i) for i in matrix] # This handles tensors from PyTorch
# Using tokenizer's decode method to convert token IDs to text
smiles = self.voc.decode(matrix)
smiles = self.restore_halogen(smiles)
return smiles
def _sanitize_parameters(self, **kwargs):
preprocess_kwargs = {}
postprocess_kwargs = {}
if "steps" in kwargs:
preprocess_kwargs["steps"] = kwargs["steps"]
if "k" in kwargs:
preprocess_kwargs["k"] = kwargs["k"]
return preprocess_kwargs, {}, postprocess_kwargs
def preprocess(self, smiles_str, steps=30, k=1):
tokenized = self.tokenize_smiles(smiles_str)
encoded = self.encode(tokenized)
return (encoded, steps, k)
def _forward(self, inputs):
encoded, steps, k = inputs
return self.unmask_partial_smiles(encoded, self.model, self.voc, steps, k)
def postprocess(self, result):
return result
def multinomial_sample(self, probs):
return torch.multinomial(probs, 1)
def unmask_partial_smiles(self, input_ids, model, voc, steps=30, k=3):
model.eval()
with torch.no_grad():
sequences = torch.tensor(input_ids, dtype=torch.long, device=self.device).unsqueeze(0)
mask_token = voc.vocab['[MASK]']
pad_token = voc.vocab.get('[PAD]', None)
total_tokens = sequences.size(1)
original_num_masked = (sequences == mask_token).sum().item()
print(f"Original number of masked tokens: {original_num_masked}")
if total_tokens == 0:
raise ValueError("Encoded SMILES has no tokens.")
if original_num_masked == 0:
print("No masked tokens found. Returning original SMILES.")
input_ids = input_ids[1:-1]
return self.decode(input_ids)
filled_tokens_confidence = {}
# --- ITERATION LOGIC: unmask one token at a time based on highest confidence ---
num_masks = (sequences == mask_token).sum().item()
for _ in range(num_masks):
# Estimate t per sequence
frac_masked = (sequences == mask_token).sum().item() / total_tokens
t = torch.tensor([frac_masked], device=self.device)
logits = model(sequences, t=t)
probs = F.softmax(logits, dim=-1)
mask_positions = (sequences == mask_token)
mask_indices = mask_positions.nonzero(as_tuple=False)
if len(mask_indices) == 0:
print("All tokens filled.")
break
masked_probs = probs[mask_positions] # (#masked, vocab)
sampled_ids = torch.multinomial(masked_probs, num_samples=1).squeeze(-1)
masked_confidence = torch.max(masked_probs, dim=-1).values
best_idx = torch.argmax(masked_confidence).item()
chosen_pos = mask_indices[best_idx] # [batch=0, position]
chosen_token = sampled_ids[best_idx].item()
chosen_conf = masked_confidence[best_idx].item()
sequences[0, chosen_pos[1]] = chosen_token
filled_tokens_confidence[chosen_pos[1].item()] = chosen_conf
print(f"Filled index {chosen_pos[1].item()} with token {chosen_token} (conf={chosen_conf:.4f})")
# Print confidence for all filled tokens
for idx, confidence in sorted(filled_tokens_confidence.items(), key=lambda x: x[0]):
print(f"Final confidence for token at index {idx}: {confidence:.3f}")
# Decode, removing <s> and </s>
decoded_tokens = [t for t in sequences[0].tolist() if t not in (voc.vocab.get('<s>'), voc.vocab.get('</s>'))]
if len(decoded_tokens) > 2:
decoded_tokens = decoded_tokens[1:-1]
decoded = self.decode(decoded_tokens).replace(" ", "")
return decoded