vit-oct-wamd / modeling.py
Kavin
Upload modeling.py with huggingface_hub
b2b25ad verified
Raw
History Blame Contribute Delete
2.23 kB
"""
Model definition for kr3131/vit-oct-wamd.
Architecture: a SigLIP vision transformer (google/siglip-so400m-patch14-384)
fine-tuned with a 2-class linear head (normal vs. wet AMD) on OCT images.
The checkpoint also contains a `siglip_loss` branch (a frozen T5-base text
encoder + projection used for an auxiliary contrastive alignment loss during
training). It is NOT used at inference time -- `forward()` only calls
`image_encoder` and `cls_head` -- but is included in the state dict for
training-time fidelity. Loading the full checkpoint therefore requires
`alignment.py` and `embedder.py` (included in this repo) even though those
weights are unused for classification.
"""
import torch
import torch.nn as nn
from transformers import SiglipVisionModel
from alignment import SigLIPLoss
MAX_TEXT_LEN = 128
IMAGE_SIZE = 384 # SigLIP input size
class SigLIPModel(nn.Module):
"""SigLIP-based classifier for OCT wet-AMD detection."""
def __init__(self, dropout_rate: float = 0.057129660535791646):
super().__init__()
self.image_encoder = SiglipVisionModel.from_pretrained(
"google/siglip-so400m-patch14-384"
)
encoder_output_dim = 1152
self.dropout = nn.Dropout(dropout_rate)
self.cls_head = nn.Linear(encoder_output_dim, 2) # 0=normal, 1=wet_amd
# Present for checkpoint compatibility; unused in forward().
self.siglip_loss = SigLIPLoss(
latent_dim=encoder_output_dim,
text_model="google-t5/t5-base",
max_txt_len=MAX_TEXT_LEN,
pool="mean",
dtype=torch.float32,
)
def forward(self, images, input_ids=None, attention_mask=None):
img_features = self.image_encoder(pixel_values=images).last_hidden_state # (B, 729, 1152)
cls_features = self.dropout(img_features[:, 0]) # CLS token, (B, 1152)
return self.cls_head(cls_features) # (B, 2)
def load_model(checkpoint_path: str, device: str = "cpu") -> SigLIPModel:
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
model = SigLIPModel()
model.load_state_dict(checkpoint["model_state_dict"])
model.to(device)
model.eval()
return model