from typing import Any, Dict, List, Optional, Union, Callable import torch from transformers import GenerationMixin, LogitsProcessorList, StoppingCriteriaList from transformers.generation.utils import GenerationConfig, GenerateOutput from transformers.utils import ModelOutput class TSGenerationMixin(GenerationMixin): @torch.no_grad() def generate( self, inputs: Optional[torch.Tensor] = None, generation_config: Optional[GenerationConfig] = None, logits_processor: Optional[LogitsProcessorList] = None, stopping_criteria: Optional[StoppingCriteriaList] = None, prefix_allowed_tokens_fn: Optional[Callable[[int, torch.Tensor], List[int]]] = None, synced_gpus: Optional[bool] = None, assistant_model: Optional["PreTrainedModel"] = None, streamer: Optional["BaseStreamer"] = None, negative_prompt_ids: Optional[torch.Tensor] = None, negative_prompt_attention_mask: Optional[torch.Tensor] = None, revin: Optional[bool] = True, num_samples: Optional[int] = 1, max_output_length: Optional[int] = 96, inference_patch_len: Optional[int] = 48, **kwargs, ) -> Union[GenerateOutput, torch.Tensor]: if len(inputs.shape) != 2: raise ValueError('Input shape must be: [batch_size, seq_len]') if revin: means = inputs.mean(dim=-1, keepdim=True) stdev = inputs.std(dim=-1, keepdim=True, unbiased=False) + 1e-5 inputs = (inputs - means) / stdev model_inputs = { "input_ids": inputs, "max_output_length": max_output_length, "revin": False, "num_samples": num_samples, "inference_patch_len": inference_patch_len, } outputs = self(**model_inputs) predictions = outputs.logits if revin: stdev = stdev.unsqueeze(1).repeat(1, num_samples, 1) means = means.unsqueeze(1).repeat(1, num_samples, 1) predictions = (predictions * stdev) + means return predictions def _update_model_kwargs_for_generation( self, outputs: ModelOutput, model_kwargs: Dict[str, Any], horizon_length: int = 1, is_encoder_decoder: bool = False, standardize_cache_format: bool = False, ) -> Dict[str, Any]: return model_kwargs