mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-08-10 14:16:50 +00:00
refactor: extract model loader initialization (#1844)
This commit is contained in:
parent
eb7f35ca49
commit
db99efdd6d
@ -696,45 +696,11 @@ public:
|
|||||||
LOG_DEBUG("loaded alphas_cumprod from model file");
|
LOG_DEBUG("loaded alphas_cumprod from model file");
|
||||||
}
|
}
|
||||||
|
|
||||||
bool init(const sd_ctx_params_t* sd_ctx_params) {
|
bool init_model_loader(ModelLoader& model_loader,
|
||||||
n_threads = sd_ctx_params->n_threads;
|
const sd_ctx_params_t* sd_ctx_params,
|
||||||
enable_mmap = sd_ctx_params->enable_mmap;
|
bool& use_tae,
|
||||||
stream_layers = sd_ctx_params->stream_layers;
|
bool& use_audio_vae,
|
||||||
eager_load = sd_ctx_params->eager_load;
|
bool& use_control_net) {
|
||||||
backend_spec = SAFE_STR(sd_ctx_params->backend);
|
|
||||||
params_backend_spec = SAFE_STR(sd_ctx_params->params_backend);
|
|
||||||
split_mode_spec = SAFE_STR(sd_ctx_params->split_mode);
|
|
||||||
auto_fit_enabled = sd_ctx_params->auto_fit;
|
|
||||||
max_vram_assignment.reset(0.f);
|
|
||||||
{
|
|
||||||
std::string error;
|
|
||||||
if (!max_vram_assignment.parse(SAFE_STR(sd_ctx_params->max_vram), &error)) {
|
|
||||||
LOG_ERROR("%s", error.c_str());
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
std::string rpc_servers_spec = SAFE_STR(sd_ctx_params->rpc_servers);
|
|
||||||
add_rpc_devices(rpc_servers_spec);
|
|
||||||
|
|
||||||
bool use_tae = false;
|
|
||||||
bool use_audio_vae = false;
|
|
||||||
bool use_control_net = false;
|
|
||||||
|
|
||||||
rng = get_rng(sd_ctx_params->rng_type);
|
|
||||||
if (sd_ctx_params->sampler_rng_type != RNG_TYPE_COUNT && sd_ctx_params->sampler_rng_type != sd_ctx_params->rng_type) {
|
|
||||||
sampler_rng = get_rng(sd_ctx_params->sampler_rng_type);
|
|
||||||
} else {
|
|
||||||
sampler_rng = rng;
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_log_set(ggml_log_callback_default, nullptr);
|
|
||||||
|
|
||||||
model_manager = std::make_shared<ModelManager>();
|
|
||||||
model_manager->set_n_threads(n_threads);
|
|
||||||
model_manager->set_enable_mmap(enable_mmap);
|
|
||||||
ModelLoader& model_loader = model_manager->loader();
|
|
||||||
|
|
||||||
if (strlen(SAFE_STR(sd_ctx_params->model_path)) > 0) {
|
if (strlen(SAFE_STR(sd_ctx_params->model_path)) > 0) {
|
||||||
LOG_INFO("loading model from '%s'", sd_ctx_params->model_path);
|
LOG_INFO("loading model from '%s'", sd_ctx_params->model_path);
|
||||||
if (!model_loader.init_from_file(sd_ctx_params->model_path)) {
|
if (!model_loader.init_from_file(sd_ctx_params->model_path)) {
|
||||||
@ -874,24 +840,69 @@ public:
|
|||||||
|
|
||||||
model_loader.convert_tensors_name();
|
model_loader.convert_tensors_name();
|
||||||
|
|
||||||
version = model_loader.get_sd_version();
|
|
||||||
if (version == VERSION_COUNT) {
|
|
||||||
LOG_ERROR("get sd version from file failed: '%s'", SAFE_STR(sd_ctx_params->model_path));
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
auto& tensor_storage_map = model_loader.get_tensor_storage_map();
|
|
||||||
|
|
||||||
LOG_INFO("Version: %s ", model_version_to_str[version]);
|
|
||||||
ggml_type wtype = sd_type_to_ggml_type(sd_ctx_params->wtype);
|
ggml_type wtype = sd_type_to_ggml_type(sd_ctx_params->wtype);
|
||||||
std::string tensor_type_rules = SAFE_STR(sd_ctx_params->tensor_type_rules);
|
std::string tensor_type_rules = SAFE_STR(sd_ctx_params->tensor_type_rules);
|
||||||
if (wtype != GGML_TYPE_COUNT || tensor_type_rules.size() > 0) {
|
if (wtype != GGML_TYPE_COUNT || tensor_type_rules.size() > 0) {
|
||||||
model_loader.set_wtype_override(wtype, tensor_type_rules);
|
model_loader.set_wtype_override(wtype, tensor_type_rules);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
bool init(const sd_ctx_params_t* sd_ctx_params) {
|
||||||
|
n_threads = sd_ctx_params->n_threads;
|
||||||
|
enable_mmap = sd_ctx_params->enable_mmap;
|
||||||
|
stream_layers = sd_ctx_params->stream_layers;
|
||||||
|
eager_load = sd_ctx_params->eager_load;
|
||||||
|
backend_spec = SAFE_STR(sd_ctx_params->backend);
|
||||||
|
params_backend_spec = SAFE_STR(sd_ctx_params->params_backend);
|
||||||
|
split_mode_spec = SAFE_STR(sd_ctx_params->split_mode);
|
||||||
|
auto_fit_enabled = sd_ctx_params->auto_fit;
|
||||||
|
max_vram_assignment.reset(0.f);
|
||||||
|
{
|
||||||
|
std::string error;
|
||||||
|
if (!max_vram_assignment.parse(SAFE_STR(sd_ctx_params->max_vram), &error)) {
|
||||||
|
LOG_ERROR("%s", error.c_str());
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
std::string rpc_servers_spec = SAFE_STR(sd_ctx_params->rpc_servers);
|
||||||
|
add_rpc_devices(rpc_servers_spec);
|
||||||
|
|
||||||
|
bool use_tae = false;
|
||||||
|
bool use_audio_vae = false;
|
||||||
|
bool use_control_net = false;
|
||||||
|
|
||||||
|
rng = get_rng(sd_ctx_params->rng_type);
|
||||||
|
if (sd_ctx_params->sampler_rng_type != RNG_TYPE_COUNT && sd_ctx_params->sampler_rng_type != sd_ctx_params->rng_type) {
|
||||||
|
sampler_rng = get_rng(sd_ctx_params->sampler_rng_type);
|
||||||
|
} else {
|
||||||
|
sampler_rng = rng;
|
||||||
|
}
|
||||||
|
|
||||||
|
ggml_log_set(ggml_log_callback_default, nullptr);
|
||||||
|
|
||||||
|
model_manager = std::make_shared<ModelManager>();
|
||||||
|
model_manager->set_n_threads(n_threads);
|
||||||
|
model_manager->set_enable_mmap(enable_mmap);
|
||||||
|
ModelLoader& model_loader = model_manager->loader();
|
||||||
|
|
||||||
|
if (!init_model_loader(model_loader, sd_ctx_params, use_tae, use_audio_vae, use_control_net)) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
version = model_loader.get_sd_version();
|
||||||
|
if (version == VERSION_COUNT) {
|
||||||
|
LOG_ERROR("get sd version from file failed: '%s'", SAFE_STR(sd_ctx_params->model_path));
|
||||||
|
return false;
|
||||||
|
} else {
|
||||||
|
LOG_INFO("Version: %s ", model_version_to_str[version]);
|
||||||
|
}
|
||||||
|
|
||||||
if (auto_fit_enabled) {
|
if (auto_fit_enabled) {
|
||||||
if (!sd::backend_fit::derive_backend_specs(model_loader,
|
if (!sd::backend_fit::derive_backend_specs(model_loader,
|
||||||
wtype,
|
sd_type_to_ggml_type(sd_ctx_params->wtype),
|
||||||
max_vram_assignment,
|
max_vram_assignment,
|
||||||
backend_spec,
|
backend_spec,
|
||||||
params_backend_spec)) {
|
params_backend_spec)) {
|
||||||
@ -946,14 +957,10 @@ public:
|
|||||||
|
|
||||||
if (sd_ctx_params->lora_apply_mode == LORA_APPLY_AUTO) {
|
if (sd_ctx_params->lora_apply_mode == LORA_APPLY_AUTO) {
|
||||||
bool have_quantized_weight = false;
|
bool have_quantized_weight = false;
|
||||||
if (wtype != GGML_TYPE_COUNT && ggml_is_quantized(wtype)) {
|
for (const auto& [type, _] : wtype_stat) {
|
||||||
have_quantized_weight = true;
|
if (ggml_is_quantized(type)) {
|
||||||
} else {
|
have_quantized_weight = true;
|
||||||
for (const auto& [type, _] : wtype_stat) {
|
break;
|
||||||
if (ggml_is_quantized(type)) {
|
|
||||||
have_quantized_weight = true;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Avoid full-model LoRA merge buffers on constrained setups.
|
// Avoid full-model LoRA merge buffers on constrained setups.
|
||||||
@ -997,6 +1004,8 @@ public:
|
|||||||
use_tae = true;
|
use_tae = true;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
auto& tensor_storage_map = model_loader.get_tensor_storage_map();
|
||||||
|
|
||||||
{
|
{
|
||||||
if (!ensure_backend_pair(SDBackendModule::TE) ||
|
if (!ensure_backend_pair(SDBackendModule::TE) ||
|
||||||
!ensure_backend_pair(SDBackendModule::DIFFUSION)) {
|
!ensure_backend_pair(SDBackendModule::DIFFUSION)) {
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user