feat: allow dflash to work with spec auto tune (#2112)
This commit is contained in:
parent
fedb48fa74
commit
3c6cbf6e2a
|
|
@ -98,6 +98,14 @@ void spec_tuner::reset_exploration() {
|
||||||
for (auto & coord : coords) {
|
for (auto & coord : coords) {
|
||||||
coord.reset_scores();
|
coord.reset_scores();
|
||||||
}
|
}
|
||||||
|
dflash_quarantined.assign(dflash_quarantined.size(), false);
|
||||||
|
dflash_probe_cursor = 0;
|
||||||
|
dflash_last_exploratory = false;
|
||||||
|
dflash_last_recovery_probe = false;
|
||||||
|
if (has_dflash_target_only_arm() && !coords.empty() && !coords[0].arms.empty()) {
|
||||||
|
coords[0].current_idx = 0;
|
||||||
|
coords[0].best_idx = 0;
|
||||||
|
}
|
||||||
n_low = 0;
|
n_low = 0;
|
||||||
cooldown = cooldown_max;
|
cooldown = cooldown_max;
|
||||||
step_ema = 0.0;
|
step_ema = 0.0;
|
||||||
|
|
@ -107,7 +115,11 @@ void spec_tuner::reset_exploration() {
|
||||||
void spec_tuner::write_best(common_params_speculative & params) const {
|
void spec_tuner::write_best(common_params_speculative & params) const {
|
||||||
for (const auto & coord : coords) {
|
for (const auto & coord : coords) {
|
||||||
float val = coord.arms[coord.best_idx].value;
|
float val = coord.arms[coord.best_idx].value;
|
||||||
if (coord.name == "n_max") params.n_max = (int32_t)val;
|
if (coord.name == "n_max") {
|
||||||
|
params.n_max = (spec_type == COMMON_SPECULATIVE_TYPE_DFLASH && (int32_t)val == 0)
|
||||||
|
? configured_n_max
|
||||||
|
: (int32_t)val;
|
||||||
|
}
|
||||||
else if (coord.name == "p_min") params.p_min = val;
|
else if (coord.name == "p_min") params.p_min = val;
|
||||||
else if (coord.name == "n_min") params.n_min = (int32_t)val;
|
else if (coord.name == "n_min") params.n_min = (int32_t)val;
|
||||||
else if (coord.name == "ngram_size_n") params.ngram_size_n = (uint16_t)val;
|
else if (coord.name == "ngram_size_n") params.ngram_size_n = (uint16_t)val;
|
||||||
|
|
@ -120,6 +132,16 @@ void spec_tuner::write_best(common_params_speculative & params) const {
|
||||||
void spec_tuner::init(common_speculative_type type, const common_params_speculative & user_params, const llama_model * model_tgt) {
|
void spec_tuner::init(common_speculative_type type, const common_params_speculative & user_params, const llama_model * model_tgt) {
|
||||||
enabled = true;
|
enabled = true;
|
||||||
spec_type = type;
|
spec_type = type;
|
||||||
|
configured_n_max = std::max(1, (int) user_params.n_max);
|
||||||
|
dflash_quarantined.clear();
|
||||||
|
dflash_probe_cursor = 0;
|
||||||
|
dflash_last_exploratory = false;
|
||||||
|
dflash_last_recovery_probe = false;
|
||||||
|
n_target_only_selections = 0;
|
||||||
|
n_dflash_selections = 0;
|
||||||
|
n_exploratory_selections = 0;
|
||||||
|
n_quarantines = 0;
|
||||||
|
n_recovery_probes = 0;
|
||||||
coords.clear();
|
coords.clear();
|
||||||
n_calls = 0;
|
n_calls = 0;
|
||||||
n_requests = 0;
|
n_requests = 0;
|
||||||
|
|
@ -137,12 +159,19 @@ void spec_tuner::init(common_speculative_type type, const common_params_speculat
|
||||||
spec_tuner_coord coord;
|
spec_tuner_coord coord;
|
||||||
coord.name = "n_max";
|
coord.name = "n_max";
|
||||||
const bool recurrent_target = model_tgt != nullptr && llama_model_has_recurrent(model_tgt);
|
const bool recurrent_target = model_tgt != nullptr && llama_model_has_recurrent(model_tgt);
|
||||||
int hi = recurrent_target ? std::max(1, (int) user_params.n_max)
|
int hi = type == COMMON_SPECULATIVE_TYPE_DFLASH
|
||||||
: std::max(16, (int) user_params.n_max);
|
? configured_n_max
|
||||||
coord.build_grid_int(1, hi, 1, user_params.n_max);
|
: (recurrent_target ? std::max(1, (int) user_params.n_max)
|
||||||
|
: std::max(16, (int) user_params.n_max));
|
||||||
|
const int lo = type == COMMON_SPECULATIVE_TYPE_DFLASH ? 0 : 1;
|
||||||
|
coord.build_grid_int(lo, hi, 1, user_params.n_max);
|
||||||
coords.push_back(std::move(coord));
|
coords.push_back(std::move(coord));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (type == COMMON_SPECULATIVE_TYPE_DFLASH) {
|
||||||
|
dflash_quarantined.assign(coords[0].arms.size(), false);
|
||||||
|
}
|
||||||
|
|
||||||
if (type == COMMON_SPECULATIVE_TYPE_DRAFT) {
|
if (type == COMMON_SPECULATIVE_TYPE_DRAFT) {
|
||||||
{
|
{
|
||||||
spec_tuner_coord coord;
|
spec_tuner_coord coord;
|
||||||
|
|
@ -212,12 +241,111 @@ void spec_tuner::init(common_speculative_type type, const common_params_speculat
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
int spec_tuner::select_dflash_arm(spec_tuner_coord & coord) {
|
||||||
|
dflash_last_exploratory = false;
|
||||||
|
dflash_last_recovery_probe = false;
|
||||||
|
|
||||||
|
if (dflash_quarantined.size() != coord.arms.size()) {
|
||||||
|
dflash_quarantined.assign(coord.arms.size(), false);
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int i = 0; i < (int) coord.arms.size(); ++i) {
|
||||||
|
if (coord.arms[i].N < dflash_min_samples_per_arm) {
|
||||||
|
dflash_last_exploratory = true;
|
||||||
|
return i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (dflash_recovery_probe_interval > 0 && n_calls > 0 &&
|
||||||
|
n_calls % dflash_recovery_probe_interval == 0) {
|
||||||
|
for (int offset = 0; offset < (int) coord.arms.size(); ++offset) {
|
||||||
|
const int i = (dflash_probe_cursor + offset) % (int) coord.arms.size();
|
||||||
|
if (dflash_quarantined[i]) {
|
||||||
|
dflash_probe_cursor = (i + 1) % (int) coord.arms.size();
|
||||||
|
dflash_last_exploratory = true;
|
||||||
|
dflash_last_recovery_probe = true;
|
||||||
|
n_recovery_probes++;
|
||||||
|
return i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
int best_idx = -1;
|
||||||
|
for (int i = 0; i < (int) coord.arms.size(); ++i) {
|
||||||
|
if (dflash_quarantined[i]) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if (best_idx < 0 || coord.arms[i].Q > coord.arms[best_idx].Q) {
|
||||||
|
best_idx = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (best_idx < 0) {
|
||||||
|
best_idx = 0;
|
||||||
|
}
|
||||||
|
coord.best_idx = best_idx;
|
||||||
|
return best_idx;
|
||||||
|
}
|
||||||
|
|
||||||
|
void spec_tuner::update_dflash_quarantine() {
|
||||||
|
if (!has_dflash_target_only_arm() || coords.empty() || coords[0].arms.empty()) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
auto & coord = coords[0];
|
||||||
|
if (dflash_quarantined.size() != coord.arms.size()) {
|
||||||
|
dflash_quarantined.assign(coord.arms.size(), false);
|
||||||
|
}
|
||||||
|
|
||||||
|
const int zero_idx = coord.find_nearest_arm(0.0f);
|
||||||
|
const auto & zero = coord.arms[zero_idx];
|
||||||
|
if (zero.N < 3 || zero.Q <= 0.0) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int i = 0; i < (int) coord.arms.size(); ++i) {
|
||||||
|
if (i == zero_idx) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
const bool should_quarantine = coord.arms[i].N >= 3 &&
|
||||||
|
coord.arms[i].Q < zero.Q * dflash_quarantine_ratio;
|
||||||
|
if (should_quarantine && !dflash_quarantined[i]) {
|
||||||
|
n_quarantines++;
|
||||||
|
}
|
||||||
|
dflash_quarantined[i] = should_quarantine;
|
||||||
|
}
|
||||||
|
|
||||||
|
int best_idx = -1;
|
||||||
|
for (int i = 0; i < (int) coord.arms.size(); ++i) {
|
||||||
|
if (dflash_quarantined[i]) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if (best_idx < 0 || coord.arms[i].Q > coord.arms[best_idx].Q) {
|
||||||
|
best_idx = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (best_idx >= 0) {
|
||||||
|
coord.best_idx = best_idx;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
void spec_tuner::propose(common_params_speculative & params) {
|
void spec_tuner::propose(common_params_speculative & params) {
|
||||||
int64_t t_start = ggml_time_us();
|
int64_t t_start = ggml_time_us();
|
||||||
|
|
||||||
// always select fresh arm for every draft call
|
// always select fresh arm for every draft call
|
||||||
for (auto & coord : coords) {
|
for (auto & coord : coords) {
|
||||||
|
if (has_dflash_target_only_arm() && coord.name == "n_max") {
|
||||||
|
coord.current_idx = select_dflash_arm(coord);
|
||||||
|
if ((int) coord.arms[coord.current_idx].value == 0) {
|
||||||
|
n_target_only_selections++;
|
||||||
|
} else {
|
||||||
|
n_dflash_selections++;
|
||||||
|
}
|
||||||
|
if (dflash_last_exploratory) {
|
||||||
|
n_exploratory_selections++;
|
||||||
|
}
|
||||||
|
} else {
|
||||||
coord.current_idx = coord.select_epsilon_greedy(epsilon);
|
coord.current_idx = coord.select_epsilon_greedy(epsilon);
|
||||||
|
}
|
||||||
|
|
||||||
float val = coord.arms[coord.current_idx].value;
|
float val = coord.arms[coord.current_idx].value;
|
||||||
if (coord.name == "n_max") params.n_max = (int32_t)val;
|
if (coord.name == "n_max") params.n_max = (int32_t)val;
|
||||||
|
|
@ -235,7 +363,8 @@ void spec_tuner::propose(common_params_speculative & params) {
|
||||||
|
|
||||||
void spec_tuner::enforce_constraints(common_params_speculative & params) {
|
void spec_tuner::enforce_constraints(common_params_speculative & params) {
|
||||||
if (params.n_min < 0) params.n_min = 0;
|
if (params.n_min < 0) params.n_min = 0;
|
||||||
if (params.n_max < 1) params.n_max = 1;
|
const int min_n_max = has_dflash_target_only_arm() ? 0 : 1;
|
||||||
|
if (params.n_max < min_n_max) params.n_max = min_n_max;
|
||||||
if (params.n_min > params.n_max) params.n_min = params.n_max;
|
if (params.n_min > params.n_max) params.n_min = params.n_max;
|
||||||
|
|
||||||
if (params.p_min < 0.0f) params.p_min = 0.0f;
|
if (params.p_min < 0.0f) params.p_min = 0.0f;
|
||||||
|
|
@ -257,6 +386,8 @@ void spec_tuner::accept_feedback(int n_accepted, int n_drafted, double step_tps)
|
||||||
coord.update(reward);
|
coord.update(reward);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
update_dflash_quarantine();
|
||||||
|
|
||||||
if (cooldown > 0) {
|
if (cooldown > 0) {
|
||||||
cooldown--;
|
cooldown--;
|
||||||
if (step_ema <= 0.0) {
|
if (step_ema <= 0.0) {
|
||||||
|
|
@ -287,11 +418,17 @@ void spec_tuner::accept_feedback(int n_accepted, int n_drafted, double step_tps)
|
||||||
<< " n_accepted=" << n_accepted
|
<< " n_accepted=" << n_accepted
|
||||||
<< " step_tps=" << std::fixed << std::setprecision(1) << step_tps
|
<< " step_tps=" << std::fixed << std::setprecision(1) << step_tps
|
||||||
<< " ema=" << std::fixed << std::setprecision(1) << step_ema;
|
<< " ema=" << std::fixed << std::setprecision(1) << step_ema;
|
||||||
|
if (has_dflash_target_only_arm()) {
|
||||||
|
oss << " configured_n_max=" << configured_n_max;
|
||||||
|
}
|
||||||
for (const auto & coord : coords) {
|
for (const auto & coord : coords) {
|
||||||
bool is_int = (coord.name != "p_min");
|
bool is_int = (coord.name != "p_min");
|
||||||
oss << " " << coord.name << "=";
|
oss << " " << coord.name << "=";
|
||||||
if (is_int) oss << (int)coord.arms[coord.current_idx].value;
|
if (is_int) oss << (int)coord.arms[coord.current_idx].value;
|
||||||
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.current_idx].value;
|
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.current_idx].value;
|
||||||
|
if (coord.name == "n_max" && spec_type == COMMON_SPECULATIVE_TYPE_DFLASH) {
|
||||||
|
oss << "(target_only=" << ((int) coord.arms[coord.current_idx].value == 0 ? "true" : "false") << ")";
|
||||||
|
}
|
||||||
oss << "→best=";
|
oss << "→best=";
|
||||||
if (is_int) oss << (int)coord.arms[coord.best_idx].value;
|
if (is_int) oss << (int)coord.arms[coord.best_idx].value;
|
||||||
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.best_idx].value;
|
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.best_idx].value;
|
||||||
|
|
@ -332,6 +469,14 @@ void spec_tuner::print_best() const {
|
||||||
|
|
||||||
if (n_resets > 0) oss << " resets=" << n_resets;
|
if (n_resets > 0) oss << " resets=" << n_resets;
|
||||||
if (n_low > 0) oss << " n_low=" << n_low;
|
if (n_low > 0) oss << " n_low=" << n_low;
|
||||||
|
if (has_dflash_target_only_arm()) {
|
||||||
|
oss << " configured_n_max=" << configured_n_max
|
||||||
|
<< " target_only=" << n_target_only_selections
|
||||||
|
<< " dflash=" << n_dflash_selections
|
||||||
|
<< " exploratory=" << n_exploratory_selections
|
||||||
|
<< " quarantines=" << n_quarantines
|
||||||
|
<< " probes=" << n_recovery_probes;
|
||||||
|
}
|
||||||
|
|
||||||
oss << " best:";
|
oss << " best:";
|
||||||
for (const auto & coord : coords) {
|
for (const auto & coord : coords) {
|
||||||
|
|
@ -345,15 +490,23 @@ void spec_tuner::print_best() const {
|
||||||
|
|
||||||
if (!coords.empty()) {
|
if (!coords.empty()) {
|
||||||
oss << " | n_max arms:";
|
oss << " | n_max arms:";
|
||||||
for (const auto & arm : coords[0].arms) {
|
for (size_t arm_idx = 0; arm_idx < coords[0].arms.size(); ++arm_idx) {
|
||||||
|
const auto & arm = coords[0].arms[arm_idx];
|
||||||
|
const bool quarantined = has_dflash_target_only_arm() && arm_idx < dflash_quarantined.size()
|
||||||
|
? dflash_quarantined[arm_idx]
|
||||||
|
: false;
|
||||||
oss << " " << (int)arm.value << "(Q=" << std::fixed << std::setprecision(2) << arm.Q
|
oss << " " << (int)arm.value << "(Q=" << std::fixed << std::setprecision(2) << arm.Q
|
||||||
<< ",N=" << arm.N << ")";
|
<< ",N=" << arm.N << ",quarantined=" << (quarantined ? "true" : "false") << ")";
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
oss << " tuner=" << std::fixed << std::setprecision(3) << t_tuner_us / 1000.0 << "ms";
|
oss << " tuner=" << std::fixed << std::setprecision(3) << t_tuner_us / 1000.0 << "ms";
|
||||||
|
if (has_dflash_target_only_arm()) {
|
||||||
|
LOG_INF("%s\n", oss.str().c_str());
|
||||||
|
} else {
|
||||||
LOG_DBG("%s\n", oss.str().c_str());
|
LOG_DBG("%s\n", oss.str().c_str());
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
{
|
{
|
||||||
std::ostringstream oss;
|
std::ostringstream oss;
|
||||||
|
|
@ -364,7 +517,11 @@ void spec_tuner::print_best() const {
|
||||||
oss << (first_kv ? ':' : ',') << coord.name << '=';
|
oss << (first_kv ? ':' : ',') << coord.name << '=';
|
||||||
first_kv = false;
|
first_kv = false;
|
||||||
|
|
||||||
if (is_int) oss << (int)coord.arms[coord.best_idx].value;
|
int reuse_value = is_int ? (int) coord.arms[coord.best_idx].value : 0;
|
||||||
|
if (coord.name == "n_max" && spec_type == COMMON_SPECULATIVE_TYPE_DFLASH && reuse_value == 0) {
|
||||||
|
reuse_value = configured_n_max;
|
||||||
|
}
|
||||||
|
if (is_int) oss << reuse_value;
|
||||||
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.best_idx].value;
|
else oss << std::fixed << std::setprecision(2) << coord.arms[coord.best_idx].value;
|
||||||
}
|
}
|
||||||
LOG_INF("%s\n", oss.str().c_str());
|
LOG_INF("%s\n", oss.str().c_str());
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,20 @@ struct spec_tuner {
|
||||||
|
|
||||||
double epsilon = 0.15; // 15% explore, 85% exploit
|
double epsilon = 0.15; // 15% explore, 85% exploit
|
||||||
|
|
||||||
|
int configured_n_max = 0;
|
||||||
|
int dflash_min_samples_per_arm = 3;
|
||||||
|
int dflash_recovery_probe_interval = 128;
|
||||||
|
double dflash_quarantine_ratio = 0.90;
|
||||||
|
std::vector<bool> dflash_quarantined;
|
||||||
|
int dflash_probe_cursor = 0;
|
||||||
|
bool dflash_last_exploratory = false;
|
||||||
|
bool dflash_last_recovery_probe = false;
|
||||||
|
uint64_t n_target_only_selections = 0;
|
||||||
|
uint64_t n_dflash_selections = 0;
|
||||||
|
uint64_t n_exploratory_selections = 0;
|
||||||
|
uint64_t n_quarantines = 0;
|
||||||
|
uint64_t n_recovery_probes = 0;
|
||||||
|
|
||||||
// task-change detection (per-call)
|
// task-change detection (per-call)
|
||||||
// If tuner goes bad for 30 consecutive calls, reset the tuner.
|
// If tuner goes bad for 30 consecutive calls, reset the tuner.
|
||||||
double step_ema = 0.0;
|
double step_ema = 0.0;
|
||||||
|
|
@ -66,4 +80,12 @@ struct spec_tuner {
|
||||||
void reset_exploration();
|
void reset_exploration();
|
||||||
|
|
||||||
void write_best(common_params_speculative & params) const;
|
void write_best(common_params_speculative & params) const;
|
||||||
|
|
||||||
|
bool has_dflash_target_only_arm() const {
|
||||||
|
return enabled && spec_type == COMMON_SPECULATIVE_TYPE_DFLASH && configured_n_max > 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
private:
|
||||||
|
int select_dflash_arm(spec_tuner_coord & coord);
|
||||||
|
void update_dflash_quarantine();
|
||||||
};
|
};
|
||||||
|
|
|
||||||
|
|
@ -1067,6 +1067,7 @@ struct common_speculative {
|
||||||
std::unique_ptr<spec_tuner> tuner;
|
std::unique_ptr<spec_tuner> tuner;
|
||||||
int last_n_drafted = 0;
|
int last_n_drafted = 0;
|
||||||
int64_t t_step_start_us = 0;
|
int64_t t_step_start_us = 0;
|
||||||
|
bool last_step_target_only = false;
|
||||||
};
|
};
|
||||||
|
|
||||||
static bool common_speculative_stage_chain_matches(
|
static bool common_speculative_stage_chain_matches(
|
||||||
|
|
@ -1474,7 +1475,7 @@ common_speculative * common_speculative_init(
|
||||||
if (actual_type != COMMON_SPECULATIVE_TYPE_NONE &&
|
if (actual_type != COMMON_SPECULATIVE_TYPE_NONE &&
|
||||||
actual_type != COMMON_SPECULATIVE_TYPE_EAGLE3) {
|
actual_type != COMMON_SPECULATIVE_TYPE_EAGLE3) {
|
||||||
result->tuner = std::make_unique<spec_tuner>();
|
result->tuner = std::make_unique<spec_tuner>();
|
||||||
result->tuner->init(actual_type, params, llama_get_model(ctx_tgt));
|
result->tuner->init(actual_type, result->configs[0].params, llama_get_model(ctx_tgt));
|
||||||
LOG_DBG("Autotune initialized for %s, tuning %zu parameters\n",
|
LOG_DBG("Autotune initialized for %s, tuning %zu parameters\n",
|
||||||
common_speculative_type_to_str(actual_type).c_str(),
|
common_speculative_type_to_str(actual_type).c_str(),
|
||||||
result->tuner->coords.size());
|
result->tuner->coords.size());
|
||||||
|
|
@ -1518,6 +1519,7 @@ llama_tokens common_speculative_draft(
|
||||||
llama_tokens result;
|
llama_tokens result;
|
||||||
|
|
||||||
spec->t_step_start_us = ggml_time_us();
|
spec->t_step_start_us = ggml_time_us();
|
||||||
|
spec->last_step_target_only = false;
|
||||||
|
|
||||||
// apply autotune proposal if enabled
|
// apply autotune proposal if enabled
|
||||||
if (spec->tuner && spec->tuner->enabled) {
|
if (spec->tuner && spec->tuner->enabled) {
|
||||||
|
|
@ -1533,8 +1535,19 @@ llama_tokens common_speculative_draft(
|
||||||
auto & impl = spec->impls[i];
|
auto & impl = spec->impls[i];
|
||||||
const auto & runtime_stage = use_runtime_stage_overrides ? runtime_stages[i] : spec->configs[i].stage;
|
const auto & runtime_stage = use_runtime_stage_overrides ? runtime_stages[i] : spec->configs[i].stage;
|
||||||
common_params_speculative impl_params = common_speculative_get_runtime_params(spec->configs[i], params, runtime_stage);
|
common_params_speculative impl_params = common_speculative_get_runtime_params(spec->configs[i], params, runtime_stage);
|
||||||
|
if (spec->tuner && spec->tuner->enabled && impl->type == COMMON_SPECULATIVE_TYPE_DFLASH) {
|
||||||
|
impl_params.n_max = params.n_max;
|
||||||
|
}
|
||||||
result.clear();
|
result.clear();
|
||||||
|
|
||||||
|
if (spec->tuner && spec->tuner->has_dflash_target_only_arm() &&
|
||||||
|
impl->type == COMMON_SPECULATIVE_TYPE_DFLASH && impl_params.n_max == 0) {
|
||||||
|
spec->curr_impl = impl.get();
|
||||||
|
spec->last_step_target_only = true;
|
||||||
|
LOG_DBG("%s: selected DFlash target-only arm\n", __func__);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
{
|
{
|
||||||
common_time_meas tm(impl->t_draft_us, !impl->gen_perf);
|
common_time_meas tm(impl->t_draft_us, !impl->gen_perf);
|
||||||
impl->draft(impl_params, prompt_tgt, id_last, draft_base_pos, draft_seq_id, result);
|
impl->draft(impl_params, prompt_tgt, id_last, draft_base_pos, draft_seq_id, result);
|
||||||
|
|
@ -1847,10 +1860,18 @@ common_speculative_draft_result common_speculative_draft_ex(
|
||||||
result.type = spec != nullptr && spec->curr_impl != nullptr
|
result.type = spec != nullptr && spec->curr_impl != nullptr
|
||||||
? spec->curr_impl->type
|
? spec->curr_impl->type
|
||||||
: COMMON_SPECULATIVE_TYPE_NONE;
|
: COMMON_SPECULATIVE_TYPE_NONE;
|
||||||
|
result.target_only = spec != nullptr && spec->last_step_target_only;
|
||||||
|
|
||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
int common_speculative_get_configured_n_max(const common_speculative * spec) {
|
||||||
|
if (spec == nullptr || spec->tuner == nullptr || !spec->tuner->has_dflash_target_only_arm()) {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
return spec->tuner->configured_n_max;
|
||||||
|
}
|
||||||
|
|
||||||
static bool common_speculative_has_target_features(const common_speculative * spec) {
|
static bool common_speculative_has_target_features(const common_speculative * spec) {
|
||||||
return common_speculative_has_type(spec, COMMON_SPECULATIVE_TYPE_MTP) ||
|
return common_speculative_has_type(spec, COMMON_SPECULATIVE_TYPE_MTP) ||
|
||||||
common_speculative_has_type(spec, COMMON_SPECULATIVE_TYPE_DFLASH);
|
common_speculative_has_type(spec, COMMON_SPECULATIVE_TYPE_DFLASH);
|
||||||
|
|
@ -2681,6 +2702,7 @@ void common_speculative_clear_sequence(
|
||||||
spec->curr_impl = nullptr;
|
spec->curr_impl = nullptr;
|
||||||
spec->last_n_drafted = 0;
|
spec->last_n_drafted = 0;
|
||||||
spec->t_step_start_us = 0;
|
spec->t_step_start_us = 0;
|
||||||
|
spec->last_step_target_only = false;
|
||||||
}
|
}
|
||||||
|
|
||||||
common_speculative_clear_sequence_hidden(spec, seq_id);
|
common_speculative_clear_sequence_hidden(spec, seq_id);
|
||||||
|
|
@ -2903,6 +2925,10 @@ void common_speculative_context_shift(
|
||||||
llama_pos kv_keep,
|
llama_pos kv_keep,
|
||||||
llama_pos kv_discard,
|
llama_pos kv_discard,
|
||||||
llama_pos kv_past) {
|
llama_pos kv_past) {
|
||||||
|
if (spec != nullptr) {
|
||||||
|
spec->last_step_target_only = false;
|
||||||
|
spec->t_step_start_us = 0;
|
||||||
|
}
|
||||||
if (auto * ctx_mtp = common_speculative_get_companion_ctx(spec); ctx_mtp != nullptr) {
|
if (auto * ctx_mtp = common_speculative_get_companion_ctx(spec); ctx_mtp != nullptr) {
|
||||||
llama_kv_cache_seq_rm (ctx_mtp, seq_id, kv_keep, kv_keep + kv_discard);
|
llama_kv_cache_seq_rm (ctx_mtp, seq_id, kv_keep, kv_keep + kv_discard);
|
||||||
llama_kv_cache_seq_add(ctx_mtp, seq_id, kv_keep + kv_discard, kv_past, -kv_discard);
|
llama_kv_cache_seq_add(ctx_mtp, seq_id, kv_keep + kv_discard, kv_past, -kv_discard);
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,7 @@ struct common_speculative_checkpoint {
|
||||||
struct common_speculative_draft_result {
|
struct common_speculative_draft_result {
|
||||||
llama_tokens tokens;
|
llama_tokens tokens;
|
||||||
common_speculative_type type = COMMON_SPECULATIVE_TYPE_NONE;
|
common_speculative_type type = COMMON_SPECULATIVE_TYPE_NONE;
|
||||||
|
bool target_only = false;
|
||||||
};
|
};
|
||||||
|
|
||||||
// comma separated list of all types
|
// comma separated list of all types
|
||||||
|
|
@ -108,6 +109,8 @@ common_speculative_draft_result common_speculative_draft_ex(
|
||||||
llama_pos draft_base_pos = -1,
|
llama_pos draft_base_pos = -1,
|
||||||
llama_seq_id draft_seq_id = 0);
|
llama_seq_id draft_seq_id = 0);
|
||||||
|
|
||||||
|
int common_speculative_get_configured_n_max(const common_speculative * spec);
|
||||||
|
|
||||||
// informs the speculative decoder that n_accepted tokens were accepted by the target model
|
// informs the speculative decoder that n_accepted tokens were accepted by the target model
|
||||||
void common_speculative_accept(common_speculative * spec, uint16_t n_accepted);
|
void common_speculative_accept(common_speculative * spec, uint16_t n_accepted);
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1001,7 +1001,7 @@ int main(int argc, char ** argv) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (!draft.empty()) {
|
if (!draft.empty() || draft_result.target_only) {
|
||||||
llama_batch verify_batch = llama_batch_init((int) draft.size() + 1, 0, 1);
|
llama_batch verify_batch = llama_batch_init((int) draft.size() + 1, 0, 1);
|
||||||
std::vector<int> verify_indices;
|
std::vector<int> verify_indices;
|
||||||
verify_indices.reserve(draft.size() + 1);
|
verify_indices.reserve(draft.size() + 1);
|
||||||
|
|
|
||||||
|
|
@ -481,6 +481,7 @@ void server_slot::reset() {
|
||||||
prompt_batch_i1 = -1;
|
prompt_batch_i1 = -1;
|
||||||
n_sent_text = 0;
|
n_sent_text = 0;
|
||||||
drafted.clear();
|
drafted.clear();
|
||||||
|
spec_target_only = false;
|
||||||
i_batch_dft.clear();
|
i_batch_dft.clear();
|
||||||
spec_prompt_warmup_failed = false;
|
spec_prompt_warmup_failed = false;
|
||||||
n_sent_token_probs = 0;
|
n_sent_token_probs = 0;
|
||||||
|
|
@ -587,6 +588,10 @@ int server_slot::get_n_draft_max() const {
|
||||||
|
|
||||||
// determine the max draft that fits the current slot state
|
// determine the max draft that fits the current slot state
|
||||||
int n_draft_max = params.speculative.get_max_stage_n_max();
|
int n_draft_max = params.speculative.get_max_stage_n_max();
|
||||||
|
const int configured_dflash_n_max = common_speculative_get_configured_n_max(spec);
|
||||||
|
if (configured_dflash_n_max > 0) {
|
||||||
|
n_draft_max = configured_dflash_n_max;
|
||||||
|
}
|
||||||
|
|
||||||
// note: slot.prompt is not yet expanded with the `id` token sampled above
|
// note: slot.prompt is not yet expanded with the `id` token sampled above
|
||||||
// also, need to leave space for 1 extra token to allow context shifts
|
// also, need to leave space for 1 extra token to allow context shifts
|
||||||
|
|
@ -613,6 +618,7 @@ void server_slot::release() {
|
||||||
state = SLOT_STATE_IDLE;
|
state = SLOT_STATE_IDLE;
|
||||||
task.reset();
|
task.reset();
|
||||||
}
|
}
|
||||||
|
spec_target_only = false;
|
||||||
llama_decode_reset();
|
llama_decode_reset();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -3476,6 +3482,7 @@ void server_context::add_sampled_tokens() {
|
||||||
if (slot.state == SLOT_STATE_IDLE) {
|
if (slot.state == SLOT_STATE_IDLE) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
slot.spec_target_only = false;
|
||||||
|
|
||||||
// generate draft tokens in speculative decoding mode
|
// generate draft tokens in speculative decoding mode
|
||||||
// TODO: rework to have a single draft llama_context shared across all slots [TAG_SERVER_SPEC_REWORK]
|
// TODO: rework to have a single draft llama_context shared across all slots [TAG_SERVER_SPEC_REWORK]
|
||||||
|
|
@ -3503,6 +3510,7 @@ void server_context::add_sampled_tokens() {
|
||||||
draft_base_pos,
|
draft_base_pos,
|
||||||
slot.id);
|
slot.id);
|
||||||
llama_tokens & draft = draft_result.tokens;
|
llama_tokens & draft = draft_result.tokens;
|
||||||
|
slot.spec_target_only = draft_result.target_only;
|
||||||
|
|
||||||
const int n_draft_max = slot.get_n_draft_max();
|
const int n_draft_max = slot.get_n_draft_max();
|
||||||
|
|
||||||
|
|
@ -3522,13 +3530,16 @@ void server_context::add_sampled_tokens() {
|
||||||
slot.cache_tokens.push_back(slot.sampled);
|
slot.cache_tokens.push_back(slot.sampled);
|
||||||
|
|
||||||
const int min_usable_draft = slot.params.speculative.get_min_usable_stage_n_min();
|
const int min_usable_draft = slot.params.speculative.get_min_usable_stage_n_min();
|
||||||
if (min_usable_draft > (int)draft.size()) {
|
if (!slot.spec_target_only && min_usable_draft > (int)draft.size()) {
|
||||||
SLT_DBG(slot, "ignoring small draft: %d < %d\n", (int)draft.size(), min_usable_draft);
|
SLT_DBG(slot, "ignoring small draft: %d < %d\n", (int)draft.size(), min_usable_draft);
|
||||||
// fallback to normal decoding
|
// fallback to normal decoding
|
||||||
slot.i_batch = slot.i_batch_dft[0];
|
slot.i_batch = slot.i_batch_dft[0];
|
||||||
slot.drafted.clear();
|
slot.drafted.clear();
|
||||||
slot.i_batch_dft.clear();
|
slot.i_batch_dft.clear();
|
||||||
} else {
|
} else {
|
||||||
|
if (slot.spec_target_only) {
|
||||||
|
SLT_DBG(slot, "%s\n", "selected DFlash target-only arm: root-only target batch");
|
||||||
|
}
|
||||||
// keep track of total number of drafted tokens tested
|
// keep track of total number of drafted tokens tested
|
||||||
slot.n_draft_total += draft.size();
|
slot.n_draft_total += draft.size();
|
||||||
if (slot.n_draft_by_depth.size() < draft.size()) {
|
if (slot.n_draft_by_depth.size() < draft.size()) {
|
||||||
|
|
@ -4221,6 +4232,7 @@ void server_context::speculative_decoding_accept() {
|
||||||
n_draft,
|
n_draft,
|
||||||
spec_pos_base,
|
spec_pos_base,
|
||||||
accepted_output_indices);
|
accepted_output_indices);
|
||||||
|
slot.spec_target_only = false;
|
||||||
|
|
||||||
for (size_t i = 0; i < ids.size(); ++i) {
|
for (size_t i = 0; i < ids.size(); ++i) {
|
||||||
completion_token_output result;
|
completion_token_output result;
|
||||||
|
|
|
||||||
|
|
@ -126,6 +126,7 @@ struct server_slot {
|
||||||
// sampling
|
// sampling
|
||||||
llama_token sampled; // in speculative mode, this is the last accepted token
|
llama_token sampled; // in speculative mode, this is the last accepted token
|
||||||
llama_tokens drafted;
|
llama_tokens drafted;
|
||||||
|
bool spec_target_only = false;
|
||||||
|
|
||||||
json json_schema;
|
json json_schema;
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue