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

