| from pathlib import Path |
| from types import SimpleNamespace |
| import re |
| import torch |
|
|
| from models import EquivariantThreeHopGINE |
| from smiles_to_npy_discretize import smiles_to_graph_with_labels |
|
|
|
|
| class SimpleGraph: |
| """The tiny subset of DGLGraph used by EquivariantThreeHopGINE.""" |
| def __init__(self, src, dst, weights, features): |
| self._src, self._dst = src, dst |
| self.edata = {"weight": weights} |
| self.ndata = {"feat": features} |
|
|
| def edges(self): |
| return self._src, self._dst |
|
|
| def num_edges(self): |
| return int(self._src.numel()) |
|
|
| def to(self, device): |
| self._src = self._src.to(device) |
| self._dst = self._dst.to(device) |
| self.edata = {k: v.to(device) for k, v in self.edata.items()} |
| self.ndata = {k: v.to(device) for k, v in self.ndata.items()} |
| return self |
|
|
|
|
| def build_args(): |
| return SimpleNamespace( |
| hidden_dim=16, |
| codebook_size=10000, |
| edge_emb_dim=32, |
| ema_decay=0.8, |
| dynamic_threshold=True, |
| epoch_at_mode_shift=0, |
| ss_max_total_latent_count=40000, |
| train_or_infer="infer", |
| use_checkpoint=True, |
| ) |
|
|
|
|
| def ensure_codebook_keys(model, state_dict): |
| cb = model.vq._codebook |
| prefixes = ["vq._codebook.", ""] |
| for prefix in prefixes: |
| ea_pat = re.compile(rf"^{re.escape(prefix)}embed_avg_(.+)$") |
| cs_pat = re.compile(rf"^{re.escape(prefix)}cluster_size_(.+)$") |
| emb_pat = re.compile(rf"^{re.escape(prefix)}embed\.(.+)$") |
| created = False |
| for key, value in state_dict.items(): |
| match = ea_pat.match(key) |
| if match and torch.is_tensor(value) and value.ndim == 2: |
| original = match.group(1) |
| k, d = value.shape |
| if not hasattr(cb, f"cluster_size_{original}"): |
| cb.register_buffer(f"cluster_size_{original}", torch.zeros(k)) |
| if not hasattr(cb, f"embed_avg_{original}"): |
| cb.register_buffer(f"embed_avg_{original}", torch.zeros(k, d)) |
| cb._get_or_create_safe_key(original, K_e=int(k), D=int(d), device="cpu") |
| created = True |
| for key, value in state_dict.items(): |
| match = cs_pat.match(key) |
| if match and torch.is_tensor(value) and value.ndim == 1: |
| original = match.group(1) |
| if not hasattr(cb, f"cluster_size_{original}"): |
| cb.register_buffer(f"cluster_size_{original}", torch.zeros_like(value, device="cpu")) |
| created = True |
| for key, value in state_dict.items(): |
| match = emb_pat.match(key) |
| if match and torch.is_tensor(value) and value.ndim == 2: |
| safe = match.group(1) |
| if safe not in cb.embed: |
| cb.embed[safe] = torch.nn.Parameter(torch.zeros_like(value, device="cpu")) |
| created = True |
| if created: |
| return |
|
|
|
|
| def make_graph(adj, features, device): |
| x = torch.as_tensor(features, dtype=torch.float32, device=device) |
| w1 = torch.as_tensor(adj, dtype=torch.float32, device=device) |
| n = x.shape[0] |
| a1 = w1 > 0 |
| a1 = a1 | a1.T |
| a1.fill_diagonal_(True) |
| w1 = w1.clone() |
| w1.fill_diagonal_(1.0) |
| a1f = a1.float() |
| a2 = (a1f @ a1f) > 0 |
| a3 = (a2.float() @ a1f) > 0 |
| two_only = a2 & ~a1 |
| three_only = a3 & ~(a1 | a2) |
| full_w = w1 + two_only.float() * 0.5 + three_only.float() * 0.3 |
| full_w.fill_diagonal_(1.0) |
| src, dst = (full_w > 0).nonzero(as_tuple=True) |
| return SimpleGraph(src, dst, full_w[src, dst], x), n |
|
|
|
|
| def make_masks(features, device): |
| masks = {} |
| for i, row in enumerate(features): |
| key = "_".join(str(int(row[j])) for j in (0, 2, 3, 4, 5, 6)) |
| masks.setdefault(key, []).append(i) |
| return {k: torch.tensor(v, dtype=torch.long, device=device) for k, v in masks.items()} |
|
|
|
|
| class VQAtomTokenizer: |
| def __init__(self, checkpoint="data/model_epoch_3.pt", device=None): |
| self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu")) |
| args = build_args() |
| self.model = EquivariantThreeHopGINE( |
| in_feats=64, hidden_feats=16, out_feats=16, args=args |
| ) |
| ckpt = torch.load(Path(checkpoint), map_location="cpu", weights_only=False) |
| state = ckpt.get("model", ckpt.get("state_dict", ckpt)) if isinstance(ckpt, dict) else ckpt |
| drop = ( |
| "vq._codebook.usage_ema_k_", "vq._codebook.split_cd_k_", |
| "vq._codebook.ever_used_k_", "vq._codebook.last_used_ep_k_", |
| ) |
| state = {k: v for k, v in state.items() if not any(k.startswith(p) for p in drop)} |
| ensure_codebook_keys(self.model, state) |
| missing, unexpected = self.model.load_state_dict(state, strict=False) |
| essential_missing = [k for k in missing if not any(x in k for x in ("usage_ema", "split_cd", "ever_used", "last_used"))] |
| if essential_missing: |
| raise RuntimeError(f"Checkpoint is missing required weights: {essential_missing[:8]}") |
| self.model.to(self.device).eval() |
|
|
| @torch.inference_mode() |
| def encode(self, smiles): |
| adj, features, _ = smiles_to_graph_with_labels(smiles, 0) |
| if features.shape[1] != 79: |
| raise RuntimeError(f"Expected 79 atom features, got {features.shape[1]}") |
| if features.shape[0] >= 100: |
| raise ValueError("This demo supports molecules with fewer than 100 heavy atoms.") |
| graph, _ = make_graph(adj, features, self.device) |
| masks = make_masks(features, self.device) |
| out = self.model( |
| graph, graph.ndata["feat"], 0, masks, None, 0, |
| None, "infer", [graph.ndata["feat"]], |
| ) |
| _, ids, _ = out |
| key_ids, cluster_ids, global_ids, id2safe = ids |
| return { |
| "tokens": global_ids.detach().cpu().long().tolist(), |
| "key_ids": key_ids.detach().cpu().long().tolist(), |
| "cluster_ids": cluster_ids.detach().cpu().long().tolist(), |
| "id2safe": id2safe, |
| } |
|
|