如何将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
相关产品推荐
相关产品推荐

