.NET ML项目中如何运行时动态生成Feature类?
解决ML.NET动态特征类的方案
核心结论
ML.NET的PredictionEngine等组件依赖带有ColumnName和LoadColumn特性的静态类型完成数据映射,因此必须保留这些特性。针对动态特征数量的场景,可行方案是运行时动态生成符合要求的静态类,再通过反射调用ML.NET的泛型API。
方案一:用Reflection.Emit动态生成类
这种方式轻量、性能接近静态类,适合大量特征的场景。
1. 动态生成Feature类的代码
using System.Reflection; using System.Reflection.Emit; using Microsoft.ML.Data; public static Type GenerateDynamicFeatureType(IEnumerable<string> featureColumnNames, string labelColumnName) { var assemblyName = new AssemblyName("DynamicMLFeatures"); var assemblyBuilder = AssemblyBuilder.DefineDynamicAssembly(assemblyName, AssemblyBuilderAccess.Run); var moduleBuilder = assemblyBuilder.DefineDynamicModule("DynamicMLFeaturesModule"); var typeBuilder = moduleBuilder.DefineType("DynamicFeature", TypeAttributes.Public | TypeAttributes.Class); // 定义Label属性及特性 var columnNameAttrCtor = typeof(ColumnNameAttribute).GetConstructor(new[] { typeof(string) }); var labelProp = typeBuilder.DefineProperty(labelColumnName, PropertyAttributes.None, typeof(float), null); var labelField = typeBuilder.DefineField($"_{labelColumnName}", typeof(float), FieldAttributes.Private); var labelColAttr = new CustomAttributeBuilder(columnNameAttrCtor, new object[] { labelColumnName }); labelProp.SetCustomAttribute(labelColAttr); // Label的Getter/Setter var labelGet = typeBuilder.DefineMethod($"get_{labelColumnName}", MethodAttributes.Public | MethodAttributes.SpecialName | MethodAttributes.HideBySig, typeof(float), Type.EmptyTypes); var labelGetIl = labelGet.GetILGenerator(); labelGetIl.Emit(OpCodes.Ldarg_0); labelGetIl.Emit(OpCodes.Ldfld, labelField); labelGetIl.Emit(OpCodes.Ret); var labelSet = typeBuilder.DefineMethod($"set_{labelColumnName}", MethodAttributes.Public | MethodAttributes.SpecialName | MethodAttributes.HideBySig, typeof(void), new[] { typeof(float) }); var labelSetIl = labelSet.GetILGenerator(); labelSetIl.Emit(OpCodes.Ldarg_0); labelSetIl.Emit(OpCodes.Ldarg_1); labelSetIl.Emit(OpCodes.Stfld, labelField); labelSetIl.Emit(OpCodes.Ret); labelProp.SetGetMethod(labelGet); labelProp.SetSetMethod(labelSet); // 定义特征属性及特性 int columnIndex = 0; foreach (var colName in featureColumnNames) { if (colName.Equals(labelColumnName, StringComparison.OrdinalIgnoreCase)) continue; var prop = typeBuilder.DefineProperty(colName, PropertyAttributes.None, typeof(float), null); var field = typeBuilder.DefineField($"_{colName}", typeof(float), FieldAttributes.Private); // 添加ColumnName和LoadColumn特性 var loadColAttrCtor = typeof(LoadColumnAttribute).GetConstructor(new[] { typeof(int) }); var loadAttr = new CustomAttributeBuilder(loadColAttrCtor, new object[] { columnIndex }); prop.SetCustomAttribute(loadAttr); var colAttr = new CustomAttributeBuilder(columnNameAttrCtor, new object[] { colName }); prop.SetCustomAttribute(colAttr); // 特征属性的Getter/Setter var getMethod = typeBuilder.DefineMethod($"get_{colName}", MethodAttributes.Public | MethodAttributes.SpecialName | MethodAttributes.HideBySig, typeof(float), Type.EmptyTypes); var getIl = getMethod.GetILGenerator(); getIl.Emit(OpCodes.Ldarg_0); getIl.Emit(OpCodes.Ldfld, field); getIl.Emit(OpCodes.Ret); var setMethod = typeBuilder.DefineMethod($"set_{colName}", MethodAttributes.Public | MethodAttributes.SpecialName | MethodAttributes.HideBySig, typeof(void), new[] { typeof(float) }); var setIl = setMethod.GetILGenerator(); setIl.Emit(OpCodes.Ldarg_0); setIl.Emit(OpCodes.Ldarg_1); setIl.Emit(OpCodes.Stfld, field); setIl.Emit(OpCodes.Ret); prop.SetGetMethod(getMethod); prop.SetSetMethod(setMethod); columnIndex++; } return typeBuilder.CreateType()!; }
2. 使用动态类调用ML.NET API
// 读取CSV并获取列名 var mlContext = new MLContext(); var dataView = mlContext.Data.LoadFromTextFile("data.csv", separatorChar: ',', hasHeader: true); var columnNames = dataView.Schema.Select(col => col.Name).ToList(); var labelColumnName = "Label"; var featureColumns = columnNames.Where(n => !n.Equals(labelColumnName, StringComparison.OrdinalIgnoreCase)).ToList(); // 生成动态Feature类型 var dynamicFeatureType = GenerateDynamicFeatureType(featureColumns, labelColumnName); // 反射调用CreatePredictionEngine(假设Prediction类是硬编码的) var predictionType = typeof(Prediction); var createPredEngineMethod = typeof(ModelOperationsCatalog) .GetMethod(nameof(ModelOperationsCatalog.CreatePredictionEngine))! .MakeGenericMethod(dynamicFeatureType, predictionType); var predictionEngine = createPredEngineMethod.Invoke(mlContext.Model, new[] { trainedModel })!;
3. 创建实例并访问属性
// 创建动态Feature实例 var featureInstance = Activator.CreateInstance(dynamicFeatureType)!; // 设置属性值(通过反射) var labelProp = dynamicFeatureType.GetProperty(labelColumnName)!; labelProp.SetValue(featureInstance, 1.0f); foreach (var colName in featureColumns) { var prop = dynamicFeatureType.GetProperty(colName)!; prop.SetValue(featureInstance, 0.5f); // 替换为实际数据值 } // 执行预测 var predictMethod = predictionEngine.GetType().GetMethod("Predict")!; var prediction = predictMethod.Invoke(predictionEngine, new[] { featureInstance })!; // 获取预测结果 var score = predictionType.GetProperty("Score")!.GetValue(prediction);
方案二:用Roslyn编译动态生成的源代码
如果需要更复杂的类结构(比如不同数据类型的特征),可以用Roslyn动态编译C#代码:
using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; using System.IO; using System.Linq; using System.Reflection; using Microsoft.ML.Data; public static Type GenerateFeatureTypeWithRoslyn(IEnumerable<string> featureColumnNames, string labelColumnName) { // 生成类的源代码 var source = $@" using Microsoft.ML.Data; public class DynamicFeature {{ [ColumnName(""{labelColumnName}"")] public float {labelColumnName} {{ get; set; }} "; foreach (var colName in featureColumnNames) { if (colName.Equals(labelColumnName, StringComparison.OrdinalIgnoreCase)) continue; source += $@" [ColumnName(""{colName}"")] [LoadColumn({featureColumnNames.ToList().IndexOf(colName)})] public float {colName} {{ get; set; }} "; } source += "}"; // 编译源代码 var syntaxTree = CSharpSyntaxTree.ParseText(source); var references = AppDomain.CurrentDomain.GetAssemblies() .Where(a => !a.IsDynamic) .Select(a => MetadataReference.CreateFromFile(a.Location)) .Append(MetadataReference.CreateFromFile(typeof(ColumnNameAttribute).Assembly.Location)); var compilation = CSharpCompilation.Create("DynamicMLFeatures") .WithOptions(new CSharpCompilationOptions(OutputKind.DynamicallyLinkedLibrary)) .AddReferences(references) .AddSyntaxTrees(syntaxTree); using var ms = new MemoryStream(); var result = compilation.Emit(ms); if (!result.Success) { var errors = result.Diagnostics.Where(d => d.Severity == DiagnosticSeverity.Error); throw new InvalidOperationException($"编译失败: {string.Join("\n", errors.Select(e => e.GetMessage()))}"); } ms.Seek(0, SeekOrigin.Begin); var assembly = Assembly.Load(ms.ToArray()); return assembly.GetType("DynamicFeature")!; }
关键注意事项
- 特性必须保留:
ColumnName用于关联CSV列名和类属性,LoadColumn用于指定列索引,ML.NET的核心组件依赖这些特性完成数据绑定。 - 数据类型匹配:如果CSV中特征的类型不是
float,需要从IDataView.Schema中获取列的实际类型,调整动态类的属性类型(比如double、int)。 - 泛型方法调用:由于动态类型无法在编译时作为泛型参数,必须通过反射的
MakeGenericMethod来调用CreatePredictionEngine等泛型API。
内容的提问来源于stack exchange,提问作者Thomas Bont
相关产品推荐
相关产品推荐

