| """ELF slide encoder: interpolate → LayerNorm → 8-head ABMIL.""" |
| from __future__ import annotations |
|
|
| from collections import OrderedDict |
| from pathlib import Path |
| from typing import Optional, Union |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
| DEFAULT_REPO_ID = "luoxd96/ELF" |
| WEIGHTS_FILE = "elf_slide_encoder.pth" |
|
|
|
|
| class BatchedABMIL(nn.Module): |
| def __init__(self, dim: int): |
| super().__init__() |
| self.attention_a = nn.Sequential(nn.Linear(dim, dim), nn.Tanh()) |
| self.attention_b = nn.Sequential(nn.Linear(dim, dim), nn.Sigmoid()) |
| self.attention_c = nn.Linear(dim, 1) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.attention_c(self.attention_a(x) * self.attention_b(x)) |
|
|
|
|
| class ELFSlideEncoder(nn.Module): |
| def __init__(self, embed_dim: int = 768, num_heads: int = 8): |
| super().__init__() |
| if embed_dim % num_heads != 0: |
| raise ValueError(f"embed_dim ({embed_dim}) must be divisible by num_heads ({num_heads})") |
| self.embed_dim = embed_dim |
| self.num_heads = num_heads |
| self.norm = nn.LayerNorm(embed_dim) |
| self.attn = nn.ModuleList( |
| [BatchedABMIL(embed_dim // num_heads) for _ in range(num_heads)] |
| ) |
|
|
| def forward(self, x: torch.Tensor, lens: Optional[torch.Tensor] = None): |
| """ |
| Args: |
| x: ``[B, N, C]`` patch features (``C`` in {768, 1024, 1280, 1536}). |
| lens: ``[B]`` native ``C`` per item; defaults to ``x.shape[-1]``. |
| |
| Returns: |
| features_dim: ``[B, C]`` — ``softmax(ᾱ)ᵀ X`` |
| features: ``[B, 768]`` — ``softmax(ᾱ)ᵀ X_768`` |
| attention: ``[B, 1, N]`` |
| """ |
| if x.ndim != 3: |
| raise ValueError(f"expected [B, N, C], got {tuple(x.shape)}") |
| batch, n_tiles, feat_dim = x.shape |
| if lens is None: |
| lens = torch.full((batch,), feat_dim, dtype=torch.long, device=x.device) |
|
|
| x768 = [] |
| for i in range(batch): |
| c = int(lens[i].item()) |
| x768.append( |
| F.interpolate(x[i, :, :c].unsqueeze(0), size=self.embed_dim, mode="linear", align_corners=True).squeeze(0) |
| ) |
| x768 = self.norm(torch.stack(x768, dim=0)) |
|
|
| head_dim = self.embed_dim // self.num_heads |
| heads = x768.view(batch, n_tiles, head_dim, self.num_heads) |
| logits = torch.stack([self.attn[h](heads[:, :, :, h]) for h in range(self.num_heads)], dim=-1) |
| attn = F.softmax(logits.mean(dim=-1).transpose(1, 2), dim=-1) |
|
|
| feat_768 = torch.bmm(attn, x768)[:, 0] |
| feat_native = torch.stack( |
| [torch.bmm(attn[i : i + 1], x[i : i + 1, :, : int(lens[i].item())])[0, 0] for i in range(batch)] |
| ) |
| return feat_native, feat_768, attn |
|
|
| @classmethod |
| def from_pretrained( |
| cls, |
| repo_id: str = DEFAULT_REPO_ID, |
| filename: str = WEIGHTS_FILE, |
| device: Union[str, torch.device] = "cpu", |
| embed_dim: int = 768, |
| num_heads: int = 8, |
| ) -> "ELFSlideEncoder": |
| from huggingface_hub import hf_hub_download |
|
|
| path = hf_hub_download(repo_id=repo_id, filename=filename) |
| return load_encoder(path, device=device, embed_dim=embed_dim, num_heads=num_heads) |
|
|
|
|
| def preprocess_patch_features(features: torch.Tensor, foundation_model: Optional[str] = None) -> torch.Tensor: |
| x = features.float() |
| if (foundation_model or "").lower() == "virchow2" and x.shape[-1] >= 2560: |
| x = 0.5 * (x[..., :1280] + x[..., 1280:2560]) |
| return x |
|
|
|
|
| def _unwrap_state_dict(raw) -> OrderedDict: |
| if isinstance(raw, dict) and "state_dict" in raw: |
| raw = raw["state_dict"] |
| return OrderedDict((k[7:] if k.startswith("module.") else k, v) for k, v in raw.items()) |
|
|
|
|
| def extract_inference_weights(state_dict: dict) -> OrderedDict: |
| keys = list(state_dict.keys()) |
| prefix = "" |
| if any(k.startswith("momentum_enc.") for k in keys): |
| prefix = "momentum_enc." |
| keep = ("norm.", "attn.") |
| out = OrderedDict( |
| (k[len(prefix) :], v) |
| for k, v in state_dict.items() |
| if k.startswith(prefix) and k[len(prefix) :].startswith(keep) |
| ) |
| if not out: |
| raise KeyError(f"no norm/attn weights found; prefixes={sorted({k.split('.')[0] for k in keys})[:12]}") |
| return out |
|
|
|
|
| def load_encoder( |
| checkpoint: Union[str, Path], |
| device: Union[str, torch.device] = "cpu", |
| embed_dim: int = 768, |
| num_heads: int = 8, |
| ) -> ELFSlideEncoder: |
| ckpt = torch.load(str(checkpoint), map_location="cpu", weights_only=False) |
| weights = extract_inference_weights(_unwrap_state_dict(ckpt)) |
| model = ELFSlideEncoder(embed_dim=embed_dim, num_heads=num_heads) |
| missing, unexpected = model.load_state_dict(weights, strict=True) |
| if missing or unexpected: |
| raise RuntimeError(f"load mismatch missing={missing} unexpected={unexpected}") |
| return model.to(device).eval() |
|
|