| import math |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from typing import List |
| from sentence_transformers import SentenceTransformer |
| from gsfm import Vocab, GSFM |
|
|
|
|
| def freeze_all(model: nn.Module): |
| for p in model.parameters(): |
| p.requires_grad = False |
|
|
|
|
| def get_gsfm_layers(model) -> list: |
| for attr in ("transformer", "encoder", "layers", "blocks"): |
| obj = getattr(model, attr, None) |
| if obj is None: |
| continue |
| if hasattr(obj, "layer"): |
| return list(obj.layer) |
| if hasattr(obj, "layers"): |
| return list(obj.layers) |
| if isinstance(obj, (nn.ModuleList, nn.Sequential)): |
| return list(obj) |
| cand = [] |
| for n, m in model.named_modules(): |
| if "layer" in n.lower() and any( |
| p.requires_grad is not None for p in m.parameters(recurse=False) |
| ): |
| cand.append(m) |
| return cand |
|
|
|
|
| def get_text_layers(st_model: SentenceTransformer) -> list: |
| try: |
| hf_model = st_model[0].auto_model |
| if (hasattr(hf_model, "encoder") or hasattr(hf_model, "enc")) and hasattr( |
| hf_model.encoder, "layer" |
| ): |
| return list(hf_model.encoder.layer) |
| except Exception: |
| pass |
| return [] |
|
|
|
|
| def unfreeze_last_n(layers: list, n: int): |
| if n <= 0: |
| return |
| for lyr in layers[-n:]: |
| for p in lyr.parameters(): |
| p.requires_grad = True |
|
|
|
|
| class VariableDepthMLP(nn.Module): |
| def __init__(self, input_dim, hidden_dim, output_dim, depth=2, dropout=0.1): |
| super().__init__() |
| assert depth >= 1 |
| layers = [] |
| if depth == 1: |
| layers.append(nn.Linear(input_dim, output_dim)) |
| else: |
| layers.extend( |
| [nn.Linear(input_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout)] |
| ) |
| for _ in range(depth - 2): |
| layers.extend( |
| [nn.Linear(hidden_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout)] |
| ) |
| layers.append(nn.Linear(hidden_dim, output_dim)) |
| layers.append(nn.LayerNorm(output_dim)) |
| self.net = nn.Sequential(*layers) |
|
|
| def forward(self, x): |
| return self.net(x) |
|
|
|
|
| class GeneSetTextEncoder(nn.Module): |
| def __init__( |
| self, |
| gsfm_model: GSFM, |
| text_model: SentenceTransformer, |
| vocab: Vocab, |
| gsfm_dim: int, |
| text_dim: int, |
| shared_dim: int = 256, |
| gene_proj_hidden: int = 512, |
| gene_proj_depth: int = 2, |
| text_proj_hidden: int = 512, |
| text_proj_depth: int = 2, |
| dropout: float = 0.1, |
| temperature_init: float = 0.07, |
| n_unfreeze_gsfm: int = 0, |
| n_unfreeze_text: int = 0, |
| ): |
| super().__init__() |
| self.gsfm = gsfm_model |
| self.text_model = text_model |
| self.vocab = vocab |
| self.n_unfreeze_gsfm = int(n_unfreeze_gsfm) |
| self.n_unfreeze_text = int(n_unfreeze_text) |
|
|
| freeze_all(self.gsfm) |
| freeze_all(self.text_model) |
|
|
| g_layers = get_gsfm_layers(self.gsfm) |
| t_layers = get_text_layers(self.text_model) |
| unfreeze_last_n(g_layers, self.n_unfreeze_gsfm) |
| unfreeze_last_n(t_layers, self.n_unfreeze_text) |
|
|
| self.gene_proj = VariableDepthMLP( |
| input_dim=gsfm_dim, |
| hidden_dim=gene_proj_hidden, |
| output_dim=shared_dim, |
| depth=gene_proj_depth, |
| dropout=dropout, |
| ) |
| self.text_proj = VariableDepthMLP( |
| input_dim=text_dim, |
| hidden_dim=text_proj_hidden, |
| output_dim=shared_dim, |
| depth=text_proj_depth, |
| dropout=dropout, |
| ) |
|
|
| self.log_temperature = nn.Parameter( |
| torch.tensor(math.log(temperature_init), dtype=torch.float32) |
| ) |
|
|
| self.gene_decoder = nn.Sequential( |
| nn.Linear(shared_dim, gene_proj_hidden), |
| nn.GELU(), |
| nn.Dropout(dropout), |
| nn.Linear(gene_proj_hidden, len(vocab)), |
| ) |
|
|
| @property |
| def temperature(self): |
| return torch.clamp(self.log_temperature.exp(), min=0.01, max=1.0) |
|
|
| def encode_genes(self, gene_ids: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: |
| out = self.gsfm.encode(gene_ids) |
| if out.dim() == 3: |
| m = mask.unsqueeze(-1).float() |
| out = (out * m).sum(dim=1) / m.sum(dim=1).clamp(min=1.0) |
| return out |
|
|
| def encode_text(self, texts: List[str], device: torch.device = None) -> torch.Tensor: |
| if device is None: |
| device = next(self.parameters()).device |
| if self.n_unfreeze_text > 0: |
| tok = self.text_model.tokenize(texts) |
| tok = {k: v.to(device) for k, v in tok.items()} |
| hf = self.text_model[0].auto_model(**tok) |
| token_emb = hf.last_hidden_state |
| att = tok["attention_mask"].unsqueeze(-1).float() |
| sent = (token_emb * att).sum(dim=1) / att.sum(dim=1).clamp(min=1.0) |
| return sent |
| with torch.no_grad(): |
| return self.text_model.encode( |
| texts, |
| convert_to_tensor=True, |
| show_progress_bar=False, |
| device=str(device), |
| ) |
|
|
| def encode_gene_symbols( |
| self, gene_symbols: List[str], device: torch.device = None |
| ) -> torch.Tensor: |
| """Encode a gene set from symbol strings → normalized CLIP embedding.""" |
| if device is None: |
| device = next(self.parameters()).device |
|
|
| ids = self.vocab(gene_symbols) |
| pad_id = getattr(self.vocab, "pad_id", 0) |
| ids = [i for i in ids if i != pad_id] |
|
|
| if len(ids) == 0: |
| raise ValueError( |
| f"All symbols mapped to pad token: {gene_symbols[:10]}" |
| ) |
|
|
| gene_ids = torch.tensor([ids], dtype=torch.long, device=device) |
| mask = torch.ones_like(gene_ids, dtype=torch.long) |
|
|
| g_raw = self.encode_genes(gene_ids, mask) |
| g = F.normalize(self.gene_proj(g_raw), dim=-1) |
| return g |
|
|
| def encode_text_normalized( |
| self, texts: List[str], device: torch.device = None |
| ) -> torch.Tensor: |
| """Encode text descriptions → normalized CLIP embedding.""" |
| if device is None: |
| device = next(self.parameters()).device |
| t_raw = self.encode_text(texts, device=device) |
| t = F.normalize(self.text_proj(t_raw), dim=-1) |
| return t |
|
|
| def forward( |
| self, |
| gene_ids: torch.Tensor, |
| mask: torch.Tensor, |
| texts: List[str], |
| device: torch.device = None, |
| ): |
| if device is None: |
| device = next(self.parameters()).device |
|
|
| if self.n_unfreeze_gsfm > 0: |
| g_raw = self.encode_genes(gene_ids, mask) |
| else: |
| with torch.no_grad(): |
| g_raw = self.encode_genes(gene_ids, mask) |
|
|
| t_raw = self.encode_text(texts, device=device) |
|
|
| g = F.normalize(self.gene_proj(g_raw), dim=-1) |
| t = F.normalize(self.text_proj(t_raw), dim=-1) |
|
|
| logits_g2t = (g @ t.T) / self.temperature |
| logits_t2g = logits_g2t.T |
| recon = self.gene_decoder(g) |
|
|
| return g, t, logits_g2t, logits_t2g, recon |
|
|
| @classmethod |
| def from_pretrained(cls, repo_path: str, device: str = "cpu"): |
| """Load model from a local directory or HuggingFace repo.""" |
| import json, os |
|
|
| try: |
| from huggingface_hub import hf_hub_download, list_repo_files |
|
|
| files = list_repo_files(repo_path) |
| local_dir = {} |
| for fn in ["config.json", "gtCLIP.pt"]: |
| if fn in files: |
| local_dir[fn] = hf_hub_download(repo_path, fn) |
| config_path = local_dir.get("config.json") |
| weights_path = local_dir.get("gtCLIP.pt") |
| except Exception: |
| config_path = os.path.join(repo_path, "config.json") |
| weights_path = os.path.join(repo_path, "gtCLIP.pt") |
|
|
| with open(config_path) as f: |
| config = json.load(f) |
|
|
| gsfm_name = config.get("gsfm_pretrained", "maayanlab/gsfm") |
| text_name = config.get("text_model_name", "pritamdeka/S-PubMedBert-MS-MARCO") |
|
|
| vocab = Vocab.from_pretrained(gsfm_name) |
| gsfm = GSFM.from_pretrained(gsfm_name).to(device) |
| gsfm.eval() |
| st = SentenceTransformer(text_name, device=device) |
| st.eval() |
|
|
| gene_proj_depth = int(config.get("gene_proj_depth", 2)) |
| text_proj_depth = int(config.get("text_proj_depth", 2)) |
|
|
| model = cls( |
| gsfm_model=gsfm, |
| text_model=st, |
| vocab=vocab, |
| gsfm_dim=int(config.get("gsfm_dim", 256)), |
| text_dim=int(config.get("text_dim", 768)), |
| shared_dim=int(config.get("shared_dim", 256)), |
| gene_proj_hidden=int(config.get("gene_proj_hidden", 256)), |
| gene_proj_depth=gene_proj_depth, |
| text_proj_hidden=int(config.get("text_proj_hidden", 256)), |
| text_proj_depth=text_proj_depth, |
| dropout=float(config.get("dropout", 0.1)), |
| temperature_init=float(config.get("temperature_init", 0.07)), |
| n_unfreeze_gsfm=int(config.get("n_unfreeze_gsfm", 0)), |
| n_unfreeze_text=int(config.get("n_unfreeze_text", 0)), |
| ).to(device) |
|
|
| ckpt = torch.load(weights_path, map_location=device) |
| state = ckpt.get("model_state_dict", ckpt) |
| model.load_state_dict(state, strict=True) |
| model.eval() |
| return model |
|
|