#ifndef MEDIA_LEARNING_IMPL_ONE_HOT_H_
#define MEDIA_LEARNING_IMPL_ONE_HOT_H_
#include <map>
#include <memory>
#include <vector>
#include "base/component_export.h"
#include "media/learning/common/labelled_example.h"
#include "media/learning/common/learning_task.h"
#include "media/learning/common/value.h"
#include "media/learning/impl/model.h"
#include "third_party/abseil-cpp/absl/types/optional.h"
namespace media {
namespace learning {
class COMPONENT_EXPORT(LEARNING_IMPL) OneHotConverter {
public:
OneHotConverter(const LearningTask& task, const TrainingData& training_data);
OneHotConverter(const OneHotConverter&) = delete;
OneHotConverter& operator=(const OneHotConverter&) = delete;
~OneHotConverter();
const LearningTask& converted_task() const { return converted_task_; }
TrainingData Convert(const TrainingData& training_data) const;
FeatureVector Convert(const FeatureVector& feature_vector) const;
private:
void ProcessOneFeature(
size_t index,
const LearningTask::ValueDescription& original_description,
const TrainingData& training_data);
LearningTask converted_task_;
using ValueVectorIndexMap = std::map<Value, size_t>;
std::vector<absl::optional<ValueVectorIndexMap>> converters_;
};
class COMPONENT_EXPORT(LEARNING_IMPL) ConvertingModel : public Model {
public:
ConvertingModel(std::unique_ptr<OneHotConverter> converter,
std::unique_ptr<Model> model);
ConvertingModel(const ConvertingModel&) = delete;
ConvertingModel& operator=(const ConvertingModel&) = delete;
~ConvertingModel() override;
TargetHistogram PredictDistribution(const FeatureVector& instance) override;
private:
std::unique_ptr<OneHotConverter> converter_;
std::unique_ptr<Model> model_;
};
}
}
#endif