Spaces:
Running on Zero
Running on Zero
Upload folder using huggingface_hub
Browse files- README.md +44 -16
- app.py +449 -266
- requirements.txt +16 -4
README.md
CHANGED
|
@@ -4,28 +4,56 @@ emoji: 🧵
|
|
| 4 |
colorFrom: indigo
|
| 5 |
colorTo: gray
|
| 6 |
sdk: gradio
|
| 7 |
-
sdk_version: 6.
|
| 8 |
app_file: app.py
|
| 9 |
-
short_description:
|
| 10 |
python_version: "3.12"
|
| 11 |
-
startup_duration_timeout:
|
|
|
|
| 12 |
---
|
| 13 |
|
| 14 |
# TimeBraid
|
| 15 |
|
| 16 |
-
Interactive demo
|
| 17 |
-
|
| 18 |
-
through Mixture-of-Transformers layers.
|
| 19 |
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
|
| 24 |
-
|
| 25 |
-
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
`trust_remote_code=False`.
|
| 29 |
|
| 30 |
-
|
| 31 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
colorFrom: indigo
|
| 5 |
colorTo: gray
|
| 6 |
sdk: gradio
|
| 7 |
+
sdk_version: 6.29.1
|
| 8 |
app_file: app.py
|
| 9 |
+
short_description: Understand and forecast time series with language
|
| 10 |
python_version: "3.12"
|
| 11 |
+
startup_duration_timeout: 1h
|
| 12 |
+
pinned: false
|
| 13 |
---
|
| 14 |
|
| 15 |
# TimeBraid
|
| 16 |
|
| 17 |
+
Interactive demo for **TimeBraid-2.5B** — *Unifying Time Series and Language for
|
| 18 |
+
Understanding and Forecasting*.
|
|
|
|
| 19 |
|
| 20 |
+
- Paper: <https://huggingface.co/papers/2609.29792>
|
| 21 |
+
- Model: [`XinyueWangg/TimeBraid-2.5B`](https://huggingface.co/XinyueWangg/TimeBraid-2.5B)
|
| 22 |
+
- Code: <https://github.com/CharonWangg/TimeBraid>
|
| 23 |
|
| 24 |
+
TimeBraid braids a Qwen3-1.7B language backbone with a TimesFM 2.5 time-series
|
| 25 |
+
expert through an interleaved Mixture-of-Transformers, so numeric series and text
|
| 26 |
+
live in the same token stream. The same numbers can therefore be read very
|
| 27 |
+
differently depending on the context you give them.
|
|
|
|
| 28 |
|
| 29 |
+
## How to use it
|
| 30 |
+
|
| 31 |
+
1. **Paste one series per line** — comma-separated numbers. One line = one
|
| 32 |
+
series; two lines = two related series the model can compare or choose
|
| 33 |
+
between.
|
| 34 |
+
2. **Write your question or instruction.** Ask for an explanation of the
|
| 35 |
+
observed data, or ask for future values.
|
| 36 |
+
3. **Set the forecast horizon.** `0` keeps the run in *understanding* mode and
|
| 37 |
+
the model answers in words. A positive `N` switches to *forecasting* and the
|
| 38 |
+
model returns `N` future values for one target series, plotted against the
|
| 39 |
+
history.
|
| 40 |
+
4. Textual context matters for forecasts: describing an upcoming promotion, a
|
| 41 |
+
heatwave, or a school term materially changes the predicted numbers.
|
| 42 |
+
|
| 43 |
+
Advanced options hold an optional system prompt, a 1-based selector for which
|
| 44 |
+
input series to forecast, and the text-generation budget.
|
| 45 |
+
|
| 46 |
+
## Implementation notes
|
| 47 |
+
|
| 48 |
+
- Runs on ZeroGPU with the full BF16 checkpoint resident on the GPU.
|
| 49 |
+
- The checkpoint vendors its own `timebraid` inference package, loaded through
|
| 50 |
+
`trust_remote_code=True`; no separate TimesFM checkpoint is required.
|
| 51 |
+
- MoT mixed attention in TimeBraid is FlashAttention-2 only, and FA3/FA4 have no
|
| 52 |
+
sm_120 kernels, so the Space installs the prebuilt Blackwell FA2 wheel.
|
| 53 |
+
- Decoding follows the released recipe: greedy (`do_sample=False`), one
|
| 54 |
+
returned sequence, one request at a time.
|
| 55 |
+
|
| 56 |
+
## Attribution
|
| 57 |
+
|
| 58 |
+
The example requests in the UI are the illustrative tasks from the TimeBraid
|
| 59 |
+
model card and the authors' `examples/inference_tasks.ipynb` (Apache-2.0).
|
app.py
CHANGED
|
@@ -1,38 +1,64 @@
|
|
| 1 |
-
"""TimeBraid:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
|
| 3 |
import os
|
| 4 |
|
| 5 |
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 6 |
|
| 7 |
-
import spaces #
|
| 8 |
-
|
|
|
|
| 9 |
|
| 10 |
-
import gradio as gr
|
| 11 |
import matplotlib
|
| 12 |
|
| 13 |
matplotlib.use("Agg")
|
| 14 |
-
import matplotlib.pyplot as plt
|
| 15 |
-
import torch
|
| 16 |
-
from matplotlib.ticker import MaxNLocator
|
| 17 |
-
from transformers import AutoModelForCausalLM, AutoProcessor
|
| 18 |
|
| 19 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
MODEL_ID = "XinyueWangg/TimeBraid-2.5B"
|
| 22 |
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 28 |
|
| 29 |
-
processor = AutoProcessor.from_pretrained(
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
model = (
|
| 31 |
AutoModelForCausalLM.from_pretrained(
|
| 32 |
MODEL_ID,
|
|
|
|
| 33 |
dtype=torch.bfloat16,
|
| 34 |
attn_implementation="flash_attention_2",
|
| 35 |
-
trust_remote_code=False,
|
| 36 |
use_safetensors=True,
|
| 37 |
)
|
| 38 |
.eval()
|
|
@@ -40,290 +66,447 @@ model = (
|
|
| 40 |
)
|
| 41 |
|
| 42 |
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 47 |
series: list[list[float]] = []
|
| 48 |
-
for
|
| 49 |
-
line =
|
| 50 |
if not line:
|
| 51 |
continue
|
| 52 |
-
|
| 53 |
-
values
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
if not values:
|
| 55 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
series.append(values)
|
|
|
|
|
|
|
| 57 |
return series
|
| 58 |
|
| 59 |
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 73 |
ax.plot(
|
| 74 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 75 |
values,
|
| 76 |
-
|
|
|
|
|
|
|
| 77 |
linewidth=2.0,
|
| 78 |
-
|
| 79 |
-
markersize=4,
|
| 80 |
-
label="Observed history" if forecast is not None else f"Series {index + 1}",
|
| 81 |
)
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
)
|
| 95 |
-
ax.scatter(future_x, forecast, color=_FORECAST_COLOR, s=24, zorder=3)
|
| 96 |
-
ax.axvline(
|
| 97 |
-
boundary,
|
| 98 |
-
color=_BOUNDARY_COLOR,
|
| 99 |
-
linewidth=1.2,
|
| 100 |
-
linestyle=":",
|
| 101 |
-
label="Forecast boundary",
|
| 102 |
-
)
|
| 103 |
-
ax.text(
|
| 104 |
-
boundary,
|
| 105 |
-
1.01,
|
| 106 |
-
"forecast starts",
|
| 107 |
-
transform=ax.get_xaxis_transform(),
|
| 108 |
ha="center",
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
color=_BOUNDARY_COLOR,
|
| 112 |
)
|
| 113 |
-
|
| 114 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
)
|
| 116 |
-
ax.set_title(
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
|
|
|
| 125 |
return fig
|
| 126 |
|
| 127 |
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 131 |
series_text: str,
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
target_series_index: 1-based index of the series to forecast (multi-series only).
|
| 144 |
-
max_new_tokens: Maximum generated text tokens.
|
| 145 |
-
Returns:
|
| 146 |
-
(answer text, forecast values text, plot, elapsed seconds text).
|
| 147 |
"""
|
| 148 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 149 |
try:
|
| 150 |
-
|
| 151 |
-
except ValueError as exc:
|
| 152 |
-
raise gr.Error(f"Could not
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
raise gr.Error("A forecast horizon needs at least one input series.")
|
| 159 |
-
if len(series) > 1:
|
| 160 |
-
if not target_series_index or target_series_index < 1:
|
| 161 |
-
raise gr.Error(
|
| 162 |
-
"With multiple series, pick which one to forecast (Target series)."
|
| 163 |
-
)
|
| 164 |
-
if int(target_series_index) > len(series):
|
| 165 |
-
raise gr.Error(
|
| 166 |
-
f"Target series is out of range: only {len(series)} series given."
|
| 167 |
-
)
|
| 168 |
-
target_i = int(target_series_index) - 1
|
| 169 |
-
|
| 170 |
-
messages = [{"role": "user", "content": prompt}]
|
| 171 |
-
inputs = processor.apply_chat_template(
|
| 172 |
-
messages,
|
| 173 |
-
timeseries=series if series else None,
|
| 174 |
-
horizon=horizon_i,
|
| 175 |
-
target_series_index=target_i,
|
| 176 |
-
add_generation_prompt=True,
|
| 177 |
-
tokenize=True,
|
| 178 |
-
return_dict=True,
|
| 179 |
-
return_tensors="pt",
|
| 180 |
-
).to(model.device)
|
| 181 |
with torch.inference_mode():
|
| 182 |
-
|
| 183 |
-
**
|
| 184 |
-
max_new_tokens=
|
| 185 |
do_sample=False,
|
| 186 |
num_beams=1,
|
| 187 |
num_return_sequences=1,
|
| 188 |
)
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
forecast_values = None
|
| 197 |
-
forecast_text = "No numeric forecast (understanding route)."
|
| 198 |
-
if ts is not None and ts.get("values"):
|
| 199 |
-
forecast_values = [float(v) for v in ts["values"]]
|
| 200 |
-
forecast_text = ", ".join(f"{v:.4g}" for v in forecast_values)
|
| 201 |
-
|
| 202 |
-
plot = render_plot(series, forecast_values, result.get("target_series_index"))
|
| 203 |
-
stats = (
|
| 204 |
-
f"route: {'forecast' if horizon_i is not None else 'understanding'} · "
|
| 205 |
-
f"prompt tokens: {result.get('prompt_tokens')} · "
|
| 206 |
-
f"completion tokens: {result.get('completion_tokens')} · "
|
| 207 |
-
f"finish: {result.get('finish_reason')} · {elapsed:.1f}s on GPU"
|
| 208 |
-
)
|
| 209 |
-
return answer, forecast_text, plot, stats
|
| 210 |
|
| 211 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
CSS = """
|
| 213 |
-
#col-container { max-width: 1150px; margin: 0 auto; }
|
| 214 |
.dark .gradio-container { color: var(--body-text-color); }
|
|
|
|
| 215 |
"""
|
| 216 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 217 |
with gr.Blocks(title="TimeBraid") as demo:
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 229 |
)
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
series_box = gr.Textbox(
|
| 238 |
-
label="Time series (one series per line)",
|
| 239 |
-
placeholder="101.2, 101.8, 102.1, 102.5, 103.0, 103.4\n200, 208, 215, ...",
|
| 240 |
-
lines=5,
|
| 241 |
-
)
|
| 242 |
-
horizon = gr.Number(
|
| 243 |
-
label="Forecast horizon (0 = no forecast)",
|
| 244 |
-
value=0,
|
| 245 |
-
precision=0,
|
| 246 |
-
minimum=0,
|
| 247 |
-
)
|
| 248 |
-
run_btn = gr.Button("Run", variant="primary")
|
| 249 |
-
with gr.Accordion("Advanced settings", open=False):
|
| 250 |
-
target_idx = gr.Number(
|
| 251 |
-
label="Target series (1-based; multi-series forecasts)",
|
| 252 |
-
value=1,
|
| 253 |
-
precision=0,
|
| 254 |
-
minimum=1,
|
| 255 |
-
)
|
| 256 |
-
max_new_tokens = gr.Slider(
|
| 257 |
-
label="Max new text tokens",
|
| 258 |
-
minimum=64,
|
| 259 |
-
maximum=1024,
|
| 260 |
-
value=256,
|
| 261 |
-
step=32,
|
| 262 |
-
)
|
| 263 |
-
with gr.Column(scale=1):
|
| 264 |
-
answer_out = gr.Textbox(label="Answer", lines=8, buttons=["copy"])
|
| 265 |
-
forecast_out = gr.Textbox(label="Forecast values (raw scale)", buttons=["copy"])
|
| 266 |
-
plot_out = gr.Plot(label="Series and forecast")
|
| 267 |
-
stats_out = gr.Markdown()
|
| 268 |
-
gr.Examples(
|
| 269 |
-
examples=[
|
| 270 |
-
[
|
| 271 |
-
"Forecast the next 8 values from the observed seasonal pattern.",
|
| 272 |
-
"100, 102, 105, 107, 103, 101, 99, 100, 103, 106, 108, 104, 102, 100, 101, 104",
|
| 273 |
-
8,
|
| 274 |
-
1,
|
| 275 |
-
256,
|
| 276 |
-
],
|
| 277 |
-
[
|
| 278 |
-
"A promotion starts at the first forecast step and is expected to lift demand above the recent seasonal baseline. Forecast the next 6 values.",
|
| 279 |
-
"200, 208, 215, 205, 198, 210, 218, 207, 201, 212, 220, 209",
|
| 280 |
-
6,
|
| 281 |
-
1,
|
| 282 |
-
256,
|
| 283 |
-
],
|
| 284 |
-
[
|
| 285 |
-
"Describe the dominant trend, major turning points, and whether the series becomes more volatile.",
|
| 286 |
-
"2.0, 2.1, 2.2, 2.4, 2.8, 3.1, 3.0, 2.9, 3.4, 3.8, 4.1, 4.3",
|
| 287 |
-
0,
|
| 288 |
-
1,
|
| 289 |
-
256,
|
| 290 |
-
],
|
| 291 |
-
[
|
| 292 |
-
"Does this series contain an isolated upward spike? Say only whether it occurs near the beginning, middle, or end; do not give a numeric index. Compare it with the neighboring level.",
|
| 293 |
-
"20.1, 20.0, 20.2, 20.1, 20.3, 34.8, 20.2, 20.1, 20.0, 20.2",
|
| 294 |
-
0,
|
| 295 |
-
1,
|
| 296 |
-
256,
|
| 297 |
-
],
|
| 298 |
-
[
|
| 299 |
-
"Series 1 is website visits and Series 2 is conversion rate. Compare their trends, volatility, and co-movement.",
|
| 300 |
-
"100, 104, 107, 111, 116, 120, 125, 129\n3.2, 3.1, 3.3, 3.4, 3.8, 3.7, 4.0, 4.2",
|
| 301 |
-
0,
|
| 302 |
-
1,
|
| 303 |
-
256,
|
| 304 |
-
],
|
| 305 |
-
[
|
| 306 |
-
"Explain in one sentence what a moving average reveals about a time series.",
|
| 307 |
-
"",
|
| 308 |
-
0,
|
| 309 |
-
1,
|
| 310 |
-
256,
|
| 311 |
-
],
|
| 312 |
-
],
|
| 313 |
-
inputs=[prompt, series_box, horizon, target_idx, max_new_tokens],
|
| 314 |
-
outputs=[answer_out, forecast_out, plot_out, stats_out],
|
| 315 |
-
fn=run,
|
| 316 |
-
cache_examples=True,
|
| 317 |
-
cache_mode="lazy",
|
| 318 |
-
label="Examples (from the authors' inference notebook)",
|
| 319 |
)
|
| 320 |
-
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
|
| 324 |
-
|
| 325 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 326 |
)
|
| 327 |
|
| 328 |
-
|
| 329 |
-
|
|
|
|
| 1 |
+
"""TimeBraid: unify time series and language for understanding and forecasting.
|
| 2 |
+
|
| 3 |
+
A single-request demo around ``XinyueWangg/TimeBraid-2.5B``. The model reads one
|
| 4 |
+
or more numeric series plus a natural-language request and either explains what
|
| 5 |
+
it sees or forecasts future values.
|
| 6 |
+
"""
|
| 7 |
|
| 8 |
import os
|
| 9 |
|
| 10 |
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 11 |
|
| 12 |
+
import spaces # noqa: F401 # must be imported before torch / transformers
|
| 13 |
+
|
| 14 |
+
from typing import Any
|
| 15 |
|
|
|
|
| 16 |
import matplotlib
|
| 17 |
|
| 18 |
matplotlib.use("Agg")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
|
| 20 |
+
import matplotlib.pyplot as plt # noqa: E402
|
| 21 |
+
import torch # noqa: E402
|
| 22 |
+
from matplotlib.ticker import MaxNLocator # noqa: E402
|
| 23 |
+
from transformers import AutoModelForCausalLM, AutoProcessor # noqa: E402
|
| 24 |
+
|
| 25 |
+
import gradio as gr # noqa: E402
|
| 26 |
|
| 27 |
MODEL_ID = "XinyueWangg/TimeBraid-2.5B"
|
| 28 |
|
| 29 |
+
MAX_SERIES = 8
|
| 30 |
+
MAX_POINTS_PER_SERIES = 2048
|
| 31 |
+
MAX_HORIZON = 64
|
| 32 |
+
|
| 33 |
+
INPUT_COLOR = "#0072B2"
|
| 34 |
+
OUTPUT_COLOR = "#D55E00"
|
| 35 |
+
BOUNDARY_COLOR = "#555555"
|
| 36 |
+
SERIES_COLORS = (
|
| 37 |
+
"#0072B2",
|
| 38 |
+
"#009E73",
|
| 39 |
+
"#CC79A7",
|
| 40 |
+
"#E69F00",
|
| 41 |
+
"#56B4E9",
|
| 42 |
+
"#F0E442",
|
| 43 |
+
"#000000",
|
| 44 |
+
"#999999",
|
| 45 |
+
)
|
| 46 |
+
|
| 47 |
+
# ---------------------------------------------------------------------------
|
| 48 |
+
# Model (loaded once, at module scope, on the GPU)
|
| 49 |
+
# ---------------------------------------------------------------------------
|
| 50 |
|
| 51 |
+
processor = AutoProcessor.from_pretrained(
|
| 52 |
+
MODEL_ID,
|
| 53 |
+
trust_remote_code=True,
|
| 54 |
+
fix_mistral_regex=False,
|
| 55 |
+
)
|
| 56 |
model = (
|
| 57 |
AutoModelForCausalLM.from_pretrained(
|
| 58 |
MODEL_ID,
|
| 59 |
+
trust_remote_code=True,
|
| 60 |
dtype=torch.bfloat16,
|
| 61 |
attn_implementation="flash_attention_2",
|
|
|
|
| 62 |
use_safetensors=True,
|
| 63 |
)
|
| 64 |
.eval()
|
|
|
|
| 66 |
)
|
| 67 |
|
| 68 |
|
| 69 |
+
# ---------------------------------------------------------------------------
|
| 70 |
+
# Input parsing / coercion helpers
|
| 71 |
+
# ---------------------------------------------------------------------------
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def _coerce_text(value: Any, default: str = "") -> str:
|
| 75 |
+
return value if isinstance(value, str) else default
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def _coerce_int(value: Any, default: int, low: int, high: int) -> int:
|
| 79 |
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
| 80 |
+
return default
|
| 81 |
+
try:
|
| 82 |
+
number = int(value)
|
| 83 |
+
except (TypeError, ValueError):
|
| 84 |
+
return default
|
| 85 |
+
return max(low, min(high, number))
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def parse_series_block(text: str) -> list[list[float]]:
|
| 89 |
+
"""Parse one series per non-empty line, comma (or tab / space) separated."""
|
| 90 |
series: list[list[float]] = []
|
| 91 |
+
for raw_line in (text or "").replace(";", "\n").replace("\t", ",").splitlines():
|
| 92 |
+
line = raw_line.strip()
|
| 93 |
if not line:
|
| 94 |
continue
|
| 95 |
+
chunks = line.replace(" ", ",") if "," not in line else line
|
| 96 |
+
values: list[float] = []
|
| 97 |
+
for chunk in chunks.split(","):
|
| 98 |
+
chunk = chunk.strip()
|
| 99 |
+
if not chunk:
|
| 100 |
+
continue
|
| 101 |
+
try:
|
| 102 |
+
number = float(chunk)
|
| 103 |
+
except ValueError as exc:
|
| 104 |
+
raise gr.Error(
|
| 105 |
+
f"Could not read {chunk!r} as a number. Use one series per "
|
| 106 |
+
"line with comma-separated values, e.g. ``1.0, 2.0, 3.0``."
|
| 107 |
+
) from exc
|
| 108 |
+
values.append(number)
|
| 109 |
if not values:
|
| 110 |
+
continue
|
| 111 |
+
if len(values) > MAX_POINTS_PER_SERIES:
|
| 112 |
+
raise gr.Error(
|
| 113 |
+
f"Each series can hold at most {MAX_POINTS_PER_SERIES} values "
|
| 114 |
+
f"(got {len(values)})."
|
| 115 |
+
)
|
| 116 |
series.append(values)
|
| 117 |
+
if len(series) > MAX_SERIES:
|
| 118 |
+
raise gr.Error(f"At most {MAX_SERIES} input series are supported.")
|
| 119 |
return series
|
| 120 |
|
| 121 |
|
| 122 |
+
# ---------------------------------------------------------------------------
|
| 123 |
+
# Chart
|
| 124 |
+
# ---------------------------------------------------------------------------
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def make_figure(
|
| 128 |
+
series: list[list[float]],
|
| 129 |
+
horizon: int,
|
| 130 |
+
target_index: int,
|
| 131 |
+
result: dict[str, Any],
|
| 132 |
+
) -> plt.Figure:
|
| 133 |
+
"""Draw the observed series and, when forecasting, the predicted values."""
|
| 134 |
+
fig, ax = plt.subplots(figsize=(7.8, 3.9), constrained_layout=False)
|
| 135 |
+
fig.patch.set_facecolor("#FFFFFF")
|
| 136 |
+
|
| 137 |
+
if not series:
|
| 138 |
+
ax.set_axis_off()
|
| 139 |
+
ax.text(
|
| 140 |
+
0.5,
|
| 141 |
+
0.5,
|
| 142 |
+
"No time series supplied\n(text-only completion)",
|
| 143 |
+
ha="center",
|
| 144 |
+
va="center",
|
| 145 |
+
fontsize=12,
|
| 146 |
+
color="#5F6368",
|
| 147 |
+
)
|
| 148 |
+
return fig
|
| 149 |
+
|
| 150 |
+
forecast = result.get("timeseries") if horizon > 0 else None
|
| 151 |
+
if forecast is not None and forecast.get("values"):
|
| 152 |
+
history = series[target_index]
|
| 153 |
+
values = list(forecast["values"])
|
| 154 |
+
x_hist = list(range(len(history)))
|
| 155 |
+
x_future = list(range(len(history), len(history) + len(values)))
|
| 156 |
ax.plot(
|
| 157 |
+
x_hist,
|
| 158 |
+
history,
|
| 159 |
+
"-o",
|
| 160 |
+
color=INPUT_COLOR,
|
| 161 |
+
markersize=3.5,
|
| 162 |
+
linewidth=2.0,
|
| 163 |
+
label=f"Observed (series {target_index + 1})",
|
| 164 |
+
)
|
| 165 |
+
ax.plot(
|
| 166 |
+
x_future,
|
| 167 |
values,
|
| 168 |
+
"-o",
|
| 169 |
+
color=OUTPUT_COLOR,
|
| 170 |
+
markersize=3.5,
|
| 171 |
linewidth=2.0,
|
| 172 |
+
label=f"Forecast (horizon {horizon})",
|
|
|
|
|
|
|
| 173 |
)
|
| 174 |
+
ax.axvline(
|
| 175 |
+
len(history) - 0.5,
|
| 176 |
+
color=BOUNDARY_COLOR,
|
| 177 |
+
linestyle="--",
|
| 178 |
+
linewidth=1.1,
|
| 179 |
+
)
|
| 180 |
+
for x_value, y_value in zip(x_future, values):
|
| 181 |
+
ax.annotate(
|
| 182 |
+
f"{y_value:.2f}",
|
| 183 |
+
(x_value, y_value),
|
| 184 |
+
textcoords="offset points",
|
| 185 |
+
xytext=(0, 7),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 186 |
ha="center",
|
| 187 |
+
fontsize=7.5,
|
| 188 |
+
color=OUTPUT_COLOR,
|
|
|
|
| 189 |
)
|
| 190 |
+
ax.set_title(
|
| 191 |
+
f"Forecast for series {target_index + 1} — {horizon} step(s), "
|
| 192 |
+
"values on the original scale",
|
| 193 |
+
fontsize=10.5,
|
| 194 |
+
color="#202124",
|
| 195 |
+
)
|
| 196 |
+
else:
|
| 197 |
+
for index, values in enumerate(series):
|
| 198 |
+
ax.plot(
|
| 199 |
+
range(len(values)),
|
| 200 |
+
values,
|
| 201 |
+
"-o",
|
| 202 |
+
color=SERIES_COLORS[index % len(SERIES_COLORS)],
|
| 203 |
+
markersize=3.2,
|
| 204 |
+
linewidth=1.9,
|
| 205 |
+
label=f"Series {index + 1}",
|
| 206 |
)
|
| 207 |
+
ax.set_title("Observed time series", fontsize=10.5, color="#202124")
|
| 208 |
+
|
| 209 |
+
ax.set_xlabel("Observation / forecast index")
|
| 210 |
+
ax.set_ylabel("Value (raw scale)")
|
| 211 |
+
ax.xaxis.set_major_locator(MaxNLocator(integer=True))
|
| 212 |
+
ax.grid(axis="y", color="#D9DCE1", linewidth=0.8, alpha=0.7)
|
| 213 |
+
ax.set_axisbelow(True)
|
| 214 |
+
ax.spines["top"].set_visible(False)
|
| 215 |
+
ax.spines["right"].set_visible(False)
|
| 216 |
+
ax.legend(loc="best", fontsize=8.5, frameon=False)
|
| 217 |
return fig
|
| 218 |
|
| 219 |
|
| 220 |
+
def format_answer(result: dict[str, Any], horizon: int) -> str:
|
| 221 |
+
"""Render the model response plus a short run summary."""
|
| 222 |
+
lines: list[str] = [(result.get("content") or "(no text returned)").strip()]
|
| 223 |
+
forecast = result.get("timeseries") if horizon > 0 else None
|
| 224 |
+
if forecast is not None and forecast.get("values"):
|
| 225 |
+
values = ", ".join(f"{value:.3f}" for value in forecast["values"])
|
| 226 |
+
lines += ["", f"Forecast (original scale, {horizon} step(s)):", values]
|
| 227 |
+
lines += [
|
| 228 |
+
"",
|
| 229 |
+
"—",
|
| 230 |
+
f"finish_reason={result.get('finish_reason')} "
|
| 231 |
+
f"prompt_tokens={result.get('prompt_tokens')} "
|
| 232 |
+
f"completion_tokens={result.get('completion_tokens')} "
|
| 233 |
+
f"decode={result.get('decode_impl')}",
|
| 234 |
+
]
|
| 235 |
+
return "\n".join(lines)
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
# ---------------------------------------------------------------------------
|
| 239 |
+
# Inference
|
| 240 |
+
# ---------------------------------------------------------------------------
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
@spaces.GPU(duration=180)
|
| 244 |
+
def run_timebraid(
|
| 245 |
series_text: str,
|
| 246 |
+
prompt: str,
|
| 247 |
+
horizon: int = 0,
|
| 248 |
+
system_prompt: str = "",
|
| 249 |
+
target_number: int = 1,
|
| 250 |
+
max_new_tokens: int = 256,
|
| 251 |
+
) -> tuple[str, Any]:
|
| 252 |
+
"""Run one TimeBraid request and return the written answer plus a chart.
|
| 253 |
+
|
| 254 |
+
Supply one series per line in ``series_text`` (comma-separated numbers).
|
| 255 |
+
Leave ``horizon`` at 0 for an explanation of the observed series; set it to
|
| 256 |
+
a positive number to forecast that many future values for the target series.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 257 |
"""
|
| 258 |
+
series_text = _coerce_text(series_text)
|
| 259 |
+
prompt = _coerce_text(prompt)
|
| 260 |
+
system_prompt = _coerce_text(system_prompt)
|
| 261 |
+
horizon = _coerce_int(horizon, 0, 0, MAX_HORIZON)
|
| 262 |
+
target_number = _coerce_int(target_number, 1, 1, MAX_SERIES)
|
| 263 |
+
max_new_tokens = _coerce_int(max_new_tokens, 256, 32, 1024)
|
| 264 |
+
|
| 265 |
+
instruction = prompt.strip()
|
| 266 |
+
if not instruction:
|
| 267 |
+
raise gr.Error("Enter a question or instruction for the model.")
|
| 268 |
+
|
| 269 |
+
series = parse_series_block(series_text)
|
| 270 |
+
|
| 271 |
+
if horizon > 0 and not series:
|
| 272 |
+
raise gr.Error(
|
| 273 |
+
"Forecasting needs at least one input series — paste comma-separated "
|
| 274 |
+
"values, one series per line."
|
| 275 |
+
)
|
| 276 |
+
|
| 277 |
+
target_index = 0
|
| 278 |
+
if horizon > 0 and len(series) > 1:
|
| 279 |
+
if target_number > len(series):
|
| 280 |
+
raise gr.Error(
|
| 281 |
+
f"Only {len(series)} series were supplied, so the forecast target "
|
| 282 |
+
f"cannot be series {target_number}."
|
| 283 |
+
)
|
| 284 |
+
target_index = target_number - 1
|
| 285 |
+
|
| 286 |
+
messages: list[dict[str, str]] = []
|
| 287 |
+
if system_prompt.strip():
|
| 288 |
+
messages.append({"role": "system", "content": system_prompt.strip()})
|
| 289 |
+
messages.append({"role": "user", "content": instruction})
|
| 290 |
+
|
| 291 |
+
processor_kwargs: dict[str, Any] = {
|
| 292 |
+
"messages": messages,
|
| 293 |
+
"timeseries": series if series else None,
|
| 294 |
+
"horizon": horizon if horizon > 0 else None,
|
| 295 |
+
"return_tensors": "pt",
|
| 296 |
+
}
|
| 297 |
+
if horizon > 0 and len(series) > 1:
|
| 298 |
+
processor_kwargs["target_series_index"] = target_index
|
| 299 |
+
|
| 300 |
try:
|
| 301 |
+
model_inputs = processor(**processor_kwargs)
|
| 302 |
+
except (ValueError, RuntimeError, TypeError) as exc:
|
| 303 |
+
raise gr.Error(f"Could not prepare the request: {exc}") from exc
|
| 304 |
+
|
| 305 |
+
device_inputs = {
|
| 306 |
+
key: value.to(model.device) if isinstance(value, torch.Tensor) else value
|
| 307 |
+
for key, value in model_inputs.items()
|
| 308 |
+
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 309 |
with torch.inference_mode():
|
| 310 |
+
output = model.generate(
|
| 311 |
+
**device_inputs,
|
| 312 |
+
max_new_tokens=max_new_tokens,
|
| 313 |
do_sample=False,
|
| 314 |
num_beams=1,
|
| 315 |
num_return_sequences=1,
|
| 316 |
)
|
| 317 |
+
try:
|
| 318 |
+
result = processor.post_process_generation(output, model_inputs=model_inputs)
|
| 319 |
+
except (ValueError, RuntimeError, TypeError) as exc:
|
| 320 |
+
raise gr.Error(f"Could not decode the model output: {exc}") from exc
|
| 321 |
+
|
| 322 |
+
figure = make_figure(series, horizon, target_index, result)
|
| 323 |
+
return format_answer(result, horizon), figure
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 324 |
|
| 325 |
|
| 326 |
+
# ---------------------------------------------------------------------------
|
| 327 |
+
# Interface
|
| 328 |
+
# ---------------------------------------------------------------------------
|
| 329 |
+
|
| 330 |
CSS = """
|
|
|
|
| 331 |
.dark .gradio-container { color: var(--body-text-color); }
|
| 332 |
+
.dark .gradio-container .prose { color: var(--body-text-color); }
|
| 333 |
"""
|
| 334 |
|
| 335 |
+
EXAMPLES = [
|
| 336 |
+
[
|
| 337 |
+
"2.0, 2.1, 2.2, 2.4, 2.8, 3.1, 3.0, 2.9, 3.4, 3.8, 4.1, 4.3",
|
| 338 |
+
"Describe the dominant trend, major turning points, and whether the "
|
| 339 |
+
"series becomes more volatile.",
|
| 340 |
+
0,
|
| 341 |
+
"",
|
| 342 |
+
],
|
| 343 |
+
[
|
| 344 |
+
"20.1, 20.0, 20.2, 20.1, 20.3, 34.8, 20.2, 20.1, 20.0, 20.2",
|
| 345 |
+
"Does this series contain an isolated upward spike? Say only whether it "
|
| 346 |
+
"occurs near the beginning, middle, or end; do not give a numeric index. "
|
| 347 |
+
"Compare it with the neighboring level.",
|
| 348 |
+
0,
|
| 349 |
+
"",
|
| 350 |
+
],
|
| 351 |
+
[
|
| 352 |
+
"100.0, 104.0, 107.0, 111.0, 116.0, 120.0, 125.0, 129.0\n"
|
| 353 |
+
"3.2, 3.1, 3.3, 3.4, 3.8, 3.7, 4.0, 4.2",
|
| 354 |
+
"Series 1 is website visits and Series 2 is conversion rate. Compare "
|
| 355 |
+
"their trends, volatility, and co-movement.",
|
| 356 |
+
0,
|
| 357 |
+
"",
|
| 358 |
+
],
|
| 359 |
+
[
|
| 360 |
+
"100.0, 102.0, 105.0, 107.0, 103.0, 101.0, 99.0, 100.0, 103.0, 106.0, "
|
| 361 |
+
"108.0, 104.0, 102.0, 100.0, 101.0, 104.0",
|
| 362 |
+
"Forecast the next 8 values from the observed seasonal pattern.",
|
| 363 |
+
8,
|
| 364 |
+
"",
|
| 365 |
+
],
|
| 366 |
+
[
|
| 367 |
+
"200.0, 208.0, 215.0, 205.0, 198.0, 210.0, 218.0, 207.0, 201.0, 212.0, "
|
| 368 |
+
"220.0, 209.0",
|
| 369 |
+
"A promotion starts at the first forecast step and is expected to lift "
|
| 370 |
+
"demand above the recent seasonal baseline. Forecast the next 6 values.",
|
| 371 |
+
6,
|
| 372 |
+
"You analyze weekly product demand.",
|
| 373 |
+
],
|
| 374 |
+
[
|
| 375 |
+
"63.0, 62.0, 64.0, 63.0, 65.0, 64.0, 63.0, 72.0, 88.0, 91.0, 86.0, "
|
| 376 |
+
"74.0, 66.0, 64.0, 63.0, 62.0",
|
| 377 |
+
"This series is a patient's daily resting heart rate in beats per minute. "
|
| 378 |
+
"The patient reported a fever starting around day 8 that lasted a few "
|
| 379 |
+
"days. Is the heart-rate pattern consistent with that report, and does "
|
| 380 |
+
"the recovery look complete by the end of the series?",
|
| 381 |
+
0,
|
| 382 |
+
"",
|
| 383 |
+
],
|
| 384 |
+
[
|
| 385 |
+
"420.0, 380.0, 395.0, 410.0, 430.0, 485.0, 560.0, 660.0, 790.0, 940.0, "
|
| 386 |
+
"1130.0, 1350.0",
|
| 387 |
+
"A new school term begins next week, which typically accelerates "
|
| 388 |
+
"transmission for several weeks. Forecast the next 6 weekly case counts.",
|
| 389 |
+
6,
|
| 390 |
+
"You analyze weekly influenza-like illness case counts for a regional "
|
| 391 |
+
"health authority.",
|
| 392 |
+
],
|
| 393 |
+
[
|
| 394 |
+
"2.10, 2.05, 2.02, 2.00, 2.04, 2.15, 2.45, 2.80, 3.00, 3.10, 3.15, "
|
| 395 |
+
"3.20, 3.25, 3.30, 3.35, 3.30, 3.40, 3.55, 3.60, 3.45, 3.15, 2.80, "
|
| 396 |
+
"2.50, 2.25, 2.12, 2.06, 2.03, 2.02, 2.06, 2.18, 2.48, 2.83, 3.02, "
|
| 397 |
+
"3.12, 3.18, 3.24, 3.28, 3.34, 3.38, 3.33, 3.44, 3.58, 3.62, 3.48, "
|
| 398 |
+
"3.18, 2.83, 2.52, 2.28",
|
| 399 |
+
"The history covers two days of hourly load in gigawatts. A heatwave "
|
| 400 |
+
"arrives tomorrow and cooling demand is expected to lift the afternoon "
|
| 401 |
+
"and evening peak. Forecast the next 24 hourly values.",
|
| 402 |
+
24,
|
| 403 |
+
"You analyze hourly electricity load for a regional grid operator.",
|
| 404 |
+
],
|
| 405 |
+
[
|
| 406 |
+
"",
|
| 407 |
+
"Explain in one sentence what a moving average reveals about a time series.",
|
| 408 |
+
0,
|
| 409 |
+
"",
|
| 410 |
+
],
|
| 411 |
+
]
|
| 412 |
+
|
| 413 |
+
|
| 414 |
with gr.Blocks(title="TimeBraid") as demo:
|
| 415 |
+
gr.Markdown(
|
| 416 |
+
"# TimeBraid\n"
|
| 417 |
+
"**Unifying time series and language for understanding and forecasting** — "
|
| 418 |
+
"[paper](https://huggingface.co/papers/2609.29792) · "
|
| 419 |
+
"[model](https://huggingface.co/XinyueWangg/TimeBraid-2.5B) · "
|
| 420 |
+
"[GitHub](https://github.com/CharonWangg/TimeBraid)\n\n"
|
| 421 |
+
"Paste one series per line (comma-separated numbers), write what you want "
|
| 422 |
+
"to know, and choose whether the model should *explain* the observed data "
|
| 423 |
+
"or *forecast* its future. TimeBraid-2.5B braids a Qwen3-1.7B language "
|
| 424 |
+
"backbone with a TimesFM 2.5 time-series expert, so the same numbers can "
|
| 425 |
+
"read very differently depending on the context you give them."
|
| 426 |
+
)
|
| 427 |
+
|
| 428 |
+
with gr.Row():
|
| 429 |
+
with gr.Column(scale=5):
|
| 430 |
+
series_box = gr.Textbox(
|
| 431 |
+
label="Time series — one series per line, comma-separated values",
|
| 432 |
+
placeholder="200.0, 208.0, 215.0, 205.0, 198.0, 210.0",
|
| 433 |
+
lines=7,
|
| 434 |
+
)
|
| 435 |
+
prompt_box = gr.Textbox(
|
| 436 |
+
label="Question or instruction",
|
| 437 |
+
placeholder=(
|
| 438 |
+
"Describe the trend and turning points in this series."
|
| 439 |
+
),
|
| 440 |
+
lines=4,
|
| 441 |
+
)
|
| 442 |
+
horizon_slider = gr.Slider(
|
| 443 |
+
minimum=0,
|
| 444 |
+
maximum=MAX_HORIZON,
|
| 445 |
+
value=0,
|
| 446 |
+
step=1,
|
| 447 |
+
precision=0,
|
| 448 |
+
label="Forecast horizon",
|
| 449 |
+
info="0 = explain the observed series; N > 0 = predict the next N values",
|
| 450 |
+
)
|
| 451 |
+
run_button = gr.Button("Run TimeBraid", variant="primary")
|
| 452 |
+
|
| 453 |
+
with gr.Column(scale=5):
|
| 454 |
+
answer_box = gr.Textbox(
|
| 455 |
+
label="Model response",
|
| 456 |
+
lines=12,
|
| 457 |
+
show_copy_button=True,
|
| 458 |
+
)
|
| 459 |
+
plot_box = gr.Plot(label="Series and forecast")
|
| 460 |
+
|
| 461 |
+
with gr.Accordion("Advanced options", open=False):
|
| 462 |
+
system_box = gr.Textbox(
|
| 463 |
+
label="System prompt (optional)",
|
| 464 |
+
placeholder="You analyze weekly product demand.",
|
| 465 |
+
lines=2,
|
| 466 |
+
)
|
| 467 |
+
target_number = gr.Number(
|
| 468 |
+
label="Forecast target series (1-based)",
|
| 469 |
+
value=1,
|
| 470 |
+
precision=0,
|
| 471 |
+
info="Used only when forecasting from more than one input series.",
|
| 472 |
)
|
| 473 |
+
max_tokens_slider = gr.Slider(
|
| 474 |
+
minimum=64,
|
| 475 |
+
maximum=1024,
|
| 476 |
+
value=256,
|
| 477 |
+
step=32,
|
| 478 |
+
precision=0,
|
| 479 |
+
label="Maximum new text tokens",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 480 |
)
|
| 481 |
+
|
| 482 |
+
inputs = [
|
| 483 |
+
series_box,
|
| 484 |
+
prompt_box,
|
| 485 |
+
horizon_slider,
|
| 486 |
+
system_box,
|
| 487 |
+
target_number,
|
| 488 |
+
max_tokens_slider,
|
| 489 |
+
]
|
| 490 |
+
outputs = [answer_box, plot_box]
|
| 491 |
+
|
| 492 |
+
run_button.click(fn=run_timebraid, inputs=inputs, outputs=outputs)
|
| 493 |
+
prompt_box.submit(fn=run_timebraid, inputs=inputs, outputs=outputs)
|
| 494 |
+
|
| 495 |
+
gr.Examples(
|
| 496 |
+
examples=EXAMPLES,
|
| 497 |
+
inputs=[series_box, prompt_box, horizon_slider, system_box],
|
| 498 |
+
fn=run_timebraid,
|
| 499 |
+
cache_examples=True,
|
| 500 |
+
cache_mode="lazy",
|
| 501 |
+
examples_per_page=3,
|
| 502 |
+
label="Example requests (from the TimeBraid repository and model card)",
|
| 503 |
+
)
|
| 504 |
+
|
| 505 |
+
gr.Markdown(
|
| 506 |
+
"All series and forecasts are plotted on their original scale. Forecasts "
|
| 507 |
+
"come from greedy decoding on a fresh prompt, exactly as in the released "
|
| 508 |
+
"inference recipe."
|
| 509 |
)
|
| 510 |
|
| 511 |
+
|
| 512 |
+
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)
|
requirements.txt
CHANGED
|
@@ -1,8 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
torch==2.11.0
|
| 2 |
-
transformers==4.57.6
|
| 3 |
-
huggingface_hub==0.36.2
|
| 4 |
-
accelerate==1.11.0
|
| 5 |
numpy==2.1.3
|
|
|
|
|
|
|
|
|
|
| 6 |
safetensors==0.5.3
|
|
|
|
| 7 |
matplotlib
|
| 8 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# TimeBraid-2.5B inference stack for HF Spaces ZeroGPU (NVIDIA RTX PRO 6000, sm_120).
|
| 2 |
+
#
|
| 3 |
+
# `spaces`, `gradio` and `huggingface_hub` are intentionally omitted: the platform
|
| 4 |
+
# preinstalls and pins them.
|
| 5 |
+
#
|
| 6 |
+
# torch is pinned to 2.11.0 to match the prebuilt FlashAttention-2 wheel cell
|
| 7 |
+
# below (`pt211-cu130-cp312`). ZeroGPU accepts 2.8.0 / 2.9.1 / 2.10.0 / 2.11.0.
|
| 8 |
torch==2.11.0
|
|
|
|
|
|
|
|
|
|
| 9 |
numpy==2.1.3
|
| 10 |
+
einops
|
| 11 |
+
accelerate==1.11.0
|
| 12 |
+
transformers==4.57.6
|
| 13 |
safetensors==0.5.3
|
| 14 |
+
tokenizers==0.22.2
|
| 15 |
matplotlib
|
| 16 |
+
|
| 17 |
+
# MoT mixed attention in TimeBraid is FlashAttention-2 only, and FA3/FA4 cannot
|
| 18 |
+
# run on sm_120, so we use the prebuilt Blackwell FA2 wheel rather than sdpa.
|
| 19 |
+
# Requires `python_version: "3.12"` in README.md to match the cp312 tag.
|
| 20 |
+
https://huggingface.co/datasets/multimodalart/zerogpu-blackwell-wheels/resolve/main/wheels/pt211-cu130-cp312/flash_attn-2.8.3-cp312-cp312-linux_x86_64.whl
|