| 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 |
|
|