PathAGG / vits /modeling_pathagg.py
luoxd96's picture
Add transformers model card layout (vitb/vits AutoModel)
b612017 verified
Raw
History Blame Contribute Delete
6.85 kB
"""PathAGG Hugging Face loaders (trust_remote_code).
Usage:
from transformers import AutoModel
model = AutoModel.from_pretrained("luoxd96/PathAGG", subfolder="vitb", trust_remote_code=True)
cls = model(images) # [B, D]
"""
from __future__ import annotations
from pathlib import Path
from typing import Any
import torch
import torch.nn as nn
from transformers import PretrainedConfig, PreTrainedModel
from transformers.utils import cached_file
class PathAGGConfig(PretrainedConfig):
model_type = "pathagg"
def __init__(
self,
variant: str = "vitb",
img_size: int = 224,
patch_size: int = 14,
embed_dim: int = 768,
depth: int = 12,
num_heads: int = 12,
num_register_tokens: int = 4,
mlp_ratio: float = 4.0,
qkv_bias: bool = True,
init_values: float | None = None,
no_embed_class: bool = False,
return_patch: bool = False,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self.variant = variant
self.img_size = img_size
self.patch_size = patch_size
self.embed_dim = embed_dim
self.depth = depth
self.num_heads = num_heads
self.num_register_tokens = num_register_tokens
self.mlp_ratio = mlp_ratio
self.qkv_bias = qkv_bias
self.init_values = init_values
self.no_embed_class = no_embed_class
self.return_patch = return_patch
def _build_backbone(config: PathAGGConfig) -> nn.Module:
from timm.models.vision_transformer import VisionTransformer
kwargs: dict[str, Any] = dict(
img_size=config.img_size,
patch_size=config.patch_size,
in_chans=3,
num_classes=0,
global_pool="",
embed_dim=config.embed_dim,
depth=config.depth,
num_heads=config.num_heads,
mlp_ratio=config.mlp_ratio,
qkv_bias=config.qkv_bias,
reg_tokens=config.num_register_tokens,
)
if config.init_values is not None:
kwargs["init_values"] = config.init_values
if config.no_embed_class:
kwargs["no_embed_class"] = True
return VisionTransformer(**kwargs)
class PathAGGModel(PreTrainedModel):
"""Pathology multi-teacher KD student. Forward returns CLS ``[B, D]`` by default."""
config_class = PathAGGConfig
base_model_prefix = "pathagg"
_no_split_modules = ["Block"]
def __init__(self, config: PathAGGConfig) -> None:
super().__init__(config)
self.return_patch = bool(config.return_patch)
self.model = _build_backbone(config)
# Do not call post_init() random re-init; weights come from checkpoint.
@property
def embed_dim(self) -> int:
return int(self.model.embed_dim)
@property
def num_register_tokens(self) -> int:
return int(getattr(self.model, "reg_tokens", getattr(self.model, "num_register_tokens", 4)))
def forward_features(self, x: torch.Tensor) -> torch.Tensor:
return self.model.forward_features(x)
def forward(
self,
pixel_values: torch.Tensor | None = None,
x: torch.Tensor | None = None,
return_patch: bool | None = None,
**kwargs: Any,
):
if pixel_values is None and x is None:
raise ValueError("provide pixel_values=... or x=...")
images = pixel_values if pixel_values is not None else x
tokens = self.forward_features(images)
cls = tokens[:, 0]
use_patch = self.return_patch if return_patch is None else bool(return_patch)
if not use_patch:
return cls
num_prefix = int(getattr(self.model, "num_prefix_tokens", 1 + self.num_register_tokens))
patch = tokens[:, num_prefix:]
return cls, patch
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
"""Load config + weights without transformers gamma↔weight remapping (LayerScale)."""
subfolder = kwargs.pop("subfolder", "")
local_files_only = kwargs.pop("local_files_only", False)
revision = kwargs.pop("revision", None)
cache_dir = kwargs.pop("cache_dir", None)
token = kwargs.pop("token", None)
if token is None:
token = kwargs.pop("use_auth_token", None)
kwargs.pop("trust_remote_code", None)
kwargs.pop("torch_dtype", None)
kwargs.pop("device_map", None)
kwargs.pop("low_cpu_mem_usage", None)
return_patch = kwargs.pop("return_patch", None)
config = kwargs.pop("config", None)
if config is None:
config = PathAGGConfig.from_pretrained(
pretrained_model_name_or_path,
subfolder=subfolder,
local_files_only=local_files_only,
revision=revision,
cache_dir=cache_dir,
token=token,
**kwargs,
)
if return_patch is not None:
config.return_patch = bool(return_patch)
model = cls(config)
weight_file = cached_file(
pretrained_model_name_or_path,
"pytorch_model.bin",
subfolder=subfolder,
local_files_only=local_files_only,
revision=revision,
cache_dir=cache_dir,
token=token,
)
if weight_file is None:
raise FileNotFoundError(
f"pytorch_model.bin not found under {pretrained_model_name_or_path!r} (subfolder={subfolder!r})"
)
try:
state = torch.load(weight_file, map_location="cpu", weights_only=True)
except TypeError:
state = torch.load(weight_file, map_location="cpu")
# transformers may have renamed LayerScale gamma→weight in some pipelines; normalize back.
fixed = {}
for k, v in state.items():
if k.endswith(".ls1.weight") or k.endswith(".ls2.weight"):
fixed[k[: -len(".weight")] + ".gamma"] = v
else:
fixed[k] = v
missing, unexpected = model.load_state_dict(fixed, strict=True)
if missing or unexpected:
raise RuntimeError(f"load failed: missing={missing}, unexpected={unexpected}")
model.eval()
return model
def get_preprocess(img_size: int = 224):
"""ImageNet normalize preprocess (PIL → tensor)."""
from torchvision import transforms
return transforms.Compose(
[
transforms.Resize(img_size, interpolation=transforms.InterpolationMode.BICUBIC),
transforms.CenterCrop(img_size),
transforms.ToTensor(),
transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),
]
)
__all__ = [
"PathAGGConfig",
"PathAGGModel",
"get_preprocess",
]