File size: 2,438 Bytes
e317359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
"""Exercise the actual Gradio model-selection callbacks and serialized plots."""
from __future__ import annotations

import json


def check_forecast_ui(client, config: dict) -> dict:
    dependencies = config["dependencies"]
    refresh = next(d for d in dependencies if d.get("api_name") == "refresh_forecasts")
    outputs = client.predict(*([None] * len(refresh["inputs"])), api_name="/refresh_forecasts")
    assert len(outputs) == 3 * len(refresh["inputs"]), "Forecast refresh output mismatch"
    selections = {d["inputs"][0]: d["api_name"] for d in dependencies
                  if str(d.get("api_name", "")).startswith("forecast_")}
    checked = []
    for index, component_id in enumerate(refresh["inputs"]):
        update = outputs[3 * index]
        choices = update.get("choices", [])
        labels = [v[1] if isinstance(v, (list, tuple)) else v for v in choices]
        if not labels:
            assert "No forecast snapshot" in outputs[3 * index + 2]
            continue
        assert update["interactive"], "Available model selector is disabled"
        api_name = selections[component_id]
        for label in dict.fromkeys([labels[0], labels[-1]]):
            payload, status = client.predict(label, api_name="/" + api_name)
            assert payload and payload["type"] == "plotly", f"Missing plot: {api_name}/{label}"
            figure = json.loads(payload["plot"])
            traces = {trace["name"]: trace for trace in figure["data"]}
            assert "Forecast (p50)" in traces and "Ground truth (actual)" in traces
            prediction = traces["Forecast (p50)"]
            truth = traces["Ground truth (actual)"]
            assert len(prediction["x"]) == len(prediction["y"]) > 1
            assert prediction["x"] == truth["x"]
            assert label in figure["layout"]["title"]["text"]
            assert "Ground truth:" in status
            # The first point anchors the line to history, not a future observation.
            checked.append({"dataset": api_name.removeprefix("forecast_"), "model": label,
                            "models_available": len(labels),
                            "target_points": len(truth["y"]) - 1,
                            "observed_targets": sum(v is not None for v in truth["y"][1:])})
    assert checked, "No usable forecast snapshots were published"
    return {"datasets_checked": len({r["dataset"] for r in checked}), "model_switches": checked}