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

ML.NET中FastForest训练器缺失Probability列的校准实现问询

ML.NET FastForest训练器缺失Probability列的修复方案

问题说明

使用FastTree训练器时代码可正常运行:

var trainer = mlContext.BinaryClassification.Trainers.FastTree();

但切换为FastForest训练器后:

var trainer = mlContext.BinaryClassification.Trainers.FastForest();

执行评估代码mlContext.BinaryClassification.Evaluate(predictions, labelColumnName: "Label")时,会抛出如下错误:

System.ArgumentOutOfRangeException: 'Probability column 'Probability' not found
Parameter name: schema'

这是因为FastForest训练器默认不会生成Probability列,而BinaryClassification.Evaluate方法需要该列计算完整评估指标,必须手动添加校准组件生成概率值。

解决方法

在训练管道中,于FastForest训练器之后添加Platt校准器(适用于大多数二分类场景),基于训练器输出的Score列生成Probability列。

修改后的完整代码

void trainModel()
{
    // Set up MLContext
    var mlContext = new MLContext();

    // Define database connection parameters
    string databaseName = "AAPL";
    string connectionString = "Server=localhost;Database=" + databaseName + ";Trusted_Connection=True;";
    var tableName = "AAPL_1day";

    // Get the column names from the database table
    var columnNames = GetColumnNamesFromDatabaseTable(connectionString, tableName);

    // Generate the list of column conditions for IS NOT NULL
    var columnConditions = columnNames.Select(column => $"[{column}] IS NOT NULL").ToList();

    // Create the SELECT query with the column conditions
    var query = $"SELECT * FROM {tableName} WHERE {string.Join(" AND ", columnConditions)}";

    // Create the database source with the dynamically generated query
    var databaseSource = new DatabaseSource(SqlClientFactory.Instance, connectionString, query);

    // Load the data from the database
    var loader = mlContext.Data.CreateDatabaseLoader<AAPLData>();
    var dataView = loader.Load(databaseSource);
    var preview = dataView.Preview();

    // Split the data into train and test sets
    var split = mlContext.Data.TrainTestSplit(dataView, testFraction: 0.35);

    // Define the data preprocessing pipeline
    var pipeline = mlContext.Transforms.Concatenate("Features", "MonthNr", "DayNr", "roc_1", "roc_2")
        .Append(mlContext.Transforms.Conversion.ConvertType("Features", "Features", DataKind.Single))
        .Append(mlContext.Transforms.NormalizeMinMax("Features"));

    // -------------------------- 修改部分 --------------------------
    // 替换为FastForest训练器,并添加Platt校准器生成Probability列
    var trainer = mlContext.BinaryClassification.Trainers.FastForest();
    var trainingPipeline = pipeline.Append(trainer)
        // 添加Platt校准器,基于Score列生成Probability列
        .Append(mlContext.BinaryClassification.Calibrators.PlattCalibrator(
            scoreColumnName: "Score",
            labelColumnName: "Label"));
    // -------------------------- 修改结束 --------------------------

    // Train the model
    var model = trainingPipeline.Fit(split.TrainSet);

    // Make predictions
    var predictions = model.Transform(split.TestSet);

    // Evaluate the model
    var metrics = mlContext.BinaryClassification.Evaluate(predictions, labelColumnName: "Label");

    // Retrieve the evaluation metrics
    var accuracy = metrics.Accuracy;
    var auc = metrics.AreaUnderRocCurve;

    MessageBox.Show($"Accuracy: {accuracy}\n" + $"AUC: {auc}");
}

// Define a class to hold your data
public class AAPLData
{
    [LoadColumn(2)] public double MonthNr;
    [LoadColumn(3)] public double DayNr;
    [LoadColumn(4)] public double roc_1;
    [LoadColumn(5)] public double roc_2;
    [LoadColumn(6), ColumnName("Label")] public bool hypo_upordown7days;
}

public static List<string> GetColumnNamesFromDatabaseTable(string connectionString, string tableName)
{
    var columnNames = new List<string>();
    using (var connection = new SqlConnection(connectionString))
    {
        connection.Open();

        // Get the schema information for the table
        var schemaTable = connection.GetSchema("Columns", new[] { null, null, tableName });
        foreach (DataRow row in schemaTable.Rows)
        {
            // Extract the column name from the schema table
            var columnName = row["COLUMN_NAME"].ToString();
            columnNames.Add(columnName);
        }
    }
    return columnNames;
}

可选优化

如果需要更精准的校准效果,可替换为IsotonicCalibrator(计算成本更高,适合大样本场景):

.Append(mlContext.BinaryClassification.Calibrators.IsotonicCalibrator(
    scoreColumnName: "Score",
    labelColumnName: "Label"))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 12:36:01