【问题标题】:ML.NET how to make input model generic?ML.NET 如何使输入模型通用?
【发布时间】:2020-05-13 00:21:52
【问题描述】:

我有 3 个用于多类分类的用例,它们的 InputModel 都不同,因为它们具有不同的列和数据结构。如何重构下面的方法,以便它可以预测任何类型的 InputModel,而无需复制和重复该方法 3 次以满足 3 种不同的输入数据结构?

    private List<MulticlassClassificationPrediction> Predict(string modelName, string testDataPath)
    {
        PredictionEngine<InputModel, MulticlassClassificationPrediction> predEngine;

        predEngine = _predEnginePool.GetPredictionEngine(modelName: modelName);

        IDataView dataView = _mlContext.Data.LoadFromTextFile<InputModel>(
                            path: testDataPath,
                            hasHeader: true,
                            separatorChar: ',',
                            allowQuoting: true,
                            allowSparse: false);

        // Use first line of dataset as model input
        // You can replace this with new test data (hardcoded or from end-user application)
        List<InputModel> testDataList = _mlContext.Data.CreateEnumerable<InputModel>(dataView, false).ToList();

        List<MulticlassClassificationPrediction> predictionList = new List<MulticlassClassificationPrediction>();
        foreach (InputModel testData in testDataList)
        {

            MulticlassClassificationPrediction result = predEngine.Predict(testData);

            predictionList.Add(result);

        }

        return predictionList;
    }

【问题讨论】:

    标签: ml.net


    【解决方案1】:

    如果我理解你的问题是正确的,你有机会尝试这样的事情吗?

    private List<MulticlassClassificationPrediction> Predict<TInputModel>(string modelName, string testDataPath) where TInputModel: class, new()
    {
        PredictionEngine<TInputModel, MulticlassClassificationPrediction> predEngine;
    
        predEngine = _predEnginePool.GetPredictionEngine(modelName: modelName);
    
        IDataView dataView = _mlContext.Data.LoadFromTextFile<TInputModel>(
                            path: testDataPath,
                            hasHeader: true,
                            separatorChar: ',',
                            allowQuoting: true,
                            allowSparse: false);
    
        // Use first line of dataset as model input
        // You can replace this with new test data (hardcoded or from end-user application)
        var testDataList = _mlContext.Data.CreateEnumerable<TInputModel>(dataView, false).ToList();
    
        List<MulticlassClassificationPrediction> predictionList = new List<MulticlassClassificationPrediction>();
        foreach (var testData in testDataList)
        {
    
            MulticlassClassificationPrediction result = predEngine.Predict(testData);
    
            predictionList.Add(result);
    
        }
    
        return predictionList;
    }
    

    【讨论】:

    • 我做了,但 _mlContext.Data.CreateEnumerable 不接受导致编译时错误的通用对象
    • 我的代码中有错字,已修复。上面的例子现在对我来说在本地编译得很好。查看 CreateEnumerable 的源代码,它采用类 new() 的类型。我希望这会有所帮助
    猜你喜欢
    • 2020-09-02
    • 2023-01-25
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2021-05-12
    • 2022-07-23
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多