feat: hot-reload ControlNet - swap without rebuilding the context (#1768)

This commit is contained in:
fszontagh
2026-07-10 17:21:38 +02:00
committed by GitHub
parent cc73429228
commit 12b6fbff28
4 changed files with 202 additions and 1 deletions

View File

@@ -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_image_generation(const sd_ctx_t* sd_ctx);
SD_API bool sd_ctx_supports_video_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 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 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); SD_API const char* sd_rng_type_name(enum rng_type_t rng_type);

View File

@@ -179,6 +179,102 @@ bool ModelManager::register_param_tensors(const std::string& desc,
return true; 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() { bool ModelManager::load_all_params_eagerly() {
std::vector<TensorState*> all_states; std::vector<TensorState*> all_states;
all_states.reserve(tensor_states_.size()); all_states.reserve(tensor_states_.size());

View File

@@ -134,6 +134,9 @@ public:
bool allow_split_buffer = false, bool allow_split_buffer = false,
bool params_follow_compute_backend = false); bool params_follow_compute_backend = false);
bool unregister_param_tensors(const std::string& desc,
size_t* registered_tensor_size = nullptr);
template <typename Runner> template <typename Runner>
bool register_runner_params(const std::string& desc, bool register_runner_params(const std::string& desc,
Runner& runner, Runner& runner,

View File

@@ -222,9 +222,13 @@ public:
std::string split_mode_spec; std::string split_mode_spec;
bool auto_fit_enabled = false; bool auto_fit_enabled = false;
bool diffusion_conv_direct = false;
bool is_using_v_parameterization = false; bool is_using_v_parameterization = false;
bool is_using_edm_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<ModelManager> model_manager;
std::shared_ptr<Denoiser> denoiser = std::make_shared<CompVisDenoiser>(); std::shared_ptr<Denoiser> denoiser = std::make_shared<CompVisDenoiser>();
@@ -494,6 +498,76 @@ public:
params_follow_runtime); 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() { bool init_backend() {
std::string error; std::string error;
if (!backend_manager.init(backend_spec.c_str(), if (!backend_manager.init(backend_spec.c_str(),
@@ -852,10 +926,12 @@ public:
model_loader.process_model_files(enable_mmap, needs_writable_mmap); model_loader.process_model_files(enable_mmap, needs_writable_mmap);
load_alphas_cumprod(model_loader); load_alphas_cumprod(model_loader);
diffusion_conv_direct = sd_ctx_params->diffusion_conv_direct;
size_t text_encoder_params_mem_size = 0; size_t text_encoder_params_mem_size = 0;
size_t unet_params_mem_size = 0; size_t unet_params_mem_size = 0;
size_t vae_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; size_t extension_params_mem_size = 0;
bool tae_preview_only = sd_ctx_params->tae_preview_only; 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); 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) { 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_ctx != nullptr && sd_ctx->sd != nullptr) {
if (sd_version_is_pid(sd_ctx->sd->version)) { if (sd_version_is_pid(sd_ctx->sd->version)) {