refactor: move model-specific args into model parsers (#1757)

This commit is contained in:
leejet
2026-07-06 23:13:18 +08:00
committed by GitHub
parent e22272ee63
commit bb84971129
9 changed files with 81 additions and 73 deletions
+21 -1
View File
@@ -6,6 +6,7 @@
#include <optional>
#include "core/tensor_ggml.hpp"
#include "core/util.h"
#include "model/te/clip.hpp"
#include "model/te/llm.hpp"
#include "model/te/t5.hpp"
@@ -1217,8 +1218,27 @@ struct T5CLIPEmbedder : public Conditioner {
bool use_mask = false,
int mask_pad = 0,
bool is_umt5 = false,
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr)
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr,
const char* model_args = nullptr)
: use_mask(use_mask), mask_pad(mask_pad), t5_tokenizer(is_umt5) {
for (const auto& [key, value] : parse_key_value_args(model_args, "model arg")) {
if (key == "chroma_use_t5_mask") {
bool parsed = false;
if (parse_strict_bool(value, parsed)) {
this->use_mask = parsed;
} else {
LOG_WARN("ignoring invalid Chroma T5 model arg '%s=%s'", key.c_str(), value.c_str());
}
} else if (key == "chroma_t5_mask_pad") {
int parsed = 0;
if (parse_strict_int(value, parsed)) {
this->mask_pad = parsed;
} else {
LOG_WARN("ignoring invalid Chroma T5 model arg '%s=%s'", key.c_str(), value.c_str());
}
}
}
bool use_t5 = false;
for (auto pair : tensor_storage_map) {
if (pair.first.find("text_encoders.t5xxl") != std::string::npos) {
+16 -6
View File
@@ -4,6 +4,7 @@
#include <memory>
#include <vector>
#include "core/util.h"
#include "model/adapter/pulid.hpp"
#include "model/common/rope.hpp"
#include "model/diffusion/dit.hpp"
@@ -1400,18 +1401,28 @@ namespace Flux {
std::vector<float> dct_vec;
sd::Tensor<float> guidance_tensor;
SDVersion version;
bool use_mask = false;
bool use_mask = true;
FluxRunner(ggml_backend_t backend,
const String2TensorStorage& tensor_storage_map = {},
const std::string prefix = "",
SDVersion version = VERSION_FLUX,
bool use_mask = false,
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr)
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr,
const char* model_args = nullptr)
: DiffusionModelRunner(backend, prefix, weight_manager),
config(FluxConfig::detect_from_weights(tensor_storage_map, prefix, version)),
version(version),
use_mask(use_mask) {
version(version) {
for (const auto& [key, value] : parse_key_value_args(model_args, "model arg")) {
if (key == "chroma_use_dit_mask") {
bool parsed = true;
if (parse_strict_bool(value, parsed)) {
use_mask = parsed;
} else {
LOG_WARN("ignoring invalid Chroma DiT model arg '%s=%s'", key.c_str(), value.c_str());
}
}
}
if (config.is_chroma) {
LOG_INFO("Using pruned modulation (Chroma)");
}
@@ -1718,7 +1729,6 @@ namespace Flux {
tensor_storage_map,
"model.diffusion_model",
VERSION_FLUX2,
false,
model_manager);
if (!model_manager->register_runner_params("Flux test",
+13 -4
View File
@@ -3,6 +3,7 @@
#include <memory>
#include "core/util.h"
#include "model/common/block.hpp"
#include "model/diffusion/dit.hpp"
#include "model/diffusion/flux.hpp"
@@ -566,12 +567,21 @@ namespace Qwen {
const String2TensorStorage& tensor_storage_map = {},
const std::string prefix = "",
SDVersion version = VERSION_QWEN_IMAGE,
bool zero_cond_t = false,
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr)
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr,
const char* model_args = nullptr)
: DiffusionModelRunner(backend, prefix, weight_manager),
config(QwenImageConfig::detect_from_weights(tensor_storage_map, prefix)),
version(version) {
config.zero_cond_t = config.zero_cond_t || zero_cond_t;
for (const auto& [key, value] : parse_key_value_args(model_args, "model arg")) {
if (key == "qwen_image_zero_cond_t") {
bool parsed = false;
if (parse_strict_bool(value, parsed)) {
config.zero_cond_t = config.zero_cond_t || parsed;
} else {
LOG_WARN("ignoring invalid Qwen Image model arg '%s=%s'", key.c_str(), value.c_str());
}
}
}
if (version == VERSION_QWEN_IMAGE_LAYERED) {
config.use_additional_t_cond = true;
}
@@ -775,7 +785,6 @@ namespace Qwen {
tensor_storage_map,
"model.diffusion_model",
VERSION_QWEN_IMAGE,
false,
model_manager);
if (!model_manager->register_runner_params("Qwen image test",
+17 -23
View File
@@ -57,11 +57,11 @@
#include "name_conversion.h"
#include "runtime/latent-preview.h"
#include <atomic>
const char* sd_vae_format_name(enum sd_vae_format_t format);
static SDVersion sd_vae_format_to_version(enum sd_vae_format_t format, SDVersion fallback);
#include <atomic>
const char* model_version_to_str[] = {
"SD 1.x",
"SD 1.x Inpaint",
@@ -921,10 +921,11 @@ public:
if (is_chroma) {
cond_stage_model = std::make_shared<T5CLIPEmbedder>(backend_for(SDBackendModule::TE),
tensor_storage_map,
sd_ctx_params->chroma_use_t5_mask,
sd_ctx_params->chroma_t5_mask_pad,
false,
model_manager);
1,
false,
model_manager,
sd_ctx_params->model_args);
} else if (version == VERSION_OVIS_IMAGE) {
cond_stage_model = std::make_shared<LLMEmbedder>(backend_for(SDBackendModule::TE),
tensor_storage_map,
@@ -941,8 +942,8 @@ public:
tensor_storage_map,
"model.diffusion_model",
version,
sd_ctx_params->chroma_use_dit_mask,
model_manager);
model_manager,
sd_ctx_params->model_args);
} else if (sd_version_is_flux2(version) || sd_version_is_sefi_image(version)) {
bool is_chroma = false;
cond_stage_model = std::make_shared<LLMEmbedder>(backend_for(SDBackendModule::TE),
@@ -955,8 +956,8 @@ public:
tensor_storage_map,
"model.diffusion_model",
version,
sd_ctx_params->chroma_use_dit_mask,
model_manager);
model_manager,
sd_ctx_params->model_args);
} else if (sd_version_is_ltxav(version)) {
cond_stage_model = std::make_shared<LTXAVEmbedder>(backend_for(SDBackendModule::TE),
tensor_storage_map,
@@ -1014,8 +1015,8 @@ public:
tensor_storage_map,
"model.diffusion_model",
version,
sd_ctx_params->qwen_image_zero_cond_t,
model_manager);
model_manager,
sd_ctx_params->model_args);
} else if (sd_version_is_longcat(version)) {
cond_stage_model = std::make_shared<LLMEmbedder>(backend_for(SDBackendModule::TE),
tensor_storage_map,
@@ -1027,8 +1028,8 @@ public:
tensor_storage_map,
"model.diffusion_model",
version,
sd_ctx_params->chroma_use_dit_mask,
model_manager);
model_manager,
sd_ctx_params->model_args);
} else if (version == VERSION_HIDREAM_O1) {
cond_stage_model = std::make_shared<HiDreamO1::HiDreamO1Conditioner>(backend_for(SDBackendModule::TE),
tensor_storage_map,
@@ -1364,7 +1365,6 @@ public:
high_noise_diffusion_model->set_flash_attention_enabled(true);
}
}
}
LOG_DEBUG("validating model metadata");
@@ -3048,15 +3048,13 @@ void sd_ctx_params_init(sd_ctx_params_t* sd_ctx_params) {
sd_ctx_params->eager_load = false;
sd_ctx_params->enable_mmap = false;
sd_ctx_params->diffusion_flash_attn = false;
sd_ctx_params->chroma_use_dit_mask = true;
sd_ctx_params->chroma_use_t5_mask = false;
sd_ctx_params->chroma_t5_mask_pad = 1;
sd_ctx_params->vae_format = SD_VAE_FORMAT_AUTO;
sd_ctx_params->backend = nullptr;
sd_ctx_params->params_backend = nullptr;
sd_ctx_params->split_mode = nullptr;
sd_ctx_params->auto_fit = false;
sd_ctx_params->rpc_servers = nullptr;
sd_ctx_params->model_args = nullptr;
sd_ctx_params->pulid_weights_path = nullptr;
}
@@ -3096,12 +3094,10 @@ char* sd_ctx_params_to_str(const sd_ctx_params_t* sd_ctx_params) {
"backend: %s\n"
"params_backend: %s\n"
"split_mode: %s\n"
"model_args: %s\n"
"auto_fit: %s\n"
"flash_attn: %s\n"
"diffusion_flash_attn: %s\n"
"chroma_use_dit_mask: %s\n"
"chroma_use_t5_mask: %s\n"
"chroma_t5_mask_pad: %d\n"
"vae_format: %s\n",
SAFE_STR(sd_ctx_params->model_path),
SAFE_STR(sd_ctx_params->clip_l_path),
@@ -3132,12 +3128,10 @@ char* sd_ctx_params_to_str(const sd_ctx_params_t* sd_ctx_params) {
SAFE_STR(sd_ctx_params->backend),
SAFE_STR(sd_ctx_params->params_backend),
SAFE_STR(sd_ctx_params->split_mode),
SAFE_STR(sd_ctx_params->model_args),
BOOL_STR(sd_ctx_params->auto_fit),
BOOL_STR(sd_ctx_params->flash_attn),
BOOL_STR(sd_ctx_params->diffusion_flash_attn),
BOOL_STR(sd_ctx_params->chroma_use_dit_mask),
BOOL_STR(sd_ctx_params->chroma_use_t5_mask),
sd_ctx_params->chroma_t5_mask_pad,
sd_vae_format_name(sd_ctx_params->vae_format));
return buf;