Download pipe.py from kelu01/Chemfuser: direct link, hf CLI and curl.
- Browser
- Download file 5.71 kB
-
https://huggingface.co/kelu01/Chemfuser/resolve/main/pipe.py
- Command line
-
hf download hf://kelu01/Chemfuser/pipe.py
-
curl -L -o pipe.py https://huggingface.co/kelu01/Chemfuser/resolve/main/pipe.py
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 |