Source code for ts1.models.pretrained.ndt_stitch.ndt_stitch_eval_trainer

import hydra
from torch_brain.transforms.container import Compose

from core.utils.util import log_param_breakdown
from pretrain.models.ndt_stitch import NDTStitch
from ts1.ts1_eval_trainer import TS1EvalTrainer


[docs] class NDTStitchEvalTrainer(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, NDTStitch): raise TypeError( f"{type(self).__name__} requires a NDTStitch 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()