mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-07-23 11:20:53 -05:00
feat: add PiD 1.5 support (#1790)
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
# How to Use
|
||||
|
||||
PiD is NVIDIA's Pixel Diffusion Decoder. It replaces the usual VAE decode or decode-then-upscale path with a pixel-space diffusion decoder conditioned on a
|
||||
source latent and text prompt.
|
||||
source latent and text prompt. Both the original PiD checkpoints and PiD 1.5 are supported.
|
||||
|
||||
In stable-diffusion.cpp, PiD currently runs as an image edit pipeline: provide a reference image with `-r`/`--ref-image`, encode that image with a matching VAE, then let the PiD diffusion model decode/upscale directly to RGB.
|
||||
|
||||
@@ -16,6 +16,7 @@ In stable-diffusion.cpp, PiD currently runs as an image edit pipeline: provide a
|
||||
- Flux / Z-Image PiD: use the Flux VAE and pass `--vae-format flux`
|
||||
- SD3 PiD: use the SD3 VAE and pass `--vae-format sd3`
|
||||
- Flux.2 PiD: use the Flux.2 VAE and pass `--vae-format flux2`
|
||||
- Qwen-Image PiD: use the Qwen-Image 2D VAE and pass `--vae-format wan`
|
||||
|
||||
The official PiD model card should be checked before use. At the time of the initial PiD release, the official weights are under the NSCLv1 non-commercial license.
|
||||
|
||||
|
||||
@@ -50,6 +50,9 @@ static sd_vae_format_t str_to_vae_format(const std::string& value) {
|
||||
if (value == "flux2") {
|
||||
return SD_VAE_FORMAT_FLUX2;
|
||||
}
|
||||
if (value == "wan") {
|
||||
return SD_VAE_FORMAT_WAN;
|
||||
}
|
||||
return SD_VAE_FORMAT_COUNT;
|
||||
}
|
||||
|
||||
@@ -401,7 +404,7 @@ ArgOptions SDContextParams::get_options() {
|
||||
&vae_path},
|
||||
{"",
|
||||
"--vae-format",
|
||||
"VAE latent format override: auto, flux, sd3, or flux2 (default: auto)",
|
||||
"VAE latent format override: auto, flux, sd3, flux2, or wan (default: auto)",
|
||||
0,
|
||||
&vae_format},
|
||||
{"",
|
||||
@@ -743,7 +746,7 @@ bool SDContextParams::validate(SDMode mode) {
|
||||
}
|
||||
|
||||
if (str_to_vae_format(vae_format) == SD_VAE_FORMAT_COUNT) {
|
||||
LOG_ERROR("error: vae_format must be 'auto', 'flux', 'sd3', or 'flux2'");
|
||||
LOG_ERROR("error: vae_format must be 'auto', 'flux', 'sd3', 'flux2', or 'wan'");
|
||||
return false;
|
||||
}
|
||||
|
||||
|
||||
@@ -180,6 +180,7 @@ enum sd_vae_format_t {
|
||||
SD_VAE_FORMAT_FLUX,
|
||||
SD_VAE_FORMAT_SD3,
|
||||
SD_VAE_FORMAT_FLUX2,
|
||||
SD_VAE_FORMAT_WAN,
|
||||
SD_VAE_FORMAT_COUNT,
|
||||
};
|
||||
|
||||
|
||||
@@ -17,30 +17,38 @@ namespace Pid {
|
||||
constexpr float PID_PI = 3.14159265358979323846f;
|
||||
|
||||
struct PixelDiTConfig {
|
||||
int64_t in_channels = 3;
|
||||
int64_t hidden_size = 1536;
|
||||
int64_t num_groups = 24;
|
||||
int64_t patch_mlp_hidden_dim = 4096;
|
||||
int64_t pixel_hidden_size = 16;
|
||||
int64_t pixel_attn_hidden_size = 1152;
|
||||
int64_t pixel_num_groups = 16;
|
||||
int64_t patch_depth = 14;
|
||||
int64_t pixel_depth = 2;
|
||||
int64_t patch_size = 16;
|
||||
int64_t txt_embed_dim = 2304;
|
||||
int64_t txt_max_length = 300;
|
||||
float text_rope_theta = 10000.f;
|
||||
int64_t lq_latent_channels = 16;
|
||||
int64_t lq_hidden_dim = 512;
|
||||
int64_t lq_num_res_blocks = 4;
|
||||
int64_t lq_interval = 2;
|
||||
int64_t lq_sr_scale = 4;
|
||||
int64_t lq_latent_down_factor = 8;
|
||||
int64_t rope_ref_grid_h = 64;
|
||||
int64_t rope_ref_grid_w = 64;
|
||||
int64_t in_channels = 3;
|
||||
int64_t hidden_size = 1536;
|
||||
int64_t num_groups = 24;
|
||||
int64_t patch_mlp_hidden_dim = 4096;
|
||||
int64_t pixel_hidden_size = 16;
|
||||
int64_t pixel_attn_hidden_size = 1152;
|
||||
int64_t pixel_num_groups = 16;
|
||||
int64_t patch_depth = 14;
|
||||
int64_t pixel_depth = 2;
|
||||
int64_t patch_size = 16;
|
||||
int64_t txt_embed_dim = 2304;
|
||||
int64_t txt_max_length = 300;
|
||||
float text_rope_theta = 10000.f;
|
||||
int64_t lq_latent_channels = 16;
|
||||
int64_t lq_hidden_dim = 512;
|
||||
int64_t lq_num_res_blocks = 4;
|
||||
int64_t lq_interval = 2;
|
||||
int64_t lq_sr_scale = 4;
|
||||
int64_t lq_latent_down_factor = 8;
|
||||
int64_t lq_latent_unpatchify_factor = 1;
|
||||
bool lq_replicate_padding = false;
|
||||
bool lq_gate_per_token = false;
|
||||
bool pit_lq_inject = false;
|
||||
int64_t rope_ref_grid_h = 64;
|
||||
int64_t rope_ref_grid_w = 64;
|
||||
|
||||
static PixelDiTConfig detect_from_weights(const String2TensorStorage& tensor_storage_map, const std::string& prefix) {
|
||||
PixelDiTConfig config;
|
||||
int64_t latent_proj_in_channels = config.lq_latent_channels;
|
||||
int64_t num_lq_gates = 0;
|
||||
const std::string lq_prefix = prefix + ".lq_proj.";
|
||||
config.pit_lq_inject = tensor_storage_map.find(lq_prefix + "pit_head.weight") != tensor_storage_map.end();
|
||||
for (const auto& [name, tensor_storage] : tensor_storage_map) {
|
||||
if (!starts_with(name, prefix)) {
|
||||
continue;
|
||||
@@ -61,20 +69,56 @@ namespace Pid {
|
||||
config.pixel_depth = std::max<int64_t>(config.pixel_depth, block_index + 1);
|
||||
}
|
||||
}
|
||||
if (name.find("lq_proj.latent_proj.0.weight") != std::string::npos) {
|
||||
config.lq_latent_channels = tensor_storage.ne[2];
|
||||
config.lq_latent_down_factor = config.lq_latent_channels >= 64 ? 16 : 8;
|
||||
if (name == lq_prefix + "latent_proj.0.weight") {
|
||||
latent_proj_in_channels = tensor_storage.ne[2];
|
||||
config.lq_hidden_dim = tensor_storage.ne[3];
|
||||
}
|
||||
if (starts_with(name, lq_prefix + "gate_modules.")) {
|
||||
auto items = split_string(name.substr(lq_prefix.size()), '.');
|
||||
if (items.size() > 1) {
|
||||
int gate_index = atoi(items[1].c_str());
|
||||
num_lq_gates = std::max<int64_t>(num_lq_gates, gate_index + 1);
|
||||
}
|
||||
}
|
||||
if (name.find("patch_blocks.0.mlp_x.w1.weight") != std::string::npos) {
|
||||
config.patch_mlp_hidden_dim = tensor_storage.ne[1];
|
||||
}
|
||||
}
|
||||
LOG_DEBUG("pid: patch_depth = %" PRId64 ", pixel_depth = %" PRId64 ", patch_mlp_hidden_dim = %" PRId64 ", lq_latent_channels = %" PRId64 ", lq_latent_down_factor = %" PRId64,
|
||||
if (num_lq_gates > 0) {
|
||||
config.lq_interval = (config.patch_depth + num_lq_gates - 1) / num_lq_gates;
|
||||
}
|
||||
if (config.pit_lq_inject) {
|
||||
if (latent_proj_in_channels == 16) {
|
||||
config.lq_latent_channels = 16;
|
||||
config.lq_latent_down_factor = 8;
|
||||
config.lq_latent_unpatchify_factor = 1;
|
||||
} else {
|
||||
GGML_ASSERT(latent_proj_in_channels == 32);
|
||||
config.lq_latent_channels = 128;
|
||||
config.lq_latent_down_factor = 16;
|
||||
config.lq_latent_unpatchify_factor = 2;
|
||||
}
|
||||
auto gate_weight = tensor_storage_map.find(lq_prefix + "gate_modules.0.content_proj.weight");
|
||||
if (gate_weight != tensor_storage_map.end()) {
|
||||
config.lq_gate_per_token = gate_weight->second.ne[1] == 1;
|
||||
}
|
||||
config.lq_replicate_padding = true;
|
||||
config.rope_ref_grid_h = 128;
|
||||
config.rope_ref_grid_w = 128;
|
||||
} else {
|
||||
config.lq_latent_channels = latent_proj_in_channels;
|
||||
config.lq_latent_down_factor = latent_proj_in_channels >= 64 ? 16 : 8;
|
||||
}
|
||||
LOG_DEBUG("pid: version = %s, patch_depth = %" PRId64 ", pixel_depth = %" PRId64 ", patch_mlp_hidden_dim = %" PRId64 ", lq_latent_channels = %" PRId64 ", lq_hidden_dim = %" PRId64 ", lq_latent_down_factor = %" PRId64 ", lq_latent_unpatchify_factor = %" PRId64 ", lq_interval = %" PRId64,
|
||||
config.pit_lq_inject ? "1.5" : "1",
|
||||
config.patch_depth,
|
||||
config.pixel_depth,
|
||||
config.patch_mlp_hidden_dim,
|
||||
config.lq_latent_channels,
|
||||
config.lq_latent_down_factor);
|
||||
config.lq_hidden_dim,
|
||||
config.lq_latent_down_factor,
|
||||
config.lq_latent_unpatchify_factor,
|
||||
config.lq_interval);
|
||||
return config;
|
||||
}
|
||||
};
|
||||
@@ -135,6 +179,18 @@ namespace Pid {
|
||||
return ggml_add(ctx, ggml_add(ctx, x, ggml_mul(ctx, x, scale)), shift);
|
||||
}
|
||||
|
||||
inline ggml_tensor* replicate_pad_2d(ggml_context* ctx, ggml_tensor* x) {
|
||||
auto left = ggml_ext_slice(ctx, x, 0, 0, 1);
|
||||
auto right = ggml_ext_slice(ctx, x, 0, x->ne[0] - 1, x->ne[0]);
|
||||
x = ggml_concat(ctx, left, x, 0);
|
||||
x = ggml_concat(ctx, x, right, 0);
|
||||
|
||||
auto top = ggml_ext_slice(ctx, x, 1, 0, 1);
|
||||
auto bottom = ggml_ext_slice(ctx, x, 1, x->ne[1] - 1, x->ne[1]);
|
||||
x = ggml_concat(ctx, top, x, 1);
|
||||
return ggml_concat(ctx, x, bottom, 1);
|
||||
}
|
||||
|
||||
struct PatchTokenEmbedder : public GGMLBlock {
|
||||
bool use_rms_norm;
|
||||
|
||||
@@ -457,9 +513,9 @@ namespace Pid {
|
||||
struct SigmaAwareGate : public GGMLBlock {
|
||||
int64_t dim;
|
||||
|
||||
SigmaAwareGate(int64_t dim)
|
||||
SigmaAwareGate(int64_t dim, bool per_token = false)
|
||||
: dim(dim) {
|
||||
blocks["content_proj"] = std::make_shared<Linear>(dim * 2, dim, true);
|
||||
blocks["content_proj"] = std::make_shared<Linear>(dim * 2, per_token ? 1 : dim, true);
|
||||
}
|
||||
|
||||
void init_params(ggml_context* ctx,
|
||||
@@ -479,16 +535,20 @@ namespace Pid {
|
||||
auto alpha = ggml_exp(ctx->ggml_ctx, params["log_alpha"]);
|
||||
auto offset = ggml_neg(ctx->ggml_ctx, ggml_mul(ctx->ggml_ctx, alpha, sigma));
|
||||
auto gate = ggml_sigmoid(ctx->ggml_ctx, ggml_add(ctx->ggml_ctx, content_logit, offset));
|
||||
return ggml_add(ctx->ggml_ctx, x, ggml_mul(ctx->ggml_ctx, gate, lq));
|
||||
return ggml_add(ctx->ggml_ctx, x, ggml_mul(ctx->ggml_ctx, lq, gate));
|
||||
}
|
||||
};
|
||||
|
||||
struct PiDResBlock : public GGMLBlock {
|
||||
PiDResBlock(int64_t channels) {
|
||||
blocks["block.0"] = std::make_shared<GroupNorm>(4, channels, 1e-5f);
|
||||
blocks["block.2"] = std::make_shared<Conv2d>(channels, channels, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, std::pair<int, int>{1, 1});
|
||||
blocks["block.3"] = std::make_shared<GroupNorm>(4, channels, 1e-5f);
|
||||
blocks["block.5"] = std::make_shared<Conv2d>(channels, channels, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, std::pair<int, int>{1, 1});
|
||||
bool replicate_padding;
|
||||
|
||||
PiDResBlock(int64_t channels, bool replicate_padding = false)
|
||||
: replicate_padding(replicate_padding) {
|
||||
std::pair<int, int> padding = replicate_padding ? std::pair<int, int>{0, 0} : std::pair<int, int>{1, 1};
|
||||
blocks["block.0"] = std::make_shared<GroupNorm>(4, channels, 1e-5f);
|
||||
blocks["block.2"] = std::make_shared<Conv2d>(channels, channels, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, padding);
|
||||
blocks["block.3"] = std::make_shared<GroupNorm>(4, channels, 1e-5f);
|
||||
blocks["block.5"] = std::make_shared<Conv2d>(channels, channels, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, padding);
|
||||
}
|
||||
|
||||
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x) {
|
||||
@@ -497,9 +557,15 @@ namespace Pid {
|
||||
auto norm2 = std::dynamic_pointer_cast<GroupNorm>(blocks["block.3"]);
|
||||
auto conv2 = std::dynamic_pointer_cast<Conv2d>(blocks["block.5"]);
|
||||
auto h = ggml_silu_inplace(ctx->ggml_ctx, norm1->forward(ctx, x));
|
||||
h = conv1->forward(ctx, h);
|
||||
h = ggml_silu_inplace(ctx->ggml_ctx, norm2->forward(ctx, h));
|
||||
h = conv2->forward(ctx, h);
|
||||
if (replicate_padding) {
|
||||
h = replicate_pad_2d(ctx->ggml_ctx, h);
|
||||
}
|
||||
h = conv1->forward(ctx, h);
|
||||
h = ggml_silu_inplace(ctx->ggml_ctx, norm2->forward(ctx, h));
|
||||
if (replicate_padding) {
|
||||
h = replicate_pad_2d(ctx->ggml_ctx, h);
|
||||
}
|
||||
h = conv2->forward(ctx, h);
|
||||
return ggml_add(ctx->ggml_ctx, x, h);
|
||||
}
|
||||
};
|
||||
@@ -509,16 +575,23 @@ namespace Pid {
|
||||
|
||||
LQProjection2D(const PixelDiTConfig& config)
|
||||
: config(config) {
|
||||
blocks["latent_proj.0"] = std::make_shared<Conv2d>(config.lq_latent_channels, config.lq_hidden_dim, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, std::pair<int, int>{1, 1});
|
||||
blocks["latent_proj.2"] = std::make_shared<Conv2d>(config.lq_hidden_dim, config.lq_hidden_dim, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, std::pair<int, int>{1, 1});
|
||||
int64_t unpatchify_area = config.lq_latent_unpatchify_factor * config.lq_latent_unpatchify_factor;
|
||||
GGML_ASSERT(config.lq_latent_channels % unpatchify_area == 0);
|
||||
int64_t latent_proj_in_channels = config.lq_latent_channels / unpatchify_area;
|
||||
std::pair<int, int> padding = config.lq_replicate_padding ? std::pair<int, int>{0, 0} : std::pair<int, int>{1, 1};
|
||||
blocks["latent_proj.0"] = std::make_shared<Conv2d>(latent_proj_in_channels, config.lq_hidden_dim, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, padding);
|
||||
blocks["latent_proj.2"] = std::make_shared<Conv2d>(config.lq_hidden_dim, config.lq_hidden_dim, std::pair<int, int>{3, 3}, std::pair<int, int>{1, 1}, padding);
|
||||
for (int i = 0; i < config.lq_num_res_blocks; ++i) {
|
||||
blocks["latent_proj." + std::to_string(3 + i)] = std::make_shared<PiDResBlock>(config.lq_hidden_dim);
|
||||
blocks["latent_proj." + std::to_string(3 + i)] = std::make_shared<PiDResBlock>(config.lq_hidden_dim, config.lq_replicate_padding);
|
||||
}
|
||||
|
||||
int num_outputs = static_cast<int>((config.patch_depth + config.lq_interval - 1) / config.lq_interval);
|
||||
for (int i = 0; i < num_outputs; ++i) {
|
||||
blocks["output_heads." + std::to_string(i)] = std::make_shared<Linear>(config.lq_hidden_dim, config.hidden_size, true);
|
||||
blocks["gate_modules." + std::to_string(i)] = std::make_shared<SigmaAwareGate>(config.hidden_size);
|
||||
blocks["gate_modules." + std::to_string(i)] = std::make_shared<SigmaAwareGate>(config.hidden_size, config.lq_gate_per_token);
|
||||
}
|
||||
if (config.pit_lq_inject) {
|
||||
blocks["pit_head"] = std::make_shared<Linear>(config.lq_hidden_dim, config.hidden_size, true);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -543,9 +616,29 @@ namespace Pid {
|
||||
ggml_tensor* lq_latent,
|
||||
int64_t target_pH,
|
||||
int64_t target_pW) {
|
||||
auto conv0 = std::dynamic_pointer_cast<Conv2d>(blocks["latent_proj.0"]);
|
||||
auto conv2 = std::dynamic_pointer_cast<Conv2d>(blocks["latent_proj.2"]);
|
||||
float z_to_patch_ratio = static_cast<float>(config.lq_sr_scale * config.lq_latent_down_factor) /
|
||||
auto conv0 = std::dynamic_pointer_cast<Conv2d>(blocks["latent_proj.0"]);
|
||||
auto conv2 = std::dynamic_pointer_cast<Conv2d>(blocks["latent_proj.2"]);
|
||||
int64_t unpatchify_factor = config.lq_latent_unpatchify_factor;
|
||||
if (unpatchify_factor > 1) {
|
||||
int64_t latent_h = lq_latent->ne[1];
|
||||
int64_t latent_w = lq_latent->ne[0];
|
||||
lq_latent = ggml_ext_cont(ctx->ggml_ctx, ggml_ext_torch_permute(ctx->ggml_ctx, lq_latent, 2, 0, 1, 3));
|
||||
lq_latent = ggml_reshape_3d(ctx->ggml_ctx,
|
||||
lq_latent,
|
||||
lq_latent->ne[0],
|
||||
lq_latent->ne[1] * lq_latent->ne[2],
|
||||
lq_latent->ne[3]);
|
||||
lq_latent = DiT::unpatchify(ctx->ggml_ctx,
|
||||
lq_latent,
|
||||
latent_h,
|
||||
latent_w,
|
||||
static_cast<int>(unpatchify_factor),
|
||||
static_cast<int>(unpatchify_factor),
|
||||
true);
|
||||
}
|
||||
|
||||
int64_t effective_down_factor = config.lq_latent_down_factor / unpatchify_factor;
|
||||
float z_to_patch_ratio = static_cast<float>(config.lq_sr_scale * effective_down_factor) /
|
||||
static_cast<float>(config.patch_size);
|
||||
GGML_ASSERT(z_to_patch_ratio >= 1.0f);
|
||||
if (lq_latent->ne[0] != target_pW || lq_latent->ne[1] != target_pH) {
|
||||
@@ -558,9 +651,15 @@ namespace Pid {
|
||||
GGML_SCALE_MODE_NEAREST);
|
||||
}
|
||||
|
||||
if (config.lq_replicate_padding) {
|
||||
lq_latent = replicate_pad_2d(ctx->ggml_ctx, lq_latent);
|
||||
}
|
||||
auto feat = conv0->forward(ctx, lq_latent);
|
||||
feat = ggml_silu_inplace(ctx->ggml_ctx, feat);
|
||||
feat = conv2->forward(ctx, feat);
|
||||
if (config.lq_replicate_padding) {
|
||||
feat = replicate_pad_2d(ctx->ggml_ctx, feat);
|
||||
}
|
||||
feat = conv2->forward(ctx, feat);
|
||||
for (int i = 0; i < config.lq_num_res_blocks; ++i) {
|
||||
auto block = std::dynamic_pointer_cast<PiDResBlock>(blocks["latent_proj." + std::to_string(3 + i)]);
|
||||
feat = block->forward(ctx, feat);
|
||||
@@ -574,11 +673,15 @@ namespace Pid {
|
||||
|
||||
int num_outputs = static_cast<int>((config.patch_depth + config.lq_interval - 1) / config.lq_interval);
|
||||
std::vector<ggml_tensor*> outputs;
|
||||
outputs.reserve(num_outputs);
|
||||
outputs.reserve(num_outputs + (config.pit_lq_inject ? 1 : 0));
|
||||
for (int i = 0; i < num_outputs; ++i) {
|
||||
auto head = std::dynamic_pointer_cast<Linear>(blocks["output_heads." + std::to_string(i)]);
|
||||
outputs.push_back(head->forward(ctx, tokens));
|
||||
}
|
||||
if (config.pit_lq_inject) {
|
||||
auto pit_head = std::dynamic_pointer_cast<Linear>(blocks["pit_head"]);
|
||||
outputs.push_back(pit_head->forward(ctx, tokens));
|
||||
}
|
||||
return outputs;
|
||||
}
|
||||
};
|
||||
@@ -606,6 +709,9 @@ namespace Pid {
|
||||
}
|
||||
blocks["final_layer"] = std::make_shared<FinalLayer>(config.pixel_hidden_size, config.in_channels);
|
||||
blocks["lq_proj"] = std::make_shared<LQProjection2D>(config);
|
||||
if (config.pit_lq_inject) {
|
||||
blocks["pit_lq_gate"] = std::make_shared<SigmaAwareGate>(config.hidden_size, config.lq_gate_per_token);
|
||||
}
|
||||
}
|
||||
|
||||
void init_params(ggml_context* ctx,
|
||||
@@ -654,6 +760,11 @@ namespace Pid {
|
||||
y_emb = ggml_add(ctx->ggml_ctx, y_emb, y_pos);
|
||||
|
||||
std::vector<ggml_tensor*> lq_features = lq_proj->forward(ctx, lq_latent, Hs, Ws);
|
||||
ggml_tensor* pit_lq_feature = nullptr;
|
||||
if (config.pit_lq_inject) {
|
||||
pit_lq_feature = lq_features.back();
|
||||
lq_features.pop_back();
|
||||
}
|
||||
|
||||
auto s = s_embedder->forward(ctx, x_patches);
|
||||
|
||||
@@ -677,6 +788,10 @@ namespace Pid {
|
||||
sd::ggml_graph_cut::mark_graph_cut(y_emb, "pid.patch_blocks." + std::to_string(i), "y");
|
||||
}
|
||||
s = ggml_silu(ctx->ggml_ctx, ggml_add(ctx->ggml_ctx, s, t_emb));
|
||||
if (pit_lq_feature != nullptr) {
|
||||
auto pit_lq_gate = std::dynamic_pointer_cast<SigmaAwareGate>(blocks["pit_lq_gate"]);
|
||||
s = pit_lq_gate->forward(ctx, s, pit_lq_feature, degrade_sigma);
|
||||
}
|
||||
|
||||
auto s_cond = ggml_reshape_2d(ctx->ggml_ctx, s, config.hidden_size, L * B);
|
||||
auto pixels = pixel_embedder->forward(ctx, x, config.patch_size, pixel_pos_full);
|
||||
|
||||
@@ -1306,12 +1306,12 @@ public:
|
||||
false,
|
||||
version,
|
||||
model_manager);
|
||||
} else if (sd_version_uses_wan_vae(version)) {
|
||||
} else if (sd_version_uses_wan_vae(vae_version)) {
|
||||
return std::make_shared<WAN::WanVAERunner>(backend_for(SDBackendModule::VAE),
|
||||
tensor_storage_map,
|
||||
"first_stage_model",
|
||||
false,
|
||||
version,
|
||||
vae_version,
|
||||
model_manager);
|
||||
} else {
|
||||
auto model = std::make_shared<AutoEncoderKL>(backend_for(SDBackendModule::VAE),
|
||||
@@ -3213,6 +3213,8 @@ const char* sd_vae_format_name(enum sd_vae_format_t format) {
|
||||
return "sd3";
|
||||
case SD_VAE_FORMAT_FLUX2:
|
||||
return "flux2";
|
||||
case SD_VAE_FORMAT_WAN:
|
||||
return "wan";
|
||||
default:
|
||||
return NONE_STR;
|
||||
}
|
||||
@@ -3226,6 +3228,8 @@ static SDVersion sd_vae_format_to_version(enum sd_vae_format_t format, SDVersion
|
||||
return VERSION_SD3;
|
||||
case SD_VAE_FORMAT_FLUX2:
|
||||
return VERSION_FLUX2;
|
||||
case SD_VAE_FORMAT_WAN:
|
||||
return VERSION_WAN2;
|
||||
case SD_VAE_FORMAT_AUTO:
|
||||
default:
|
||||
return fallback;
|
||||
|
||||
Reference in New Issue
Block a user