kelu01 commited on
Commit
4792a76
·
verified ·
1 Parent(s): 0b570c8

Update pipe.py

Browse files
Files changed (1) hide show
  1. pipe.py +26 -29
pipe.py CHANGED
@@ -80,52 +80,49 @@ class SmilesDiffusionPipe(Pipeline):
80
 
81
  def unmask_partial_smiles(self, input_ids, model, voc, steps, k):
82
  """
83
- Iteratively unmask a SMILES sequence using a trained diffusion model.
84
- Always fills the masked token with highest confidence first, until all masks are gone.
85
  """
86
- model.eval()
87
- with torch.no_grad():
88
- sequences = torch.tensor(
89
- input_ids, dtype=torch.long, device=self.device
90
- ).unsqueeze(0)
91
-
92
- mask_token = voc.vocab["[MASK]"]
93
- pad_token = voc.vocab.get("[PAD]", None)
94
 
95
- while True:
96
- mask_positions = sequences == mask_token
97
- num_masked = mask_positions.sum().item()
98
- if num_masked == 0:
99
- break # stop only when all masks are filled
100
 
101
- # ---- estimate diffusion time t ----
 
 
102
  if pad_token is not None:
103
- valid = sequences != pad_token
104
- frac_masked = mask_positions.sum().float() / valid.sum()
105
  else:
106
- frac_masked = mask_positions.sum().float() / sequences.size(1)
107
- t = torch.tensor([frac_masked], device=self.device)
108
 
109
- # ---- forward pass ----
110
- logits = model(sequences, t=t)
111
  probs = F.softmax(logits, dim=-1)
112
 
113
- mask_indices = mask_positions.nonzero(as_tuple=False)
114
  masked_probs = probs[mask_positions]
115
 
116
- # ---- pick the mask with highest confidence ----
 
 
 
117
  masked_confidence = masked_probs.max(dim=-1).values
118
  best_idx = torch.argmax(masked_confidence)
119
-
120
- sampled_id = torch.multinomial(masked_probs[best_idx], num_samples=1).item()
121
  pos = mask_indices[best_idx]
122
 
123
- # ---- fill that mask ----
124
- sequences[pos[0], pos[1]] = sampled_id
 
125
 
126
  # ---- decode ----
127
  decoded = [
128
- t for t in sequences[0].tolist()
129
  if t not in (pad_token, voc.vocab.get("<s>"), voc.vocab.get("</s>"))
130
  ]
131
  if len(decoded) > 2:
 
80
 
81
  def unmask_partial_smiles(self, input_ids, model, voc, steps, k):
82
  """
83
+ Iteratively unmask a single SMILES string using a diffusion model.
84
+ Always fills the masked token with highest confidence first until done.
85
  """
86
+ pad_token = voc.vocab["[PAD]"]
87
+ mask_token = voc.vocab["[MASK]"]
 
 
 
 
 
 
88
 
89
+ # Encode input SMILES
90
+ encoded = voc.encode(voc.tokenize(input_ids))
91
+ unmasked = torch.tensor(encoded, dtype=torch.long).unsqueeze(0).to(device)
 
 
92
 
93
+ with torch.no_grad():
94
+ while (unmasked == mask_token).any():
95
+ # Estimate diffusion time t
96
  if pad_token is not None:
97
+ valid_mask = unmasked != pad_token
98
+ frac_masked = ((unmasked == mask_token).sum().float() / valid_mask.sum())
99
  else:
100
+ frac_masked = ((unmasked == mask_token).sum().float() / unmasked.size(1))
101
+ t_val = torch.tensor([frac_masked], device=device)
102
 
103
+ # Forward pass
104
+ logits = model.model(unmasked, t=t_val)
105
  probs = F.softmax(logits, dim=-1)
106
 
107
+ mask_positions = (unmasked == mask_token)
108
  masked_probs = probs[mask_positions]
109
 
110
+ if masked_probs.size(0) == 0:
111
+ break # safety check
112
+
113
+ # Pick the mask with highest confidence
114
  masked_confidence = masked_probs.max(dim=-1).values
115
  best_idx = torch.argmax(masked_confidence)
116
+ mask_indices = mask_positions.nonzero(as_tuple=False)
 
117
  pos = mask_indices[best_idx]
118
 
119
+ # Unmask the selected token
120
+ sampled_id = torch.multinomial(masked_probs[best_idx], num_samples=1).item()
121
+ unmasked[0, pos[1]] = sampled_id
122
 
123
  # ---- decode ----
124
  decoded = [
125
+ t for t in unmasked[0].tolist()
126
  if t not in (pad_token, voc.vocab.get("<s>"), voc.vocab.get("</s>"))
127
  ]
128
  if len(decoded) > 2: