ccloud0525 commited on
Commit
b0d5f8a
·
2 Parent(s): c73792af2bff8a

Merge branch 'main' of hf.co:Ccloud0525/FLAME

Browse files
Files changed (1) hide show
  1. README.md +95 -0
README.md ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ ---
4
+
5
+ ![image/png](https://cdn-uploads.huggingface.co/production/uploads/66276727368ec2a0b933772c/dPvMCbnSDjhZG_ddCw1kk.png)
6
+
7
+
8
+
9
+ [![Python](https://img.shields.io/badge/Python-3.10%2B-blue)](https://www.python.org/) [![PyTorch](https://img.shields.io/badge/PyTorch-2.6.0-blue)](https://pytorch.org/)
10
+
11
+ # FLAME: Flow Enhanced Legendre Memory Models for General Time Series Forecasting
12
+
13
+ This is the official repository of **FLAME**: Flow Enhanced Legendre Memory Models for General Time Series Forecasting.
14
+
15
+
16
+
17
+ ## Introduction
18
+ FLAME is a family of extremely **lightweight** and highly capable time series foundation models. Based on the normalization-based forecasting head, it can support both the **deterministic** and **probabilistic** forecasting.
19
+
20
+ To our best knowldege, FLAME is the first time series foundation model possessing both lightweight backbones and generative prediction capabilities!
21
+
22
+
23
+
24
+ ![image/png](https://cdn-uploads.huggingface.co/production/uploads/66276727368ec2a0b933772c/OHKXxjrwYR3n3LO6UeVRS.png)
25
+
26
+ ## Architecture
27
+
28
+ FLAME adopts the Channel-Independent pretraining paradigm, and each variable is preprocessed through Instance Normalization to mitigate the value discrepancy. FLAME utilizes the Re-Norm to further mitigate the statistical differences between inputs and forecasts, and its backbone mainly consists of three modules: 1) Encoding, including Time Series Tokenization, **Local-Perception**, and MSA-Encoder, which tokenize the time series and enhance them through fusing the local environmental information with LegT; 2) Decoding, including **LegS based SSD-Decoder** and MCA-Enhancer, which utilize the SSD layers and MCA layers to make long-term inference ; 3) **Flow-based Head**, which leverages the Normalization Flow to support generative probabilistic forecasting, with both efficiency and accuracy.
29
+
30
+
31
+ ![image/png](https://cdn-uploads.huggingface.co/production/uploads/66276727368ec2a0b933772c/LchwMIvetMu10EjWjcW5N.png)
32
+
33
+
34
+ ## Quickstart
35
+
36
+ We release all three versions of FLAME in different branches:
37
+ ```shell
38
+ FLAME Small (2M) -- branch main & FLAME_Small
39
+ FLAME Base (6M) -- branch FLAME_Base
40
+ FLAME Large (10M) -- branch FLAME_Large
41
+ ```
42
+
43
+
44
+
45
+ You need to install the following packages:
46
+
47
+ ```shell
48
+ # pip install transformers[torch]
49
+
50
+ # pip install mamba-ssm[causal-conv1d]
51
+
52
+ # pip install zuko
53
+ ```
54
+
55
+ To make deterministic or probabilistic forecasts, just follow:
56
+
57
+ ```python
58
+ from transformers import AutoModel, AutoConfig
59
+ import torch
60
+
61
+ model_path = "path/to/your/model"
62
+ config_path = "path/to/your/config"
63
+
64
+ config = AutoConfig.from_pretrained(config_path)
65
+
66
+ model = AutoModel.from_pretrained(model_path, config=config)
67
+ model.eval()
68
+
69
+ # The inputs need to be [batch_size, seq_len]. If multivariate, transform the inputs to [batch_size * n_vars, seq_len]
70
+ inputs = torch.randn(batch_size, seq_length)
71
+
72
+ # deterministic forecasting
73
+ with torch.no_grad():
74
+ # output shape: [batch_size, 1, seq_len]
75
+ outputs = model.generate(
76
+ inputs=inputs,
77
+ max_length=96,
78
+ revin=True,
79
+ num_samples=1,
80
+ inference_patch_len=48 # recommend to input the period length
81
+ )
82
+
83
+
84
+ # probabilistic forecasting
85
+ with torch.no_grad():
86
+ # output shape: [batch_size, 100, seq_len]
87
+ outputs = model.generate(
88
+ inputs=inputs,
89
+ max_length=96,
90
+ revin=True,
91
+ num_samples=100
92
+ )
93
+
94
+ ```
95
+