under test, not sure no errors
This commit is contained in:
343
core/runtime/py_attention_metadata.cpp
Normal file
343
core/runtime/py_attention_metadata.cpp
Normal file
@@ -0,0 +1,343 @@
|
||||
/* 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
|
||||
100
core/runtime/py_attention_metadata.h
Normal file
100
core/runtime/py_attention_metadata.h
Normal file
@@ -0,0 +1,100 @@
|
||||
/* Adapted from xLLM commit 78aa2a85 (PR #2258).
|
||||
Adds dp_token_counts / dp_is_decode fields to PyAttentionMetadataView
|
||||
so the Python attention backend can partition KV cache by DP group.
|
||||
|
||||
Original: xllm/core/runtime/py_attention_metadata.h
|
||||
Scope: Qwen3.5 data-parallel support in project_6.
|
||||
==============================================================================*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <torch/torch.h>
|
||||
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <vector>
|
||||
|
||||
/* Forward declarations — project_6 keeps these in its own layer namespace. */
|
||||
namespace project6::layer {
|
||||
struct AttentionMetadata;
|
||||
struct ExpandedDecodeMetadata;
|
||||
} // namespace project6::layer
|
||||
|
||||
namespace project6 {
|
||||
|
||||
struct ModelInputParams;
|
||||
|
||||
void register_attention_metadata_views(pybind11::module_& module);
|
||||
|
||||
class PyExpandedDecodeMetadataView final {
|
||||
public:
|
||||
explicit PyExpandedDecodeMetadataView(
|
||||
std::shared_ptr<layer::AttentionMetadata> metadata);
|
||||
|
||||
bool enabled() const;
|
||||
pybind11::object kv_seq_lens() const;
|
||||
pybind11::object block_table() const;
|
||||
pybind11::object paged_kv_indptr() const;
|
||||
pybind11::object paged_kv_indices() const;
|
||||
pybind11::object paged_kv_last_page_len() const;
|
||||
pybind11::object paged_attention_tiling_data() const;
|
||||
pybind11::object kv_seq_lens_host() const;
|
||||
const std::vector<int32_t>& kv_seq_lens_host_values() const;
|
||||
|
||||
private:
|
||||
const layer::ExpandedDecodeMetadata& metadata() const;
|
||||
|
||||
std::shared_ptr<layer::AttentionMetadata> metadata_;
|
||||
};
|
||||
|
||||
class PyAttentionMetadataView final {
|
||||
public:
|
||||
explicit PyAttentionMetadataView(
|
||||
std::shared_ptr<layer::AttentionMetadata> metadata);
|
||||
PyAttentionMetadataView(std::shared_ptr<layer::AttentionMetadata> metadata,
|
||||
const ModelInputParams& params);
|
||||
|
||||
const torch::Tensor& slot_mapping() const;
|
||||
const torch::Tensor& paged_kv_indptr() const;
|
||||
const torch::Tensor& paged_kv_indices() const;
|
||||
const torch::Tensor& paged_kv_last_page_len() const;
|
||||
pybind11::object qo_indptr() const;
|
||||
pybind11::object q_cu_seq_lens() const;
|
||||
pybind11::object kv_cu_seq_lens() const;
|
||||
pybind11::object kv_seq_lens_host() const;
|
||||
const std::vector<int32_t>& kv_seq_lens_host_values() const;
|
||||
pybind11::object q_seq_lens_host() const;
|
||||
pybind11::object block_table() const;
|
||||
pybind11::object kv_seq_lens() const;
|
||||
pybind11::object linear_state_indices() const;
|
||||
pybind11::object has_initial_state() const;
|
||||
|
||||
/* ---- DP fields (added by PR #2258) ---------------------------------- */
|
||||
const std::vector<int32_t>& dp_token_counts() const;
|
||||
const std::vector<int32_t>& dp_is_decode() const;
|
||||
/* --------------------------------------------------------------------- */
|
||||
|
||||
pybind11::object q_seq_lens() const;
|
||||
PyExpandedDecodeMetadataView expanded_decode_metadata() const;
|
||||
bool is_prefill() const;
|
||||
bool is_chunked_prefill() const;
|
||||
|
||||
private:
|
||||
static torch::Tensor make_host_int32_view(
|
||||
const std::shared_ptr<layer::AttentionMetadata>& metadata,
|
||||
std::vector<int32_t>& host_vec);
|
||||
static pybind11::object optional_tensor(const torch::Tensor& tensor);
|
||||
|
||||
std::shared_ptr<layer::AttentionMetadata> metadata_;
|
||||
torch::Tensor kv_seq_lens_host_;
|
||||
torch::Tensor q_seq_lens_host_;
|
||||
torch::Tensor linear_state_indices_;
|
||||
|
||||
/* ---- DP fields (added by PR #2258) ---------------------------------- */
|
||||
std::vector<int32_t> dp_token_counts_;
|
||||
std::vector<int32_t> dp_is_decode_;
|
||||
/* --------------------------------------------------------------------- */
|
||||
};
|
||||
|
||||
} // namespace project6
|
||||
Reference in New Issue
Block a user