Team Ai
Apppublic

YuXT/18_Multi_cycle_inventory_prediction_algorithm

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py241 linesDownload Raw Back to root
1import time2 3import matplotlib.pyplot as plt4import numpy as np5import pandas as pd6from keras.layers.core import Dense, Activation, Dropout7from keras.layers import LSTM8from keras.models import Sequential9from sklearn.metrics import explained_variance_score, mean_absolute_error, mean_squared_error, median_absolute_error, \10    r2_score11import gradio as gr12 13examples = [14    '1天'15]16 17 18class Conf:19    # epochs20    EPOCHS = 5021    # 时间序列长度22    SEQ_LEN = 5023    # 预测步数24    PREDICT_STEP = 725    # 测试训练集比例26    TRAIN_DATA_RATE = 0.827    # 批大小28    BATCH_SIZE = 5029    # 网络形状30    LAYERS = [1, 50, 100, 1]31 32 33def model_evaluation(y_true, y_pred):34    metrics = _cal_metrics(y_true, y_pred)35    for (k, v) in metrics.items():36        print(k + ": " + str(v))37 38 39def _cal_metrics(y_true, y_pred):40    """41    计算各个指标的值42    """43    re = _calc_re(y_true, y_pred)44    metrics = {45        "explained_variance_score":46            explained_variance_score(y_true, y_pred),47        "mean_absolute_error":48            mean_absolute_error(y_true, y_pred),49        "mean_squared_error":50            mean_squared_error(y_true, y_pred),51        "median_absolute_error":52            median_absolute_error(y_true, y_pred),53        "r2_score":54            r2_score(y_true, y_pred),55        "sum_relative_error":56            re[0],57        "mean_relative_error":58            re[1]59    }60 61    return metrics62 63 64def _calc_re(y_true, y_pred):65    """66    计算相对误差(Sum/Mean Relative Error)67    """68    return [((y_true - y_pred) / y_pred).sum().values, ((y_true - y_pred) / y_pred).mean().values]69 70 71def load_data(filename):72    """73    数据准备74    """75    data = pd.read_csv(filename).values76 77    result_0 = []78    for index in range(len(data) - Conf.SEQ_LEN - 1):79        result_0.append(data[index: index + Conf.SEQ_LEN + 1])80    # 数据标准化81 82    result = normalise_windows(result_0)83 84    result = np.array(result)85    result_0 = np.array(result_0)86 87    row = round(result.shape[0] * Conf.TRAIN_DATA_RATE)88    train = result[:int(row), :]89    np.random.shuffle(train)90 91    _X_train = train[:, :-1]92    _y_train = train[:, -1]93    _X_test = result[int(row):, :-1]94    _y_test = result[int(row):, -1]95    _y_test_p_0 = result_0[int(row):, 0]96 97    # 增加一列98    _X_train = _X_train[:, :, np.newaxis]99    _X_test = _X_test[:, :, np.newaxis]100 101    print(_X_train.shape)102    print(_X_test.shape)103    return [_X_train, _y_train, _X_test, _y_test, _y_test_p_0]104 105 106def normalise_windows(window_data):107    """108    对原始数据做标准化:n_i = (p_i/p)0 - 1)109    对应的反标准化公式为:p_i = p_0(n_i + 1)110    """111    normalised_data = []112    for window in window_data:113        normalised_window = [((float(p) / float(window[0])) - 1) for p in window]114        normalised_data.append(normalised_window)115    return normalised_data116 117 118def anti_normalise_windows(p_data, normalised_data):119    """120    对原始数据做标准化:n_i = (p_i/p_0 - 1)121    对应的反标准化公式为:p_i = p_0(n_i + 1)122    """123    anti_normalised_data = []124    for (n, p_0) in zip(normalised_data, p_data):125        anti_normalised_window = p_0 * (n + 1)126        anti_normalised_data.append(anti_normalised_window)127    return anti_normalised_data128 129 130def build_model(layers):131    """132    模型定义133    """134    model = Sequential()135 136    model.add(LSTM(units=layers[1], input_shape=(layers[1], layers[0]), return_sequences=True))137    model.add(Dropout(0.2))138 139    model.add(LSTM(layers[2], return_sequences=False))140    model.add(Dropout(0.2))141 142    model.add(Dense(units=layers[3]))143    model.add(Activation("tanh"))144 145    start = time.time()146    model.compile(loss="mse", optimizer="rmsprop")147    print("> Compilation Time : ", time.time() - start)148    return model149 150 151def predict_point_by_point(model, data):152    """153    每次预测1步154    """155    predict = model.predict(data)156    predict = np.reshape(predict, (len(predict),))157    return predict158 159 160def predict_by_len(model, data, predict_len):161    """162    预测predict_len步163    """164    predicted = []165    for i in range(predict_len):166        predicted.append(model.predict(data[np.newaxis, :, :])[0, 0])167        data = data[1:]168        data = np.insert(data, -1, predicted[-1], axis=0)169    return predicted170 171 172def plot_results(y_true, y_pred):173    fig = plt.figure(facecolor='white')174    ax = fig.add_subplot(111)175    ax.plot(y_true, label='True Data')176    plt.plot(y_pred, label='Prediction')177    plt.legend()178    return fig179 180 181global_start_time = time.time()182 183print('> Loading data... ')184 185X_train, y_train, X_test, y_test, y_test_p_0 = load_data('price.csv')186 187print('> Data Loaded. Compiling...')188 189model = build_model(Conf.LAYERS)190 191model.fit(X_train, y_train, batch_size=Conf.BATCH_SIZE, epochs=Conf.EPOCHS, validation_split=0.05)192 193# 预测一步194predicted = predict_point_by_point(model, X_test)195 196# 预测7天197predicted_len = predict_by_len(model, X_test[-1], 7)198predicted_len_0 = anti_normalise_windows(y_test_p_0, predicted_len)199 200# 7天绘图201fig = plt.figure(facecolor='white')202ax = fig.add_subplot(111)203ax.plot(predicted_len_0, label='Prediction')204plt.legend()205 206print('Training duration (s) : ', time.time() - global_start_time)207 208y_test_0 = anti_normalise_windows(y_test_p_0, y_test)209 210# 预测一步211predicted_0 = anti_normalise_windows(y_test_p_0, predicted)212# 预测一步绘图213fig1 = plot_results(y_test_0, predicted_0)214 215# 模型评估216model_evaluation(pd.DataFrame(y_test_0), pd.DataFrame(predicted_0))217 218 219def process(choice):220    """221    整个后端的输入输出都在这里体现,也就是模型的推理部分。222    """223    if choice == "1天":224        output = fig1225    else:226        output = fig227    # 对推理结果做特定处理,方便展示228    return output  # 返回处理后的结果229 230 231iface = gr.Interface(  # 定义这个gr的接口232    fn=process,  # 实现输入输出的函数233    inputs=gr.Dropdown(["1天", "7天"], label="选择天数"),  # 输入本文框234    outputs=gr.Plot(label="Plot"),  # 输出结果以什么样的形式展示,随意替换235    # outputs=gr.Textbox(),236    title="多周期库存算法",237    examples=examples,  # gradio直接给设计了examples这个属性238    # allow_flagging="never",239)240iface.launch(debug=True, share=False)  # 写好对象直接launch就行241