mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-08-06 02:00:41 -05:00
refactor: extract model loader initialization (#1844)
This commit is contained in:
@@ -696,45 +696,11 @@ public:
|
||||
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<ModelManager>();
|
||||
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 @@ public:
|
||||
|
||||
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<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 (!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 @@ public:
|
||||
|
||||
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 @@ public:
|
||||
use_tae = true;
|
||||
}
|
||||
|
||||
auto& tensor_storage_map = model_loader.get_tensor_storage_map();
|
||||
|
||||
{
|
||||
if (!ensure_backend_pair(SDBackendModule::TE) ||
|
||||
!ensure_backend_pair(SDBackendModule::DIFFUSION)) {
|
||||
|
||||
Reference in New Issue
Block a user