FLAME / ts_generation_mixin.py
ccloud0525
feat: 'main'
e11caaf
Raw
History Blame Contribute Delete
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):
@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