|
#include "query_classifier/AstFeatureVectorizer.hpp" |
|
#include "query_classifier/AstNode.hpp" |
|
#include "query_classifier/BitnessFeatureVectorizer.hpp" |
|
#include "query_classifier/FeatureVector.hpp" |
|
#include "query_classifier/MemoryPressureFeatureVectorizer.hpp" |
|
#include "query_classifier/OperatingSystem.hpp" |
|
#include "query_classifier/OperatingSystemFeatureVectorizer.hpp" |
|
#include "query_classifier/QueryContext.hpp" |
|
#include "query_classifier/QueryLabel.hpp" |
|
#include "query_classifier/WorkFeatureVectorizer.hpp" |
|
#include "query_classifier/LogisticRegression.hpp" |
|
|
|
#include <iostream> |
|
#include <vector> |
|
|
|
using query_classifier::AstNode; |
|
using query_classifier::AstNodeType; |
|
using query_classifier::FeatureVector; |
|
using query_classifier::AstFeatureVectorizer; |
|
using query_classifier::BitnessFeatureVectorizer; |
|
using query_classifier::LogisticRegression; |
|
using query_classifier::MemoryPressureFeatureVectorizer; |
|
using query_classifier::OperatingSystemFeatureVectorizer; |
|
using query_classifier::WorkFeatureVectorizer; |
|
using query_classifier::OperatingSystem; |
|
using query_classifier::QueryLabel; |
|
using query_classifier::QueryContext; |
|
|
|
int main() { |
|
const AstNode fast_query{AstNodeType::Aggregate, |
|
{{AstNodeType::Selector, {{AstNodeType::Metric, {}}}}}}; |
|
const AstNode slow_query{AstNodeType::Join, |
|
{{AstNodeType::RangeScan, |
|
{{AstNodeType::Selector, {{AstNodeType::Metric, {}}}}}}}}; |
|
struct Sample { |
|
AstNode query; |
|
QueryContext context; |
|
QueryLabel label; |
|
}; |
|
const std::vector<Sample> samples = { |
|
{fast_query, {OperatingSystem::Mac, false, true, 2.0, 10}, QueryLabel::Ok}, |
|
{fast_query, {OperatingSystem::Linux, false, true, 5.0, 20}, QueryLabel::Ok}, |
|
{slow_query, {OperatingSystem::Mac, true, true, 30.0, 100}, QueryLabel::Slow}, |
|
{slow_query, {OperatingSystem::Linux, true, true, 60.0, 200}, QueryLabel::Slow}, |
|
}; |
|
|
|
AstFeatureVectorizer ast_vectorizer; |
|
OperatingSystemFeatureVectorizer operating_system_vectorizer; |
|
BitnessFeatureVectorizer bitness_vectorizer; |
|
MemoryPressureFeatureVectorizer memory_pressure_vectorizer; |
|
WorkFeatureVectorizer work_vectorizer; |
|
std::vector<FeatureVector> training_features; |
|
std::vector<int> training_labels; |
|
for (const auto& sample : samples) { |
|
training_features.push_back( |
|
ast_vectorizer.vectorize(sample.query) + |
|
operating_system_vectorizer.vectorize(sample.context) + |
|
bitness_vectorizer.vectorize(sample.context) + |
|
memory_pressure_vectorizer.vectorize(sample.context) + |
|
work_vectorizer.vectorize(sample.context)); |
|
training_labels.push_back(static_cast<int>(sample.label)); |
|
} |
|
|
|
LogisticRegression model(0.1, 1000, 0.001); |
|
model.train(training_features, training_labels); |
|
|
|
// end training |
|
|
|
const QueryContext query_context{OperatingSystem::Linux, false, true, 8.0, 20}; |
|
const auto query_features = |
|
ast_vectorizer.vectorize(slow_query) + |
|
operating_system_vectorizer.vectorize(query_context) + |
|
bitness_vectorizer.vectorize(query_context) + |
|
memory_pressure_vectorizer.vectorize(query_context) + |
|
work_vectorizer.vectorize(query_context); |
|
std::cout << "slow probability: " |
|
<< model.predict_probability(query_features) << "\n"; |
|
std::cout << "prediction: " << (model.predict(query_features) ? "slow" : "ok") |
|
<< "\n"; |
|
return 0; |
|
} |