ITADN
BayesWitnesses/m2cgen · 文件 下载 ZIP
文件最后提交记录最后更新时间
README.md
以下内容由 AI 翻译,如有问题请点此提交 issue 反馈

m2cgen

GitHub Actions Status Coverage Status License: MIT Python Versions PyPI Version Downloads

m2cgen (Model 2 Code Generator) - 是一个轻量级库,提供了一种将训练好的统计模型转译为原生代码(Python、C、Java、Go、JavaScript、Visual Basic、C#、PowerShell、R、PHP、Dart、Haskell、Ruby、F#、Rust、Elixir)的便捷方式。

安装

支持的 Python 版本为 >= 3.7

pip install m2cgen

开发

在提交 PR 之前,请确保以下命令能成功运行:

make pre-pr

或者,您可以运行该命令的 Docker 版本:

make docker-build docker-pre-pr

支持的语言

  • C
  • C#
  • Dart
  • F#
  • Go
  • Haskell
  • Java
  • JavaScript
  • PHP
  • PowerShell
  • Python
  • R
  • Ruby
  • Rust
  • Visual Basic (VBA-compatible)
  • Elixir

支持的模型

分类回归
线性
  • scikit-learn
    • LogisticRegression
    • LogisticRegressionCV
    • PassiveAggressiveClassifier
    • Perceptron
    • RidgeClassifier
    • RidgeClassifierCV
    • SGDClassifier
  • lightning
    • AdaGradClassifier
    • CDClassifier
    • FistaClassifier
    • SAGAClassifier
    • SAGClassifier
    • SDCAClassifier
    • SGDClassifier
  • scikit-learn
    • ARDRegression
    • BayesianRidge
    • ElasticNet
    • ElasticNetCV
    • GammaRegressor
    • HuberRegressor
    • Lars
    • LarsCV
    • Lasso
    • LassoCV
    • LassoLars
    • LassoLarsCV
    • LassoLarsIC
    • LinearRegression
    • OrthogonalMatchingPursuit
    • OrthogonalMatchingPursuitCV
    • PassiveAggressiveRegressor
    • PoissonRegressor
    • RANSACRegressor(only supported regression estimators can be used as a base estimator)
    • Ridge
    • RidgeCV
    • SGDRegressor
    • TheilSenRegressor
    • TweedieRegressor
  • StatsModels
    • Generalized Least Squares (GLS)
    • Generalized Least Squares with AR Errors (GLSAR)
    • Generalized Linear Models (GLM)
    • Ordinary Least Squares (OLS)
    • [Gaussian] Process Regression Using Maximum Likelihood-based Estimation (ProcessMLE)
    • Quantile Regression (QuantReg)
    • Weighted Least Squares (WLS)
  • lightning
    • AdaGradRegressor
    • CDRegressor
    • FistaRegressor
    • SAGARegressor
    • SAGRegressor
    • SDCARegressor
    • SGDRegressor
SVM
  • scikit-learn
    • LinearSVC
    • NuSVC
    • OneClassSVM
    • SVC
  • lightning
    • KernelSVC
    • LinearSVC
  • scikit-learn
    • LinearSVR
    • NuSVR
    • SVR
  • lightning
    • LinearSVR
Tree
  • DecisionTreeClassifier
  • ExtraTreeClassifier
  • DecisionTreeRegressor
  • ExtraTreeRegressor
Random Forest
  • ExtraTreesClassifier
  • LGBMClassifier(rf booster only)
  • RandomForestClassifier
  • XGBRFClassifier
  • ExtraTreesRegressor
  • LGBMRegressor(rf booster only)
  • RandomForestRegressor
  • XGBRFRegressor
Boosting
  • LGBMClassifier(gbdt/dart/goss booster only)
  • XGBClassifier(gbtree(including boosted forests)/gblinear booster only)
    • LGBMRegressor(gbdt/dart/goss booster only)
    • XGBRegressor(gbtree(including boosted forests)/gblinear booster only)

    您可以在此处]找到通过 CI 测试保证兼容性的软件包版本。 其他版本也可能受支持,但未经测试。

    分类输出

    线性 / 线性 SVM / 核 SVM

    二分类

    标量值;样本到第二类超平面的有符号距离。

    多分类

    向量值;样本到每个类别超平面的有符号距离。

    注释

    输出与 LinearClassifierMixin.decision_function 的输出一致。

    SVM

    异常检测

    标量值;样本到分离超平面的有符号距离:内部点为正,异常点为负。

    二分类

    标量值;样本到第二类超平面的有符号距离。

    多分类

    向量值;每个类别的一对一得分,形状为 (n_samples, n_classes * (n_classes-1) / 2)。

    注释

    decision_function_shape 设置为 ovo 时,输出与 BaseSVC.decision_function 的输出一致。

    树 / 随机森林 / 提升

    二分类

    向量值;类别概率。

    多分类

    向量值;类别概率。

    注释

    输出与 DecisionTreeClassifier / ExtraTreeClassifier / ExtraTreesClassifier / RandomForestClassifier / XGBRFClassifier / XGBClassifier / LGBMClassifierpredict_proba 方法的输出一致。

    用法

    以下是一个简单示例,展示了在 Python 环境中训练的线性模型如何在 Java 代码中表示:

    from sklearn.datasets import load_diabetes
    from sklearn import linear_model
    import m2cgen as m2c
    
    X, y = load_diabetes(return_X_y=True)
    
    estimator = linear_model.LinearRegression()
    estimator.fit(X, y)
    
    code = m2c.export_to_java(estimator)

    生成的 Java 代码:

    public class Model {
        public static double score(double[] input) {
            return ((((((((((152.1334841628965) + ((input[0]) * (-10.012197817470472))) + ((input[1]) * (-239.81908936565458))) + ((input[2]) * (519.8397867901342))) + ((input[3]) * (324.39042768937657))) + ((input[4]) * (-792.1841616283054))) + ((input[5]) * (476.74583782366153))) + ((input[6]) * (101.04457032134408))) + ((input[7]) * (177.06417623225025))) + ((input[8]) * (751.2793210873945))) + ((input[9]) * (67.62538639104406));
        }
    }

    您可以在此处]找到不同模型/语言的更多生成代码示例。

    CLI

    m2cgen 可以作为 CLI 工具,使用序列化模型对象(pickle 协议)生成代码:

    $ m2cgen <pickle_file> --language <language> [--indent <indent>] [--function_name <function_name>]
             [--class_name <class_name>] [--module_name <module_name>] [--package_name <package_name>]
             [--namespace <namespace>] [--recursion-limit <recursion_limit>]
    

    请记住,对于反序列化模型对象,它们的类必须在反序列化环境中可导入模块的顶层定义。

    也支持管道:

    $ cat <pickle_file> | m2cgen --language <language>
    

    FAQ

    Q: 生成时出现 RecursionError: maximum recursion depth exceeded 错误。

    A: 如果在使用集成模型生成代码时出现此错误,请尝试减少该模型中训练估计器的数量。或者,您可以使用 sys.setrecursionlimit(<new_depth>) 增加最大递归深度。

    Q: 在从序列化模型对象转译模型时,生成失败并出现 ImportError: No module named <module_name_here> 错误。

    A: 此错误表明 pickle 协议无法反序列化模型对象。对于反序列化模型对象,要求它们的类必须在反序列化环境中可导入模块的顶层定义。因此,安装提供模型类定义的包应该能解决此问题。

    Q: 由 m2cgen 生成的代码对于某些输入与原始 Python 模型(代码从中获取)的结果不同。

    A: 某些模型在其原生 Python 库的预测阶段会强制输入数据为特定类型。目前,m2cgen 仅支持 float64 (double) 数据类型。您可以尝试手动将输入数据转换为其他类型,然后再次检查结果。此外,由于目标语言中浮点运算的具体实现,可能会出现一些细微差异。