File size: 2,484 Bytes
e11caaf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
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