mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-09-18 16:55:14 +02:00
sd: sync to master-529-630ee03 (#2040)
This commit is contained in:
@@ -603,87 +603,6 @@ inline std::vector<int> generate_scm_mask(
|
||||
return mask;
|
||||
}
|
||||
|
||||
inline std::vector<int> get_scm_preset(const std::string& preset, int total_steps) {
|
||||
struct Preset {
|
||||
std::vector<int> compute_bins;
|
||||
std::vector<int> cache_bins;
|
||||
};
|
||||
|
||||
Preset slow = {{8, 3, 3, 2, 1, 1}, {1, 2, 2, 2, 3}};
|
||||
Preset medium = {{6, 2, 2, 2, 2, 1}, {1, 3, 3, 3, 3}};
|
||||
Preset fast = {{6, 1, 1, 1, 1, 1}, {1, 3, 4, 5, 4}};
|
||||
Preset ultra = {{4, 1, 1, 1, 1}, {2, 5, 6, 7}};
|
||||
|
||||
Preset* p = nullptr;
|
||||
if (preset == "slow" || preset == "s" || preset == "S")
|
||||
p = &slow;
|
||||
else if (preset == "medium" || preset == "m" || preset == "M")
|
||||
p = &medium;
|
||||
else if (preset == "fast" || preset == "f" || preset == "F")
|
||||
p = &fast;
|
||||
else if (preset == "ultra" || preset == "u" || preset == "U")
|
||||
p = &ultra;
|
||||
else
|
||||
return {};
|
||||
|
||||
if (total_steps != 28 && total_steps > 0) {
|
||||
float scale = static_cast<float>(total_steps) / 28.0f;
|
||||
std::vector<int> scaled_compute, scaled_cache;
|
||||
|
||||
for (int v : p->compute_bins) {
|
||||
scaled_compute.push_back(std::max(1, static_cast<int>(v * scale + 0.5f)));
|
||||
}
|
||||
for (int v : p->cache_bins) {
|
||||
scaled_cache.push_back(std::max(1, static_cast<int>(v * scale + 0.5f)));
|
||||
}
|
||||
|
||||
return generate_scm_mask(scaled_compute, scaled_cache, total_steps);
|
||||
}
|
||||
|
||||
return generate_scm_mask(p->compute_bins, p->cache_bins, total_steps);
|
||||
}
|
||||
|
||||
inline float get_preset_threshold(const std::string& preset) {
|
||||
if (preset == "slow" || preset == "s" || preset == "S")
|
||||
return 0.20f;
|
||||
if (preset == "medium" || preset == "m" || preset == "M")
|
||||
return 0.25f;
|
||||
if (preset == "fast" || preset == "f" || preset == "F")
|
||||
return 0.30f;
|
||||
if (preset == "ultra" || preset == "u" || preset == "U")
|
||||
return 0.34f;
|
||||
return 0.08f;
|
||||
}
|
||||
|
||||
inline int get_preset_warmup(const std::string& preset) {
|
||||
if (preset == "slow" || preset == "s" || preset == "S")
|
||||
return 8;
|
||||
if (preset == "medium" || preset == "m" || preset == "M")
|
||||
return 6;
|
||||
if (preset == "fast" || preset == "f" || preset == "F")
|
||||
return 6;
|
||||
if (preset == "ultra" || preset == "u" || preset == "U")
|
||||
return 4;
|
||||
return 8;
|
||||
}
|
||||
|
||||
inline int get_preset_Fn(const std::string& preset) {
|
||||
if (preset == "slow" || preset == "s" || preset == "S")
|
||||
return 8;
|
||||
if (preset == "medium" || preset == "m" || preset == "M")
|
||||
return 8;
|
||||
if (preset == "fast" || preset == "f" || preset == "F")
|
||||
return 6;
|
||||
if (preset == "ultra" || preset == "u" || preset == "U")
|
||||
return 4;
|
||||
return 8;
|
||||
}
|
||||
|
||||
inline int get_preset_Bn(const std::string& preset) {
|
||||
(void)preset;
|
||||
return 0;
|
||||
}
|
||||
|
||||
inline void parse_dbcache_options(const std::string& opts, DBCacheConfig& cfg) {
|
||||
if (opts.empty())
|
||||
return;
|
||||
|
||||
@@ -1052,7 +1052,6 @@ struct SDGenerationParams {
|
||||
|
||||
std::string cache_mode;
|
||||
std::string cache_option;
|
||||
std::string cache_preset;
|
||||
std::string scm_mask;
|
||||
bool scm_policy_dynamic = true;
|
||||
sd_cache_params_t cache_params{};
|
||||
@@ -1466,21 +1465,6 @@ struct SDGenerationParams {
|
||||
return 1;
|
||||
};
|
||||
|
||||
auto on_cache_preset_arg = [&](int argc, const char** argv, int index) {
|
||||
if (++index >= argc) {
|
||||
return -1;
|
||||
}
|
||||
cache_preset = argv_to_utf8(index, argv);
|
||||
if (cache_preset != "slow" && cache_preset != "s" && cache_preset != "S" &&
|
||||
cache_preset != "medium" && cache_preset != "m" && cache_preset != "M" &&
|
||||
cache_preset != "fast" && cache_preset != "f" && cache_preset != "F" &&
|
||||
cache_preset != "ultra" && cache_preset != "u" && cache_preset != "U") {
|
||||
fprintf(stderr, "error: invalid cache preset '%s', must be 'slow'/'s', 'medium'/'m', 'fast'/'f', or 'ultra'/'u'\n", cache_preset.c_str());
|
||||
return -1;
|
||||
}
|
||||
return 1;
|
||||
};
|
||||
|
||||
options.manual_options = {
|
||||
{"-s",
|
||||
"--seed",
|
||||
@@ -1518,16 +1502,12 @@ struct SDGenerationParams {
|
||||
on_ref_image_arg},
|
||||
{"",
|
||||
"--cache-mode",
|
||||
"caching method: 'easycache' (DiT), 'ucache' (UNET), 'dbcache'/'taylorseer'/'cache-dit' (DiT block-level)",
|
||||
"caching method: 'easycache' (DiT), 'ucache' (UNET), 'dbcache'/'taylorseer'/'cache-dit' (DiT block-level), 'spectrum' (UNET/DiT Chebyshev+Taylor forecasting)",
|
||||
on_cache_mode_arg},
|
||||
{"",
|
||||
"--cache-option",
|
||||
"named cache params (key=value format, comma-separated). easycache/ucache: threshold=,start=,end=,decay=,relative=,reset=; dbcache/taylorseer/cache-dit: Fn=,Bn=,threshold=,warmup=. Examples: \"threshold=0.25\" or \"threshold=1.5,reset=0\"",
|
||||
"named cache params (key=value format, comma-separated). easycache/ucache: threshold=,start=,end=,decay=,relative=,reset=; dbcache/taylorseer/cache-dit: Fn=,Bn=,threshold=,warmup=; spectrum: w=,m=,lam=,window=,flex=,warmup=,stop=. Examples: \"threshold=0.25\" or \"threshold=1.5,reset=0\"",
|
||||
on_cache_option_arg},
|
||||
{"",
|
||||
"--cache-preset",
|
||||
"cache-dit preset: 'slow'/'s', 'medium'/'m', 'fast'/'f', 'ultra'/'u'",
|
||||
on_cache_preset_arg},
|
||||
{"",
|
||||
"--scm-mask",
|
||||
"SCM steps mask for cache-dit: comma-separated 0/1 (e.g., \"1,1,1,0,0,1,0,0,1,0\") - 1=compute, 0=can cache",
|
||||
@@ -1580,7 +1560,6 @@ struct SDGenerationParams {
|
||||
load_if_exists("negative_prompt", negative_prompt);
|
||||
load_if_exists("cache_mode", cache_mode);
|
||||
load_if_exists("cache_option", cache_option);
|
||||
load_if_exists("cache_preset", cache_preset);
|
||||
load_if_exists("scm_mask", scm_mask);
|
||||
|
||||
load_if_exists("clip_skip", clip_skip);
|
||||
@@ -1815,48 +1794,17 @@ struct SDGenerationParams {
|
||||
|
||||
if (!cache_mode.empty()) {
|
||||
if (cache_mode == "easycache") {
|
||||
cache_params.mode = SD_CACHE_EASYCACHE;
|
||||
cache_params.reuse_threshold = 0.2f;
|
||||
cache_params.start_percent = 0.15f;
|
||||
cache_params.end_percent = 0.95f;
|
||||
cache_params.error_decay_rate = 1.0f;
|
||||
cache_params.use_relative_threshold = true;
|
||||
cache_params.reset_error_on_compute = true;
|
||||
cache_params.mode = SD_CACHE_EASYCACHE;
|
||||
} else if (cache_mode == "ucache") {
|
||||
cache_params.mode = SD_CACHE_UCACHE;
|
||||
cache_params.reuse_threshold = 1.0f;
|
||||
cache_params.start_percent = 0.15f;
|
||||
cache_params.end_percent = 0.95f;
|
||||
cache_params.error_decay_rate = 1.0f;
|
||||
cache_params.use_relative_threshold = true;
|
||||
cache_params.reset_error_on_compute = true;
|
||||
cache_params.mode = SD_CACHE_UCACHE;
|
||||
} else if (cache_mode == "dbcache") {
|
||||
cache_params.mode = SD_CACHE_DBCACHE;
|
||||
cache_params.Fn_compute_blocks = 8;
|
||||
cache_params.Bn_compute_blocks = 0;
|
||||
cache_params.residual_diff_threshold = 0.08f;
|
||||
cache_params.max_warmup_steps = 8;
|
||||
cache_params.mode = SD_CACHE_DBCACHE;
|
||||
} else if (cache_mode == "taylorseer") {
|
||||
cache_params.mode = SD_CACHE_TAYLORSEER;
|
||||
cache_params.Fn_compute_blocks = 8;
|
||||
cache_params.Bn_compute_blocks = 0;
|
||||
cache_params.residual_diff_threshold = 0.08f;
|
||||
cache_params.max_warmup_steps = 8;
|
||||
cache_params.mode = SD_CACHE_TAYLORSEER;
|
||||
} else if (cache_mode == "cache-dit") {
|
||||
cache_params.mode = SD_CACHE_CACHE_DIT;
|
||||
cache_params.Fn_compute_blocks = 8;
|
||||
cache_params.Bn_compute_blocks = 0;
|
||||
cache_params.residual_diff_threshold = 0.08f;
|
||||
cache_params.max_warmup_steps = 8;
|
||||
cache_params.mode = SD_CACHE_CACHE_DIT;
|
||||
} else if (cache_mode == "spectrum") {
|
||||
cache_params.mode = SD_CACHE_SPECTRUM;
|
||||
cache_params.spectrum_w = 0.40f;
|
||||
cache_params.spectrum_m = 3;
|
||||
cache_params.spectrum_lam = 1.0f;
|
||||
cache_params.spectrum_window_size = 2;
|
||||
cache_params.spectrum_flex_window = 0.50f;
|
||||
cache_params.spectrum_warmup_steps = 4;
|
||||
cache_params.spectrum_stop_percent = 0.9f;
|
||||
cache_params.mode = SD_CACHE_SPECTRUM;
|
||||
}
|
||||
|
||||
if (!cache_option.empty()) {
|
||||
|
||||
@@ -831,8 +831,6 @@ static void parse_cache_options(sd_cache_params_t & params, const std::string& c
|
||||
params.mode = SD_CACHE_EASYCACHE;
|
||||
} else if (cache_mode == "ucache") {
|
||||
params.mode = SD_CACHE_UCACHE;
|
||||
// this is the only difference from the defaults right now
|
||||
params.reuse_threshold = 1.0f;
|
||||
} else if (cache_mode == "dbcache") {
|
||||
params.mode = SD_CACHE_DBCACHE;
|
||||
} else if (cache_mode == "taylorseer") {
|
||||
|
||||
@@ -100,6 +100,19 @@ void suppress_pp(int step, int steps, float time, void* data) {
|
||||
return;
|
||||
}
|
||||
|
||||
static float get_cache_reuse_threshold(const sd_cache_params_t& params) {
|
||||
float reuse_threshold = params.reuse_threshold;
|
||||
if (reuse_threshold == INFINITY) {
|
||||
if (params.mode == SD_CACHE_EASYCACHE) {
|
||||
reuse_threshold = 0.2;
|
||||
}
|
||||
else if (params.mode == SD_CACHE_UCACHE) {
|
||||
reuse_threshold = 1.0;
|
||||
}
|
||||
}
|
||||
return std::max(0.0f, reuse_threshold);
|
||||
}
|
||||
|
||||
/*=============================================== StableDiffusionGGML ================================================*/
|
||||
|
||||
class StableDiffusionGGML {
|
||||
@@ -1900,7 +1913,7 @@ public:
|
||||
} else {
|
||||
EasyCacheConfig easycache_config;
|
||||
easycache_config.enabled = true;
|
||||
easycache_config.reuse_threshold = std::max(0.0f, cache_params->reuse_threshold);
|
||||
easycache_config.reuse_threshold = get_cache_reuse_threshold(*cache_params);
|
||||
easycache_config.start_percent = cache_params->start_percent;
|
||||
easycache_config.end_percent = cache_params->end_percent;
|
||||
easycache_state.init(easycache_config, denoiser.get());
|
||||
@@ -1921,7 +1934,7 @@ public:
|
||||
} else {
|
||||
UCacheConfig ucache_config;
|
||||
ucache_config.enabled = true;
|
||||
ucache_config.reuse_threshold = std::max(0.0f, cache_params->reuse_threshold);
|
||||
ucache_config.reuse_threshold = get_cache_reuse_threshold(*cache_params);
|
||||
ucache_config.start_percent = cache_params->start_percent;
|
||||
ucache_config.end_percent = cache_params->end_percent;
|
||||
ucache_config.error_decay_rate = std::max(0.0f, std::min(1.0f, cache_params->error_decay_rate));
|
||||
@@ -1982,9 +1995,9 @@ public:
|
||||
}
|
||||
}
|
||||
} else if (cache_params->mode == SD_CACHE_SPECTRUM) {
|
||||
bool spectrum_supported = sd_version_is_unet(version);
|
||||
bool spectrum_supported = sd_version_is_unet(version) || sd_version_is_dit(version);
|
||||
if (!spectrum_supported) {
|
||||
LOG_WARN("Spectrum requested but not supported for this model type (only UNET models)");
|
||||
LOG_WARN("Spectrum requested but not supported for this model type (only UNET and DiT models)");
|
||||
} else {
|
||||
SpectrumConfig spectrum_config;
|
||||
spectrum_config.w = cache_params->spectrum_w;
|
||||
@@ -3188,7 +3201,7 @@ enum lora_apply_mode_t str_to_lora_apply_mode(const char* str) {
|
||||
void sd_cache_params_init(sd_cache_params_t* cache_params) {
|
||||
*cache_params = {};
|
||||
cache_params->mode = SD_CACHE_DISABLED;
|
||||
cache_params->reuse_threshold = 1.0f;
|
||||
cache_params->reuse_threshold = INFINITY;
|
||||
cache_params->start_percent = 0.15f;
|
||||
cache_params->end_percent = 0.95f;
|
||||
cache_params->error_decay_rate = 1.0f;
|
||||
@@ -3434,7 +3447,7 @@ char* sd_img_gen_params_to_str(const sd_img_gen_params_t* sd_img_gen_params) {
|
||||
snprintf(buf + strlen(buf), 4096 - strlen(buf),
|
||||
"cache: %s (threshold=%.3f, start=%.2f, end=%.2f)\n",
|
||||
cache_mode_str,
|
||||
sd_img_gen_params->cache.reuse_threshold,
|
||||
get_cache_reuse_threshold(sd_img_gen_params->cache),
|
||||
sd_img_gen_params->cache.start_percent,
|
||||
sd_img_gen_params->cache.end_percent);
|
||||
free(sample_params_str);
|
||||
|
||||
Reference in New Issue
Block a user