Source code for ts1.models.pretrained.neds.neds_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 pretrain.models.neds import NEDS, NEDSMasker
from ts1.ts1_eval_trainer import TS1EvalTrainer


[docs] class NEDSEvalTrainer(TS1EvalTrainer):
[docs] def setup_model(self, ckpt: dict | None): modalities = ["spikes", self.cfg.task] self.model = hydra.utils.instantiate( self.cfg.model, finetune_enable=self.cfg.finetuning.enable, ) if not isinstance(self.model, NEDS): raise TypeError( f"{type(self).__name__} requires a NEDS model, got {type(self.model).__name__}" ) self.model.modalities = modalities 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=None, target_mask=None): model_inputs = X["model_inputs"] keep_masks = X["keep_masks"] modality_masks = self.masker(model_inputs["spikes"]) preds = self.model(model_inputs, keep_masks, modality_masks) return preds[self.task]