You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将Jupyter Notebook中训练的XGBoost模型保存并在前端JavaScript中加载调用,实现用户输入预测碳强度功能?

如何将Jupyter Notebook中训练的XGBoost模型保存并在前端JavaScript中加载调用,实现用户输入预测碳强度功能?

我来一步步帮你解决这个问题,你已经完成了模型训练和保存的核心步骤,接下来只需要把模型适配到JavaScript环境,再完成前端交互逻辑就行:


一、准备可在JS中加载的模型文件

你已经用xgb_model.save_model("xgb_model.json")保存了XGBoost的原生JSON模型,目前有两种主流方式让JS能识别并调用它:

方案1:直接用xgboost-js加载原生JSON模型

xgboost-js是官方推出的JavaScript版本XGBoost,能直接读取你保存的JSON模型,无需额外转换。

如果是前端项目,可通过npm安装依赖:

npm install xgboost-js

如果是纯HTML项目,也可以用CDN直接引入:

<script src="https://unpkg.com/xgboost-js@latest/dist/xgboost.min.js"></script>

方案2:转换为ONNX格式(兼容性更强)

如果需要更好的跨框架兼容性,或者后续要和其他AI工具联动,可以把XGB模型转成ONNX格式,用onnxruntime-web加载。在Python中执行以下代码完成转换:

import onnxmltools
from onnxmltools.convert.common.data_types import FloatTensorType
import xgboost as xgb

# 加载已保存的XGB模型
xgb_model = xgb.Booster()
xgb_model.load_model("xgb_model.json")

# 定义输入张量形状:对应单条22个特征的数据([样本数, 特征数])
initial_type = [('float_input', FloatTensorType([1, 22]))]
onnx_model = onnxmltools.convert.convert_xgboost(xgb_model, initial_types=initial_type)

# 保存ONNX模型文件
onnxmltools.utils.save_model(onnx_model, "xgb_model.onnx")

二、前端实现用户输入与预测功能

这里以纯HTML+JavaScript为例,实现一个简单的交互页面:用户输入22个特征值,点击按钮后返回预测的碳强度。

示例代码(使用xgboost-js)

<!DOCTYPE html>
<html>
<head>
    <meta charset="UTF-8">
    <title>碳强度预测工具</title>
    <!-- 引入xgboost-js -->
    <script src="https://unpkg.com/xgboost-js@latest/dist/xgboost.min.js"></script>
</head>
<body>
    <h3>请输入22个特征值(用逗号分隔):</h3>
    <input type="text" id="featureInput" placeholder="例如:242.86,19.78,4.59,..." style="width: 80%; padding: 8px;">
    <button onclick="predictCarbonIntensity()" style="margin-left: 10px; padding: 8px 16px;">预测碳强度</button>
    <p style="margin-top: 20px;">预测结果:<strong><span id="result">-</span></strong></p>

    <script>
        let xgbModel;

        // 页面加载时先异步加载模型
        async function loadModel() {
            try {
                // 加载本地的xgb_model.json(需放在HTML同目录或前端可访问的静态资源路径)
                xgbModel = await xgboost.loadModel('xgb_model.json');
                console.log("模型加载成功!");
            } catch (error) {
                console.error("模型加载失败:", error);
                alert("模型加载失败,请检查文件路径是否正确");
            }
        }

        // 核心预测函数
        async function predictCarbonIntensity() {
            if (!xgbModel) {
                alert("模型还在加载中,请稍等几秒!");
                return;
            }

            // 处理用户输入
            const inputStr = document.getElementById('featureInput').value.trim();
            const features = inputStr.split(',').map(val => parseFloat(val));
            
            // 校验输入有效性
            if (features.length !== 22 || features.some(isNaN)) {
                alert("请输入22个有效的数值,用英文逗号分隔!");
                return;
            }

            // 转换为XGBoost JS要求的输入格式:二维数组([样本数, 特征数])
            const inputData = [features];
            // 执行预测
            const prediction = await xgbModel.predict(inputData);
            // 显示结果(保留4位小数)
            document.getElementById('result').textContent = prediction[0].toFixed(4);
        }

        // 页面加载完成后自动加载模型
        window.onload = loadModel;
    </script>
</body>
</html>

示例代码(使用ONNX格式)

如果选择ONNX方案,先引入onnxruntime-web:

<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@latest/dist/ort.min.js"></script>

然后修改加载和预测的JS逻辑:

let ortSession;

async function loadModel() {
    try {
        ortSession = await ort.InferenceSession.create('xgb_model.onnx');
        console.log("ONNX模型加载成功!");
    } catch (error) {
        console.error("模型加载失败:", error);
        alert("模型加载失败,请检查文件路径是否正确");
    }
}

async function predictCarbonIntensity() {
    if (!ortSession) {
        alert("模型还在加载中,请稍等几秒!");
        return;
    }

    const inputStr = document.getElementById('featureInput').value.trim();
    const features = inputStr.split(',').map(val => parseFloat(val));
    if (features.length !== 22 || features.some(isNaN)) {
        alert("请输入22个有效的数值,用英文逗号分隔!");
        return;
    }

    // 转换为ONNX要求的输入格式:Float32Array,形状为[1,22]
    const inputTensor = new ort.Tensor('float32', features, [1, 22]);
    const feeds = { float_input: inputTensor }; // 键名需和转换时定义的initial_type名称一致

    const results = await ortSession.run(feeds);
    // 提取预测结果并显示
    document.getElementById('result').textContent = results.sequential_1[0].toFixed(4);
}

window.onload = loadModel;

三、关键注意事项

  • 特征顺序必须严格匹配:用户输入的22个特征顺序,必须和你训练时sorted_columns的顺序完全一致(即boiler_features + turbine_features + power_features + coal_features + carbon_emission_features去掉Carbon Intensity和Carbon Emission后的顺序),否则预测结果会完全错误。
  • 模型文件路径:确保模型文件(xgb_model.json或xgb_model.onnx)放在前端页面能访问的路径下,比如和HTML文件同目录,或者项目的静态资源文件夹中。
  • 数据类型校验:必须将用户输入的字符串转换为浮点型数组,避免因类型不匹配导致预测失败。

备注:内容来源于stack exchange,提问作者Ryan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 17:55:28