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

时间序列预测中XGBoost类别变量的序列化与推理问题

解决XGBoost类别变量在MLflow序列化与实时推理中的类型丢失问题

针对你遇到的两个核心问题——MLflow日志模型时类别型特征被转回object类型、实时推理端点因类别类型不匹配触发错误,即使已设置enable_categorical=True,可以按以下步骤解决:

1. 训练阶段确保类别特征与模型配置正确

首先要保证训练数据的类别列是pandas Categorical类型,而非默认的object,同时XGBoost模型明确启用类别支持:

转换类别列类型

import pandas as pd

# 显式转换为Categorical,保留训练时的类别集合
df["MACHINE_SERIAL_NUMBER"] = df["MACHINE_SERIAL_NUMBER"].astype("category")
df["RECIPE"] = df["RECIPE"].astype("category")

初始化XGBoost模型的正确参数

必须设置enable_categorical=True,且指定tree_method="hist"或"gpu_hist"(XGBoost的类别特征支持仅对这两种树方法生效):

import xgboost as xgb

model = xgb.XGBRegressor(
    enable_categorical=True,
    tree_method="hist",  # 关键:无此参数类别支持不生效
    objective="reg:squarederror",
    # 其他模型参数...
)

# 注意:X必须是pandas DataFrame,不能是numpy数组(否则丢失类别元数据)
model.fit(X_train, y_train)

2. MLflow日志模型时保留类别元数据

直接使用mlflow.xgboost.log_model时,需显式传递enable_categorical=True,并创建清晰的模型签名,避免MLflow自动推断类型时出错:

创建模型签名

手动定义输入输出的 schema,确保类别列被识别为字符串(因为JSON输入以字符串传递):

from mlflow.models.signature import ModelSignature
from mlflow.types.schema import Schema, ColSpec

# 定义输入schema:类别列用string类型,对应JSON输入的字符串格式
input_schema = Schema([
    ColSpec("string", "MACHINE_SERIAL_NUMBER"),
    ColSpec("string", "RECIPE"),
    ColSpec("datetime", "Date"),
])
output_schema = Schema([ColSpec("double", "Volume")])
signature = ModelSignature(inputs=input_schema, outputs=output_schema)

日志模型

import mlflow

mlflow.xgboost.log_model(
    model,
    artifact_path="xgb_time_series_model",
    signature=signature,
    input_example=X_train.head(),  # 提供示例输入,帮助MLflow识别类型
    enable_categorical=True  # 显式告知MLflow启用类别支持
)

3. 实时推理阶段处理类别转换

因为JSON输入的类别值是字符串,直接传入模型会被识别为object类型,需自定义PyFunc模型,在推理前将字符串映射回训练时的Categorical类型:

自定义PyFunc模型

import mlflow.pyfunc

class XGBoostCatInferModel(mlflow.pyfunc.PythonModel):
    def load_context(self, context):
        import xgboost as xgb
        # 加载训练好的XGBoost模型
        self.model = xgb.Booster()
        self.model.load_model(context.artifacts["model_path"])
        # 加载训练时保存的类别映射
        self.category_maps = context.artifacts["category_maps"]
    
    def predict(self, context, model_input):
        # 将输入字符串转换为训练时的Categorical类型
        for col, categories in self.category_maps.items():
            model_input[col] = pd.Categorical(
                model_input[col],
                categories=categories,
                ordered=False
            )
        # 创建DMatrix并启用类别支持
        dmatrix = xgb.DMatrix(model_input, enable_categorical=True)
        return self.model.predict(dmatrix)

保存类别映射并日志自定义模型

# 保存训练时的类别集合
category_maps = {
    "MACHINE_SERIAL_NUMBER": df["MACHINE_SERIAL_NUMBER"].cat.categories.tolist(),
    "RECIPE": df["RECIPE"].cat.categories.tolist()
}
mlflow.log_dict(category_maps, "category_maps.json")

# 日志自定义PyFunc模型
mlflow.pyfunc.log_model(
    artifact_path="xgb_cat_infer_model",
    python_model=XGBoostCatInferModel(),
    artifacts={
        "model_path": mlflow.xgboost.get_model_uri("xgb_time_series_model"),
        "category_maps": "category_maps.json"
    },
    signature=signature,
    input_example=X_train.head()
)

4. 额外注意事项

  • 升级依赖:确保使用最新版本的MLflow、XGBoost和pandas,旧版本对类别特征的支持存在bug:
    pip install --upgrade mlflow xgboost pandas
    
  • 避免使用numpy数组作为训练输入:XGBoost仅能从pandas DataFrame中读取类别元数据,numpy数组会丢失该信息。

内容的提问来源于stack exchange,提问作者Raju Natra

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 03:06:04