Files
project_6/core/runtime/py_attention_metadata.cpp
2026-09-01 10:24:14 +00:00

344 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