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