ML.NET多分类训练报错:Label列类型不匹配求助
ML.NET多分类训练标签类型不匹配错误排查
错误信息
System.ArgumentOutOfRangeException: 'Schema mismatch for label column 'Label': expected Key
, got Single Arg_ParamName_Name'
初始代码
using Microsoft.ML.Data; using System; using System.Collections.Generic; using System.Linq; using System.Text; using System.Threading.Tasks; namespace MLPredictor1 { internal class TicketInputDataModel { [LoadColumn(8)] public float state { get; set; } [LoadColumn(1)] public string? short_description { get; set; } [LoadColumn(2)] public string? description { get; set; } [LoadColumn(3)] public string? email { get; set; } [LoadColumn(4)] public float priority { get; set; } [LoadColumn(5)] public bool active { get; set; } [LoadColumn(6)] public DateTime opened_at { get; set; } [LoadColumn(7)] public float child_incidents { get; set; } [LoadColumn(0), ColumnName("Label")] public float num_of_days_com { get; set; } } } using Microsoft.ML.Data; using System; using System.Collections.Generic; using System.Linq; using System.Text; using System.Threading.Tasks; namespace MLPredictor1 { internal class TicketOutputDataModel { [ColumnName("PredictedLabel")] public float num_of_days_com { get; set; } } } // See https://aka.ms/new-console-template for more information using Microsoft.ML; using MLPredictor1; Console.WriteLine("Hello, World!"); List<TicketInputDataModel> ticketInputModel = new List<TicketInputDataModel>(); ticketInputModel.Add(new TicketInputDataModel() { active = true, child_incidents = 4, description = "Great to hear that this will be resolved soon", email = "john.doe@telefonicatech.uk", num_of_days_com = 5, opened_at = DateTime.Now, priority = 5, short_description = "will be resolved soon", state = 2 }); ticketInputModel.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); ticketInputModel.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); ticketInputModel.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); List<TicketInputDataModel> ticketInputModel2 = new List<TicketInputDataModel>(); ticketInputModel2.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); ticketInputModel2.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); ticketInputModel2.Add(new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", num_of_days_com = 3, opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }); MLContext mlContext = new MLContext(); IDataView? trainingData = mlContext.Data.LoadFromEnumerable<TicketInputDataModel>(ticketInputModel); var pipeline = mlContext.Transforms.Concatenate("Features", "state", "child_incidents", "priority"); var multiclassTrainer = pipeline.Append(mlContext.MulticlassClassification.Trainers .NaiveBayes(labelColumnName: "Label", featureColumnName: "Features")).Append(mlContext.Transforms.Categorical.OneHotEncoding(outputColumnName: "short_descriptionEncoded", inputColumnName: nameof(TicketInputDataModel.short_description))); ITransformer model = multiclassTrainer.Fit(trainingData); // --- ERROR var predictor = mlContext.Model.CreatePredictionEngine<TicketInputDataModel, TicketOutputDataModel>(model); var batchData = mlContext.Data.LoadFromEnumerable<TicketInputDataModel>(ticketInputModel2); IDataView predictions = model.Transform(batchData); ITransformer model2 = multiclassTrainer.Fit(predictions); var predictor2 = mlContext.Model.CreatePredictionEngine<TicketInputDataModel, TicketOutputDataModel>(model2); TicketInputDataModel tt = new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }; TicketOutputDataModel ttt = predictor2.Predict(tt);
更新后的代码
MLContext mlContext = new MLContext(); IDataView? trainingData = mlContext.Data.LoadFromEnumerable<TicketInputDataModel>(ticketInputModel); var multiclassTrainer = mlContext.Transforms.Concatenate("FeaturesText", new[] { "description", "short_description", "email", }) .Append(mlContext.Transforms.Text.FeaturizeText("Features", "FeaturesText")).Append(mlContext.MulticlassClassification.Trainers .NaiveBayes(labelColumnName: "Label", featureColumnName: "Features")); ITransformer model = multiclassTrainer.Fit(trainingData); // -- ERROR var predictor = mlContext.Model.CreatePredictionEngine<TicketInputDataModel, TicketOutputDataModel>(model); var batchData = mlContext.Data.LoadFromEnumerable<TicketInputDataModel>(ticketInputModel2); IDataView predictions = model.Transform(batchData); ITransformer model2 = multiclassTrainer.Fit(predictions); var predictor2 = mlContext.Model.CreatePredictionEngine<TicketInputDataModel, TicketOutputDataModel>(model2); TicketInputDataModel tt = new TicketInputDataModel() { active = true, child_incidents = 6, description = "This is taking a bit of time but making progress", email = "john.doe@telefonicatech.uk", opened_at = DateTime.Now, priority = 1, short_description = "progress being made", state = 3 }; TicketOutputDataModel ttt = predictor2.Predict(tt);
问题原因与解决方法
原因
ML.NET的多分类训练器要求标签列必须是Key类型(代表离散的类别标识),但当前代码中作为Label的num_of_days_com是float(Single)类型,属于连续数值类型,不符合多分类任务的标签格式要求,因此触发Schema不匹配错误。
解决方法
根据实际业务需求,分两种情况处理:
情况1:确实是多分类任务(预测离散的天数类别)
如果num_of_days_com是离散的类别(比如天数区间:1-3天、4-6天等),需要做以下修改:
- 修改输入模型的标签类型:将
num_of_days_com的类型改为uint或int,并添加[KeyType]属性标注类别数量(根据实际类别数调整参数):
internal class TicketInputDataModel { // 其他属性保持不变 [LoadColumn(0), ColumnName("Label"), KeyType(5)] // 假设共有5个类别,按需修改 public uint num_of_days_com { get; set; } }
- 修改输出模型:
PredictedLabel的类型要与输入标签一致,还可以添加Score属性获取各类别的概率:
internal class TicketOutputDataModel { [ColumnName("PredictedLabel")] public uint num_of_days_com { get; set; } [ColumnName("Score")] public float[]? CategoryProbabilities { get; set; } // 可选,获取每个类别的概率值 }
情况2:实际是回归任务(预测连续的天数数值)
如果num_of_days_com是连续的天数,你应该使用回归训练器而非多分类训练器,修改训练器部分代码:
// 替换多分类训练器为回归训练器 var regressionTrainer = mlContext.Transforms.Concatenate("FeaturesText", new[] { "description", "short_description", "email", }) .Append(mlContext.Transforms.Text.FeaturizeText("Features", "FeaturesText")) .Append(mlContext.Regression.Trainers.Sdca(labelColumnName: "Label", featureColumnName: "Features")); // 后续训练和预测逻辑保持一致,输出模型的PredictedLabel类型仍为float即可
内容的提问来源于stack exchange,提问作者redoc01
相关产品推荐
相关产品推荐

