Feature Extraction
Transformers
Safetensors
mettle
computational-pathology
histopathology
foundation-model
scanner-robustness
custom_code
Instructions to use slideflow-labs/Mettle with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use slideflow-labs/Mettle with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="slideflow-labs/Mettle", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("slideflow-labs/Mettle", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 5,952 Bytes
0e83a2b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 | """Hugging Face model implementation for the Mettle tile encoder."""
from dataclasses import dataclass
from typing import Optional
import timm
import torch
import torch.nn as nn
from transformers import PreTrainedModel
from transformers.utils import ModelOutput
try:
from .configuration_mettle import MettleConfig
except ImportError:
from configuration_mettle import MettleConfig
class MettleRefineHead(nn.Module):
"""Content-conditional refinement of the pooled CLS embedding."""
def __init__(self, dim, num_atoms, rank, hidden_size):
super().__init__()
self.P = nn.Parameter(torch.zeros(num_atoms, dim, rank))
self.Q = nn.Parameter(torch.zeros(num_atoms, dim, rank))
self.mix = nn.Sequential(
nn.Linear(dim, hidden_size),
nn.GELU(),
nn.Linear(hidden_size, num_atoms),
)
self.shift = nn.Sequential(
nn.Linear(dim, hidden_size),
nn.GELU(),
nn.Linear(hidden_size, dim),
)
def forward(self, embedding):
mixture = self.mix(embedding)
shift = self.shift(embedding)
projected = torch.einsum("bd,kdr->bkr", embedding, self.P)
delta = torch.einsum("bkr,kdr->bkd", projected, self.Q)
return embedding + (mixture.unsqueeze(-1) * delta).sum(1) + shift
@dataclass
class MettleModelOutput(ModelOutput):
"""Outputs from :class:`MettleModel`.
``pooler_output`` follows the requested feature view. ``cls_embedding`` is
the headed 1536-dimensional CLS representation. ``cls_mean_embedding`` is
the concatenation of headed CLS and the unheaded mean spatial patch token.
"""
last_hidden_state: Optional[torch.Tensor] = None
pooler_output: Optional[torch.Tensor] = None
cls_embedding: Optional[torch.Tensor] = None
mean_patch_embedding: Optional[torch.Tensor] = None
cls_mean_embedding: Optional[torch.Tensor] = None
class MettleModel(PreTrainedModel):
"""Mettle scanner-resilient pathology tile encoder."""
config_class = MettleConfig
base_model_prefix = "backbone"
main_input_name = "pixel_values"
def __init__(self, config):
super().__init__(config)
self.backbone = timm.create_model(
config.backbone_name,
pretrained=False,
num_classes=0,
init_values=1e-5,
dynamic_img_size=False,
img_size=config.image_size,
)
if int(self.backbone.num_features) != config.hidden_size:
raise ValueError(
f"{config.backbone_name} has width {self.backbone.num_features}, "
f"but the config declares {config.hidden_size}"
)
self.head = None
if config.head_enabled:
self.head = MettleRefineHead(
dim=config.hidden_size,
num_atoms=config.head_num_atoms,
rank=config.head_rank,
hidden_size=config.head_hidden_size,
)
self.post_init()
def _embeddings(self, pixel_values):
tokens = self.backbone.forward_features(pixel_values)
expected_prefix = 1 + self.config.num_register_tokens
actual_prefix = int(self.backbone.num_prefix_tokens)
if actual_prefix != expected_prefix:
raise RuntimeError(
f"expected {expected_prefix} prefix tokens, found {actual_prefix}"
)
patches = tokens[:, actual_prefix:]
if patches.shape[1] == 0:
raise RuntimeError("the backbone returned no spatial patch tokens")
cls_embedding = tokens[:, 0].to(torch.float32)
if self.head is not None:
cls_embedding = self.head(cls_embedding)
mean_patch_embedding = patches.to(torch.float32).mean(dim=1)
cls_mean_embedding = torch.cat(
(cls_embedding, mean_patch_embedding),
dim=-1,
)
return tokens, cls_embedding, mean_patch_embedding, cls_mean_embedding
def forward(
self,
pixel_values,
feature_view=None,
return_dict=None,
):
"""Encode normalized 224-pixel RGB tiles.
Args:
pixel_values: Tensor with shape ``(batch, 3, 224, 224)`` after the
preprocessing specified by ``preprocessor_config.json``.
feature_view: ``"cls"`` for a 1536-dimensional embedding or
``"cls_mean"`` for headed CLS concatenated with mean spatial
patch pooling (3072 dimensions).
return_dict: Follow the standard Transformers return convention.
"""
feature_view = feature_view or self.config.default_feature_view
if feature_view not in {"cls", "cls_mean"}:
raise ValueError(
"feature_view must be 'cls' or 'cls_mean', got "
f"{feature_view!r}"
)
return_dict = (
self.config.use_return_dict if return_dict is None else return_dict
)
tokens, cls_embedding, mean_patch, cls_mean = self._embeddings(
pixel_values
)
selected = cls_embedding if feature_view == "cls" else cls_mean
if not return_dict:
return tokens, selected, cls_embedding, mean_patch, cls_mean
return MettleModelOutput(
last_hidden_state=tokens,
pooler_output=selected,
cls_embedding=cls_embedding,
mean_patch_embedding=mean_patch,
cls_mean_embedding=cls_mean,
)
def encode(self, pixel_values, feature_view=None):
"""Return only the selected tile embedding."""
return self(
pixel_values=pixel_values,
feature_view=feature_view,
return_dict=True,
).pooler_output
def get_input_embeddings(self):
"""Return the ViT patch projection for Transformers integrations."""
return self.backbone.patch_embed.proj
|