From 6d1551b669676f002c682d59b0ace4e28f5480d1 Mon Sep 17 00:00:00 2001 From: leejet Date: Sun, 2 Aug 2026 16:42:47 +0800 Subject: [PATCH] refactor: extract model loader initialization --- src/stable-diffusion.cpp | 123 +++++++++++++++++++++------------------ 1 file changed, 66 insertions(+), 57 deletions(-) diff --git a/src/stable-diffusion.cpp b/src/stable-diffusion.cpp index 9153b52b8..5c34d4042 100644 --- a/src/stable-diffusion.cpp +++ b/src/stable-diffusion.cpp @@ -696,45 +696,11 @@ class StableDiffusionGGML { LOG_DEBUG("loaded alphas_cumprod from model file"); } - 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(); - model_manager->set_n_threads(n_threads); - model_manager->set_enable_mmap(enable_mmap); - ModelLoader& model_loader = model_manager->loader(); - + bool init_model_loader(ModelLoader& model_loader, + const sd_ctx_params_t* sd_ctx_params, + bool& use_tae, + bool& use_audio_vae, + bool& use_control_net) { if (strlen(SAFE_STR(sd_ctx_params->model_path)) > 0) { LOG_INFO("loading model from '%s'", sd_ctx_params->model_path); if (!model_loader.init_from_file(sd_ctx_params->model_path)) { @@ -874,24 +840,69 @@ class StableDiffusionGGML { 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); std::string tensor_type_rules = SAFE_STR(sd_ctx_params->tensor_type_rules); if (wtype != GGML_TYPE_COUNT || tensor_type_rules.size() > 0) { 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(); + 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 (!sd::backend_fit::derive_backend_specs(model_loader, - wtype, + sd_type_to_ggml_type(sd_ctx_params->wtype), max_vram_assignment, backend_spec, params_backend_spec)) { @@ -946,14 +957,10 @@ class StableDiffusionGGML { if (sd_ctx_params->lora_apply_mode == LORA_APPLY_AUTO) { bool have_quantized_weight = false; - if (wtype != GGML_TYPE_COUNT && ggml_is_quantized(wtype)) { - have_quantized_weight = true; - } else { - for (const auto& [type, _] : wtype_stat) { - if (ggml_is_quantized(type)) { - have_quantized_weight = true; - break; - } + for (const auto& [type, _] : wtype_stat) { + if (ggml_is_quantized(type)) { + have_quantized_weight = true; + break; } } // Avoid full-model LoRA merge buffers on constrained setups. @@ -997,6 +1004,8 @@ class StableDiffusionGGML { use_tae = true; } + auto& tensor_storage_map = model_loader.get_tensor_storage_map(); + { if (!ensure_backend_pair(SDBackendModule::TE) || !ensure_backend_pair(SDBackendModule::DIFFUSION)) {