"""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)