kashif/probabilistic-forecast
3
1import gradio as gr2import pandas as pd3from gluonts.dataset.pandas import PandasDataset4from gluonts.dataset.split import split5from gluonts.torch.model.deepar import DeepAREstimator6import matplotlib7 8matplotlib.use("Agg")9import matplotlib.pyplot as plt10 11 12def fn(upload_data):13 df = pd.read_csv(upload_data.name, index_col=0, parse_dates=True)14 dataset = PandasDataset(df, target=df.columns[0])15 training_data, test_gen = split(dataset, offset=-36)16 17 model = DeepAREstimator(18 prediction_length=12,19 freq=dataset.freq,20 trainer_kwargs=dict(max_epochs=10),21 ).train(22 training_data=training_data,23 )24 25 test_data = test_gen.generate_instances(prediction_length=12, windows=3)26 forecasts = list(model.predict(test_data.input))27 28 fig = plt.figure()29 df["#Passengers"].plot(color="black")30 for forecast, color in zip(forecasts, ["green", "blue", "purple"]):31 forecast.plot(color=f"tab:{color}")32 plt.legend(["True values"], loc="upper left", fontsize="xx-large")33 return fig34 35 36with gr.Blocks() as demo:37 plot = gr.Plot()38 upload_btn = gr.UploadButton()39 40 upload_btn.upload(fn, inputs=upload_btn, outputs=plot)41 42if __name__ == "__main__":43 demo.launch()44 