multimodalart HF Staff commited on
Commit
84ab377
·
verified ·
1 Parent(s): 7006eaf

Upload folder using huggingface_hub

Browse files
Files changed (3) hide show
  1. README.md +44 -16
  2. app.py +449 -266
  3. 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.17.3
8
  app_file: app.py
9
- short_description: Time series + language model for analysis and forecasting
10
  python_version: "3.12"
11
- startup_duration_timeout: 30m
 
12
  ---
13
 
14
  # TimeBraid
15
 
16
- Interactive demo of **TimeBraid-2.5B** — a unified time-series + language model that
17
- interleaves a Qwen3-1.7B language backbone with TimesFM 2.5 200M time-series experts
18
- through Mixture-of-Transformers layers.
19
 
20
- Paste numeric series (one series per line, values comma- or space-separated) and ask a
21
- question, or set a forecast horizon to get a numeric point forecast plotted against the
22
- observed history.
23
 
24
- - Model: [XinyueWangg/TimeBraid-2.5B](https://huggingface.co/XinyueWangg/TimeBraid-2.5B) (Apache-2.0)
25
- - Code: [CharonWangg/TimeBraid](https://github.com/CharonWangg/TimeBraid)
26
- - The `timebraid/` package in this Space is the authors' inference package, shipped
27
- verbatim so `AutoModelForCausalLM` resolves the custom architecture with
28
- `trust_remote_code=False`.
29
 
30
- Runs on ZeroGPU. Example inputs are drawn from the authors'
31
- `examples/inference_tasks.ipynb` notebook (Apache-2.0).
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: unified time series + language demo on ZeroGPU."""
 
 
 
 
 
2
 
3
  import os
4
 
5
  os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
6
 
7
- import spaces # MUST come before any torch / CUDA-touching import
8
- import time
 
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
- from timebraid import TimeBraidProcessor # noqa: F401 - registers HF AutoClasses
 
 
 
 
 
20
 
21
  MODEL_ID = "XinyueWangg/TimeBraid-2.5B"
22
 
23
- # Chart conventions from the authors' notebook (examples/inference_tasks.ipynb).
24
- _INPUT_COLOR = "#0072B2"
25
- _FORECAST_COLOR = "#D55E00"
26
- _BOUNDARY_COLOR = "#555555"
27
- _SERIES_COLORS = ("#0072B2", "#009E73", "#CC79A7", "#E69F00")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
28
 
29
- processor = AutoProcessor.from_pretrained(MODEL_ID, trust_remote_code=False)
 
 
 
 
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
- def parse_series(text: str) -> list[list[float]]:
44
- """Parse one series per line; values comma- or whitespace-separated."""
45
- if text is None:
46
- return []
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
47
  series: list[list[float]] = []
48
- for line in text.splitlines():
49
- line = line.strip().rstrip(",")
50
  if not line:
51
  continue
52
- tokens = line.replace(",", " ").split()
53
- values = [float(token) for token in tokens]
 
 
 
 
 
 
 
 
 
 
 
 
54
  if not values:
55
- raise ValueError("A series line parsed to zero values.")
 
 
 
 
 
56
  series.append(values)
 
 
57
  return series
58
 
59
 
60
- def render_plot(
61
- history: list[list[float]],
62
- forecast: list[float] | None,
63
- target_index: int | None,
64
- ) -> matplotlib.figure.Figure:
65
- """Draw observed history and, when present, the forecast continuation."""
66
- n = len(history)
67
- fig, axes = plt.subplots(
68
- n, 1, figsize=(9.5, 2.9 * n), constrained_layout=True, squeeze=False
69
- )
70
- for index, (ax, values) in enumerate(zip(axes[:, 0], history)):
71
- color = _INPUT_COLOR if forecast is not None else _SERIES_COLORS[index % 4]
72
- x = list(range(len(values)))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
73
  ax.plot(
74
- x,
 
 
 
 
 
 
 
 
 
75
  values,
76
- color=color,
 
 
77
  linewidth=2.0,
78
- marker="o" if len(values) <= 64 else None,
79
- markersize=4,
80
- label="Observed history" if forecast is not None else f"Series {index + 1}",
81
  )
82
- title = f"INPUT TIMESERIES · Series {index + 1}"
83
- if forecast is not None and target_index == index:
84
- future_x = list(range(len(values), len(values) + len(forecast)))
85
- boundary = len(values) - 0.5
86
- ax.axvspan(boundary, future_x[-1] + 0.5, color=_FORECAST_COLOR, alpha=0.06)
87
- ax.plot(
88
- (len(values) - 1, *future_x),
89
- (values[-1], *forecast),
90
- color=_FORECAST_COLOR,
91
- linewidth=2.0,
92
- linestyle="--",
93
- label="Point forecast",
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
- va="bottom",
110
- fontsize=8,
111
- color=_BOUNDARY_COLOR,
112
  )
113
- title = (
114
- f"INPUT HISTORY + OUTPUT FORECAST · Series {index + 1} (forecast target)"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
  )
116
- ax.set_title(title, loc="left", fontsize=11, fontweight="bold")
117
- ax.set_xlabel("Observation / forecast index" if forecast is not None else "Observation index")
118
- ax.set_ylabel("Value (raw scale)")
119
- ax.xaxis.set_major_locator(MaxNLocator(integer=True))
120
- ax.grid(axis="y", color="#D9DCE1", linewidth=0.8, alpha=0.7)
121
- ax.set_axisbelow(True)
122
- ax.spines["top"].set_visible(False)
123
- ax.spines["right"].set_visible(False)
124
- ax.legend(loc="upper left", frameon=False, ncol=3, fontsize=9)
 
125
  return fig
126
 
127
 
128
- @spaces.GPU(duration=120)
129
- def run(
130
- prompt: str,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
131
  series_text: str,
132
- horizon: float,
133
- target_series_index: float,
134
- max_new_tokens: float,
135
- progress=gr.Progress(track_tqdm=True),
136
- ):
137
- """Analyze or forecast a time series with TimeBraid.
138
-
139
- Args:
140
- prompt: Question or instruction about the series (or a plain text question).
141
- series_text: Numeric series, one per line; values comma- or space-separated.
142
- horizon: Number of future points to forecast; 0 = no forecast (understanding only).
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
- started = time.perf_counter()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
149
  try:
150
- series = parse_series(series_text)
151
- except ValueError as exc:
152
- raise gr.Error(f"Could not parse the time series: {exc}")
153
-
154
- horizon_i = int(horizon) if horizon and horizon > 0 else None
155
- target_i: int | None = None
156
- if horizon_i is not None:
157
- if not series:
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
- outputs = model.generate(
183
- **inputs,
184
- max_new_tokens=int(max_new_tokens),
185
  do_sample=False,
186
  num_beams=1,
187
  num_return_sequences=1,
188
  )
189
- result = processor.post_process_generation(outputs, model_inputs=inputs)
190
- elapsed = time.perf_counter() - started
191
-
192
- answer = (result.get("content") or "").strip()
193
- if not answer and horizon_i is None:
194
- answer = "(empty response)"
195
- ts = result.get("timeseries")
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
- with gr.Column(elem_id="col-container"):
219
- gr.Markdown(
220
- """
221
- # 🧵 TimeBraid: Unifying Time Series and Language
222
- Ask questions about numeric series, or request a point forecast.
223
- One series per line, values comma- or space-separated. Leave the series
224
- box empty for a pure text question. Set **Forecast horizon** to 0 for
225
- analysis-only, or a positive number to get numeric forecast values.
226
- **Model:** [XinyueWangg/TimeBraid-2.5B](https://huggingface.co/XinyueWangg/TimeBraid-2.5B) ·
227
- **Code:** [GitHub](https://github.com/CharonWangg/TimeBraid) (Apache-2.0)
228
- """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
229
  )
230
- with gr.Row():
231
- with gr.Column(scale=1):
232
- prompt = gr.Textbox(
233
- label="Prompt",
234
- placeholder="e.g. Describe the dominant trend and turning points, or forecast the next 8 values.",
235
- lines=3,
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
- run_btn.click(
321
- run,
322
- inputs=[prompt, series_box, horizon, target_idx, max_new_tokens],
323
- outputs=[answer_out, forecast_out, plot_out, stats_out],
324
- api_name="run",
325
- concurrency_limit=1,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
326
  )
327
 
328
- if __name__ == "__main__":
329
- demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)
 
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
- 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
 
 
 
 
 
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