ML.NET多项式回归中CustomMapping如何设置运行时指定的泛型输入输出类型
泛型化自定义多项式映射实现方案
方案1:接口约束泛型参数(编译时类型安全,推荐)
先定义公共接口约束输入输出类的必要属性,保证转换逻辑可以正常访问所需字段:
// 输入类约束接口 public interface IPolynomialInput { float Level { get; } float Salary { get; } } // 输出类约束接口 public interface IPolynomialOutput { float[] Features { get; set; } float Salary { get; set; } }
将PolynomialFeatures改造为泛型类,添加接口约束:
[CustomMappingFactoryAttribute("Features")] public class PolynomialFeatures<TInput, TOutput> : CustomMappingFactory<TInput, TOutput> where TInput : class, IPolynomialInput, new() where TOutput : class, IPolynomialOutput, new() { private readonly int _degree; public PolynomialFeatures(int degree) { _degree = degree; } public void Transform(TInput input, TOutput output) { output.Features = Enumerable.Range(0, _degree + 1).Select(i => (float)(Math.Pow(input.Level, i))).ToArray(); output.Salary = input.Salary; } public override Action<TInput, TOutput> GetMapping() { return Transform; } }
调用方式和原逻辑基本一致,仅需在实例化时指定具体的输入输出类型即可:
// 实例化时指定你自定义的输入、输出类 var polyFeatures = new PolynomialFeatures<你的输入类, 你的输出类>(degree: 2); var dataProcessPipeline = mLContext.Transforms.CustomMapping(polyFeatures.GetMapping(), contractName: null,inputSchemaDefinition: schema) .Append(mLContext.Transforms.Concatenate("Features", new[] { "Features" })); var trainer = mLContext.Regression.Trainers.Ols(featureColumnName: "Features", labelColumnName: "Salary"); var trainingPipeline = dataProcessPipeline.Append(trainer);
方案2:反射实现完全运行时动态类型(无接口约束,灵活度更高)
如果需要完全不限制输入输出类的结构,可通过反射动态读写属性,无需提前定义接口,还可支持运行时指定要计算的特征列、标签列名称:
[CustomMappingFactoryAttribute("Features")] public class PolynomialFeatures<TInput, TOutput> : CustomMappingFactory<TInput, TOutput> where TInput : class, new() where TOutput : class, new() { private readonly int _degree; private readonly string _inputFeatureCol; private readonly string _labelCol; private readonly string _outputFeatureCol; public PolynomialFeatures(int degree, string inputFeatureCol = "Level", string labelCol = "Salary", string outputFeatureCol = "Features") { _degree = degree; _inputFeatureCol = inputFeatureCol; _labelCol = labelCol; _outputFeatureCol = outputFeatureCol; } public void Transform(TInput input, TOutput output) { // 反射读取输入属性 var inputType = typeof(TInput); var levelValue = (float)inputType.GetProperty(_inputFeatureCol).GetValue(input); var salaryValue = (float)inputType.GetProperty(_labelCol).GetValue(input); // 反射写入输出属性 var outputType = typeof(TOutput); var polyFeatures = Enumerable.Range(0, _degree + 1).Select(i => (float)(Math.Pow(levelValue, i))).ToArray(); outputType.GetProperty(_outputFeatureCol).SetValue(output, polyFeatures); outputType.GetProperty(_labelCol).SetValue(output, salaryValue); } public override Action<TInput, TOutput> GetMapping() { return Transform; } }
该方案无需输入输出类实现任何接口,仅要求类中存在对应名称的属性即可,适合对灵活度要求高的场景,缺点是反射会带来极小的性能损耗。
注意事项
ML.NET的自定义映射要求输入、输出类型必须为公共类,且用到的属性必须具备对应的get/set访问器,否则会出现Schema识别异常。
内容的提问来源于stack exchange,提问作者hussein shaib
相关产品推荐
相关产品推荐

