Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Tighten ZeroGPU duration to measured 30s; clearer empty-content message
2a3e13a verified Download app.py from hugging-apps/timebraid: direct link, hf CLI and curl.
- Browser
- Download file 18 kB
-
https://huggingface.co/spaces/hugging-apps/timebraid/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/timebraid/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/timebraid/resolve/main/app.py
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 | |
| # --------------------------------------------------------------------------- | |
| 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) | |