Skip to content

Instantly share code, notes, and snippets.

@trevershick
Created September 4, 2026 02:06
Show Gist options
  • Select an option

  • Save trevershick/8fd180be8ab42a45a6b96141bfe15a90 to your computer and use it in GitHub Desktop.

Select an option

Save trevershick/8fd180be8ab42a45a6b96141bfe15a90 to your computer and use it in GitHub Desktop.
C++17 query AST logistic-regression classifier with composable feature vectorizers

Query Classifier

A small C++17 example for classifying query ASTs as ok or slow with logistic regression.

Build and test

cmake -S . -B build
cmake --build build
ctest --test-dir build --output-on-failure
./build/query-classifier

Feature design

AstNode uses the numeric AstNodeType enum instead of strings. AstFeatureVectorizer produces a fixed-size hashed AST vector from those values:

  • node=<kind> counts AST node kinds.
  • path=<ancestor>...<node> counts bounded root-to-node paths. The default maximum path length is four, which limits vocabulary growth.

Context features are produced independently by OperatingSystemFeatureVectorizer, BitnessFeatureVectorizer, MemoryPressureFeatureVectorizer, and WorkFeatureVectorizer. The latter produces log1p-scaled duration, steps, and duration-times-steps values. Compose these four vectors with the AST vector using operator+ before training or inference. Each vectorizer has a fixed size and requires no fitting.

The example uses labels 0 = ok and 1 = slow. In production, choose the threshold from an explicit latency SLO or a validation-set metric rather than assuming 0.5. A CSV loader and a real query parser can be added at the boundary where your database produces ASTs and measurements.

#include "query_classifier/AstFeatureVectorizer.hpp"
#include <functional>
#include <stdexcept>
namespace query_classifier {
AstFeatureVectorizer::AstFeatureVectorizer(std::size_t max_path_length,
std::size_t bucket_count)
: max_path_length_(max_path_length), bucket_count_(bucket_count) {
if (max_path_length_ == 0) {
throw std::invalid_argument("max_path_length must be greater than zero");
}
if (bucket_count_ == 0) {
throw std::invalid_argument("bucket_count must be greater than zero");
}
names_.reserve(bucket_count_);
for (std::size_t i = 0; i < bucket_count_; ++i) {
names_.push_back("hash_bucket=" + std::to_string(i));
}
}
FeatureVector AstFeatureVectorizer::vectorize(const AstNode& root) const {
FeatureVector result(bucket_count_, 0.0);
collect_features(root, result);
return result;
}
std::size_t AstFeatureVectorizer::feature_count() const { return names_.size(); }
const std::vector<std::string>& AstFeatureVectorizer::feature_names() const {
return names_;
}
void AstFeatureVectorizer::collect_features(const AstNode& root,
FeatureVector& features) const {
std::function<void(const AstNode&, std::vector<std::uint64_t>&)> visit =
[&](const AstNode& node, std::vector<std::uint64_t>& path) {
const auto node_type = static_cast<std::uint64_t>(node.type);
path.push_back(node_type);
add_hashed_feature(features, combine_hash(0x4e4f4445ULL, node_type));
std::uint64_t path_hash = 0x50415448ULL;
const auto first = path.size() > max_path_length_
? path.size() - max_path_length_
: 0;
for (std::size_t i = first; i < path.size(); ++i) {
path_hash = combine_hash(path_hash, path[i]);
}
add_hashed_feature(features, path_hash);
for (const auto& child : node.children) {
visit(child, path);
}
path.pop_back();
};
std::vector<std::uint64_t> path;
visit(root, path);
}
void AstFeatureVectorizer::add_hashed_feature(FeatureVector& features,
std::uint64_t hash) const {
features[static_cast<std::size_t>(hash % bucket_count_)] += 1.0;
}
std::uint64_t AstFeatureVectorizer::combine_hash(std::uint64_t seed,
std::uint64_t value) {
return seed ^ (value + 0x9e3779b97f4a7c15ULL + (seed << 6) + (seed >> 2));
}
} // namespace query_classifier
#pragma once
#include "query_classifier/AstNode.hpp"
#include "query_classifier/FeatureVector.hpp"
#include <cstddef>
#include <cstdint>
#include <string>
#include <vector>
namespace query_classifier {
class AstFeatureVectorizer {
public:
explicit AstFeatureVectorizer(std::size_t max_path_length = 4,
std::size_t bucket_count = 256);
FeatureVector vectorize(const AstNode& root) const;
std::size_t feature_count() const;
const std::vector<std::string>& feature_names() const;
private:
void collect_features(const AstNode& node, FeatureVector& features) const;
void add_hashed_feature(FeatureVector& features, std::uint64_t hash) const;
static std::uint64_t combine_hash(std::uint64_t seed, std::uint64_t value);
std::size_t max_path_length_;
std::size_t bucket_count_;
std::vector<std::string> names_;
};
} // namespace query_classifier
#pragma once
#include "query_classifier/AstNodeType.hpp"
#include <vector>
namespace query_classifier {
struct AstNode {
AstNodeType type;
std::vector<AstNode> children;
};
} // namespace query_classifier
#pragma once
#include <cstdint>
namespace query_classifier {
enum class AstNodeType : std::uint16_t {
Root,
Aggregate,
Selector,
Metric,
Join,
RangeScan,
Filter
};
} // namespace query_classifier
#include "query_classifier/AstFeatureVectorizer.hpp"
#include "query_classifier/AstNode.hpp"
#include "query_classifier/BitnessFeatureVectorizer.hpp"
#include "query_classifier/MemoryPressureFeatureVectorizer.hpp"
#include "query_classifier/OperatingSystemFeatureVectorizer.hpp"
#include "query_classifier/QueryContext.hpp"
#include "query_classifier/WorkFeatureVectorizer.hpp"
#include "query_classifier/LogisticRegression.hpp"
#include <algorithm>
#include <chrono>
#include <cstddef>
#include <iostream>
#include <numeric>
#include <vector>
namespace {
double percentile(const std::vector<double>& sorted_samples, double percent) {
const double rank = percent * static_cast<double>(sorted_samples.size() - 1);
const auto lower = static_cast<std::size_t>(rank);
const auto upper = std::min(lower + 1, sorted_samples.size() - 1);
const double fraction = rank - static_cast<double>(lower);
return sorted_samples[lower] +
fraction * (sorted_samples[upper] - sorted_samples[lower]);
}
} // namespace
int main() {
using namespace query_classifier;
const AstNode query{AstNodeType::Join,
{{AstNodeType::RangeScan,
{{AstNodeType::Selector, {{AstNodeType::Metric, {}}}}}}}};
const QueryContext context{OperatingSystem::Linux, true, true, 30.0, 100};
AstFeatureVectorizer ast_vectorizer;
OperatingSystemFeatureVectorizer operating_system_vectorizer;
BitnessFeatureVectorizer bitness_vectorizer;
MemoryPressureFeatureVectorizer memory_pressure_vectorizer;
WorkFeatureVectorizer work_vectorizer;
const FeatureVector template_features =
ast_vectorizer.vectorize(query) +
operating_system_vectorizer.vectorize(context) +
bitness_vectorizer.vectorize(context) +
memory_pressure_vectorizer.vectorize(context) +
work_vectorizer.vectorize(context);
LogisticRegression model;
model.load(FeatureVector(template_features.size(), 0.01), 0.0);
constexpr std::size_t warmup_iterations = 10'000;
constexpr std::size_t iterations = 1'000'000;
volatile int sink = 0;
for (std::size_t i = 0; i < warmup_iterations; ++i) {
sink += model.predict(template_features);
}
std::vector<double> samples;
samples.reserve(iterations);
for (std::size_t i = 0; i < iterations; ++i) {
QueryContext query_context = context;
query_context.steps = 1 + (i % 1000);
const auto start = std::chrono::steady_clock::now();
const FeatureVector query_features =
ast_vectorizer.vectorize(query) +
operating_system_vectorizer.vectorize(query_context) +
bitness_vectorizer.vectorize(query_context) +
memory_pressure_vectorizer.vectorize(query_context) +
work_vectorizer.vectorize(query_context);
sink += model.predict(query_features);
const auto elapsed = std::chrono::steady_clock::now() - start;
samples.push_back(static_cast<double>(
std::chrono::duration_cast<std::chrono::nanoseconds>(elapsed).count()));
}
std::sort(samples.begin(), samples.end());
const double total_nanoseconds =
std::accumulate(samples.begin(), samples.end(), 0.0);
std::cout << "iterations: " << iterations << "\n"
<< "mean nanoseconds/query: " << total_nanoseconds / iterations << "\n"
<< "median nanoseconds/query: " << percentile(samples, 0.50) << "\n"
<< "p90 nanoseconds/query: " << percentile(samples, 0.90) << "\n"
<< "p95 nanoseconds/query: " << percentile(samples, 0.95) << "\n"
<< "p99 nanoseconds/query: " << percentile(samples, 0.99) << "\n"
<< "min nanoseconds/query: " << samples.front() << "\n"
<< "max nanoseconds/query: " << samples.back() << "\n"
<< "sink: " << sink << "\n";
}
#include "query_classifier/BitnessFeatureVectorizer.hpp"
namespace query_classifier {
FeatureVector BitnessFeatureVectorizer::vectorize(
const QueryContext& context) const {
FeatureVector result(feature_count(), 0.0);
result[context.is_64_bit ? 1 : 0] = 1.0;
return result;
}
std::size_t BitnessFeatureVectorizer::feature_count() const {
return 2;
}
} // namespace query_classifier
#pragma once
#include "query_classifier/FeatureVector.hpp"
#include "query_classifier/QueryContext.hpp"
namespace query_classifier {
class BitnessFeatureVectorizer {
public:
FeatureVector vectorize(const QueryContext& context) const;
std::size_t feature_count() const;
};
} // namespace query_classifier
cmake_minimum_required(VERSION 3.16)
project(query_classifier LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CXX_EXTENSIONS OFF)
add_library(query_classifier_lib
src/AstFeatureVectorizer.cpp
src/OperatingSystemFeatureVectorizer.cpp
src/BitnessFeatureVectorizer.cpp
src/MemoryPressureFeatureVectorizer.cpp
src/WorkFeatureVectorizer.cpp
src/FeatureVector.cpp
src/LogisticRegression.cpp
)
target_include_directories(query_classifier_lib PUBLIC include)
target_compile_options(query_classifier_lib PRIVATE -Wall -Wextra -Wpedantic)
add_executable(query-classifier src/main.cpp)
target_link_libraries(query-classifier PRIVATE query_classifier_lib)
add_executable(query-classifier-benchmark bench_predict.cpp)
target_link_libraries(query-classifier-benchmark PRIVATE query_classifier_lib)
add_executable(query-classifier-tests tests/test_classifier.cpp)
target_link_libraries(query-classifier-tests PRIVATE query_classifier_lib)
enable_testing()
add_test(NAME query-classifier-tests COMMAND query-classifier-tests)
#include "query_classifier/FeatureVector.hpp"
namespace query_classifier {
FeatureVector operator+(const FeatureVector& left, const FeatureVector& right) {
FeatureVector result;
result.reserve(left.size() + right.size());
result.insert(result.end(), left.begin(), left.end());
result.insert(result.end(), right.begin(), right.end());
return result;
}
} // namespace query_classifier
#pragma once
#include <vector>
namespace query_classifier {
class FeatureVector : public std::vector<double> {
public:
using std::vector<double>::vector;
};
FeatureVector operator+(const FeatureVector& left, const FeatureVector& right);
} // namespace query_classifier
#include "query_classifier/LogisticRegression.hpp"
#include <cmath>
#include <stdexcept>
namespace query_classifier {
LogisticRegression::LogisticRegression(double learning_rate, std::size_t epochs,
double l2)
: learning_rate_(learning_rate), epochs_(epochs), l2_(l2) {}
void LogisticRegression::train(const std::vector<FeatureVector>& features,
const std::vector<int>& labels) {
if (features.empty() || features.size() != labels.size()) {
throw std::invalid_argument("features and labels must have equal nonzero size");
}
const std::size_t dimensions = features.front().size();
weights_.assign(dimensions, 0.0);
bias_ = 0.0;
for (const auto& row : features) {
if (row.size() != dimensions) {
throw std::invalid_argument("feature vectors must have equal dimensions");
}
}
for (std::size_t epoch = 0; epoch < epochs_; ++epoch) {
for (std::size_t row = 0; row < features.size(); ++row) {
double score = bias_;
for (std::size_t column = 0; column < dimensions; ++column) {
score += weights_[column] * features[row][column];
}
const double probability = 1.0 / (1.0 + std::exp(-score));
const double error = probability - static_cast<double>(labels[row]);
bias_ -= learning_rate_ * error;
for (std::size_t column = 0; column < dimensions; ++column) {
weights_[column] -= learning_rate_ *
(error * features[row][column] + l2_ * weights_[column]);
}
}
}
}
void LogisticRegression::load(const FeatureVector& weights, double bias) {
weights_ = weights;
bias_ = bias;
}
double LogisticRegression::predict_probability(const FeatureVector& features) const {
if (features.size() != weights_.size()) {
throw std::invalid_argument("feature vector dimension does not match model");
}
double score = bias_;
for (std::size_t column = 0; column < features.size(); ++column) {
score += weights_[column] * features[column];
}
return 1.0 / (1.0 + std::exp(-score));
}
int LogisticRegression::predict(const FeatureVector& features, double threshold) const {
return predict_probability(features) >= threshold ? 1 : 0;
}
} // namespace query_classifier
#pragma once
#include "query_classifier/FeatureVector.hpp"
#include <cstddef>
#include <vector>
namespace query_classifier {
class LogisticRegression {
public:
explicit LogisticRegression(double learning_rate = 0.1,
std::size_t epochs = 1000,
double l2 = 0.001);
void train(const std::vector<FeatureVector>& features,
const std::vector<int>& labels);
void load(const FeatureVector& weights, double bias);
double predict_probability(const FeatureVector& features) const;
int predict(const FeatureVector& features, double threshold = 0.5) const;
private:
double learning_rate_;
std::size_t epochs_;
double l2_;
double bias_ = 0.0;
std::vector<double> weights_;
};
} // namespace query_classifier
#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;
}
#include "query_classifier/MemoryPressureFeatureVectorizer.hpp"
namespace query_classifier {
FeatureVector MemoryPressureFeatureVectorizer::vectorize(
const QueryContext& context) const {
FeatureVector result(feature_count(), 0.0);
result[context.memory_pressure ? 1 : 0] = 1.0;
return result;
}
std::size_t MemoryPressureFeatureVectorizer::feature_count() const {
return 2;
}
} // namespace query_classifier
#pragma once
#include "query_classifier/FeatureVector.hpp"
#include "query_classifier/QueryContext.hpp"
namespace query_classifier {
class MemoryPressureFeatureVectorizer {
public:
FeatureVector vectorize(const QueryContext& context) const;
std::size_t feature_count() const;
};
} // namespace query_classifier
#pragma once
namespace query_classifier {
enum class OperatingSystem { Linux, Mac, Windows, Solaris, Aix };
} // namespace query_classifier
#include "query_classifier/OperatingSystemFeatureVectorizer.hpp"
namespace query_classifier {
FeatureVector OperatingSystemFeatureVectorizer::vectorize(
const QueryContext& context) const {
FeatureVector result(feature_count(), 0.0);
result[static_cast<std::size_t>(context.operating_system)] = 1.0;
return result;
}
std::size_t OperatingSystemFeatureVectorizer::feature_count() const {
return operating_system_count;
}
} // namespace query_classifier
#pragma once
#include "query_classifier/FeatureVector.hpp"
#include "query_classifier/QueryContext.hpp"
#include <cstddef>
namespace query_classifier {
class OperatingSystemFeatureVectorizer {
public:
FeatureVector vectorize(const QueryContext& context) const;
std::size_t feature_count() const;
private:
static constexpr std::size_t operating_system_count = 5;
};
} // namespace query_classifier
#pragma once
#include "query_classifier/OperatingSystem.hpp"
#include <cstddef>
namespace query_classifier {
struct QueryContext {
OperatingSystem operating_system = OperatingSystem::Linux;
bool memory_pressure = false;
bool is_64_bit = true;
double duration_seconds = 0.0;
std::size_t steps = 1;
};
} // namespace query_classifier
#pragma once
namespace query_classifier {
enum class QueryLabel { Ok, Slow };
} // namespace query_classifier
#include "query_classifier/AstFeatureVectorizer.hpp"
#include "query_classifier/BitnessFeatureVectorizer.hpp"
#include "query_classifier/FeatureVector.hpp"
#include "query_classifier/MemoryPressureFeatureVectorizer.hpp"
#include "query_classifier/OperatingSystemFeatureVectorizer.hpp"
#include "query_classifier/OperatingSystem.hpp"
#include "query_classifier/QueryContext.hpp"
#include "query_classifier/WorkFeatureVectorizer.hpp"
#include "query_classifier/LogisticRegression.hpp"
#include <cassert>
#include <cmath>
#include <vector>
using namespace query_classifier;
int main() {
const AstNode tree{AstNodeType::Root,
{{AstNodeType::Filter, {{AstNodeType::Metric, {}}}}}};
const QueryContext context{OperatingSystem::Linux, true, false, 20.0, 10};
AstFeatureVectorizer ast_vectorizer(3);
OperatingSystemFeatureVectorizer operating_system_vectorizer;
BitnessFeatureVectorizer bitness_vectorizer;
MemoryPressureFeatureVectorizer memory_pressure_vectorizer;
WorkFeatureVectorizer work_vectorizer;
const AstNode other_tree{AstNodeType::Root, {{AstNodeType::Metric, {}}}};
const QueryContext other_context{OperatingSystem::Mac, false, true, 1.0, 2};
const auto features =
ast_vectorizer.vectorize(tree) +
operating_system_vectorizer.vectorize(context) +
bitness_vectorizer.vectorize(context) +
memory_pressure_vectorizer.vectorize(context) +
work_vectorizer.vectorize(context);
const auto other_features =
ast_vectorizer.vectorize(other_tree) +
operating_system_vectorizer.vectorize(other_context) +
bitness_vectorizer.vectorize(other_context) +
memory_pressure_vectorizer.vectorize(other_context) +
work_vectorizer.vectorize(other_context);
assert(features.size() == ast_vectorizer.feature_count() +
operating_system_vectorizer.feature_count() +
bitness_vectorizer.feature_count() +
memory_pressure_vectorizer.feature_count() +
work_vectorizer.feature_count());
assert(ast_vectorizer.feature_count() >= 6);
assert(std::fabs(features[0]) >= 0.0);
assert(other_features.size() == features.size());
LogisticRegression model(0.2, 500, 0.001);
model.train({features, other_features}, {1, 0});
assert(model.predict(features) == 1);
assert(model.predict(other_features) == 0);
LogisticRegression loaded_model;
loaded_model.load(FeatureVector{2.0, -1.0}, 0.5);
const FeatureVector loaded_features{1.0, 2.0};
assert(std::fabs(loaded_model.predict_probability(loaded_features) -
0.6224593312) < 1e-6);
return 0;
}
#include "query_classifier/WorkFeatureVectorizer.hpp"
#include <algorithm>
#include <cmath>
namespace query_classifier {
FeatureVector WorkFeatureVectorizer::vectorize(
const QueryContext& context) const {
FeatureVector result(feature_count(), 0.0);
const double duration = std::max(0.0, context.duration_seconds);
const double steps = static_cast<double>(std::max<std::size_t>(context.steps, 1));
result[0] = std::log1p(duration);
result[1] = std::log1p(steps);
result[2] = std::log1p(duration * steps);
return result;
}
std::size_t WorkFeatureVectorizer::feature_count() const {
return 3;
}
} // namespace query_classifier
#pragma once
#include "query_classifier/FeatureVector.hpp"
#include "query_classifier/QueryContext.hpp"
namespace query_classifier {
class WorkFeatureVectorizer {
public:
FeatureVector vectorize(const QueryContext& context) const;
std::size_t feature_count() const;
};
} // namespace query_classifier
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment