diff --git a/README.md b/README.md
index 340ba5b2d..31e5251ed 100644
--- a/README.md
+++ b/README.md
@@ -54,6 +54,7 @@ API and command-line option may change frequently.***
- [ERNIE-Image](./docs/ernie_image.md)
- [Boogu Image](./docs/boogu_image.md)
- [Krea2](./docs/krea2.md)
+ - [Mage-Flow](./docs/mage_flow.md)
- [SeFi-Image](./docs/sefi_image.md)
- [HiDream-O1-Image](./docs/hidream_o1_image.md)
- [Ideogram4](./docs/ideogram4.md)
@@ -62,6 +63,7 @@ API and command-line option may change frequently.***
- [Qwen Image Edit series](./docs/qwen_image_edit.md)
- [LongCat Image Edit](./docs/longcat_image.md)
- [Boogu Image Edit](./docs/boogu_image.md)
+ - [Mage-Flow-Edit](./docs/mage_flow.md#image-editing)
- Video Models
- [Wan2.1/Wan2.2](./docs/wan.md)
- [LTX-2.3](./docs/ltx2.md)
diff --git a/assets/mage_flow/edit_example.png b/assets/mage_flow/edit_example.png
new file mode 100644
index 000000000..a53fef516
Binary files /dev/null and b/assets/mage_flow/edit_example.png differ
diff --git a/assets/mage_flow/example.png b/assets/mage_flow/example.png
new file mode 100644
index 000000000..695265f19
Binary files /dev/null and b/assets/mage_flow/example.png differ
diff --git a/docs/edit.md b/docs/edit.md
index 0b38667ba..9791d046d 100644
--- a/docs/edit.md
+++ b/docs/edit.md
@@ -16,6 +16,7 @@ Depending on the architecture, different models handle reference images differen
| [**Flux.2 [Dev] / Flux.2 [Klein]**](./flux2.md) | `flux2` |
| [**Boogu Image Edit**](./boogu_image.md) | `z_image_omni` |
| **Krea2 (Community Edit LoRAs)** | `krea2_ostris_edit` |
+| [**Mage-Flow-Edit**](./mage_flow.md#image-editing) | `mage_flow` |
| **Anima (Community Edit LoRAs)** | `cosmos_reference` |
Stable-diffusion.spp also supports basic Unet-based editing models like instruct-pix2pix or CosXL-Edit. This document is not about those.
@@ -48,6 +49,7 @@ The `--ref-image-args` argument accepts a comma-separated list of key-value pair
| `qwen_layered` | Qwen Image Layered |
| `z_image_omni` | Boogu, Z-Image Omni |
| `krea2_ostris_edit` | Most Krea2 Community edit LoRAs (trained with Ostris script) |
+| `mage_flow` | Mage-Flow-Edit |
| `krea2_edit` | Specifically for [lbouaraba/krea2edit](https://huggingface.co/conradlocke/krea2-identity-edit). (or similar) |
| `cosmos_reference` | For Anima |
| `default` | Uses the automatic detection based on model architecture. |
@@ -83,6 +85,7 @@ For a technical overview of how each preset is configured, see the table below.
| `flux2` | No | `increase` | `none` | |
| `qwen` | Yes | `increase` | `area` | |
| `qwen_layered` | Yes | `decrease` | `area` | |
+| `mage_flow` | Yes | `increase` | `longest` | `vlm_max_size = 384`, VAE input resized to target |
| `z_image_omni` | Yes | `fixed` | `area` | |
| `krea2_ostris_edit`| Yes | `increase` | `area` | `force_ref_timestep_zero = true` |
| `krea2_edit` | Yes | `increase` | `longest` | `vlm_size = 768` |
diff --git a/docs/mage_flow.md b/docs/mage_flow.md
new file mode 100644
index 000000000..1d8c9a4cd
--- /dev/null
+++ b/docs/mage_flow.md
@@ -0,0 +1,45 @@
+# Mage-Flow
+
+[Mage-Flow](https://github.com/microsoft/Mage) uses a 4B native-resolution multimodal diffusion transformer, Qwen3-VL for text and image conditioning, and the 128-channel Mage-VAE. Both text-to-image and instruction-based image editing checkpoints are supported.
+
+## Download weights
+
+- Download Mage-Flow
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow/tree/main/transformer
+- Download Mage-Flow-Base
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow-Base/tree/main/transformer
+- Download Mage-Flow-Turbo
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow-Turbo/tree/main/transformer
+- Download Mage-Flow-Edit
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow-Edit/tree/main/transformer
+- Download Mage-Flow-Edit-Turbo
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow-Edit-Turbo/tree/main/transformer
+- Download Mage-Flow-Edit-Base
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow-Edit-Base/tree/main/transformer
+- Download Mage-Flow vae
+ - safetensors: https://huggingface.co/microsoft/Mage-Flow/tree/main/vae
+- Download Qwen3-VL 4B
+ - safetensors: https://huggingface.co/Comfy-Org/Krea-2/tree/main/text_encoders
+ - gguf: https://huggingface.co/Qwen/Qwen3-VL-4B-Instruct-GGUF/tree/main
+
+## Text-to-image
+
+Use 30 steps for Base models and 4 steps with `--cfg-scale 1` for Turbo models. Image dimensions must be multiples of 16; the official checkpoints are trained for native resolutions from 512 to 2048 pixels.
+
+```bash
+.\bin\Release\sd-cli.exe --diffusion-model ..\models\diffusion_models\Mage-Flow-Turbo.safetensors --llm ..\models\text_encoders\Qwen3-VL-4B-Instruct-Q4_K_M.gguf --vae ..\models\vae\mage_vae.safetensors -p "a lovely cat holding a sign says 'mage.cpp'" --cfg-scale 1.0 --steps 4 --diffusion-fa -v --offload-to-cpu
+```
+
+
+
+## Image editing
+
+Mage-Flow-Edit accepts one or more reference images. The default `mage_flow` reference preset sends each image to both Qwen3-VL and the diffusion transformer, caps the VLM copy's longest edge at 384 pixels, and keeps the VAE copy at the requested output resolution.
+
+For the Turbo edit checkpoint, use 4 steps and `--cfg-scale 1`.
+
+```bash
+.\bin\Release\sd-cli.exe --diffusion-model ..\models\diffusion_models\Mage-Flow-Edit.safetensors --llm ..\models\text_encoders\Qwen3-VL-4B-Instruct-Q4_K_M.gguf --llm_vision ..\models\text_encoders\Qwen3-VL-4B-Instruct-mmproj-BF16.gguf --vae ..\models\vae\mage_vae.safetensors -r ..\assets\flux\flux1-dev-q8_0.png -p "change 'flux.cpp' to 'mage.cpp'" --cfg-scale 4.0 --sampling-method euler -v --diffusion-fa --offload-to-cpu
+```
+
+
diff --git a/examples/CMakeLists.txt b/examples/CMakeLists.txt
index 2dcd1d53a..dc8150a9f 100644
--- a/examples/CMakeLists.txt
+++ b/examples/CMakeLists.txt
@@ -1,4 +1,4 @@
include_directories(${CMAKE_CURRENT_SOURCE_DIR})
add_subdirectory(cli)
-add_subdirectory(server)
\ No newline at end of file
+add_subdirectory(server)
diff --git a/src/conditioning/conditioner.hpp b/src/conditioning/conditioner.hpp
index a0bc49b37..e84e86297 100644
--- a/src/conditioning/conditioner.hpp
+++ b/src/conditioning/conditioner.hpp
@@ -1800,7 +1800,12 @@ struct LLMEmbedder : public Conditioner {
arch = LLM::LLMArch::GPT_OSS_20B;
} else if (sd_version_is_pid(version)) {
arch = LLM::LLMArch::GEMMA2_2B;
- } else if (sd_version_is_lingbot_video(version) || sd_version_is_ideogram4(version) || sd_version_is_boogu_image(version) || sd_version_is_sefi_image(version) || sd_version_is_krea2(version)) {
+ } else if (sd_version_is_lingbot_video(version) ||
+ sd_version_is_ideogram4(version) ||
+ sd_version_is_boogu_image(version) ||
+ sd_version_is_sefi_image(version) ||
+ sd_version_is_krea2(version) ||
+ sd_version_is_mage_flow(version)) {
arch = LLM::LLMArch::QWEN3_VL;
} else if (sd_version_is_z_image(version) || version == VERSION_OVIS_IMAGE || version == VERSION_FLUX2_KLEIN) {
arch = LLM::LLMArch::QWEN3;
@@ -2205,22 +2210,22 @@ struct LLMEmbedder : public Conditioner {
prompt += conditioner_params.text;
prompt_attn_range = {0, 0};
prompt += "<|im_end|>\n<|im_start|>assistant\n";
- } else if (sd_version_is_qwen_image(version)) {
+ } else if (sd_version_is_qwen_image(version) || sd_version_is_mage_flow(version)) {
if (llm->enable_vision && conditioner_params.ref_images != nullptr && !conditioner_params.ref_images->empty()) {
- LOG_INFO("QwenImageEditPlusPipeline");
+ LOG_INFO("%s", sd_version_is_mage_flow(version) ? "MageFlowEditPipeline" : "QwenImageEditPlusPipeline");
prompt_template_encode_start_idx = 64;
int image_embed_idx = 64 + 6;
int min_pixels = conditioner_params.ref_image_params.vlm_min_size;
if (min_pixels <= 0) {
- min_pixels = 384;
- if (resize_mode == RefImageResizeMode::AREA) {
+ min_pixels = sd_version_is_mage_flow(version) ? -1 : 384;
+ if (min_pixels > 0 && resize_mode == RefImageResizeMode::AREA) {
min_pixels *= min_pixels;
}
}
int max_pixels = conditioner_params.ref_image_params.vlm_max_size;
if (max_pixels <= 0) {
- max_pixels = 560;
+ max_pixels = sd_version_is_mage_flow(version) ? 384 : 560;
if (resize_mode == RefImageResizeMode::AREA) {
max_pixels *= max_pixels;
}
@@ -2248,7 +2253,7 @@ struct LLMEmbedder : public Conditioner {
image_embeds.emplace_back(image_embed_idx, image_embed);
image_embed_idx += 1 + static_cast(image_embed.shape()[1]) + 6;
- img_prompt += "Picture " + std::to_string(i + 1) + ": <|vision_start|>"; // [24669, 220, index, 25, 220, 151652]
+ img_prompt += (sd_version_is_mage_flow(version) ? "Image " : "Picture ") + std::to_string(i + 1) + ": <|vision_start|>";
int64_t num_image_tokens = image_embed.shape()[1];
img_prompt.reserve(num_image_tokens * placeholder.size());
for (int j = 0; j < num_image_tokens; j++) {
@@ -2276,6 +2281,9 @@ struct LLMEmbedder : public Conditioner {
prompt += "<|im_end|>\n<|im_start|>assistant\n";
}
+ if (sd_version_is_mage_flow(version)) {
+ max_length = 2048 + prompt_template_encode_start_idx;
+ }
} else if (sd_version_is_boogu_image(version)) {
prompt_template_encode_start_idx = 0;
diff --git a/src/core/ggml_extend.hpp b/src/core/ggml_extend.hpp
index 36067ce3c..18f1645bf 100644
--- a/src/core/ggml_extend.hpp
+++ b/src/core/ggml_extend.hpp
@@ -3609,18 +3609,19 @@ class Conv2d_grouped : public UnaryBlock {
if (groups == in_channels && groups == out_channels) {
ggml_tensor* res;
if (ctx->conv2d_direct_enabled) {
- res = ggml_conv_2d_dw_direct(ctx->ggml_ctx, x, w,
+ res = ggml_conv_2d_dw_direct(ctx->ggml_ctx, w, x,
stride.second, stride.first,
padding.second, padding.first,
dilation.second, dilation.first);
} else {
- res = ggml_conv_2d_dw(ctx->ggml_ctx, x, w,
+ res = ggml_conv_2d_dw(ctx->ggml_ctx, w, x,
stride.second, stride.first,
padding.second, padding.first,
dilation.second, dilation.first);
}
if (b) {
- res = ggml_add(ctx->ggml_ctx, res, b);
+ b = ggml_reshape_4d(ctx->ggml_ctx, b, 1, 1, b->ne[0], 1);
+ res = ggml_add_inplace(ctx->ggml_ctx, res, b);
}
return res;
}
diff --git a/src/model.h b/src/model.h
index 67523656a..b0cee896b 100644
--- a/src/model.h
+++ b/src/model.h
@@ -55,6 +55,7 @@ enum SDVersion {
VERSION_IDEOGRAM4,
VERSION_SEFI_IMAGE,
VERSION_KREA2,
+ VERSION_MAGE_FLOW,
VERSION_ESRGAN,
VERSION_COUNT,
};
@@ -227,6 +228,10 @@ static inline bool sd_version_is_krea2(SDVersion version) {
return false;
}
+static inline bool sd_version_is_mage_flow(SDVersion version) {
+ return version == VERSION_MAGE_FLOW;
+}
+
static inline bool sd_version_uses_flux_vae(SDVersion version) {
if (sd_version_is_flux(version) || sd_version_is_z_image(version) || sd_version_is_boogu_image(version) || sd_version_is_longcat(version)) {
return true;
@@ -283,7 +288,8 @@ static inline bool sd_version_is_dit(SDVersion version) {
sd_version_is_pid(version) ||
sd_version_is_ideogram4(version) ||
sd_version_is_sefi_image(version) ||
- sd_version_is_krea2(version)) {
+ sd_version_is_krea2(version) ||
+ sd_version_is_mage_flow(version)) {
return true;
}
return false;
diff --git a/src/model/common/rope.hpp b/src/model/common/rope.hpp
index 100f96ee0..c36cc7e9c 100644
--- a/src/model/common/rope.hpp
+++ b/src/model/common/rope.hpp
@@ -654,6 +654,43 @@ namespace Rope {
return embed_nd(ids, bs, static_cast(theta), axes_dim, wrap_dims);
}
+ __STATIC_INLINE__ std::vector gen_mage_flow_pe(int h,
+ int w,
+ int bs,
+ int context_len,
+ const std::vector& ref_latents,
+ int theta,
+ const std::vector& axes_dim) {
+ const int axes_dim_num = static_cast(axes_dim.size());
+ auto make_image_ids = [=](int image_h, int image_w, int image_index) {
+ std::vector> image_ids(static_cast(bs) * image_h * image_w,
+ std::vector(axes_dim_num, 0.f));
+ int h_start = -(image_h - image_h / 2);
+ int w_start = -(image_w - image_w / 2);
+ for (int b = 0; b < bs; ++b) {
+ for (int y = 0; y < image_h; ++y) {
+ for (int x = 0; x < image_w; ++x) {
+ auto& id = image_ids[static_cast(b) * image_h * image_w + y * image_w + x];
+ id[0] = static_cast(image_index);
+ id[1] = static_cast(h_start + y);
+ id[2] = static_cast(w_start + x);
+ }
+ }
+ }
+ return image_ids;
+ };
+ auto ids = gen_flux_txt_ids(bs, context_len, axes_dim_num, {});
+ auto img_ids = make_image_ids(h, w, 0);
+ ids = concat_ids(ids, img_ids, bs);
+ for (size_t i = 0; i < ref_latents.size(); ++i) {
+ auto ref_ids = make_image_ids(static_cast(ref_latents[i]->ne[1]),
+ static_cast(ref_latents[i]->ne[0]),
+ static_cast(i + 1));
+ ids = concat_ids(ids, ref_ids, bs);
+ }
+ return embed_nd(ids, bs, static_cast(theta), axes_dim);
+ }
+
__STATIC_INLINE__ std::vector> gen_lens_ids(int h,
int w,
int bs,
diff --git a/src/model/diffusion/mage_flow.hpp b/src/model/diffusion/mage_flow.hpp
new file mode 100644
index 000000000..6ac2d2aa5
--- /dev/null
+++ b/src/model/diffusion/mage_flow.hpp
@@ -0,0 +1,162 @@
+#ifndef __SD_MODEL_DIFFUSION_MAGE_FLOW_HPP__
+#define __SD_MODEL_DIFFUSION_MAGE_FLOW_HPP__
+
+#include
+#include
+
+#include "model/diffusion/qwen_image.hpp"
+
+namespace MageFlow {
+ constexpr int MAGE_FLOW_GRAPH_SIZE = 20480;
+
+ // Mage-Flow was trained with BF16-rounded timestep frequencies; using Qwen's F32 projection degrades generation quality.
+ struct MageFlowTimestepProjEmbeddings : public Qwen::QwenTimestepProjEmbeddings {
+ static constexpr int TIMESTEP_DIM = 256;
+ static constexpr int HALF_DIM = TIMESTEP_DIM / 2;
+
+ std::vector frequencies;
+ std::vector timesteps_proj;
+
+ explicit MageFlowTimestepProjEmbeddings(int64_t embedding_dim)
+ : QwenTimestepProjEmbeddings(embedding_dim), frequencies(HALF_DIM) {
+ for (int i = 0; i < HALF_DIM; ++i) {
+ float frequency = std::exp(-std::log(10000.f) * static_cast(i) / HALF_DIM);
+ frequencies[i] = ggml_bf16_to_fp32(ggml_fp32_to_bf16(frequency));
+ }
+ }
+
+ void prepare(const sd::Tensor& timesteps) {
+ size_t num_timesteps = static_cast(timesteps.numel());
+ timesteps_proj.resize(static_cast(TIMESTEP_DIM) * num_timesteps);
+ for (size_t b = 0; b < num_timesteps; ++b) {
+ float sigma = ggml_bf16_to_fp32(ggml_fp32_to_bf16(timesteps.values()[b] / 1000.f));
+ for (int i = 0; i < HALF_DIM; ++i) {
+ float argument = sigma * frequencies[i] * 1000.f;
+ timesteps_proj[b * TIMESTEP_DIM + i] =
+ ggml_bf16_to_fp32(ggml_fp32_to_bf16(std::cos(argument)));
+ timesteps_proj[b * TIMESTEP_DIM + HALF_DIM + i] =
+ ggml_bf16_to_fp32(ggml_fp32_to_bf16(std::sin(argument)));
+ }
+ }
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx,
+ ggml_tensor* timesteps,
+ ggml_tensor* addition_t_cond = nullptr) override {
+ GGML_ASSERT(addition_t_cond == nullptr);
+ GGML_ASSERT(timesteps_proj.size() ==
+ static_cast(TIMESTEP_DIM * ggml_nelements(timesteps)));
+ auto projection = ggml_new_tensor_2d(ctx->ggml_ctx,
+ GGML_TYPE_F32,
+ TIMESTEP_DIM,
+ ggml_nelements(timesteps));
+ ctx->bind_backend_tensor_data(projection, timesteps_proj.data());
+ auto timestep_embedder = std::dynamic_pointer_cast(blocks["timestep_embedder"]);
+ return timestep_embedder->forward(ctx, projection);
+ }
+ };
+
+ struct MageFlowRunner : public DiffusionModelRunner {
+ public:
+ Qwen::QwenImageConfig config;
+ Qwen::QwenImageModel mage_flow;
+ std::shared_ptr time_text_embed;
+ std::vector pe_vec;
+
+ MageFlowRunner(ggml_backend_t backend,
+ const String2TensorStorage& tensor_storage_map = {},
+ const std::string prefix = "",
+ std::shared_ptr weight_manager = nullptr)
+ : DiffusionModelRunner(backend, prefix, weight_manager) {
+ config.patch_size = 1;
+ config.in_channels = 128;
+ config.out_channels = 128;
+ config.num_layers = 12;
+ config.attention_head_dim = 128;
+ config.num_attention_heads = 24;
+ config.joint_attention_dim = 2560;
+ config.theta = 10000;
+ config.axes_dim = {16, 56, 56};
+ config.axes_dim_sum = 128;
+ time_text_embed = std::make_shared(
+ config.num_attention_heads * config.attention_head_dim);
+ mage_flow = Qwen::QwenImageModel(config, time_text_embed);
+ mage_flow.init(params_ctx, tensor_storage_map, prefix);
+ }
+
+ std::string get_desc() override {
+ return "mage_flow";
+ }
+
+ void get_param_tensors(std::map& tensors, const std::string& prefix) override {
+ mage_flow.get_param_tensors(tensors, prefix);
+ }
+
+ ggml_cgraph* build_graph(const sd::Tensor& x_tensor,
+ const sd::Tensor& timesteps_tensor,
+ const sd::Tensor& context_tensor,
+ const std::vector>& ref_latents_tensor = {}) {
+ ggml_cgraph* gf = new_graph_custom(MAGE_FLOW_GRAPH_SIZE);
+ ggml_tensor* x = make_input(x_tensor);
+ ggml_tensor* timesteps = make_input(timesteps_tensor);
+ GGML_ASSERT(x->ne[3] == 1);
+ GGML_ASSERT(!context_tensor.empty());
+ ggml_tensor* context = make_input(context_tensor);
+
+ std::vector ref_latents;
+ ref_latents.reserve(ref_latents_tensor.size());
+ for (const auto& ref_latent_tensor : ref_latents_tensor) {
+ ref_latents.push_back(make_input(ref_latent_tensor));
+ }
+
+ int batch_size = static_cast(x->ne[3]);
+ pe_vec = Rope::gen_mage_flow_pe(static_cast(x->ne[1]),
+ static_cast(x->ne[0]),
+ batch_size,
+ static_cast(context->ne[1]),
+ ref_latents,
+ config.theta,
+ config.axes_dim);
+ int pos_len = static_cast(pe_vec.size() / config.axes_dim_sum / 2);
+ auto pe = ggml_new_tensor_4d(compute_ctx, GGML_TYPE_F32, 2, 2, config.axes_dim_sum / 2, pos_len);
+ set_backend_tensor_data(pe, pe_vec.data());
+
+ time_text_embed->prepare(timesteps_tensor);
+ auto runner_ctx = get_context();
+ auto out = mage_flow.forward(&runner_ctx,
+ x,
+ timesteps,
+ nullptr,
+ context,
+ pe,
+ ref_latents);
+ ggml_build_forward_expand(gf, out);
+ return gf;
+ }
+
+ sd::Tensor compute(int n_threads,
+ const sd::Tensor& x,
+ const sd::Tensor& timesteps,
+ const sd::Tensor& context,
+ const std::vector>& ref_latents = {}) {
+ auto get_graph = [&]() -> ggml_cgraph* {
+ return build_graph(x, timesteps, context, ref_latents);
+ };
+ return restore_trailing_singleton_dims(GGMLRunner::compute(get_graph, n_threads, false, false, false), x.dim());
+ }
+
+ sd::Tensor compute(int n_threads,
+ const DiffusionParams& diffusion_params) override {
+ GGML_ASSERT(diffusion_params.x != nullptr);
+ GGML_ASSERT(diffusion_params.timesteps != nullptr);
+ static const std::vector> empty_ref_latents;
+ return compute(n_threads,
+ *diffusion_params.x,
+ *diffusion_params.timesteps,
+ tensor_or_empty(diffusion_params.context),
+ diffusion_params.ref_latents && diffusion_params.ref_image_params.pass_to_dit ? *diffusion_params.ref_latents : empty_ref_latents);
+ }
+ };
+} // namespace MageFlow
+
+#endif // __SD_MODEL_DIFFUSION_MAGE_FLOW_HPP__
diff --git a/src/model/diffusion/model.hpp b/src/model/diffusion/model.hpp
index d8ee3d26f..31c504833 100644
--- a/src/model/diffusion/model.hpp
+++ b/src/model/diffusion/model.hpp
@@ -26,6 +26,7 @@ struct RefImageParams {
RefImageResizeMode vlm_resize_mode = RefImageResizeMode::AREA;
int vlm_min_size = -1;
int vlm_max_size = -1;
+ bool resize_vae_to_target = false;
};
const std::unordered_map REF_IMAGE_PRESETS = {
@@ -34,6 +35,7 @@ const std::unordered_map REF_IMAGE_PRESETS = {
{"flux2", {false, true, Rope::RefIndexMode::INCREASE, false, true, -1, RefImageResizeMode::NONE, -1, -1}},
{"qwen", {true, true, Rope::RefIndexMode::INCREASE, false, true, -1, RefImageResizeMode::AREA, -1, -1}},
{"qwen_layered", {true, true, Rope::RefIndexMode::DECREASE, false, true, -1, RefImageResizeMode::AREA, -1, -1}},
+ {"mage_flow", {true, true, Rope::RefIndexMode::INCREASE, false, true, -1, RefImageResizeMode::LONGEST_SIDE, -1, 384, true}},
{"z_image_omni", {true, true, Rope::RefIndexMode::FIXED, false, true, -1, RefImageResizeMode::AREA, -1, -1}},
{"krea2_ostris_edit", {true, true, Rope::RefIndexMode::INCREASE, true, true, -1, RefImageResizeMode::AREA, -1, -1}},
{"krea2_edit", {true, true, Rope::RefIndexMode::INCREASE, false, true, -1, RefImageResizeMode::LONGEST_SIDE, 768, 768}},
diff --git a/src/model/diffusion/qwen_image.hpp b/src/model/diffusion/qwen_image.hpp
index c6c3622f9..6f08d6180 100644
--- a/src/model/diffusion/qwen_image.hpp
+++ b/src/model/diffusion/qwen_image.hpp
@@ -103,9 +103,9 @@ namespace Qwen {
}
}
- ggml_tensor* forward(GGMLRunnerContext* ctx,
- ggml_tensor* timesteps,
- ggml_tensor* addition_t_cond = nullptr) {
+ virtual ggml_tensor* forward(GGMLRunnerContext* ctx,
+ ggml_tensor* timesteps,
+ ggml_tensor* addition_t_cond = nullptr) {
// timesteps: [N,]
// return: [N, embedding_dim]
auto timestep_embedder = std::dynamic_pointer_cast(blocks["timestep_embedder"]);
@@ -416,10 +416,14 @@ namespace Qwen {
public:
QwenImageModel() {}
- QwenImageModel(QwenImageConfig config)
+ QwenImageModel(QwenImageConfig config,
+ std::shared_ptr time_text_embed = nullptr)
: config(config) {
- int64_t inner_dim = config.num_attention_heads * config.attention_head_dim;
- blocks["time_text_embed"] = std::shared_ptr(new QwenTimestepProjEmbeddings(inner_dim, config.use_additional_t_cond));
+ int64_t inner_dim = config.num_attention_heads * config.attention_head_dim;
+ if (time_text_embed == nullptr) {
+ time_text_embed = std::make_shared(inner_dim, config.use_additional_t_cond);
+ }
+ blocks["time_text_embed"] = std::move(time_text_embed);
blocks["txt_norm"] = std::shared_ptr(new RMSNorm(config.joint_attention_dim, 1e-6f));
blocks["img_in"] = std::shared_ptr(new Linear(config.in_channels, inner_dim));
blocks["txt_in"] = std::shared_ptr(new Linear(config.joint_attention_dim, inner_dim));
diff --git a/src/model/vae/mage_vae.hpp b/src/model/vae/mage_vae.hpp
new file mode 100644
index 000000000..39075a836
--- /dev/null
+++ b/src/model/vae/mage_vae.hpp
@@ -0,0 +1,521 @@
+#ifndef __SD_MODEL_VAE_MAGE_VAE_HPP__
+#define __SD_MODEL_VAE_MAGE_VAE_HPP__
+
+#include "model/diffusion/dit.hpp"
+#include "model/vae/vae.hpp"
+
+namespace MageVAE {
+ constexpr int MAGE_VAE_GRAPH_SIZE = 327680;
+ constexpr int HIDDEN_SIZE = 384;
+ constexpr int LATENT_CHANNELS = 128;
+ constexpr int PATCH_SIZE = 16;
+
+ struct LayerNorm2d : public UnaryBlock {
+ int64_t channels;
+ bool affine;
+ std::string prefix;
+
+ void init_params(ggml_context* ctx,
+ const String2TensorStorage& tensor_storage_map = {},
+ const std::string prefix = "") override {
+ this->prefix = prefix;
+ if (affine) {
+ params["weight"] = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, channels);
+ params["bias"] = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, channels);
+ }
+ }
+
+ LayerNorm2d(int64_t channels, bool affine = true)
+ : channels(channels), affine(affine) {}
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x) override {
+ ggml_tensor* weight = affine ? params["weight"] : nullptr;
+ ggml_tensor* bias = affine ? params["bias"] : nullptr;
+ if (affine && ctx->weight_adapter) {
+ weight = ctx->weight_adapter->patch_weight(ctx->ggml_ctx, ctx->backend, weight, prefix + "weight");
+ bias = ctx->weight_adapter->patch_weight(ctx->ggml_ctx, ctx->backend, bias, prefix + "bias");
+ }
+ // [N, C, H, W] -> [N, H, W, C] so layer norm reduces over channels.
+ x = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, x, 1, 2, 0, 3));
+ x = ggml_ext_layer_norm(ctx->ggml_ctx, x, weight, bias, 1e-6f);
+ return ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, x, 2, 0, 1, 3));
+ }
+ };
+
+ inline ggml_tensor* modulate_2d(ggml_context* ctx,
+ ggml_tensor* x,
+ ggml_tensor* shift,
+ ggml_tensor* scale) {
+ shift = ggml_reshape_4d(ctx, shift, 1, 1, shift->ne[0], shift->ne[1]);
+ scale = ggml_reshape_4d(ctx, scale, 1, 1, scale->ne[0], scale->ne[1]);
+ return ggml_add(ctx, ggml_mul(ctx, x, ggml_add(ctx, scale, ggml_ext_ones(ctx, 1, 1, 1, 1))), shift);
+ }
+
+ inline ggml_tensor* channel_attention(GGMLRunnerContext* ctx,
+ ggml_tensor* x,
+ Conv2d* projection) {
+ auto pooled = ggml_reshape_3d(ctx->ggml_ctx, x, x->ne[0] * x->ne[1], x->ne[2], x->ne[3]);
+ pooled = ggml_mean(ctx->ggml_ctx, pooled);
+ pooled = ggml_reshape_4d(ctx->ggml_ctx, pooled, 1, 1, x->ne[2], x->ne[3]);
+ pooled = ggml_sigmoid(ctx->ggml_ctx, projection->forward(ctx, pooled));
+ return ggml_mul(ctx->ggml_ctx, x, pooled);
+ }
+
+ struct TimestepEmbedder : public GGMLBlock {
+ TimestepEmbedder() {
+ blocks["mlp.0"] = std::make_shared(256, HIDDEN_SIZE);
+ blocks["mlp.2"] = std::make_shared(HIDDEN_SIZE, HIDDEN_SIZE);
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* timestep) {
+ auto linear_0 = std::dynamic_pointer_cast(blocks["mlp.0"]);
+ auto linear_2 = std::dynamic_pointer_cast(blocks["mlp.2"]);
+ auto x = ggml_ext_timestep_embedding(ctx->ggml_ctx, timestep, 256, 10000, 1.f);
+ x = linear_0->forward(ctx, x);
+ x = ggml_silu_inplace(ctx->ggml_ctx, x);
+ return linear_2->forward(ctx, x);
+ }
+ };
+
+ struct EncoderDiCoBlock : public UnaryBlock {
+ explicit EncoderDiCoBlock(int64_t channels) {
+ blocks["conv1"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["conv2"] = std::make_shared(channels, channels, static_cast(channels), std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ blocks["conv3"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["ca.1"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["conv4"] = std::make_shared(channels, channels * 4, std::pair{1, 1});
+ blocks["conv5"] = std::make_shared(channels * 4, channels, std::pair{1, 1});
+ blocks["norm1"] = std::make_shared(channels);
+ blocks["norm2"] = std::make_shared(channels);
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* input) override {
+ auto conv1 = std::dynamic_pointer_cast(blocks["conv1"]);
+ auto conv2 = std::dynamic_pointer_cast(blocks["conv2"]);
+ auto conv3 = std::dynamic_pointer_cast(blocks["conv3"]);
+ auto ca = std::dynamic_pointer_cast(blocks["ca.1"]);
+ auto conv4 = std::dynamic_pointer_cast(blocks["conv4"]);
+ auto conv5 = std::dynamic_pointer_cast(blocks["conv5"]);
+ auto norm1 = std::dynamic_pointer_cast(blocks["norm1"]);
+ auto norm2 = std::dynamic_pointer_cast(blocks["norm2"]);
+
+ auto x = norm1->forward(ctx, input);
+ x = conv1->forward(ctx, x);
+ x = conv2->forward(ctx, x);
+ x = ggml_gelu(ctx->ggml_ctx, x);
+ x = channel_attention(ctx, x, ca.get());
+ x = conv3->forward(ctx, x);
+ x = ggml_add(ctx->ggml_ctx, input, x);
+ auto h = norm2->forward(ctx, x);
+ h = conv4->forward(ctx, h);
+ h = ggml_gelu(ctx->ggml_ctx, h);
+ h = conv5->forward(ctx, h);
+ return ggml_add(ctx->ggml_ctx, x, h);
+ }
+ };
+
+ struct DiCoBlock : public GGMLBlock {
+ explicit DiCoBlock(int64_t channels) {
+ blocks["conv1"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["conv2"] = std::make_shared(channels, channels, static_cast(channels), std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ blocks["conv3"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["ca.1"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["conv4"] = std::make_shared(channels, channels * 4, std::pair{1, 1});
+ blocks["conv5"] = std::make_shared(channels * 4, channels, std::pair{1, 1});
+ blocks["norm1"] = std::make_shared(channels, false);
+ blocks["norm2"] = std::make_shared(channels, false);
+ blocks["adaLN_modulation.1"] = std::make_shared(channels, channels * 6);
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* input, ggml_tensor* condition) {
+ auto conv1 = std::dynamic_pointer_cast(blocks["conv1"]);
+ auto conv2 = std::dynamic_pointer_cast(blocks["conv2"]);
+ auto conv3 = std::dynamic_pointer_cast(blocks["conv3"]);
+ auto ca = std::dynamic_pointer_cast(blocks["ca.1"]);
+ auto conv4 = std::dynamic_pointer_cast(blocks["conv4"]);
+ auto conv5 = std::dynamic_pointer_cast(blocks["conv5"]);
+ auto norm1 = std::dynamic_pointer_cast(blocks["norm1"]);
+ auto norm2 = std::dynamic_pointer_cast(blocks["norm2"]);
+ auto ada = std::dynamic_pointer_cast(blocks["adaLN_modulation.1"]);
+
+ auto params = ada->forward(ctx, ggml_silu(ctx->ggml_ctx, condition));
+ auto chunks = ggml_ext_chunk(ctx->ggml_ctx, params, 6, 0);
+ auto x = norm1->forward(ctx, input);
+ x = modulate_2d(ctx->ggml_ctx, x, chunks[0], chunks[1]);
+ x = conv1->forward(ctx, x);
+ x = conv2->forward(ctx, x);
+ x = ggml_gelu(ctx->ggml_ctx, x);
+ x = channel_attention(ctx, x, ca.get());
+ x = conv3->forward(ctx, x);
+ auto gate_1 = ggml_reshape_4d(ctx->ggml_ctx, chunks[2], 1, 1, chunks[2]->ne[0], chunks[2]->ne[1]);
+ x = ggml_add(ctx->ggml_ctx, input, ggml_mul(ctx->ggml_ctx, x, gate_1));
+
+ auto h = norm2->forward(ctx, x);
+ h = modulate_2d(ctx->ggml_ctx, h, chunks[3], chunks[4]);
+ h = conv4->forward(ctx, h);
+ h = ggml_gelu(ctx->ggml_ctx, h);
+ h = conv5->forward(ctx, h);
+ auto gate_2 = ggml_reshape_4d(ctx->ggml_ctx, chunks[5], 1, 1, chunks[5]->ne[0], chunks[5]->ne[1]);
+ return ggml_add(ctx->ggml_ctx, x, ggml_mul(ctx->ggml_ctx, h, gate_2));
+ }
+ };
+
+ struct MageResnetBlock : public UnaryBlock {
+ explicit MageResnetBlock(int64_t channels) {
+ blocks["norm1"] = std::make_shared(channels);
+ blocks["conv1"] = std::make_shared(channels, channels, std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ blocks["norm2"] = std::make_shared(channels);
+ blocks["conv2"] = std::make_shared(channels, channels, std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* input) override {
+ auto norm1 = std::dynamic_pointer_cast(blocks["norm1"]);
+ auto conv1 = std::dynamic_pointer_cast(blocks["conv1"]);
+ auto norm2 = std::dynamic_pointer_cast(blocks["norm2"]);
+ auto conv2 = std::dynamic_pointer_cast(blocks["conv2"]);
+ auto x = conv1->forward(ctx, ggml_silu(ctx->ggml_ctx, norm1->forward(ctx, input)));
+ x = conv2->forward(ctx, ggml_silu(ctx->ggml_ctx, norm2->forward(ctx, x)));
+ return ggml_add(ctx->ggml_ctx, input, x);
+ }
+ };
+
+ inline ggml_tensor* replicate_pad_right_bottom(ggml_context* ctx,
+ ggml_tensor* x,
+ int pad_w,
+ int pad_h) {
+ if (pad_w > 0) {
+ auto edge = ggml_ext_slice(ctx, x, 0, x->ne[0] - 1, x->ne[0]);
+ edge = ggml_repeat_4d(ctx, edge, pad_w, x->ne[1], x->ne[2], x->ne[3]);
+ x = ggml_concat(ctx, x, edge, 0);
+ }
+ if (pad_h > 0) {
+ auto edge = ggml_ext_slice(ctx, x, 1, x->ne[1] - 1, x->ne[1]);
+ edge = ggml_repeat_4d(ctx, edge, x->ne[0], pad_h, x->ne[2], x->ne[3]);
+ x = ggml_concat(ctx, x, edge, 1);
+ }
+ return x;
+ }
+
+ struct MageAttnBlock : public UnaryBlock {
+ int64_t channels;
+ int patch_size;
+
+ MageAttnBlock(int64_t channels, int patch_size = 32)
+ : channels(channels), patch_size(patch_size) {
+ blocks["norm"] = std::make_shared(channels);
+ blocks["q"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["k"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["v"] = std::make_shared(channels, channels, std::pair{1, 1});
+ blocks["proj_out"] = std::make_shared(channels, channels, std::pair{1, 1});
+ }
+
+ ggml_tensor* to_patches(ggml_context* ctx, ggml_tensor* x) {
+ x = DiT::patchify(ctx, x, patch_size, patch_size);
+ x = ggml_reshape_4d(ctx, x, patch_size * patch_size, channels, x->ne[1], x->ne[2]);
+ // [N, np, C, P] -> [N, np, P, C] for attention over P pixels.
+ x = ggml_ext_cont(ctx, ggml_permute(ctx, x, 1, 0, 2, 3));
+ return ggml_reshape_3d(ctx, x, channels, patch_size * patch_size, x->ne[2] * x->ne[3]);
+ }
+
+ ggml_tensor* from_patches(ggml_context* ctx,
+ ggml_tensor* x,
+ int64_t patch_count,
+ int64_t batch_size,
+ int64_t h_patches,
+ int64_t w_patches) {
+ x = ggml_reshape_4d(ctx, x, channels, patch_size * patch_size, patch_count, batch_size);
+ // [N, np, P, C] -> [N, np, C, P] before spatial unpatchify.
+ x = ggml_ext_cont(ctx, ggml_permute(ctx, x, 1, 0, 2, 3));
+ x = ggml_reshape_3d(ctx, x, patch_size * patch_size * channels, patch_count, batch_size);
+ return DiT::unpatchify(ctx, x, h_patches, w_patches, patch_size, patch_size);
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* input) override {
+ auto norm = std::dynamic_pointer_cast(blocks["norm"]);
+ auto q_proj = std::dynamic_pointer_cast(blocks["q"]);
+ auto k_proj = std::dynamic_pointer_cast(blocks["k"]);
+ auto v_proj = std::dynamic_pointer_cast(blocks["v"]);
+ auto proj_out = std::dynamic_pointer_cast(blocks["proj_out"]);
+
+ int64_t width = input->ne[0];
+ int64_t height = input->ne[1];
+ int64_t batch = input->ne[3];
+ int pad_w = (patch_size - static_cast(width % patch_size)) % patch_size;
+ int pad_h = (patch_size - static_cast(height % patch_size)) % patch_size;
+ int64_t wp = (width + pad_w) / patch_size;
+ int64_t hp = (height + pad_h) / patch_size;
+ int64_t np = wp * hp;
+
+ auto h = norm->forward(ctx, input);
+ auto q = replicate_pad_right_bottom(ctx->ggml_ctx, q_proj->forward(ctx, h), pad_w, pad_h);
+ auto k = replicate_pad_right_bottom(ctx->ggml_ctx, k_proj->forward(ctx, h), pad_w, pad_h);
+ auto v = replicate_pad_right_bottom(ctx->ggml_ctx, v_proj->forward(ctx, h), pad_w, pad_h);
+ q = to_patches(ctx->ggml_ctx, q);
+ k = to_patches(ctx->ggml_ctx, k);
+ v = to_patches(ctx->ggml_ctx, v);
+ h = ggml_ext_attention_ext(ctx->ggml_ctx, ctx->backend, q, k, v, 1, nullptr, false, ctx->flash_attn_enabled);
+ h = from_patches(ctx->ggml_ctx, h, np, batch, hp, wp);
+ if (pad_h > 0) {
+ h = ggml_ext_slice(ctx->ggml_ctx, h, 1, 0, height);
+ }
+ if (pad_w > 0) {
+ h = ggml_ext_slice(ctx->ggml_ctx, h, 0, 0, width);
+ }
+ return ggml_add(ctx->ggml_ctx, input, proj_out->forward(ctx, h));
+ }
+ };
+
+ struct Decoder : public UnaryBlock {
+ Decoder() {
+ blocks["conv_in"] = std::make_shared(LATENT_CHANNELS, HIDDEN_SIZE, std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ blocks["block.0"] = std::make_shared(HIDDEN_SIZE);
+ blocks["block.1"] = std::make_shared(HIDDEN_SIZE);
+ blocks["block.2"] = std::make_shared(HIDDEN_SIZE);
+ blocks["block.3"] = std::make_shared(HIDDEN_SIZE);
+ blocks["block.4"] = std::make_shared(HIDDEN_SIZE);
+ blocks["norm_out"] = std::make_shared(HIDDEN_SIZE);
+ blocks["conv_out"] = std::make_shared(HIDDEN_SIZE, HIDDEN_SIZE, std::pair{3, 3}, std::pair{1, 1}, std::pair{1, 1});
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x) override {
+ x = std::dynamic_pointer_cast(blocks["conv_in"])->forward(ctx, x);
+ for (int i = 0; i < 5; ++i) {
+ x = std::dynamic_pointer_cast(blocks["block." + std::to_string(i)])->forward(ctx, x);
+ }
+ x = std::dynamic_pointer_cast(blocks["norm_out"])->forward(ctx, x);
+ x = ggml_silu(ctx->ggml_ctx, x);
+ return std::dynamic_pointer_cast(blocks["conv_out"])->forward(ctx, x);
+ }
+ };
+
+ struct DConvEncoder : public UnaryBlock {
+ DConvEncoder() {
+ blocks["patch_cond_embed"] = std::make_shared(3, 768, std::pair{PATCH_SIZE, PATCH_SIZE}, std::pair{PATCH_SIZE, PATCH_SIZE});
+ for (int i = 0; i < 2; ++i) {
+ blocks["head_blocks." + std::to_string(i)] = std::make_shared(768);
+ }
+ blocks["proj_down"] = std::make_shared(768, HIDDEN_SIZE, std::pair{1, 1});
+ blocks["z_proj"] = std::make_shared(LATENT_CHANNELS, HIDDEN_SIZE, std::pair{1, 1});
+ blocks["fuse_proj"] = std::make_shared(HIDDEN_SIZE * 2, HIDDEN_SIZE, std::pair{1, 1});
+ blocks["t_embedder"] = std::make_shared();
+ for (int i = 0; i < 21; ++i) {
+ blocks["blocks." + std::to_string(i)] = std::make_shared(HIDDEN_SIZE);
+ }
+ blocks["norm_out"] = std::make_shared(HIDDEN_SIZE);
+ blocks["proj_out"] = std::make_shared(HIDDEN_SIZE, LATENT_CHANNELS * 2, std::pair{1, 1});
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* image) override {
+ auto cond = std::dynamic_pointer_cast(blocks["patch_cond_embed"])->forward(ctx, image);
+ for (int i = 0; i < 2; ++i) {
+ cond = std::dynamic_pointer_cast(blocks["head_blocks." + std::to_string(i)])->forward(ctx, cond);
+ }
+ cond = std::dynamic_pointer_cast(blocks["proj_down"])->forward(ctx, cond);
+ auto z = ggml_ext_zeros(ctx->ggml_ctx, cond->ne[0], cond->ne[1], LATENT_CHANNELS, cond->ne[3]);
+ z = std::dynamic_pointer_cast(blocks["z_proj"])->forward(ctx, z);
+ z = ggml_concat(ctx->ggml_ctx, cond, z, 2);
+ z = std::dynamic_pointer_cast(blocks["fuse_proj"])->forward(ctx, z);
+ auto t = ggml_ext_zeros(ctx->ggml_ctx, image->ne[3], 1, 1, 1);
+ t = ggml_reshape_1d(ctx->ggml_ctx, t, image->ne[3]);
+ auto c = std::dynamic_pointer_cast(blocks["t_embedder"])->forward(ctx, t);
+ for (int i = 0; i < 21; ++i) {
+ z = std::dynamic_pointer_cast(blocks["blocks." + std::to_string(i)])->forward(ctx, z, c);
+ }
+ z = std::dynamic_pointer_cast(blocks["norm_out"])->forward(ctx, z);
+ return std::dynamic_pointer_cast(blocks["proj_out"])->forward(ctx, z);
+ }
+ };
+
+ struct MLPResBlock : public GGMLBlock {
+ MLPResBlock() {
+ blocks["in_ln"] = std::make_shared(32, 1e-6f);
+ blocks["mlp.0"] = std::make_shared(32, 32);
+ blocks["mlp.2"] = std::make_shared(32, 32);
+ blocks["adaLN_modulation.1"] = std::make_shared(32, 96);
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, ggml_tensor* condition) {
+ auto params = std::dynamic_pointer_cast(blocks["adaLN_modulation.1"])->forward(ctx, ggml_silu(ctx->ggml_ctx, condition));
+ auto chunks = ggml_ext_chunk(ctx->ggml_ctx, params, 3, 0);
+ auto h = std::dynamic_pointer_cast(blocks["in_ln"])->forward(ctx, x);
+ h = ggml_add(ctx->ggml_ctx, ggml_mul(ctx->ggml_ctx, h, ggml_add(ctx->ggml_ctx, chunks[1], ggml_ext_ones(ctx->ggml_ctx, 1, 1, 1, 1))), chunks[0]);
+ h = std::dynamic_pointer_cast(blocks["mlp.0"])->forward(ctx, h);
+ h = ggml_silu(ctx->ggml_ctx, h);
+ h = std::dynamic_pointer_cast(blocks["mlp.2"])->forward(ctx, h);
+ return ggml_add(ctx->ggml_ctx, x, ggml_mul(ctx->ggml_ctx, chunks[2], h));
+ }
+ };
+
+ struct DConvDenoiser : public GGMLBlock {
+ DConvDenoiser() {
+ blocks["t_embedder"] = std::make_shared();
+ blocks["y_embedder_x"] = std::make_shared(HIDDEN_SIZE, 32 * PATCH_SIZE * PATCH_SIZE, std::pair{1, 1});
+ blocks["x_embedder.embedder.0"] = std::make_shared(3 + 32 + 64, 32);
+ blocks["s_embedder.proj1"] = std::make_shared(3, LATENT_CHANNELS, std::pair{PATCH_SIZE, PATCH_SIZE}, std::pair{PATCH_SIZE, PATCH_SIZE}, std::pair{0, 0}, std::pair{1, 1}, false);
+ blocks["s_embedder.proj2"] = std::make_shared(LATENT_CHANNELS + HIDDEN_SIZE, HIDDEN_SIZE, std::pair{1, 1});
+ for (int i = 0; i < 21; ++i) {
+ blocks["blocks." + std::to_string(i)] = std::make_shared(HIDDEN_SIZE);
+ }
+ blocks["dec_net.cond_embed"] = std::make_shared(HIDDEN_SIZE, PATCH_SIZE * PATCH_SIZE * 32);
+ blocks["dec_net.input_proj"] = std::make_shared(32, 32);
+ for (int i = 0; i < 3; ++i) {
+ blocks["dec_net.res_blocks." + std::to_string(i)] = std::make_shared();
+ }
+ blocks["final_layer.norm"] = std::make_shared(32);
+ blocks["final_layer.linear"] = std::make_shared(32, 3);
+ blocks["y_embedder.decoder"] = std::make_shared();
+ }
+
+ ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* latent, ggml_tensor* dct) {
+ auto cond = std::dynamic_pointer_cast(blocks["y_embedder.decoder"])->forward(ctx, latent);
+ int64_t w = cond->ne[0];
+ int64_t h = cond->ne[1];
+ int64_t n = cond->ne[3];
+ int64_t length = w * h;
+
+ auto image = ggml_ext_zeros(ctx->ggml_ctx, w * PATCH_SIZE, h * PATCH_SIZE, 3, n);
+ auto t = ggml_ext_zeros(ctx->ggml_ctx, n, 1, 1, 1);
+ t = ggml_reshape_1d(ctx->ggml_ctx, t, n);
+ auto c = std::dynamic_pointer_cast(blocks["t_embedder"])->forward(ctx, t);
+
+ auto s0 = std::dynamic_pointer_cast(blocks["s_embedder.proj1"])->forward(ctx, image);
+ s0 = ggml_concat(ctx->ggml_ctx, s0, cond, 2);
+ auto s = std::dynamic_pointer_cast(blocks["s_embedder.proj2"])->forward(ctx, s0);
+ for (int i = 0; i < 21; ++i) {
+ s = std::dynamic_pointer_cast(blocks["blocks." + std::to_string(i)])->forward(ctx, s, c);
+ }
+ // [N, C, H, W] -> [N*H*W, C].
+ s = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, s, 1, 2, 0, 3));
+ s = ggml_reshape_2d(ctx->ggml_ctx, s, HIDDEN_SIZE, length * n);
+
+ auto y = std::dynamic_pointer_cast(blocks["y_embedder_x"])->forward(ctx, cond);
+ // Split 32*P channels as [32, P], then produce [N*L, P, 32].
+ y = ggml_reshape_4d(ctx->ggml_ctx, y, length, PATCH_SIZE * PATCH_SIZE, 32, n);
+ y = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, y, 2, 1, 0, 3));
+ y = ggml_reshape_3d(ctx->ggml_ctx, y, 32, PATCH_SIZE * PATCH_SIZE, length * n);
+ auto zeros = ggml_ext_zeros(ctx->ggml_ctx, 3, PATCH_SIZE * PATCH_SIZE, length * n, 1);
+ dct = ggml_repeat_4d(ctx->ggml_ctx, dct, 64, PATCH_SIZE * PATCH_SIZE, length * n, 1);
+ auto x = ggml_concat(ctx->ggml_ctx, zeros, y, 0);
+ x = ggml_concat(ctx->ggml_ctx, x, dct, 0);
+ x = std::dynamic_pointer_cast(blocks["x_embedder.embedder.0"])->forward(ctx, x);
+ x = std::dynamic_pointer_cast(blocks["dec_net.input_proj"])->forward(ctx, x);
+
+ auto dec_cond = std::dynamic_pointer_cast(blocks["dec_net.cond_embed"])->forward(ctx, s);
+ dec_cond = ggml_reshape_3d(ctx->ggml_ctx, dec_cond, 32, PATCH_SIZE * PATCH_SIZE, length * n);
+ for (int i = 0; i < 3; ++i) {
+ x = std::dynamic_pointer_cast(blocks["dec_net.res_blocks." + std::to_string(i)])->forward(ctx, x, dec_cond);
+ }
+ x = std::dynamic_pointer_cast(blocks["final_layer.norm"])->forward(ctx, x);
+ x = std::dynamic_pointer_cast(blocks["final_layer.linear"])->forward(ctx, x);
+ // [N*L, P, 3] -> [N, L, 3*P] for fold/unpatchify.
+ x = ggml_reshape_4d(ctx->ggml_ctx, x, 3, PATCH_SIZE * PATCH_SIZE, length, n);
+ x = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, x, 1, 0, 2, 3));
+ x = ggml_reshape_3d(ctx->ggml_ctx, x, 3 * PATCH_SIZE * PATCH_SIZE, length, n);
+ return DiT::unpatchify(ctx->ggml_ctx, x, h, w, PATCH_SIZE, PATCH_SIZE);
+ }
+ };
+
+ struct MageVAEModel : public GGMLBlock {
+ MageVAEModel() {
+ blocks["student.dconv_encoder"] = std::make_shared();
+ blocks["pipeline"] = std::make_shared();
+ }
+
+ ggml_tensor* encode(GGMLRunnerContext* ctx, ggml_tensor* image) {
+ return std::dynamic_pointer_cast(blocks["student.dconv_encoder"])->forward(ctx, image);
+ }
+
+ ggml_tensor* decode(GGMLRunnerContext* ctx, ggml_tensor* latent, ggml_tensor* dct) {
+ return std::dynamic_pointer_cast(blocks["pipeline"])->forward(ctx, latent, dct);
+ }
+ };
+
+ struct MageVAERunner : public VAE {
+ MageVAEModel model;
+ std::vector dct_vec;
+
+ MageVAERunner(ggml_backend_t backend,
+ const String2TensorStorage& tensor_storage_map,
+ const std::string& prefix,
+ std::shared_ptr weight_manager = nullptr)
+ : VAE(VERSION_MAGE_FLOW, backend, prefix, weight_manager) {
+ model = MageVAEModel();
+ model.init(params_ctx, tensor_storage_map, prefix);
+ dct_vec.resize(64 * PATCH_SIZE * PATCH_SIZE);
+ constexpr float pi = 3.14159265358979323846f;
+ for (int py = 0; py < PATCH_SIZE; ++py) {
+ float y = static_cast(py) / static_cast(PATCH_SIZE - 1);
+ for (int px = 0; px < PATCH_SIZE; ++px) {
+ float x = static_cast(px) / static_cast(PATCH_SIZE - 1);
+ int pos = py * PATCH_SIZE + px;
+ for (int fy = 0; fy < 8; ++fy) {
+ for (int fx = 0; fx < 8; ++fx) {
+ int freq = fx * 8 + fy;
+ float freq_x = static_cast(fx) * 8.f / 7.f;
+ float freq_y = static_cast(fy) * 8.f / 7.f;
+ float coeff = 1.f / (1.f + freq_x * freq_y);
+ dct_vec[freq + 64 * pos] = std::cos(x * freq_x * pi) *
+ std::cos(y * freq_y * pi) * coeff;
+ }
+ }
+ }
+ }
+ }
+
+ std::string get_desc() override {
+ return "mage_vae";
+ }
+
+ void get_param_tensors(std::map& tensors) override {
+ model.get_param_tensors(tensors, weight_prefix);
+ }
+
+ ggml_cgraph* build_graph(const sd::Tensor& input_tensor, bool decode_graph) {
+ ggml_cgraph* gf = new_graph_custom(MAGE_VAE_GRAPH_SIZE);
+ auto input = make_input(input_tensor);
+ auto runner_ctx = get_context();
+ ggml_tensor* dct = nullptr;
+ if (decode_graph) {
+ dct = ggml_new_tensor_3d(compute_ctx, GGML_TYPE_F32, 64, PATCH_SIZE * PATCH_SIZE, 1);
+ set_backend_tensor_data(dct, dct_vec.data());
+ }
+ auto out = decode_graph ? model.decode(&runner_ctx, input, dct) : model.encode(&runner_ctx, input);
+ ggml_build_forward_expand(gf, out);
+ return gf;
+ }
+
+ sd::Tensor _compute(const int n_threads,
+ const sd::Tensor& input,
+ bool decode_graph) override {
+ auto get_graph = [&]() -> ggml_cgraph* {
+ return build_graph(input, decode_graph);
+ };
+ return restore_trailing_singleton_dims(GGMLRunner::compute(get_graph, n_threads, false, false, false), input.dim());
+ }
+
+ int get_encoder_output_channels(int input_channels) override {
+ SD_UNUSED(input_channels);
+ return LATENT_CHANNELS * 2;
+ }
+
+ sd::Tensor vae_output_to_latents(const sd::Tensor& vae_output, std::shared_ptr rng) override {
+ const auto chunks = sd::ops::chunk(vae_output, 2, 2);
+ const auto& mean = chunks[0];
+ const auto& logvar = chunks[1];
+ sd::Tensor stddev = sd::ops::exp(0.5f * sd::ops::clamp(logvar, -20.0f, 10.0f));
+ sd::Tensor noise = sd::Tensor::randn_like(mean, rng);
+ sd::Tensor latents = mean + stddev * noise;
+ return latents;
+ }
+
+ sd::Tensor diffusion_to_vae_latents(const sd::Tensor& latents) override {
+ return latents;
+ }
+
+ sd::Tensor vae_to_diffusion_latents(const sd::Tensor& latents) override {
+ return latents;
+ }
+ };
+} // namespace MageVAE
+
+#endif // __SD_MODEL_VAE_MAGE_VAE_HPP__
diff --git a/src/model/vae/vae.hpp b/src/model/vae/vae.hpp
index a462c8728..e8b4eed38 100644
--- a/src/model/vae/vae.hpp
+++ b/src/model/vae/vae.hpp
@@ -74,7 +74,7 @@ struct VAE : public GGMLRunner {
int scale_factor = 8;
if (version == VERSION_LTXAV) {
scale_factor = 32;
- } else if (version == VERSION_WAN2_2_TI2V || sd_version_is_hunyuan_video(version)) {
+ } else if (version == VERSION_WAN2_2_TI2V || sd_version_is_hunyuan_video(version) || sd_version_is_mage_flow(version)) {
scale_factor = 16;
} else if (sd_version_uses_flux2_vae(version)) {
scale_factor = 16;
diff --git a/src/model_loader.cpp b/src/model_loader.cpp
index 2f9e68cac..ccd1e7d96 100644
--- a/src/model_loader.cpp
+++ b/src/model_loader.cpp
@@ -498,6 +498,10 @@ SDVersion ModelLoader::get_sd_version() {
return VERSION_MINIT2I;
}
if (tensor_storage.name.find("model.diffusion_model.transformer_blocks.0.img_mod.1.weight") != std::string::npos) {
+ auto img_in = tensor_storage_map.find("model.diffusion_model.img_in.weight");
+ if (img_in != tensor_storage_map.end() && img_in->second.ne[0] == 128) {
+ return VERSION_MAGE_FLOW;
+ }
if (tensor_storage_map.find("model.diffusion_model.time_text_embed.addition_t_embedding.weight") != tensor_storage_map.end()) {
return VERSION_QWEN_IMAGE_LAYERED;
}
diff --git a/src/name_conversion.cpp b/src/name_conversion.cpp
index 904e12e86..2b41a5576 100644
--- a/src/name_conversion.cpp
+++ b/src/name_conversion.cpp
@@ -1048,7 +1048,7 @@ std::string convert_diffusers_to_original_wan_vae(std::string name) {
}
std::string convert_first_stage_model_name(std::string name, std::string prefix, SDVersion version) {
- if (sd_version_is_hunyuan_video(version)) {
+ if (sd_version_is_hunyuan_video(version) || sd_version_is_mage_flow(version)) {
return name;
}
if (sd_version_uses_wan_vae(version)) {
@@ -1414,8 +1414,14 @@ std::string convert_tensor_name(std::string name, SDVersion version) {
replace_with_prefix_map(name, prefix_map);
- if ((sd_version_is_boogu_image(version) || sd_version_is_krea2(version)) && starts_with(name, "text_encoders.llm.visual.")) {
- name = convert_qwen3_vl_vision_name(std::move(name));
+ if (sd_version_is_boogu_image(version) || sd_version_is_krea2(version) || sd_version_is_mage_flow(version)) {
+ const std::string hf_vision_prefix = "text_encoders.llm.model.visual.";
+ if (starts_with(name, hf_vision_prefix)) {
+ name = "text_encoders.llm.visual." + name.substr(hf_vision_prefix.size());
+ }
+ if (starts_with(name, "text_encoders.llm.visual.")) {
+ name = convert_qwen3_vl_vision_name(std::move(name));
+ }
}
// diffusion model
diff --git a/src/stable-diffusion.cpp b/src/stable-diffusion.cpp
index 0bddba2bb..81f82fbc8 100644
--- a/src/stable-diffusion.cpp
+++ b/src/stable-diffusion.cpp
@@ -36,6 +36,7 @@
#include "model/diffusion/lens.hpp"
#include "model/diffusion/lingbot_video.hpp"
#include "model/diffusion/ltxv.hpp"
+#include "model/diffusion/mage_flow.hpp"
#include "model/diffusion/minit2i.hpp"
#include "model/diffusion/mmdit.hpp"
#include "model/diffusion/model.hpp"
@@ -50,6 +51,7 @@
#include "model/vae/hunyuan_vae.hpp"
#include "model/vae/ltx_audio_vae.hpp"
#include "model/vae/ltx_vae.hpp"
+#include "model/vae/mage_vae.hpp"
#include "model/vae/tae.hpp"
#include "model/vae/vae.hpp"
#include "model/vae/wan_vae.hpp"
@@ -115,6 +117,7 @@ const char* model_version_to_str[] = {
"Ideogram 4",
"SeFi-Image",
"Krea2",
+ "Mage Flow",
"ESRGAN",
};
@@ -145,6 +148,7 @@ static bool sd_version_supports_ref_latent_img_cfg(SDVersion version) {
return version == VERSION_FLUX ||
sd_version_is_flux2(version) ||
sd_version_is_qwen_image(version) ||
+ sd_version_is_mage_flow(version) ||
sd_version_is_longcat(version) ||
sd_version_is_z_image(version) ||
sd_version_is_boogu_image(version);
@@ -1147,6 +1151,17 @@ class StableDiffusionGGML {
version,
model_manager,
sd_ctx_params->model_args);
+ } else if (sd_version_is_mage_flow(version)) {
+ cond_stage_model = std::make_shared(backend_for(SDBackendModule::TE),
+ tensor_storage_map,
+ version,
+ "",
+ true,
+ model_manager);
+ diffusion_model = std::make_shared(backend_for(SDBackendModule::DIFFUSION),
+ tensor_storage_map,
+ "model.diffusion_model",
+ model_manager);
} else if (sd_version_is_longcat(version)) {
cond_stage_model = std::make_shared(backend_for(SDBackendModule::TE),
tensor_storage_map,
@@ -1321,6 +1336,11 @@ class StableDiffusionGGML {
false,
version,
model_manager);
+ } else if (sd_version_is_mage_flow(vae_version)) {
+ return std::make_shared(backend_for(SDBackendModule::VAE),
+ tensor_storage_map,
+ "first_stage_model",
+ model_manager);
} else if (sd_version_uses_hunyuan_video_vae(vae_version)) {
return std::make_shared(backend_for(SDBackendModule::VAE),
tensor_storage_map,
@@ -1642,6 +1662,7 @@ class StableDiffusionGGML {
sd_version_is_hunyuan_video(version) ||
sd_version_is_lingbot_video(version) ||
sd_version_is_qwen_image(version) ||
+ sd_version_is_mage_flow(version) ||
version == VERSION_HIDREAM_O1 ||
sd_version_is_anima(version) ||
sd_version_is_ernie_image(version) ||
@@ -1662,6 +1683,8 @@ class StableDiffusionGGML {
default_flow_shift = 1.0f;
} else if (sd_version_is_boogu_image(version)) {
default_flow_shift = 3.16f;
+ } else if (sd_version_is_mage_flow(version)) {
+ default_flow_shift = 6.f;
} else {
default_flow_shift = 3.f;
}
@@ -2762,6 +2785,8 @@ class StableDiffusionGGML {
latent_channel = 144;
} else if (sd_version_uses_flux2_vae(version)) {
latent_channel = 128;
+ } else if (sd_version_is_mage_flow(version)) {
+ latent_channel = 128;
} else {
latent_channel = 16;
}
@@ -2910,6 +2935,8 @@ class StableDiffusionGGML {
return "qwen_layered";
} else if (sd_version_is_qwen_image(version)) {
return "qwen";
+ } else if (sd_version_is_mage_flow(version)) {
+ return "mage_flow";
} else if (sd_version_is_z_image(version) || sd_version_is_boogu_image(version)) {
return "z_image_omni";
} else if (sd_version_is_krea2(version)) {
@@ -4773,10 +4800,17 @@ static std::optional prepare_image_generation_latents(sd
sd::Tensor ref_latent;
if (ref_image_params.resize_before_vae && !sd_version_is_pid(sd_ctx->sd->version)) {
LOG_DEBUG("auto resize ref images");
- int target_pixels = ref_image_params.vae_input_max_pixels > 0 ? ref_image_params.vae_input_max_pixels : 1024 * 1024;
- int vae_image_size = std::min(target_pixels, request->width * request->height);
- double vae_width = sqrt(vae_image_size * ref_images[i].shape()[0] / ref_images[i].shape()[1]);
- double vae_height = vae_width * ref_images[i].shape()[1] / ref_images[i].shape()[0];
+ double vae_width;
+ double vae_height;
+ if (ref_image_params.resize_vae_to_target) {
+ vae_width = request->width;
+ vae_height = request->height;
+ } else {
+ int target_pixels = ref_image_params.vae_input_max_pixels > 0 ? ref_image_params.vae_input_max_pixels : 1024 * 1024;
+ int vae_image_size = std::min(target_pixels, request->width * request->height);
+ vae_width = sqrt(vae_image_size * ref_images[i].shape()[0] / ref_images[i].shape()[1]);
+ vae_height = vae_width * ref_images[i].shape()[1] / ref_images[i].shape()[0];
+ }
int factor = sd_version_is_qwen_image(sd_ctx->sd->version) ? 32 : 16;
vae_height = round(vae_height / factor) * factor;