Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
123 changes: 66 additions & 57 deletions src/stable-diffusion.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<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)) {
Expand Down Expand Up @@ -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<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)) {
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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)) {
Expand Down
Loading