Source code for ts1.models.pretrained.poyo.poyo_eval_trainer

import hydra
from torch_brain.transforms.container import Compose

from core.utils.util import get_num_params, get_num_trainable_params, human_readable
from ibl_bwb_eval.tasks import TargetResolution
from pretrain.models.poyo import POYO
from ts1.ts1_eval_trainer import TS1EvalTrainer


[docs] class POYOEvalTrainer(TS1EvalTrainer):
[docs] def setup_model(self, ckpt: dict | None): self.model = hydra.utils.instantiate( self.cfg.model, finetune_enable=self.cfg.finetuning.enable, ) if not isinstance(self.model, POYO): raise TypeError( f"{type(self).__name__} requires a POYO model, got {type(self.model).__name__}" ) self.add_checkpoint_items(model=self.model) self.logger.info(f"Precision: {self.precision}") self.logger.info(f"Model: {self.model.__class__}") self.model.train()
[docs] def predict(self, X, target_timestamps, *args): output = self.model( **X["model_inputs"], output_timestamps=target_timestamps, output_session_index=X["session_index"], ) if self.readout_spec.target_resolution == TargetResolution.TIMESTEP: # queries are chained, so restore the batch axis of the target: (B, T, D) return output.view(*target_timestamps.shape, -1) # sequence-level: one query per sample, (B, D) -> (B, 1, D), to make it compatible # with TS1TestMixin and TS1EvalTrainer.loss(). return output.unsqueeze(1)