GemmaSight / models /feature_extractor.py
TouseefAhmad's picture
Upload folder using huggingface_hub
de2e2e5 verified
Raw
History Blame Contribute Delete
6.03 kB
import torch
import torch.nn as nn
import numpy as np
from PIL import Image
from transformers import AutoModel, AutoProcessor, AutoImageProcessor
import torchvision.models as models
from gemmasight.config import PATH_FOUNDATION_ID, MEDSIGLIP_ID, DIM_PATH, DIM_SIGLIP, DIM_FUSED, FORCE_SIMULATION
class DualEncoderFeatureExtractor(nn.Module):
def __init__(self, force_simulation=FORCE_SIMULATION):
super().__init__()
self.force_simulation = force_simulation
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if self.force_simulation:
print("GemmaSight Feature Extractor: Initializing in Simulation/Fallback Mode...")
self._init_simulation()
else:
try:
print("GemmaSight Feature Extractor: Loading live models from Hugging Face...")
# Encoder A: Google Path Foundation (requires login/access if gated)
self.path_processor = AutoImageProcessor.from_pretrained(PATH_FOUNDATION_ID)
self.path_encoder = AutoModel.from_pretrained(PATH_FOUNDATION_ID, trust_remote_code=True)
self.path_encoder.to(self.device).eval()
# Encoder B: MedSigLIP (bfloat16 recommended)
self.siglip_processor = AutoProcessor.from_pretrained(MEDSIGLIP_ID)
self.siglip_encoder = AutoModel.from_pretrained(
MEDSIGLIP_ID,
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
trust_remote_code=True
)
self.siglip_encoder.to(self.device).eval()
# Freeze all parameters
for param in self.parameters():
param.requires_grad = False
print("GemmaSight Feature Extractor: Dual Encoders loaded successfully.")
self.is_simulation = False
except Exception as e:
print(f"GemmaSight Feature Extractor Error: Failed to load live models: {e}")
print("GemmaSight Feature Extractor: Falling back to Simulation/Fallback Mode.")
self._init_simulation()
def _init_simulation(self):
self.is_simulation = True
# Use torchvision ResNet18 as a backbone for deterministic visual-feature mapping
self.backbone = models.resnet18(pretrained=True)
self.backbone.eval()
# Freeze backbone
for param in self.backbone.parameters():
param.requires_grad = False
# Linear layers to project features to DIM_PATH (384) and DIM_SIGLIP (1152)
# ResNet18 outputs 1000 classes or 512 average pooled features. Let's project from 1000.
self.proj_path = nn.Linear(1000, DIM_PATH)
self.proj_siglip = nn.Linear(1000, DIM_SIGLIP)
# Freeze projection layers to maintain frozen behavior
for param in self.proj_path.parameters():
param.requires_grad = False
for param in self.proj_siglip.parameters():
param.requires_grad = False
self.to(self.device)
def forward(self, pil_image: Image.Image) -> torch.Tensor:
"""
Extracts and concatenates histopathology and medical embeddings.
Returns a fused tensor of shape (1, 1536) on the correct device.
"""
if self.is_simulation:
# Preprocess using torchvision transforms
import torchvision.transforms as T
transforms = T.Compose([
T.Resize((224, 224)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
img_tensor = transforms(pil_image).unsqueeze(0).to(self.device)
with torch.no_grad():
logits = self.backbone(img_tensor)
# Deterministic projection
feat_path = self.proj_path(logits) # (1, 384)
feat_siglip = self.proj_siglip(logits) # (1, 1152)
# Normalize features
feat_path = nn.functional.normalize(feat_path, p=2, dim=1)
feat_siglip = nn.functional.normalize(feat_siglip, p=2, dim=1)
fused = torch.cat([feat_path, feat_siglip], dim=1) # (1, 1536)
return fused
else:
# Live feature extraction
with torch.no_grad():
# Path Foundation
path_inputs = self.path_processor(images=pil_image, return_tensors="pt").to(self.device)
path_outputs = self.path_encoder(**path_inputs)
# Mean pool or take pooler output
if hasattr(path_outputs, "pooler_output") and path_outputs.pooler_output is not None:
feat_path = path_outputs.pooler_output
else:
feat_path = path_outputs.last_hidden_state.mean(dim=1)
# MedSigLIP
siglip_inputs = self.siglip_processor(images=pil_image, return_tensors="pt").to(self.device)
if torch.cuda.is_available():
siglip_inputs = {k: v.to(torch.bfloat16) if v.dtype == torch.float32 else v for k, v in siglip_inputs.items()}
siglip_outputs = self.siglip_encoder.get_image_features(**siglip_inputs)
feat_siglip = siglip_outputs
# Align types and compute
feat_path = feat_path.to(torch.float32)
feat_siglip = feat_siglip.to(torch.float32)
# L2 normalize
feat_path = nn.functional.normalize(feat_path, p=2, dim=-1)
feat_siglip = nn.functional.normalize(feat_siglip, p=2, dim=-1)
fused = torch.cat([feat_path, feat_siglip], dim=-1)
return fused