349 lines
13 KiB
C++
349 lines
13 KiB
C++
/* Adapted from xLLM commit 78aa2a85 (PR #2258).
|
|
Adds dp_token_counts / dp_is_decode to the pybind11-exported
|
|
AttentionMetadataView so Python model executors (Qwen3.5 MoE layers,
|
|
decode graph runners) can read per-DP-rank token counts and decide
|
|
between padded vs compact all-gather.
|
|
|
|
Original: xllm/core/runtime/py_attention_metadata.cpp
|
|
Scope: Qwen3.5 data-parallel support in project_6.
|
|
==============================================================================*/
|
|
|
|
#include "core/runtime/py_attention_metadata.h"
|
|
|
|
#include <pybind11/stl.h>
|
|
#include <torch/extension.h>
|
|
|
|
#include <utility>
|
|
|
|
/*
|
|
* NOTE: The upstream xLLM implementation #includes
|
|
* "core/framework/model/model_input_params.h"
|
|
* "core/layers/common/attention_metadata.h"
|
|
* Those headers are part of xLLM's internal C++ framework and are NOT
|
|
* open-sourced in project_6. The stub types below satisfy the build so
|
|
* the DP-specific logic compiles; the real integration will link against
|
|
* the xLLM shared libraries that provide the concrete structs.
|
|
*/
|
|
|
|
namespace project6::layer {
|
|
|
|
struct ExpandedDecodeMetadata {
|
|
bool enabled = false;
|
|
torch::Tensor kv_seq_lens;
|
|
torch::Tensor block_table;
|
|
torch::Tensor paged_kv_indptr;
|
|
torch::Tensor paged_kv_indices;
|
|
torch::Tensor paged_kv_last_page_len;
|
|
torch::Tensor paged_attention_tiling_data;
|
|
torch::Tensor kv_seq_lens_host;
|
|
std::vector<int32_t> kv_seq_lens_host_vec;
|
|
};
|
|
|
|
struct AttentionMetadata {
|
|
torch::Tensor slot_mapping;
|
|
torch::Tensor paged_kv_indptr;
|
|
torch::Tensor paged_kv_indices;
|
|
torch::Tensor paged_kv_last_page_len;
|
|
std::optional<torch::Tensor> qo_indptr;
|
|
torch::Tensor q_cu_seq_lens;
|
|
torch::Tensor kv_cu_seq_lens;
|
|
torch::Tensor block_table;
|
|
torch::Tensor kv_seq_lens;
|
|
torch::Tensor q_seq_lens;
|
|
torch::Tensor has_initial_states;
|
|
std::vector<int32_t> kv_seq_lens_vec;
|
|
std::vector<int32_t> q_seq_lens_vec;
|
|
bool is_prefill = false;
|
|
bool is_chunked_prefill = false;
|
|
ExpandedDecodeMetadata expanded_decode;
|
|
};
|
|
|
|
} // namespace project6::layer
|
|
|
|
namespace project6 {
|
|
|
|
/* Minimal stub so the two-arg constructor compiles. */
|
|
struct ModelInputParams {
|
|
struct {
|
|
std::vector<int32_t> raw_dp_global_token_nums;
|
|
std::vector<int32_t> dp_global_token_nums;
|
|
std::vector<int32_t> dp_is_decode;
|
|
} parallel;
|
|
struct {
|
|
torch::Tensor linear_state_indices;
|
|
} embedding;
|
|
};
|
|
|
|
namespace py = pybind11;
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// pybind11 registration
|
|
// ---------------------------------------------------------------------------
|
|
|
|
void register_attention_metadata_views(py::module_& module) {
|
|
py::class_<PyExpandedDecodeMetadataView>(module, "ExpandedDecodeMetadataView")
|
|
.def_property_readonly("enabled", &PyExpandedDecodeMetadataView::enabled)
|
|
.def_property_readonly("kv_seq_lens",
|
|
&PyExpandedDecodeMetadataView::kv_seq_lens)
|
|
.def_property_readonly("block_table",
|
|
&PyExpandedDecodeMetadataView::block_table)
|
|
.def_property_readonly("paged_kv_indptr",
|
|
&PyExpandedDecodeMetadataView::paged_kv_indptr)
|
|
.def_property_readonly("paged_kv_indices",
|
|
&PyExpandedDecodeMetadataView::paged_kv_indices)
|
|
.def_property_readonly(
|
|
"paged_kv_last_page_len",
|
|
&PyExpandedDecodeMetadataView::paged_kv_last_page_len)
|
|
.def_property_readonly(
|
|
"paged_attention_tiling_data",
|
|
&PyExpandedDecodeMetadataView::paged_attention_tiling_data)
|
|
.def_property_readonly("kv_seq_lens_host",
|
|
&PyExpandedDecodeMetadataView::kv_seq_lens_host)
|
|
.def_property_readonly(
|
|
"kv_seq_lens_host_values",
|
|
&PyExpandedDecodeMetadataView::kv_seq_lens_host_values);
|
|
|
|
py::class_<PyAttentionMetadataView>(module, "AttentionMetadataView")
|
|
.def_property_readonly("slot_mapping",
|
|
&PyAttentionMetadataView::slot_mapping)
|
|
.def_property_readonly("paged_kv_indptr",
|
|
&PyAttentionMetadataView::paged_kv_indptr)
|
|
.def_property_readonly("paged_kv_indices",
|
|
&PyAttentionMetadataView::paged_kv_indices)
|
|
.def_property_readonly("paged_kv_last_page_len",
|
|
&PyAttentionMetadataView::paged_kv_last_page_len)
|
|
.def_property_readonly("qo_indptr", &PyAttentionMetadataView::qo_indptr)
|
|
.def_property_readonly("q_cu_seq_lens",
|
|
&PyAttentionMetadataView::q_cu_seq_lens)
|
|
.def_property_readonly("kv_cu_seq_lens",
|
|
&PyAttentionMetadataView::kv_cu_seq_lens)
|
|
.def_property_readonly("kv_seq_lens_host",
|
|
&PyAttentionMetadataView::kv_seq_lens_host)
|
|
.def_property_readonly("kv_seq_lens_host_values",
|
|
&PyAttentionMetadataView::kv_seq_lens_host_values)
|
|
.def_property_readonly("q_seq_lens_host",
|
|
&PyAttentionMetadataView::q_seq_lens_host)
|
|
.def_property_readonly("block_table",
|
|
&PyAttentionMetadataView::block_table)
|
|
.def_property_readonly("kv_seq_lens",
|
|
&PyAttentionMetadataView::kv_seq_lens)
|
|
.def_property_readonly("linear_state_indices",
|
|
&PyAttentionMetadataView::linear_state_indices)
|
|
.def_property_readonly("has_initial_state",
|
|
&PyAttentionMetadataView::has_initial_state)
|
|
/* ---- DP fields (added by PR #2258) ------------------------------ */
|
|
.def_property_readonly("dp_token_counts",
|
|
&PyAttentionMetadataView::dp_token_counts)
|
|
.def_property_readonly("dp_is_decode",
|
|
&PyAttentionMetadataView::dp_is_decode)
|
|
/* ----------------------------------------------------------------- */
|
|
.def_property_readonly("q_seq_lens", &PyAttentionMetadataView::q_seq_lens)
|
|
.def_property_readonly("expanded_decode_metadata",
|
|
&PyAttentionMetadataView::expanded_decode_metadata)
|
|
.def_property_readonly("is_prefill", &PyAttentionMetadataView::is_prefill)
|
|
.def_property_readonly("is_chunked_prefill",
|
|
&PyAttentionMetadataView::is_chunked_prefill);
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// PyExpandedDecodeMetadataView
|
|
// ---------------------------------------------------------------------------
|
|
|
|
PyExpandedDecodeMetadataView::PyExpandedDecodeMetadataView(
|
|
std::shared_ptr<layer::AttentionMetadata> metadata)
|
|
: metadata_(std::move(metadata)) {}
|
|
|
|
bool PyExpandedDecodeMetadataView::enabled() const {
|
|
return metadata().enabled;
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::kv_seq_lens() const {
|
|
return metadata().kv_seq_lens.defined() ? py::cast(metadata().kv_seq_lens)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::block_table() const {
|
|
return metadata().block_table.defined() ? py::cast(metadata().block_table)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::paged_kv_indptr() const {
|
|
return metadata().paged_kv_indptr.defined()
|
|
? py::cast(metadata().paged_kv_indptr)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::paged_kv_indices() const {
|
|
return metadata().paged_kv_indices.defined()
|
|
? py::cast(metadata().paged_kv_indices)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::paged_kv_last_page_len() const {
|
|
return metadata().paged_kv_last_page_len.defined()
|
|
? py::cast(metadata().paged_kv_last_page_len)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::paged_attention_tiling_data() const {
|
|
return metadata().paged_attention_tiling_data.defined()
|
|
? py::cast(metadata().paged_attention_tiling_data)
|
|
: py::none();
|
|
}
|
|
|
|
py::object PyExpandedDecodeMetadataView::kv_seq_lens_host() const {
|
|
return metadata().kv_seq_lens_host.defined()
|
|
? py::cast(metadata().kv_seq_lens_host)
|
|
: py::none();
|
|
}
|
|
|
|
const std::vector<int32_t>&
|
|
PyExpandedDecodeMetadataView::kv_seq_lens_host_values() const {
|
|
return metadata().kv_seq_lens_host_vec;
|
|
}
|
|
|
|
const layer::ExpandedDecodeMetadata& PyExpandedDecodeMetadataView::metadata()
|
|
const {
|
|
return metadata_->expanded_decode;
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// PyAttentionMetadataView
|
|
// ---------------------------------------------------------------------------
|
|
|
|
PyAttentionMetadataView::PyAttentionMetadataView(
|
|
std::shared_ptr<layer::AttentionMetadata> metadata)
|
|
: metadata_(std::move(metadata)),
|
|
kv_seq_lens_host_(
|
|
make_host_int32_view(metadata_, metadata_->kv_seq_lens_vec)),
|
|
q_seq_lens_host_(
|
|
make_host_int32_view(metadata_, metadata_->q_seq_lens_vec)) {}
|
|
|
|
PyAttentionMetadataView::PyAttentionMetadataView(
|
|
std::shared_ptr<layer::AttentionMetadata> metadata,
|
|
const ModelInputParams& params)
|
|
: PyAttentionMetadataView(std::move(metadata)) {
|
|
linear_state_indices_ = params.embedding.linear_state_indices;
|
|
|
|
/* ---- DP fields (added by PR #2258) ---------------------------------- */
|
|
dp_token_counts_ = params.parallel.raw_dp_global_token_nums.empty()
|
|
? params.parallel.dp_global_token_nums
|
|
: params.parallel.raw_dp_global_token_nums;
|
|
dp_is_decode_ = params.parallel.dp_is_decode;
|
|
/* --------------------------------------------------------------------- */
|
|
}
|
|
|
|
const torch::Tensor& PyAttentionMetadataView::slot_mapping() const {
|
|
return metadata_->slot_mapping;
|
|
}
|
|
|
|
const torch::Tensor& PyAttentionMetadataView::paged_kv_indptr() const {
|
|
return metadata_->paged_kv_indptr;
|
|
}
|
|
|
|
const torch::Tensor& PyAttentionMetadataView::paged_kv_indices() const {
|
|
return metadata_->paged_kv_indices;
|
|
}
|
|
|
|
const torch::Tensor& PyAttentionMetadataView::paged_kv_last_page_len() const {
|
|
return metadata_->paged_kv_last_page_len;
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::qo_indptr() const {
|
|
if (!metadata_->qo_indptr.has_value() || !metadata_->qo_indptr->defined()) {
|
|
return py::none();
|
|
}
|
|
return py::cast(*metadata_->qo_indptr);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::q_cu_seq_lens() const {
|
|
return optional_tensor(metadata_->q_cu_seq_lens);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::kv_cu_seq_lens() const {
|
|
return optional_tensor(metadata_->kv_cu_seq_lens);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::kv_seq_lens_host() const {
|
|
return optional_tensor(kv_seq_lens_host_);
|
|
}
|
|
|
|
const std::vector<int32_t>& PyAttentionMetadataView::kv_seq_lens_host_values()
|
|
const {
|
|
return metadata_->kv_seq_lens_vec;
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::block_table() const {
|
|
return optional_tensor(metadata_->block_table);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::kv_seq_lens() const {
|
|
return optional_tensor(metadata_->kv_seq_lens);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::linear_state_indices() const {
|
|
return optional_tensor(linear_state_indices_);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::has_initial_state() const {
|
|
return optional_tensor(metadata_->has_initial_states);
|
|
}
|
|
|
|
/* ---- DP fields (added by PR #2258) ------------------------------------ */
|
|
const std::vector<int32_t>& PyAttentionMetadataView::dp_token_counts() const {
|
|
return dp_token_counts_;
|
|
}
|
|
|
|
const std::vector<int32_t>& PyAttentionMetadataView::dp_is_decode() const {
|
|
return dp_is_decode_;
|
|
}
|
|
/* ----------------------------------------------------------------------- */
|
|
|
|
py::object PyAttentionMetadataView::q_seq_lens() const {
|
|
return optional_tensor(metadata_->q_seq_lens);
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::q_seq_lens_host() const {
|
|
return optional_tensor(q_seq_lens_host_);
|
|
}
|
|
|
|
PyExpandedDecodeMetadataView PyAttentionMetadataView::expanded_decode_metadata()
|
|
const {
|
|
return PyExpandedDecodeMetadataView(metadata_);
|
|
}
|
|
|
|
bool PyAttentionMetadataView::is_prefill() const {
|
|
return metadata_->is_prefill;
|
|
}
|
|
|
|
bool PyAttentionMetadataView::is_chunked_prefill() const {
|
|
return metadata_->is_chunked_prefill;
|
|
}
|
|
|
|
torch::Tensor PyAttentionMetadataView::make_host_int32_view(
|
|
const std::shared_ptr<layer::AttentionMetadata>& metadata,
|
|
std::vector<int32_t>& host_vec) {
|
|
if (host_vec.empty()) {
|
|
return torch::Tensor();
|
|
}
|
|
|
|
std::shared_ptr<layer::AttentionMetadata> owner = metadata;
|
|
return torch::from_blob(
|
|
host_vec.data(),
|
|
{static_cast<int64_t>(host_vec.size())},
|
|
[owner = std::move(owner)](void*) mutable { owner.reset(); },
|
|
torch::TensorOptions().dtype(torch::kInt32).device(torch::kCPU));
|
|
}
|
|
|
|
py::object PyAttentionMetadataView::optional_tensor(
|
|
const torch::Tensor& tensor) {
|
|
return tensor.defined() ? py::cast(tensor) : py::none();
|
|
}
|
|
|
|
} // namespace project6
|
|
|
|
PYBIND11_MODULE(py_attention_metadata, m) {
|
|
m.doc() = "DP-aware attention metadata (project6, ported from xLLM PR #2258)";
|
|
project6::register_attention_metadata_views(m);
|
|
}
|