mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-07-23 11:20:53 -05:00
feat: hot-reload ControlNet - swap without rebuilding the context (#1768)
This commit is contained in:
@@ -428,6 +428,11 @@ SD_API const char* sd_get_system_info();
|
||||
SD_API bool sd_ctx_supports_image_generation(const sd_ctx_t* sd_ctx);
|
||||
SD_API bool sd_ctx_supports_video_generation(const sd_ctx_t* sd_ctx);
|
||||
|
||||
// ControlNet hot-swap APIs are not safe to call while generation is in flight.
|
||||
SD_API bool sd_ctx_load_control_net(sd_ctx_t* sd_ctx, const char* path);
|
||||
SD_API bool sd_ctx_unload_control_net(sd_ctx_t* sd_ctx);
|
||||
SD_API bool sd_ctx_has_control_net(const sd_ctx_t* sd_ctx);
|
||||
|
||||
SD_API const char* sd_type_name(enum sd_type_t type);
|
||||
SD_API enum sd_type_t str_to_sd_type(const char* str);
|
||||
SD_API const char* sd_rng_type_name(enum rng_type_t rng_type);
|
||||
|
||||
@@ -179,6 +179,102 @@ bool ModelManager::register_param_tensors(const std::string& desc,
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ModelManager::unregister_param_tensors(const std::string& desc, size_t* registered_tensor_size) {
|
||||
if (desc.empty()) {
|
||||
return true;
|
||||
}
|
||||
|
||||
std::unordered_set<TensorState*> target_states;
|
||||
size_t released_size = 0;
|
||||
for (auto& state : tensor_states_) {
|
||||
if (state == nullptr || state->desc != desc) {
|
||||
continue;
|
||||
}
|
||||
if (state->active_prepare_count > 0) {
|
||||
LOG_ERROR("model manager cannot unregister active %s tensor '%s'",
|
||||
desc.c_str(),
|
||||
state->name.c_str());
|
||||
return false;
|
||||
}
|
||||
target_states.insert(state.get());
|
||||
if (state->tensor != nullptr) {
|
||||
released_size += ggml_nbytes(state->tensor);
|
||||
}
|
||||
}
|
||||
|
||||
if (target_states.empty()) {
|
||||
return true;
|
||||
}
|
||||
|
||||
release_compute_staging_blocks(false);
|
||||
|
||||
std::vector<ParamsStorageBlock*> storage_blocks_to_release;
|
||||
std::unordered_set<TensorState*> affected_storage_states;
|
||||
for (const auto& block : params_storage_blocks_) {
|
||||
if (block == nullptr) {
|
||||
continue;
|
||||
}
|
||||
bool has_target_state = false;
|
||||
for (TensorState* state : block->states) {
|
||||
if (state != nullptr && target_states.count(state) > 0) {
|
||||
has_target_state = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (!has_target_state) {
|
||||
continue;
|
||||
}
|
||||
storage_blocks_to_release.push_back(block.get());
|
||||
for (TensorState* state : block->states) {
|
||||
if (state != nullptr) {
|
||||
affected_storage_states.insert(state);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for (TensorState* state : affected_storage_states) {
|
||||
if (state == nullptr) {
|
||||
continue;
|
||||
}
|
||||
if (state->active_prepare_count > 0 || state->staged_to_compute_backend) {
|
||||
LOG_ERROR("model manager cannot unregister %s while tensor '%s' is active",
|
||||
desc.c_str(),
|
||||
state->name.c_str());
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
for (ParamsStorageBlock* block : storage_blocks_to_release) {
|
||||
if (block != nullptr) {
|
||||
free_params_storage_block(*block);
|
||||
erase_params_storage_block(block);
|
||||
}
|
||||
}
|
||||
|
||||
for (auto it = tensor_states_by_name_.begin(); it != tensor_states_by_name_.end();) {
|
||||
if (target_states.count(it->second) > 0) {
|
||||
it = tensor_states_by_name_.erase(it);
|
||||
} else {
|
||||
++it;
|
||||
}
|
||||
}
|
||||
tensor_states_.erase(std::remove_if(tensor_states_.begin(),
|
||||
tensor_states_.end(),
|
||||
[&](const std::unique_ptr<TensorState>& s) {
|
||||
return s == nullptr || target_states.count(s.get()) > 0;
|
||||
}),
|
||||
tensor_states_.end());
|
||||
|
||||
if (registered_tensor_size != nullptr) {
|
||||
if (released_size > *registered_tensor_size) {
|
||||
*registered_tensor_size = 0;
|
||||
} else {
|
||||
*registered_tensor_size -= released_size;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ModelManager::load_all_params_eagerly() {
|
||||
std::vector<TensorState*> all_states;
|
||||
all_states.reserve(tensor_states_.size());
|
||||
|
||||
@@ -134,6 +134,9 @@ public:
|
||||
bool allow_split_buffer = false,
|
||||
bool params_follow_compute_backend = false);
|
||||
|
||||
bool unregister_param_tensors(const std::string& desc,
|
||||
size_t* registered_tensor_size = nullptr);
|
||||
|
||||
template <typename Runner>
|
||||
bool register_runner_params(const std::string& desc,
|
||||
Runner& runner,
|
||||
|
||||
@@ -222,9 +222,13 @@ public:
|
||||
std::string split_mode_spec;
|
||||
bool auto_fit_enabled = false;
|
||||
|
||||
bool diffusion_conv_direct = false;
|
||||
|
||||
bool is_using_v_parameterization = false;
|
||||
bool is_using_edm_v_parameterization = false;
|
||||
|
||||
size_t control_net_params_mem_size = 0;
|
||||
|
||||
std::shared_ptr<ModelManager> model_manager;
|
||||
|
||||
std::shared_ptr<Denoiser> denoiser = std::make_shared<CompVisDenoiser>();
|
||||
@@ -494,6 +498,76 @@ public:
|
||||
params_follow_runtime);
|
||||
}
|
||||
|
||||
bool unload_control_net() {
|
||||
if (control_net == nullptr) {
|
||||
return true;
|
||||
}
|
||||
if (model_manager != nullptr) {
|
||||
if (!model_manager->unregister_param_tensors("ControlNet", &control_net_params_mem_size)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
control_net.reset();
|
||||
control_net_params_mem_size = 0;
|
||||
return true;
|
||||
}
|
||||
|
||||
bool load_control_net_from_file(const std::string& path) {
|
||||
if (path.empty()) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: empty path");
|
||||
return false;
|
||||
}
|
||||
if (model_manager == nullptr) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: model_manager not initialized");
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!unload_control_net()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
ModelLoader& shared_loader = model_manager->loader();
|
||||
if (!shared_loader.init_from_file(path)) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: failed to load '%s'", path.c_str());
|
||||
return false;
|
||||
}
|
||||
shared_loader.convert_tensors_name();
|
||||
|
||||
if (!ensure_backend_pair(SDBackendModule::CONTROL_NET)) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: control_net backend unavailable");
|
||||
return false;
|
||||
}
|
||||
|
||||
control_net = std::make_shared<ControlNet>(backend_for(SDBackendModule::CONTROL_NET),
|
||||
params_backend_for(SDBackendModule::CONTROL_NET),
|
||||
shared_loader.get_tensor_storage_map(),
|
||||
version,
|
||||
"",
|
||||
model_manager);
|
||||
if (diffusion_conv_direct) {
|
||||
LOG_INFO("Using Conv2d direct in the control net");
|
||||
control_net->set_conv2d_direct_enabled(true);
|
||||
}
|
||||
if (!register_runner_params("ControlNet",
|
||||
control_net,
|
||||
SDBackendModule::CONTROL_NET,
|
||||
&control_net_params_mem_size)) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: register_runner_params failed");
|
||||
control_net.reset();
|
||||
control_net_params_mem_size = 0;
|
||||
return false;
|
||||
}
|
||||
if (!model_manager->validate_registered_tensors()) {
|
||||
LOG_ERROR("sd_ctx_load_control_net: registered tensors validation failed");
|
||||
unload_control_net();
|
||||
return false;
|
||||
}
|
||||
LOG_INFO("sd_ctx_load_control_net: loaded '%s' (%.2f MB)",
|
||||
path.c_str(),
|
||||
control_net_params_mem_size / 1024.0 / 1024.0);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool init_backend() {
|
||||
std::string error;
|
||||
if (!backend_manager.init(backend_spec.c_str(),
|
||||
@@ -852,10 +926,12 @@ public:
|
||||
model_loader.process_model_files(enable_mmap, needs_writable_mmap);
|
||||
load_alphas_cumprod(model_loader);
|
||||
|
||||
diffusion_conv_direct = sd_ctx_params->diffusion_conv_direct;
|
||||
|
||||
size_t text_encoder_params_mem_size = 0;
|
||||
size_t unet_params_mem_size = 0;
|
||||
size_t vae_params_mem_size = 0;
|
||||
size_t control_net_params_mem_size = 0;
|
||||
control_net_params_mem_size = 0;
|
||||
size_t extension_params_mem_size = 0;
|
||||
|
||||
bool tae_preview_only = sd_ctx_params->tae_preview_only;
|
||||
@@ -3429,6 +3505,27 @@ SD_API bool sd_ctx_supports_video_generation(const sd_ctx_t* sd_ctx) {
|
||||
return sd_version_supports_video_generation(sd_ctx->sd->version);
|
||||
}
|
||||
|
||||
SD_API bool sd_ctx_load_control_net(sd_ctx_t* sd_ctx, const char* path) {
|
||||
if (sd_ctx == nullptr || sd_ctx->sd == nullptr || path == nullptr) {
|
||||
return false;
|
||||
}
|
||||
return sd_ctx->sd->load_control_net_from_file(path);
|
||||
}
|
||||
|
||||
SD_API bool sd_ctx_unload_control_net(sd_ctx_t* sd_ctx) {
|
||||
if (sd_ctx == nullptr || sd_ctx->sd == nullptr) {
|
||||
return false;
|
||||
}
|
||||
return sd_ctx->sd->unload_control_net();
|
||||
}
|
||||
|
||||
SD_API bool sd_ctx_has_control_net(const sd_ctx_t* sd_ctx) {
|
||||
if (sd_ctx == nullptr || sd_ctx->sd == nullptr) {
|
||||
return false;
|
||||
}
|
||||
return sd_ctx->sd->control_net != nullptr;
|
||||
}
|
||||
|
||||
enum sample_method_t sd_get_default_sample_method(const sd_ctx_t* sd_ctx) {
|
||||
if (sd_ctx != nullptr && sd_ctx->sd != nullptr) {
|
||||
if (sd_version_is_pid(sd_ctx->sd->version)) {
|
||||
|
||||
Reference in New Issue
Block a user