Files
wehub-resource-sync 925e56bb5f
Unit tests / build (t4_gpu) (push) Has been cancelled
Unit tests / build (ubuntu-latest) (push) Has been cancelled
Unit tests / build (windows-latest) (push) Has been cancelled
Test CLI scripts / build (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:24:56 +08:00

54 lines
1.7 KiB
Python

from typing import Optional
import torch
from surya.common.load import ModelLoader
from surya.detection.processor import SegformerImageProcessor
from surya.detection.model.config import EfficientViTConfig
from surya.detection.model.encoderdecoder import EfficientViTForSemanticSegmentation
from surya.logging import get_logger
from surya.settings import settings
logger = get_logger()
class DetectionModelLoader(ModelLoader):
def __init__(self, checkpoint: Optional[str] = None):
super().__init__(checkpoint)
if self.checkpoint is None:
self.checkpoint = settings.DETECTOR_MODEL_CHECKPOINT
def model(
self,
device: Optional[torch.device | str] = None,
dtype: Optional[torch.dtype | str] = None,
attention_implementation: Optional[str] = None,
) -> EfficientViTForSemanticSegmentation:
if device is None:
device = settings.TORCH_DEVICE_MODEL
if dtype is None:
dtype = settings.MODEL_DTYPE
config = EfficientViTConfig.from_pretrained(self.checkpoint)
model = EfficientViTForSemanticSegmentation.from_pretrained(
self.checkpoint,
dtype=dtype,
config=config,
)
model = model.to(device)
model = model.eval()
logger.debug(
f"Loaded detection model {self.checkpoint} from {EfficientViTForSemanticSegmentation.get_local_path(self.checkpoint)} onto device {device} with dtype {dtype}"
)
return model
def processor(
self,
device: Optional[torch.device | str] = None,
dtype: Optional[torch.dtype | str] = None,
) -> SegformerImageProcessor:
return SegformerImageProcessor.from_pretrained(self.checkpoint)