timebraid / app.py
multimodalart's picture
multimodalart HF Staff
Tighten ZeroGPU duration to measured 30s; clearer empty-content message
2a3e13a verified
Raw History Blame Contribute Delete
18 kB
"""TimeBraid: unify time series and language for understanding and forecasting.
A single-request demo around ``XinyueWangg/TimeBraid-2.5B``. The model reads one
or more numeric series plus a natural-language request and either explains what
it sees or forecasts future values.
"""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: F401 # must be imported before torch / transformers
from typing import Any
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
import torch # noqa: E402
from matplotlib.ticker import MaxNLocator # noqa: E402
from transformers import AutoModelForCausalLM, AutoProcessor # noqa: E402
import gradio as gr # noqa: E402
MODEL_ID = "XinyueWangg/TimeBraid-2.5B"
MAX_SERIES = 8
MAX_POINTS_PER_SERIES = 2048
MAX_HORIZON = 64
INPUT_COLOR = "#0072B2"
OUTPUT_COLOR = "#D55E00"
BOUNDARY_COLOR = "#555555"
SERIES_COLORS = (
"#0072B2",
"#009E73",
"#CC79A7",
"#E69F00",
"#56B4E9",
"#F0E442",
"#000000",
"#999999",
)
# ---------------------------------------------------------------------------
# Model (loaded once, at module scope, on the GPU)
# ---------------------------------------------------------------------------
processor = AutoProcessor.from_pretrained(
MODEL_ID,
trust_remote_code=True,
fix_mistral_regex=False,
)
model = (
AutoModelForCausalLM.from_pretrained(
MODEL_ID,
trust_remote_code=True,
dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
use_safetensors=True,
)
.eval()
.to("cuda")
)
# ---------------------------------------------------------------------------
# Input parsing / coercion helpers
# ---------------------------------------------------------------------------
def _coerce_text(value: Any, default: str = "") -> str:
return value if isinstance(value, str) else default
def _coerce_int(value: Any, default: int, low: int, high: int) -> int:
if isinstance(value, bool) or not isinstance(value, (int, float)):
return default
try:
number = int(value)
except (TypeError, ValueError):
return default
return max(low, min(high, number))
def parse_series_block(text: str) -> list[list[float]]:
"""Parse one series per non-empty line, comma (or tab / space) separated."""
series: list[list[float]] = []
for raw_line in (text or "").replace(";", "\n").replace("\t", ",").splitlines():
line = raw_line.strip()
if not line:
continue
chunks = line.replace(" ", ",") if "," not in line else line
values: list[float] = []
for chunk in chunks.split(","):
chunk = chunk.strip()
if not chunk:
continue
try:
number = float(chunk)
except ValueError as exc:
raise gr.Error(
f"Could not read {chunk!r} as a number. Use one series per "
"line with comma-separated values, e.g. ``1.0, 2.0, 3.0``."
) from exc
values.append(number)
if not values:
continue
if len(values) > MAX_POINTS_PER_SERIES:
raise gr.Error(
f"Each series can hold at most {MAX_POINTS_PER_SERIES} values "
f"(got {len(values)})."
)
series.append(values)
if len(series) > MAX_SERIES:
raise gr.Error(f"At most {MAX_SERIES} input series are supported.")
return series
# ---------------------------------------------------------------------------
# Chart
# ---------------------------------------------------------------------------
def make_figure(
series: list[list[float]],
horizon: int,
target_index: int,
result: dict[str, Any],
) -> plt.Figure:
"""Draw the observed series and, when forecasting, the predicted values."""
fig, ax = plt.subplots(figsize=(7.8, 3.9), constrained_layout=False)
fig.patch.set_facecolor("#FFFFFF")
if not series:
ax.set_axis_off()
ax.text(
0.5,
0.5,
"No time series supplied\n(text-only completion)",
ha="center",
va="center",
fontsize=12,
color="#5F6368",
)
return fig
forecast = result.get("timeseries") if horizon > 0 else None
if forecast is not None and forecast.get("values"):
history = series[target_index]
values = list(forecast["values"])
x_hist = list(range(len(history)))
x_future = list(range(len(history), len(history) + len(values)))
ax.plot(
x_hist,
history,
"-o",
color=INPUT_COLOR,
markersize=3.5,
linewidth=2.0,
label=f"Observed (series {target_index + 1})",
)
ax.plot(
x_future,
values,
"-o",
color=OUTPUT_COLOR,
markersize=3.5,
linewidth=2.0,
label=f"Forecast (horizon {horizon})",
)
ax.axvline(
len(history) - 0.5,
color=BOUNDARY_COLOR,
linestyle="--",
linewidth=1.1,
)
for x_value, y_value in zip(x_future, values):
ax.annotate(
f"{y_value:.2f}",
(x_value, y_value),
textcoords="offset points",
xytext=(0, 7),
ha="center",
fontsize=7.5,
color=OUTPUT_COLOR,
)
ax.set_title(
f"Forecast for series {target_index + 1} — {horizon} step(s), "
"values on the original scale",
fontsize=10.5,
color="#202124",
)
else:
for index, values in enumerate(series):
ax.plot(
range(len(values)),
values,
"-o",
color=SERIES_COLORS[index % len(SERIES_COLORS)],
markersize=3.2,
linewidth=1.9,
label=f"Series {index + 1}",
)
ax.set_title("Observed time series", fontsize=10.5, color="#202124")
ax.set_xlabel("Observation / forecast index")
ax.set_ylabel("Value (raw scale)")
ax.xaxis.set_major_locator(MaxNLocator(integer=True))
ax.grid(axis="y", color="#D9DCE1", linewidth=0.8, alpha=0.7)
ax.set_axisbelow(True)
ax.spines["top"].set_visible(False)
ax.spines["right"].set_visible(False)
ax.legend(loc="best", fontsize=8.5, frameon=False)
return fig
def format_answer(result: dict[str, Any], horizon: int) -> str:
"""Render the model response plus a short run summary."""
forecast = result.get("timeseries") if horizon > 0 else None
content = (result.get("content") or "").strip()
if not content:
if forecast is not None and forecast.get("values"):
content = (
"TimeBraid returned the forecast values below without an "
"accompanying written answer, which is normal for a pure "
"forecasting request."
)
else:
content = "(the model returned no text)"
lines: list[str] = [content]
if forecast is not None and forecast.get("values"):
values = ", ".join(f"{value:.3f}" for value in forecast["values"])
lines += ["", f"Forecast (original scale, {horizon} step(s)):", values]
lines += [
"",
"—",
f"finish_reason={result.get('finish_reason')} "
f"prompt_tokens={result.get('prompt_tokens')} "
f"completion_tokens={result.get('completion_tokens')} "
f"decode={result.get('decode_impl')}",
]
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------------
@spaces.GPU(duration=30)
def run_timebraid(
series_text: str,
prompt: str,
horizon: int = 0,
system_prompt: str = "",
target_number: int = 1,
max_new_tokens: int = 256,
) -> tuple[str, Any]:
"""Run one TimeBraid request and return the written answer plus a chart.
Supply one series per line in ``series_text`` (comma-separated numbers).
Leave ``horizon`` at 0 for an explanation of the observed series; set it to
a positive number to forecast that many future values for the target series.
"""
series_text = _coerce_text(series_text)
prompt = _coerce_text(prompt)
system_prompt = _coerce_text(system_prompt)
horizon = _coerce_int(horizon, 0, 0, MAX_HORIZON)
target_number = _coerce_int(target_number, 1, 1, MAX_SERIES)
max_new_tokens = _coerce_int(max_new_tokens, 256, 32, 512)
instruction = prompt.strip()
if not instruction:
raise gr.Error("Enter a question or instruction for the model.")
series = parse_series_block(series_text)
if horizon > 0 and not series:
raise gr.Error(
"Forecasting needs at least one input series — paste comma-separated "
"values, one series per line."
)
target_index = 0
if horizon > 0 and len(series) > 1:
if target_number > len(series):
raise gr.Error(
f"Only {len(series)} series were supplied, so the forecast target "
f"cannot be series {target_number}."
)
target_index = target_number - 1
messages: list[dict[str, str]] = []
if system_prompt.strip():
messages.append({"role": "system", "content": system_prompt.strip()})
messages.append({"role": "user", "content": instruction})
processor_kwargs: dict[str, Any] = {
"messages": messages,
"timeseries": series if series else None,
"horizon": horizon if horizon > 0 else None,
"return_tensors": "pt",
}
if horizon > 0 and len(series) > 1:
processor_kwargs["target_series_index"] = target_index
try:
model_inputs = processor(**processor_kwargs)
except (ValueError, RuntimeError, TypeError) as exc:
raise gr.Error(f"Could not prepare the request: {exc}") from exc
device_inputs = {
key: value.to(model.device) if isinstance(value, torch.Tensor) else value
for key, value in model_inputs.items()
}
with torch.inference_mode():
output = model.generate(
**device_inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
num_beams=1,
num_return_sequences=1,
)
try:
result = processor.post_process_generation(output, model_inputs=model_inputs)
except (ValueError, RuntimeError, TypeError) as exc:
raise gr.Error(f"Could not decode the model output: {exc}") from exc
figure = make_figure(series, horizon, target_index, result)
return format_answer(result, horizon), figure
# ---------------------------------------------------------------------------
# Interface
# ---------------------------------------------------------------------------
CSS = """
.dark .gradio-container { color: var(--body-text-color); }
.dark .gradio-container .prose { color: var(--body-text-color); }
"""
EXAMPLES = [
[
"2.0, 2.1, 2.2, 2.4, 2.8, 3.1, 3.0, 2.9, 3.4, 3.8, 4.1, 4.3",
"Describe the dominant trend, major turning points, and whether the "
"series becomes more volatile.",
0,
"",
],
[
"20.1, 20.0, 20.2, 20.1, 20.3, 34.8, 20.2, 20.1, 20.0, 20.2",
"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.",
0,
"",
],
[
"100.0, 104.0, 107.0, 111.0, 116.0, 120.0, 125.0, 129.0\n"
"3.2, 3.1, 3.3, 3.4, 3.8, 3.7, 4.0, 4.2",
"Series 1 is website visits and Series 2 is conversion rate. Compare "
"their trends, volatility, and co-movement.",
0,
"",
],
[
"100.0, 102.0, 105.0, 107.0, 103.0, 101.0, 99.0, 100.0, 103.0, 106.0, "
"108.0, 104.0, 102.0, 100.0, 101.0, 104.0",
"Forecast the next 8 values from the observed seasonal pattern.",
8,
"",
],
[
"200.0, 208.0, 215.0, 205.0, 198.0, 210.0, 218.0, 207.0, 201.0, 212.0, "
"220.0, 209.0",
"A promotion starts at the first forecast step and is expected to lift "
"demand above the recent seasonal baseline. Forecast the next 6 values.",
6,
"You analyze weekly product demand.",
],
[
"63.0, 62.0, 64.0, 63.0, 65.0, 64.0, 63.0, 72.0, 88.0, 91.0, 86.0, "
"74.0, 66.0, 64.0, 63.0, 62.0",
"This series is a patient's daily resting heart rate in beats per minute. "
"The patient reported a fever starting around day 8 that lasted a few "
"days. Is the heart-rate pattern consistent with that report, and does "
"the recovery look complete by the end of the series?",
0,
"",
],
[
"420.0, 380.0, 395.0, 410.0, 430.0, 485.0, 560.0, 660.0, 790.0, 940.0, "
"1130.0, 1350.0",
"A new school term begins next week, which typically accelerates "
"transmission for several weeks. Forecast the next 6 weekly case counts.",
6,
"You analyze weekly influenza-like illness case counts for a regional "
"health authority.",
],
[
"2.10, 2.05, 2.02, 2.00, 2.04, 2.15, 2.45, 2.80, 3.00, 3.10, 3.15, "
"3.20, 3.25, 3.30, 3.35, 3.30, 3.40, 3.55, 3.60, 3.45, 3.15, 2.80, "
"2.50, 2.25, 2.12, 2.06, 2.03, 2.02, 2.06, 2.18, 2.48, 2.83, 3.02, "
"3.12, 3.18, 3.24, 3.28, 3.34, 3.38, 3.33, 3.44, 3.58, 3.62, 3.48, "
"3.18, 2.83, 2.52, 2.28",
"The history covers two days of hourly load in gigawatts. A heatwave "
"arrives tomorrow and cooling demand is expected to lift the afternoon "
"and evening peak. Forecast the next 24 hourly values.",
24,
"You analyze hourly electricity load for a regional grid operator.",
],
[
"",
"Explain in one sentence what a moving average reveals about a time series.",
0,
"",
],
]
with gr.Blocks(title="TimeBraid", theme=gr.themes.Citrus(), css=CSS) as demo:
gr.Markdown(
"# TimeBraid\n"
"**Unifying time series and language for understanding and forecasting** — "
"[paper](https://huggingface.co/papers/2609.29792) · "
"[model](https://huggingface.co/XinyueWangg/TimeBraid-2.5B) · "
"[GitHub](https://github.com/CharonWangg/TimeBraid)\n\n"
"Paste one series per line (comma-separated numbers), write what you want "
"to know, and choose whether the model should *explain* the observed data "
"or *forecast* its future. TimeBraid-2.5B braids a Qwen3-1.7B language "
"backbone with a TimesFM 2.5 time-series expert, so the same numbers can "
"read very differently depending on the context you give them."
)
with gr.Row():
with gr.Column(scale=5):
series_box = gr.Textbox(
label="Time series — one series per line, comma-separated values",
placeholder="200.0, 208.0, 215.0, 205.0, 198.0, 210.0",
lines=7,
)
prompt_box = gr.Textbox(
label="Question or instruction",
placeholder=(
"Describe the trend and turning points in this series."
),
lines=4,
)
horizon_slider = gr.Slider(
minimum=0,
maximum=MAX_HORIZON,
value=0,
step=1,
precision=0,
label="Forecast horizon",
info="0 = explain the observed series; N > 0 = predict the next N values",
)
run_button = gr.Button("Run TimeBraid", variant="primary")
with gr.Column(scale=5):
answer_box = gr.Textbox(
label="Model response",
lines=12,
show_copy_button=True,
)
plot_box = gr.Plot(label="Series and forecast")
with gr.Accordion("Advanced options", open=False):
system_box = gr.Textbox(
label="System prompt (optional)",
placeholder="You analyze weekly product demand.",
lines=2,
)
target_number = gr.Number(
label="Forecast target series (1-based)",
value=1,
precision=0,
info="Used only when forecasting from more than one input series.",
)
max_tokens_slider = gr.Slider(
minimum=64,
maximum=512,
value=256,
step=32,
precision=0,
label="Maximum new text tokens",
)
inputs = [
series_box,
prompt_box,
horizon_slider,
system_box,
target_number,
max_tokens_slider,
]
outputs = [answer_box, plot_box]
run_button.click(fn=run_timebraid, inputs=inputs, outputs=outputs)
prompt_box.submit(fn=run_timebraid, inputs=inputs, outputs=outputs)
gr.Examples(
examples=EXAMPLES,
inputs=[series_box, prompt_box, horizon_slider, system_box],
outputs=outputs,
fn=run_timebraid,
cache_examples=True,
cache_mode="lazy",
examples_per_page=3,
label="Example requests (from the TimeBraid repository and model card)",
)
gr.Markdown(
"All series and forecasts are plotted on their original scale. Forecasts "
"come from greedy decoding on a fresh prompt, exactly as in the released "
"inference recipe."
)
demo.launch(mcp_server=True)