Download ts_generation_mixin.py from DecisionIntelligence/FLAME: direct link, hf CLI and curl.
- Browser
- Download file 2.48 kB
-
https://huggingface.co/DecisionIntelligence/FLAME/resolve/main/ts_generation_mixin.py
- Command line
-
hf download hf://DecisionIntelligence/FLAME/ts_generation_mixin.py
-
curl -L -o ts_generation_mixin.py https://huggingface.co/DecisionIntelligence/FLAME/resolve/main/ts_generation_mixin.py
2.48 kB
| 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): | |
| 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 | |