gtCLIP / gtCLIP.py
DaveLab's picture
Upload folder using huggingface_hub
46130c3 verified
Raw
History Blame Contribute Delete
9.57 kB
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