#pragma once
#include <cstdint>
#include <memory>
#include <vector>
#include "proxy_dmatrix.h"
#include "xgboost/data.h"
#include "xgboost/span.h"
namespace xgboost::common {
class HistogramCuts;
}
namespace xgboost::data {
class QuantileDMatrix : public DMatrix {
template <typename Page>
static auto InvalidTreeMethod() {
LOG(FATAL) << "Only `hist` tree method can use `QuantileDMatrix`.";
return BatchSet<Page>(BatchIterator<Page>(nullptr));
}
public:
DMatrix *Slice(common::Span<std::int32_t const>) final {
LOG(FATAL) << "Slicing DMatrix is not supported for external memory.";
return nullptr;
}
DMatrix *SliceCol(std::int32_t, std::int32_t) final {
LOG(FATAL) << "Slicing DMatrix columns is not supported for external memory.";
return nullptr;
}
[[nodiscard]] bool SparsePageExists() const final { return false; }
BatchSet<SparsePage> GetRowBatches() final {
LOG(FATAL) << "Not implemented for `QuantileDMatrix`.";
return BatchSet<SparsePage>(BatchIterator<SparsePage>(nullptr));
}
BatchSet<CSCPage> GetColumnBatches(Context const *) final { return InvalidTreeMethod<CSCPage>(); }
BatchSet<SortedCSCPage> GetSortedColumnBatches(Context const *) final {
return InvalidTreeMethod<SortedCSCPage>();
}
[[nodiscard]] MetaInfo &Info() final { return info_; }
[[nodiscard]] MetaInfo const &Info() const final { return info_; }
[[nodiscard]] Context const *Ctx() const final { return &fmat_ctx_; }
protected:
Context fmat_ctx_;
MetaInfo info_;
};
void GetCutsFromRef(Context const *ctx, std::shared_ptr<DMatrix> ref, bst_feature_t n_features,
BatchParam p, common::HistogramCuts *p_cuts);
void GetCutsFromEllpack(EllpackPage const &page, common::HistogramCuts *cuts);
namespace cpu_impl {
void SyncFeatureType(Context const *ctx, std::vector<FeatureType> *p_h_ft);
void GetDataShape(Context const *ctx, DMatrixProxy *proxy,
DataIterProxy<DataIterResetCallback, XGDMatrixCallbackNext> iter, float missing,
ExternalDataInfo *p_info);
void MakeSketches(Context const *ctx,
DataIterProxy<DataIterResetCallback, XGDMatrixCallbackNext> *iter,
DMatrixProxy *proxy, std::shared_ptr<DMatrix> ref, float missing,
common::HistogramCuts *cuts, BatchParam const &p, MetaInfo const &info,
ExternalDataInfo const &ext_info, std::vector<FeatureType> *p_h_ft);
}
namespace cuda_impl {
void MakeSketches(Context const *ctx,
DataIterProxy<DataIterResetCallback, XGDMatrixCallbackNext> *iter,
DMatrixProxy *proxy, std::shared_ptr<DMatrix> ref, BatchParam const &p,
float missing, std::shared_ptr<common::HistogramCuts> cuts, MetaInfo const &info,
std::int64_t max_quantile_blocks, ExternalDataInfo *p_ext_info);
} }