refactor: unify model source and weight lifecycle management (#1956)

This commit is contained in:
leejet
2026-09-10 23:53:48 +08:00
committed by GitHub
parent d04e8950c1
commit 6b47fec013
36 changed files with 2548 additions and 1462 deletions
+1 -39
View File
@@ -2,8 +2,6 @@
#define __SD_MODEL_DIFFUSION_CONTROL_HPP__
#include "model/common/block.hpp"
#include "model_loader.h"
#include "model_manager.h"
// Match main UNet's MAX_GRAPH_SIZE so SDXL ControlNet (transformer_depth={1,2,10}) fits.
#define CONTROL_NET_GRAPH_SIZE MAX_GRAPH_SIZE
@@ -317,20 +315,17 @@ struct ControlNet : public GGMLRunner {
ggml_tensor* guided_hint_output_ggml = nullptr;
std::vector<sd::Tensor<float>> controls;
bool guided_hint_cached = false;
std::shared_ptr<ModelManager> owned_model_manager;
ggml_backend_t params_backend = nullptr;
static const char* guided_hint_cache_name() {
return "controlnet.guided_hint";
}
ControlNet(ggml_backend_t backend,
ggml_backend_t params_backend_,
const String2TensorStorage& tensor_storage_map = {},
SDVersion version = VERSION_SD1,
const std::string& prefix = "",
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr)
: GGMLRunner(backend, weight_manager), version(version), control_net(version), weight_prefix(prefix), params_backend(params_backend_) {
: GGMLRunner(backend, weight_manager), version(version), control_net(version), weight_prefix(prefix) {
control_net.init(params_ctx, tensor_storage_map, prefix);
}
@@ -445,39 +440,6 @@ struct ControlNet : public GGMLRunner {
guided_hint_cached = get_cache_tensor_by_name(guided_hint_cache_name()) != nullptr;
return controls;
}
bool load_from_file(const std::string& file_path, int n_threads) {
LOG_INFO("loading control net from '%s'", file_path.c_str());
std::map<std::string, ggml_tensor*> tensors;
control_net.get_param_tensors(tensors);
auto manager = std::dynamic_pointer_cast<ModelManager>(residency_manager.lock());
if (manager == nullptr) {
owned_model_manager = std::make_shared<ModelManager>();
residency_manager = owned_model_manager;
manager = owned_model_manager;
}
ModelLoader& model_loader = manager->loader();
if (!model_loader.init_from_file_and_convert_name(file_path)) {
LOG_ERROR("init control net model loader from file failed: '%s'", file_path.c_str());
return false;
}
manager->set_n_threads(n_threads);
if (!manager->register_param_tensors("ControlNet",
std::move(tensors),
ModelManager::ResidencyMode::ParamBackend,
runtime_backend,
params_backend) ||
!manager->validate_registered_tensors()) {
LOG_ERROR("register control net tensors with model manager failed");
return false;
}
LOG_INFO("control net model loaded");
return true;
}
};
#endif // __SD_MODEL_DIFFUSION_CONTROL_HPP__
+4 -3
View File
@@ -1714,8 +1714,8 @@ namespace Flux {
ggml_backend_t backend = sd_backend_cpu_init();
ggml_type model_data_type = GGML_TYPE_COUNT;
auto model_manager = std::make_shared<ModelManager>();
ModelLoader& model_loader = model_manager->loader();
auto model_manager = std::make_shared<ModelManager>();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(file_path, "model.diffusion_model.")) {
LOG_ERROR("init model loader from file failed: '%s'", file_path.c_str());
return;
@@ -1736,7 +1736,8 @@ namespace Flux {
VERSION_FLUX2,
model_manager);
if (!model_manager->register_runner_params("Flux test",
if (!model_manager->set_loader(model_loader) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*flux,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,
+4 -3
View File
@@ -2087,8 +2087,8 @@ namespace LTXV {
ggml_backend_t backend = sd_backend_cpu_init();
LOG_INFO("loading ltxav from '%s'", model_path.c_str());
auto model_manager = std::make_shared<ModelManager>();
ModelLoader& model_loader = model_manager->loader();
auto model_manager = std::make_shared<ModelManager>();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(model_path, "model.diffusion_model.")) {
LOG_ERROR("init model loader from file failed: '%s'", model_path.c_str());
return;
@@ -2107,7 +2107,8 @@ namespace LTXV {
"model.diffusion_model",
model_manager);
if (!model_manager->register_runner_params("LTXAV test",
if (!model_manager->set_loader(model_loader) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*ltxav,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,
+3 -2
View File
@@ -1064,13 +1064,14 @@ struct MMDiTRunner : public DiffusionModelRunner {
{
LOG_INFO("loading from '%s'", file_path.c_str());
ModelLoader& model_loader = model_manager->loader();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(file_path)) {
LOG_ERROR("init model loader from file failed: '%s'", file_path.c_str());
return;
}
if (!model_manager->register_runner_params("MMDiT test",
if (!model_manager->set_loader(std::move(model_loader)) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*mmdit,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,
+4 -3
View File
@@ -773,8 +773,8 @@ namespace Qwen {
ggml_backend_t backend = sd_backend_cpu_init();
ggml_type model_data_type = GGML_TYPE_Q8_0;
auto model_manager = std::make_shared<ModelManager>();
ModelLoader& model_loader = model_manager->loader();
auto model_manager = std::make_shared<ModelManager>();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(file_path, "model.diffusion_model.")) {
LOG_ERROR("init model loader from file failed: '%s'", file_path.c_str());
return;
@@ -793,7 +793,8 @@ namespace Qwen {
VERSION_QWEN_IMAGE,
model_manager);
if (!model_manager->register_runner_params("Qwen image test",
if (!model_manager->set_loader(model_loader) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*qwen_image,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,
+4 -3
View File
@@ -1020,8 +1020,8 @@ namespace WAN {
ggml_type model_data_type = GGML_TYPE_F16;
LOG_INFO("loading from '%s'", file_path.c_str());
auto model_manager = std::make_shared<ModelManager>();
ModelLoader& model_loader = model_manager->loader();
auto model_manager = std::make_shared<ModelManager>();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(file_path, "model.diffusion_model.")) {
LOG_ERROR("init model loader from file failed: '%s'", file_path.c_str());
return;
@@ -1040,7 +1040,8 @@ namespace WAN {
VERSION_WAN2_2_TI2V,
model_manager);
if (!model_manager->register_runner_params("Wan test",
if (!model_manager->set_loader(model_loader) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*wan,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,
+4 -3
View File
@@ -706,8 +706,8 @@ namespace ZImage {
ggml_backend_t backend = sd_backend_cpu_init();
ggml_type model_data_type = GGML_TYPE_Q8_0;
auto model_manager = std::make_shared<ModelManager>();
ModelLoader& model_loader = model_manager->loader();
auto model_manager = std::make_shared<ModelManager>();
ModelLoader model_loader;
if (!model_loader.init_from_file_and_convert_name(file_path, "model.diffusion_model.")) {
LOG_ERROR("init model loader from file failed: '%s'", file_path.c_str());
return;
@@ -728,7 +728,8 @@ namespace ZImage {
VERSION_QWEN_IMAGE,
model_manager);
if (!model_manager->register_runner_params("ZImage test",
if (!model_manager->set_loader(model_loader) ||
!model_manager->register_runner_params(ModelComponent::Diffusion,
*z_image,
"model.diffusion_model",
ModelManager::ResidencyMode::ParamBackend,