from __future__ import annotations from typing import Any import torch from transformers import AutoModel from .base import Scorer from .objectives import build_objective class EncoderScorer(Scorer): def __init__( self, encoder_model_id: str | None = None, *, encoder: Any | None = None, encoder_call: str = "predict_fmri", objective: str | Any = "indices_mean", device: str = "cuda", trust_remote_code: bool = True, **encoder_kwargs, ) -> None: if encoder is None: if encoder_model_id is None: raise ValueError("encoder_model_id is required when encoder is not provided.") encoder = AutoModel.from_pretrained( encoder_model_id, trust_remote_code=trust_remote_code, **encoder_kwargs, ) self.encoder = encoder self.encoder_call = encoder_call self.objective = build_objective(objective) self.device = device if hasattr(self.encoder, "to"): self.encoder.to(device) if hasattr(self.encoder, "eval"): self.encoder.eval() def score(self, videos: torch.Tensor, target: Any, **kwargs) -> list[float]: videos = videos.to(self.device) with torch.no_grad(): if self.encoder_call: fn = getattr(self.encoder, self.encoder_call) predictions = fn(videos, **kwargs) else: predictions = self.encoder(videos, **kwargs) scores = self.objective(predictions, target) return [float(x) for x in scores.detach().cpu().reshape(-1)]