From 6ab7b37f2824c0db4ebcafa4b01c658a7875514f Mon Sep 17 00:00:00 2001 From: ChrisGadek1 Date: Wed, 22 Jul 2026 18:44:30 +0200 Subject: [PATCH] Add evaluation mode and inference context to embedding methods --- toolbox/models/embedding/embedder/esm2_embedder.py | 4 +++- toolbox/models/embedding/embedder/esmc_embedder.py | 10 ++++++---- toolbox/models/embedding/embedder/glm2_embedder.py | 1 + 3 files changed, 10 insertions(+), 5 deletions(-) diff --git a/toolbox/models/embedding/embedder/esm2_embedder.py b/toolbox/models/embedding/embedder/esm2_embedder.py index 6a73050..732d949 100644 --- a/toolbox/models/embedding/embedder/esm2_embedder.py +++ b/toolbox/models/embedding/embedder/esm2_embedder.py @@ -16,6 +16,8 @@ def __init__(self, device=None, batch_size=1000, model_name='esm2_t33_650M_UR50D def get_embedding(self, prot_id, prot_seq): inputs = self.tokenizer(prot_seq, return_tensors="pt") inputs = {k: v.to(self.device) for k, v in inputs.items()} - outputs = self.model(**inputs, output_hidden_states=True) + self.model.eval() + with torch.inference_mode(): + outputs = self.model(**inputs, output_hidden_states=True) embeddings = outputs.hidden_states[-1] return embeddings[0,:].to('cpu').detach().to(torch.float32).numpy() diff --git a/toolbox/models/embedding/embedder/esmc_embedder.py b/toolbox/models/embedding/embedder/esmc_embedder.py index 097bcaf..ffacc99 100644 --- a/toolbox/models/embedding/embedder/esmc_embedder.py +++ b/toolbox/models/embedding/embedder/esmc_embedder.py @@ -11,8 +11,10 @@ def __init__(self, device=None, batch_size=1000, model_name="esmc_600m"): def get_embedding(self, prot_id, prot_seq): protein = ESMProtein(sequence=prot_seq) - protein_tensor = self.model.encode(protein) - logits_output = self.model.logits( - protein_tensor, LogitsConfig(sequence=True, return_embeddings=True) - ) + self.model.eval() + with torch.inference_mode(): + protein_tensor = self.model.encode(protein) + logits_output = self.model.logits( + protein_tensor, LogitsConfig(sequence=True, return_embeddings=True) + ) return logits_output.embeddings[0,:,:].to('cpu').detach().to(torch.float32).numpy() diff --git a/toolbox/models/embedding/embedder/glm2_embedder.py b/toolbox/models/embedding/embedder/glm2_embedder.py index a369c0b..b66bf77 100644 --- a/toolbox/models/embedding/embedder/glm2_embedder.py +++ b/toolbox/models/embedding/embedder/glm2_embedder.py @@ -27,6 +27,7 @@ def get_embedding(self, prot_id, prot_seq): sequence = PREP_SIGN + prot_seq inputs = self.tokenizer([sequence], return_tensors="pt") inputs = {k: v.to(self.device) for k, v in inputs.items()} + self.model.eval() outputs = self.model(inputs["input_ids"], output_hidden_states=True) embeddings = outputs.last_hidden_state[0] return embeddings.to("cpu").detach().to(torch.float32).numpy()