Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions docs/paged_attention_engine.md
Original file line number Diff line number Diff line change
Expand Up @@ -851,6 +851,13 @@ Graph buffers are allocated once at configured limits and reshaped as static vie

Prefill and mixed-token steps use graph id `-1`, which tells the CUDA execution provider to run eagerly.

An Engine-hosted DFlash 2 drafter also forces eager target execution. Its target decoder output is
the packed auxiliary hidden-state tensor named by `model.dflash2.main_aux_hidden_states`; that
variable-size output does not have persistent graph buffers. Engine construction validates its
rank, element type, and static width against the drafter input before allocating cache resources.
The drafter run is synchronous because its packed inputs and outputs are owned by one proposal
call, so `model.dflash2.run_options` cannot disable execution-provider synchronization.

## Backpressure and fairness

Continuous batching does not mean every pending request runs on every step.
Expand Down
137 changes: 137 additions & 0 deletions src/config.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -459,6 +459,8 @@ struct DecoderOutputs_Element : JSON::Element {
v_.state_update_recurrent_capsule_names = JSON::Get<std::string_view>(value);
} else if (name == "hidden_states") {
v_.hidden_states = JSON::Get<std::string_view>(value);
} else if (name == "aux_hidden_states") {
v_.aux_hidden_states = JSON::Get<std::string_view>(value);
} else if (name == "outputs") {
v_.outputs = JSON::Get<std::string_view>(value);
} else if (name == "lstm_hidden_state") {
Expand Down Expand Up @@ -1093,6 +1095,137 @@ struct Mtp_Element : JSON::Element {
SharedInitializers_Element shared_initializers_{v_.shared_initializers};
};

struct Dflash2Inputs_Element : JSON::Element {
explicit Dflash2Inputs_Element(Config::Model::Dflash2::Inputs& v) : v_{v} {}

void OnValue(std::string_view name, JSON::Value value) override {
if (name == "aux_hidden_states") {
v_.aux_hidden_states = JSON::Get<std::string_view>(value);
} else if (name == "input_ids") {
v_.input_ids = JSON::Get<std::string_view>(value);
} else if (name == "q_row_map") {
v_.q_row_map = JSON::Get<std::string_view>(value);
} else if (name == "qkv_row_map") {
v_.qkv_row_map = JSON::Get<std::string_view>(value);
} else if (name == "block_row_index") {
v_.block_row_index = JSON::Get<std::string_view>(value);
} else if (name == "cumulative_sequence_lengths") {
v_.cumulative_sequence_lengths = JSON::Get<std::string_view>(value);
} else if (name == "past_sequence_lengths") {
v_.past_sequence_lengths = JSON::Get<std::string_view>(value);
} else if (name == "block_table") {
v_.block_table = JSON::Get<std::string_view>(value);
} else if (name == "attention_metadata") {
v_.attention_metadata = JSON::Get<std::string_view>(value);
} else if (name == "past_key_names") {
v_.past_key_names = JSON::Get<std::string_view>(value);
} else if (name == "past_value_names") {
v_.past_value_names = JSON::Get<std::string_view>(value);
} else {
throw JSON::unknown_value_error{};
}
}

private:
Config::Model::Dflash2::Inputs& v_;
};

struct Dflash2Outputs_Element : JSON::Element {
explicit Dflash2Outputs_Element(Config::Model::Dflash2::Outputs& v) : v_{v} {}

void OnValue(std::string_view name, JSON::Value value) override {
if (name == "candidate_ids") {
v_.candidate_ids = JSON::Get<std::string_view>(value);
} else if (name == "scores") {
v_.scores = JSON::Get<std::string_view>(value);
} else if (name == "present_key_names") {
v_.present_key_names = JSON::Get<std::string_view>(value);
} else if (name == "present_value_names") {
v_.present_value_names = JSON::Get<std::string_view>(value);
} else {
throw JSON::unknown_value_error{};
}
}

private:
Config::Model::Dflash2::Outputs& v_;
};

struct Dflash2_Element : JSON::Element {
explicit Dflash2_Element(Config::Model::Dflash2& v) : v_{v} {}

void OnValue(std::string_view name, JSON::Value value) override {
if (name == "filename") {
v_.filename = JSON::Get<std::string_view>(value);
} else if (name == "num_hidden_layers") {
v_.num_hidden_layers = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.num_hidden_layers <= 0) throw std::out_of_range("num_hidden_layers must be > 0");
} else if (name == "num_key_value_heads") {
v_.num_key_value_heads = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.num_key_value_heads <= 0) throw std::out_of_range("num_key_value_heads must be > 0");
} else if (name == "head_size") {
v_.head_size = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.head_size <= 0) throw std::out_of_range("head_size must be > 0");
} else if (name == "block_size") {
v_.block_size = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.block_size <= 1) throw std::out_of_range("block_size must be > 1");
} else if (name == "num_draft_tokens") {
v_.num_draft_tokens = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.num_draft_tokens <= 0) throw std::out_of_range("num_draft_tokens must be > 0");
} else if (name == "selector_top_k") {
v_.selector_top_k = SafeDoubleToInt(JSON::Get<double>(value), name);
if (v_.selector_top_k <= 0) throw std::out_of_range("selector_top_k must be > 0");
} else if (name == "mask_token_id") {
v_.mask_token_id = SafeDoubleToInt(JSON::Get<double>(value), name);
} else if (name == "sliding_window") {
v_.sliding_window = SafeDoubleToInt(JSON::Get<double>(value), name);
} else if (name == "main_aux_hidden_states") {
v_.main_aux_hidden_states = JSON::Get<std::string_view>(value);
} else {
throw JSON::unknown_value_error{};
}
}

Element& OnObject(std::string_view name) override {
if (name == "session_options") {
v_.session_options = Config::SessionOptions{};
session_options_ = std::make_unique<SessionOptions_Element>(*v_.session_options);
return *session_options_;
}
if (name == "run_options") {
v_.run_options = Config::RunOptions{};
run_options_ = std::make_unique<RunOptions_Element>(*v_.run_options);
return *run_options_;
}
if (name == "inputs") {
return inputs_;
}
if (name == "outputs") {
return outputs_;
}
throw JSON::unknown_value_error{};
}

Element& OnArray(std::string_view name) override {
if (name == "shared_initializers") {
return shared_initializers_;
}
if (name == "aux_hidden_state_layers") {
return aux_hidden_state_layers_;
}
throw JSON::unknown_value_error{};
}

private:
Config::Model::Dflash2& v_;
std::unique_ptr<SessionOptions_Element> session_options_;
std::unique_ptr<RunOptions_Element> run_options_;
Dflash2Inputs_Element inputs_{v_.inputs};
Dflash2Outputs_Element outputs_{v_.outputs};
SharedInitializers_Element shared_initializers_{v_.shared_initializers};
IntArray_Element aux_hidden_state_layers_{v_.aux_hidden_state_layers};
};

struct VisionInputs_Element : JSON::Element {
explicit VisionInputs_Element(Config::Model::Vision::Inputs& v) : v_{v} {}

Expand Down Expand Up @@ -1652,6 +1785,9 @@ struct Model_Element : JSON::Element {
if (name == "mtp") {
return mtp_;
}
if (name == "dflash2") {
return dflash2_;
}
throw JSON::unknown_value_error{};
}

Expand All @@ -1668,6 +1804,7 @@ struct Model_Element : JSON::Element {
Joiner_Element joiner_{v_.joiner};
VAD_Element vad_{v_.vad};
Mtp_Element mtp_{v_.mtp};
Dflash2_Element dflash2_{v_.dflash2};
};

// Throws std::runtime_error (rather than std::overflow_error/std::invalid_argument) on failure.
Expand Down
48 changes: 48 additions & 0 deletions src/config.h
Original file line number Diff line number Diff line change
Expand Up @@ -488,6 +488,9 @@ struct Config {
std::string state_update_conv_value_names{Defaults::StateUpdateConvValueName};
std::string state_update_recurrent_capsule_names{Defaults::StateUpdateRecurrentCapsuleName};
std::string hidden_states; // Last hidden state output (when exported with include_hidden_states; e.g. fed to the MTP head)
// Residual streams tapped at model.dflash2.aux_hidden_state_layers, concatenated on the
// last axis. Empty unless the model was exported with aux_hidden_state_layers.
std::string aux_hidden_states;

// RNNT decoder outputs
std::string outputs;
Expand Down Expand Up @@ -557,6 +560,51 @@ struct Config {
} outputs;
} mtp;

// DFlash 2 block-drafter metadata. Unlike MTP the drafter is not decoder-shaped: it reads the
// main model's auxiliary hidden states, predicts a whole block of tokens at once, and returns
// a candidate lattice (top-k ids per slot plus the pairwise edge scores) that the engine walks
// greedily. The Engine drives its session directly rather than through a Model.
struct Dflash2 {
std::string filename; // e.g. "dflash2.onnx"
std::optional<SessionOptions> session_options;
std::optional<RunOptions> run_options;
std::vector<SharedInitializer> shared_initializers;

int num_hidden_layers{};
int num_key_value_heads{};
int head_size{};
int block_size{}; // Query rows per request: the anchor token plus one mask per draft.
int num_draft_tokens{}; // block_size - 1
int selector_top_k{};
int mask_token_id{};
int sliding_window{-1};
std::vector<int> aux_hidden_state_layers;

// Name of the main decoder's auxiliary hidden-states output that feeds the drafter.
std::string main_aux_hidden_states{"aux_hidden_states"};

struct Inputs {
std::string aux_hidden_states{"aux_hidden_states"};
std::string input_ids{Defaults::InputIdsName};
std::string q_row_map{"q_row_map"};
std::string qkv_row_map{"qkv_row_map"};
std::string block_row_index{"block_row_index"};
std::string cumulative_sequence_lengths{Defaults::CumulativeSequenceLengthsName};
std::string past_sequence_lengths{Defaults::PastSequenceLengthsName};
std::string block_table{Defaults::BlockTableName};
std::string attention_metadata{Defaults::AttentionMetadataName};
std::string past_key_names{Defaults::PastKeyName};
std::string past_value_names{Defaults::PastValueName};
} inputs;

struct Outputs {
std::string candidate_ids{"draft_candidate_ids"};
std::string scores{"draft_scores"};
std::string present_key_names{Defaults::PresentKeyName};
std::string present_value_names{Defaults::PresentValueName};
} outputs;
} dflash2;

std::optional<Decoder> draft;

} model;
Expand Down
Loading
Loading