Source code for ts1.models.pretrained.poyo_plus.poyo_plus_eval_trainer

import hydra
from torch_brain.transforms import Compose

from core.utils.util import get_num_params, get_num_trainable_params, human_readable
from pretrain.models.poyo_plus import POYOPlus
from ts1.ts1_eval_trainer import TS1EvalTrainer


[docs] class POYOPlusEvalTrainer(TS1EvalTrainer):
[docs] def predict(self, X, target_timestamps, *args): if "output_timestamps" in X["model_inputs"]: del X["model_inputs"]["output_timestamps"] if "output_session_index" in X["model_inputs"]: del X["model_inputs"]["output_session_index"] # return_dict=False + unflatten_output=True (the model's default) routes through # POYOPlus's unflatten logic, which yields (B, N, D) for all predictions # (e.g., N=1 for single-task sequence-level readouts). return self.model( **X["model_inputs"], output_timestamps=target_timestamps, output_session_index=X["session_index"], return_dict=False, )