refactor: define VAE tile dimensions in image pixels (#2059)

This commit is contained in:
leejet
2026-09-25 18:20:32 +08:00
committed by GitHub
parent 39ada0863b
commit 19bbbca1c7
14 changed files with 298 additions and 174 deletions
+30 -13
View File
@@ -478,7 +478,12 @@ namespace sd::backend_fit {
return true;
}
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params, bool prefer_temporal_tiling, ggml_status status) {
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params,
bool prefer_temporal_tiling,
ggml_status status,
int latent_tile_size_w,
int latent_tile_size_h,
int scale_factor) {
// Execution failures can leave the device unusable; tiling only helps with allocation failures.
if (status != GGML_STATUS_ALLOC_FAILED) {
return false;
@@ -487,19 +492,31 @@ namespace sd::backend_fit {
if (prefer_temporal_tiling && !tiling_params.temporal_tiling) {
tiling_params.temporal_tiling = true;
retry_mode = tiling_params.enabled ? "spatial+temporal" : "temporal";
} else if (!tiling_params.enabled) {
tiling_params.enabled = true;
tiling_params.rel_size_x = 0.5f;
tiling_params.rel_size_y = 0.5f;
if (tiling_params.tile_size_x <= 0) {
tiling_params.tile_size_x = 256;
}
if (tiling_params.tile_size_y <= 0) {
tiling_params.tile_size_y = 256;
}
retry_mode = tiling_params.temporal_tiling ? "spatial+temporal" : "spatial";
} else {
return false;
if (latent_tile_size_w <= 0 || latent_tile_size_h <= 0 || scale_factor <= 0) {
return false;
}
auto smaller_tile = [&](int size) {
int next_size = size / 2;
if (!tiling_params.enabled) {
next_size = std::min(next_size, 256 / scale_factor);
}
return std::min(size, std::max(4, next_size));
};
const int tile_size_w = smaller_tile(latent_tile_size_w);
const int tile_size_h = smaller_tile(latent_tile_size_h);
if (tile_size_w == latent_tile_size_w && tile_size_h == latent_tile_size_h) {
return false;
}
tiling_params.enabled = true;
tiling_params.rel_size_w = 0.0f;
tiling_params.rel_size_h = 0.0f;
tiling_params.tile_size_w = tile_size_w * scale_factor;
tiling_params.tile_size_h = tile_size_h * scale_factor;
retry_mode = tiling_params.temporal_tiling ? "spatial+temporal" : "spatial";
LOG_WARN("Reducing VAE decode tiles from %dx%d to %dx%d image pixels",
latent_tile_size_w * scale_factor, latent_tile_size_h * scale_factor,
tiling_params.tile_size_w, tiling_params.tile_size_h);
}
LOG_WARN("VAE decode ran out of memory; retrying with %s tiling",
+4 -1
View File
@@ -17,7 +17,10 @@ namespace sd::backend_fit {
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params,
bool prefer_temporal_tiling,
ggml_status status);
ggml_status status,
int latent_tile_size_w,
int latent_tile_size_h,
int scale_factor);
} // namespace sd::backend_fit
+12 -6
View File
@@ -556,12 +556,18 @@ namespace MiniMaxH3VAE {
tensor.shape()[3]});
}
static sd_tiling_params_t h3_tiling(sd_tiling_params_t params) {
sd_tiling_params_t resolve_tiling_params(sd_tiling_params_t params) const override {
if (!params.enabled) {
params.target_overlap = 0.25f;
}
if (params.tile_size_w == 0 && params.rel_size_w == 0.f) {
params.tile_size_w = 256;
}
if (params.tile_size_h == 0 && params.rel_size_h == 0.f) {
params.tile_size_h = 256;
}
params.enabled = true;
params.temporal_tiling = false;
params.tile_size_x = 16;
params.tile_size_y = 16;
params.target_overlap = 0.25f;
return params;
}
@@ -605,7 +611,7 @@ namespace MiniMaxH3VAE {
bool circular_x = false,
bool circular_y = false) override {
auto input = ensure_video_shape(x);
auto tiling = h3_tiling(tiling_params);
auto tiling = resolve_tiling_params(tiling_params);
if (input.shape()[2] == 1) {
auto encoded = VAE::encode(n_threads, input, tiling, circular_x, circular_y);
if (!encoded.empty() && encoded.shape()[2] > 1) {
@@ -646,7 +652,7 @@ namespace MiniMaxH3VAE {
bool circular_y = false,
bool silent = false) override {
auto input = ensure_video_shape(x);
auto tiling = h3_tiling(tiling_params);
auto tiling = resolve_tiling_params(tiling_params);
if (input.shape()[2] == 1) {
auto decoded = VAE::decode(n_threads,
input,
+82 -50
View File
@@ -1,6 +1,9 @@
#ifndef __SD_MODEL_VAE_VAE_HPP__
#define __SD_MODEL_VAE_VAE_HPP__
#include <cmath>
#include <limits>
#include "core/tensor_ggml.hpp"
#include "model/common/block.hpp"
#include "model/vae/vae_tiling.hpp"
@@ -117,8 +120,8 @@ protected:
int output_width,
int output_height,
int scale,
int p_tile_size_x,
int p_tile_size_y,
int p_tile_size_w,
int p_tile_size_h,
float tile_overlap_factor,
bool circular_x,
bool circular_y,
@@ -138,17 +141,28 @@ protected:
}
return output_tile;
};
return ::process_tiles_2d(input,
output_width,
output_height,
scale,
p_tile_size_x,
p_tile_size_y,
tile_overlap_factor,
circular_x,
circular_y,
on_processing,
silent);
const bool original_circular_x = circular_x_enabled;
const bool original_circular_y = circular_y_enabled;
const int64_t latent_width = decode_graph ? input.shape()[0] : output_width;
const int64_t latent_height = decode_graph ? input.shape()[1] : output_height;
circular_x = circular_x || original_circular_x;
circular_y = circular_y || original_circular_y;
// Full-width axes wrap in convolutions; split axes wrap between tiles.
set_circular_axes(circular_x && p_tile_size_w >= latent_width,
circular_y && p_tile_size_h >= latent_height);
auto output = ::process_tiles_2d(input,
output_width,
output_height,
scale,
p_tile_size_w,
p_tile_size_h,
tile_overlap_factor,
circular_x && p_tile_size_w < latent_width,
circular_y && p_tile_size_h < latent_height,
on_processing,
silent);
set_circular_axes(original_circular_x, original_circular_y);
return output;
}
public:
@@ -178,33 +192,48 @@ public:
return supports_temporal_tiling(VAETemporalDirection::DECODE);
}
void get_tile_sizes(int& tile_size_x,
int& tile_size_y,
virtual sd_tiling_params_t resolve_tiling_params(sd_tiling_params_t params) const {
return params;
}
bool get_tile_sizes(int& tile_size_w,
int& tile_size_h,
float& tile_overlap,
const sd_tiling_params_t& params,
int64_t latent_x,
int64_t latent_y,
float encoding_factor = 1.0f) {
tile_overlap = std::max(std::min(params.target_overlap, 0.5f), 0.0f);
auto get_tile_size = [&](int requested_size, float factor, int64_t latent_size) {
const int default_tile_size = 32;
const int min_tile_dimension = 4;
int tile_size = default_tile_size;
// factor <= 1 means simple fraction of the latent dimension
// factor > 1 means number of tiles across that dimension
if (factor > 0.f) {
if (factor > 1.0)
factor = 1 / (factor - factor * tile_overlap + tile_overlap);
tile_size = static_cast<int>(std::round(latent_size * factor));
} else if (requested_size >= min_tile_dimension) {
tile_size = requested_size;
int64_t latent_w,
int64_t latent_h) {
const auto tiling = resolve_tiling_params(params);
if (latent_w <= 0 || latent_h <= 0 ||
latent_w > std::numeric_limits<int>::max() || latent_h > std::numeric_limits<int>::max() ||
!std::isfinite(tiling.target_overlap)) {
LOG_ERROR("invalid VAE tiling dimensions or overlap");
return false;
}
const int scale_factor = get_scale_factor();
tile_overlap = std::max(std::min(tiling.target_overlap, 0.5f), 0.0f);
auto get_tile_size = [&](int requested_size, double factor, int64_t latent_size, int& tile_size) {
if (requested_size < 0 || !std::isfinite(factor) || factor < 0.0) {
LOG_ERROR("VAE tile sizes and relative sizes must be finite and non-negative");
return false;
}
tile_size = static_cast<int>(tile_size * encoding_factor);
return std::max(std::min(tile_size, static_cast<int>(latent_size)), min_tile_dimension);
const int min_tile_dimension = std::min(4, static_cast<int>(latent_size));
double size = (requested_size > 0 ? requested_size : 256) / scale_factor;
if (factor > 0.0) {
if (factor > 1.0) {
factor = 1.0 / (factor * (1.0 - tile_overlap) + tile_overlap);
}
size = std::floor(static_cast<double>(latent_size) * factor);
}
if (size < min_tile_dimension && (requested_size > 0 || factor > 0.0)) {
LOG_ERROR("VAE tile size must be at least %d image pixels on this axis", min_tile_dimension * scale_factor);
return false;
}
tile_size = static_cast<int>(std::min(static_cast<double>(latent_size), std::max<double>(min_tile_dimension, size)));
return true;
};
tile_size_x = get_tile_size(params.tile_size_x, params.rel_size_x, latent_x);
tile_size_y = get_tile_size(params.tile_size_y, params.rel_size_y, latent_y);
return get_tile_size(tiling.tile_size_w, tiling.rel_size_w, latent_w, tile_size_w) &&
get_tile_size(tiling.tile_size_h, tiling.rel_size_h, latent_h, tile_size_h);
}
virtual sd::Tensor<float> encode(int n_threads,
@@ -213,6 +242,7 @@ public:
bool circular_x = false,
bool circular_y = false) {
int64_t t0 = ggml_time_ms();
tiling_params = resolve_tiling_params(tiling_params);
sd::Tensor<float> input = x;
sd::Tensor<float> output;
if (scale_input) {
@@ -224,21 +254,19 @@ public:
int64_t W = input.shape()[0] / scale_factor;
int64_t H = input.shape()[1] / scale_factor;
float tile_overlap;
int tile_size_x, tile_size_y;
// Image VAE encode is more sensitive to tile boundary context than decode.
// Keep the smaller legacy factor for video VAEs, but default image encode
// tiles to 64 latent pixels so a 512px SD image is encoded as one tile.
const float encode_tile_factor = sd_version_is_minimax_h3(version) ? 1.f : (sd_version_is_wan(version) || sd_version_is_hunyuan_video(version) || sd_version_is_ltxav(version)) ? 1.30539f
: 2.0f;
get_tile_sizes(tile_size_x, tile_size_y, tile_overlap, tiling_params, W, H, encode_tile_factor);
LOG_VERBOSE("VAE Tile size: %dx%d", tile_size_x, tile_size_y);
int tile_size_w, tile_size_h;
if (!get_tile_sizes(tile_size_w, tile_size_h, tile_overlap, tiling_params, W, H)) {
return {};
}
LOG_VERBOSE("VAE encode tile size: %dx%d pixels (%dx%d latent)",
tile_size_w * scale_factor, tile_size_h * scale_factor, tile_size_w, tile_size_h);
output = tiled_compute(input,
n_threads,
static_cast<int>(W),
static_cast<int>(H),
scale_factor,
tile_size_x,
tile_size_y,
tile_size_w,
tile_size_h,
tile_overlap,
circular_x,
circular_y,
@@ -271,6 +299,7 @@ public:
bool circular_y = false,
bool silent = false) {
int64_t t0 = ggml_time_ms();
tiling_params = resolve_tiling_params(tiling_params);
sd::Tensor<float> input = x;
sd::Tensor<float> output;
@@ -279,10 +308,13 @@ public:
int64_t W = input.shape()[0] * scale_factor;
int64_t H = input.shape()[1] * scale_factor;
float tile_overlap;
int tile_size_x, tile_size_y;
get_tile_sizes(tile_size_x, tile_size_y, tile_overlap, tiling_params, input.shape()[0], input.shape()[1]);
int tile_size_w, tile_size_h;
if (!get_tile_sizes(tile_size_w, tile_size_h, tile_overlap, tiling_params, input.shape()[0], input.shape()[1])) {
return {};
}
if (!silent) {
LOG_VERBOSE("VAE Tile size: %dx%d", tile_size_x, tile_size_y);
LOG_VERBOSE("VAE decode tile size: %dx%d pixels (%dx%d latent)",
tile_size_w * scale_factor, tile_size_h * scale_factor, tile_size_w, tile_size_h);
}
output = tiled_compute(
input,
@@ -290,8 +322,8 @@ public:
static_cast<int>(W),
static_cast<int>(H),
scale_factor,
tile_size_x,
tile_size_y,
tile_size_w,
tile_size_h,
tile_overlap,
circular_x,
circular_y,
+19 -7
View File
@@ -2883,14 +2883,26 @@ sd::Tensor<float> StableDiffusionGGML::decode_first_stage(const sd::Tensor<float
return sd::ops::clamp((x + 1.f) * 0.5f, 0.0f, 1.0f);
}
auto latents = first_stage_model->diffusion_to_vae_latents(x);
auto decoded = first_stage_model->decode(n_threads, latents, vae_tiling_params, decode_video, circular_x, circular_y);
const bool prefer_temporal_tiling = decode_video && first_stage_model->can_temporal_tile_decode();
while (decoded.empty() &&
sd::backend_fit::prepare_vae_decode_retry_tiling(vae_tiling_params, prefer_temporal_tiling,
first_stage_model->last_compute_status())) {
decoded = first_stage_model->decode(n_threads, latents, vae_tiling_params, decode_video, circular_x, circular_y);
auto tiling_params = first_stage_model->resolve_tiling_params(vae_tiling_params);
const bool prefer_temporal_tiling = decode_video && latents.dim() == 5 && latents.shape()[2] > 1 &&
first_stage_model->can_temporal_tile_decode();
for (;;) {
int tile_size_w = static_cast<int>(latents.shape()[0]);
int tile_size_h = static_cast<int>(latents.shape()[1]);
float tile_overlap;
if (tiling_params.enabled &&
!first_stage_model->get_tile_sizes(tile_size_w, tile_size_h, tile_overlap, tiling_params,
latents.shape()[0], latents.shape()[1])) {
return {};
}
auto decoded = first_stage_model->decode(n_threads, latents, tiling_params, decode_video, circular_x, circular_y);
if (!decoded.empty() ||
!sd::backend_fit::prepare_vae_decode_retry_tiling(tiling_params, prefer_temporal_tiling,
first_stage_model->last_compute_status(),
tile_size_w, tile_size_h, first_stage_model->get_scale_factor())) {
return decoded;
}
}
return decoded;
}
sd::Tensor<float> StableDiffusionGGML::normalize_ltx_video_latents(const sd::Tensor<float>& x) {
+15 -13
View File
@@ -35,19 +35,21 @@ namespace sd::pipeline {
return original_axes;
}
int tile_size_x, tile_size_y;
int tile_size_w, tile_size_h;
float overlap;
int latent_size_x = request.width / request.vae_scale_factor;
int latent_size_y = request.height / request.vae_scale_factor;
sd->first_stage_model->get_tile_sizes(tile_size_x,
tile_size_y,
overlap,
sd_img_gen_params->vae_tiling_params,
latent_size_x,
latent_size_y);
int latent_size_w = request.width / request.vae_scale_factor;
int latent_size_h = request.height / request.vae_scale_factor;
if (!sd->first_stage_model->get_tile_sizes(tile_size_w,
tile_size_h,
overlap,
sd_img_gen_params->vae_tiling_params,
latent_size_w,
latent_size_h)) {
return original_axes;
}
sd->circular_x = sd->circular_x && (tile_size_x >= latent_size_x);
sd->circular_y = sd->circular_y && (tile_size_y >= latent_size_y);
sd->circular_x = sd->circular_x && (tile_size_w >= latent_size_w);
sd->circular_y = sd->circular_y && (tile_size_h >= latent_size_h);
if (sd->first_stage_model) {
sd->first_stage_model->set_circular_axes(sd->circular_x, sd->circular_y);
@@ -56,8 +58,8 @@ namespace sd::pipeline {
sd->preview_vae->set_circular_axes(sd->circular_x, sd->circular_y);
}
sd->circular_x = original_axes.circular_x && (tile_size_x < latent_size_x);
sd->circular_y = original_axes.circular_y && (tile_size_y < latent_size_y);
sd->circular_x = original_axes.circular_x && (tile_size_w < latent_size_w);
sd->circular_y = original_axes.circular_y && (tile_size_h < latent_size_h);
return original_axes;
}
+24 -24
View File
@@ -142,8 +142,8 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
int output_width,
int output_height,
int scale,
int p_tile_size_x,
int p_tile_size_y,
int p_tile_size_w,
int p_tile_size_h,
float tile_overlap_factor,
bool circular_x,
bool circular_y,
@@ -168,28 +168,28 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
int num_tiles_x;
float tile_overlap_factor_x;
sd_tiling_calc_tiles(num_tiles_x, tile_overlap_factor_x, small_width, p_tile_size_x, tile_overlap_factor, circular_x);
sd_tiling_calc_tiles(num_tiles_x, tile_overlap_factor_x, small_width, p_tile_size_w, tile_overlap_factor, circular_x);
int num_tiles_y;
float tile_overlap_factor_y;
sd_tiling_calc_tiles(num_tiles_y, tile_overlap_factor_y, small_height, p_tile_size_y, tile_overlap_factor, circular_y);
sd_tiling_calc_tiles(num_tiles_y, tile_overlap_factor_y, small_height, p_tile_size_h, tile_overlap_factor, circular_y);
int tile_overlap_x = static_cast<int32_t>(p_tile_size_x * tile_overlap_factor_x);
int non_tile_overlap_x = p_tile_size_x - tile_overlap_x;
int tile_overlap_y = static_cast<int32_t>(p_tile_size_y * tile_overlap_factor_y);
int non_tile_overlap_y = p_tile_size_y - tile_overlap_y;
int tile_size_x = p_tile_size_x < small_width ? p_tile_size_x : small_width;
int tile_size_y = p_tile_size_y < small_height ? p_tile_size_y : small_height;
int input_tile_size_x = tile_size_x;
int input_tile_size_y = tile_size_y;
int output_tile_size_x = tile_size_x;
int output_tile_size_y = tile_size_y;
int tile_overlap_x = static_cast<int32_t>(p_tile_size_w * tile_overlap_factor_x);
int non_tile_overlap_x = p_tile_size_w - tile_overlap_x;
int tile_overlap_y = static_cast<int32_t>(p_tile_size_h * tile_overlap_factor_y);
int non_tile_overlap_y = p_tile_size_h - tile_overlap_y;
int tile_size_w = p_tile_size_w < small_width ? p_tile_size_w : small_width;
int tile_size_h = p_tile_size_h < small_height ? p_tile_size_h : small_height;
int input_tile_size_w = tile_size_w;
int input_tile_size_h = tile_size_h;
int output_tile_size_w = tile_size_w;
int output_tile_size_h = tile_size_h;
if (decode) {
output_tile_size_x *= scale;
output_tile_size_y *= scale;
output_tile_size_w *= scale;
output_tile_size_h *= scale;
} else {
input_tile_size_x *= scale;
input_tile_size_y *= scale;
input_tile_size_w *= scale;
input_tile_size_h *= scale;
}
int num_tiles = num_tiles_x * num_tiles_y;
@@ -205,9 +205,9 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
}
for (int y = 0; y < small_height && !last_y; y += non_tile_overlap_y) {
int dy = 0;
if (!circular_y && y + tile_size_y >= small_height) {
if (!circular_y && y + tile_size_h >= small_height) {
int original_y = y;
y = small_height - tile_size_y;
y = small_height - tile_size_h;
dy = original_y - y;
if (decode) {
dy *= scale;
@@ -216,9 +216,9 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
}
for (int x = 0; x < small_width && !last_x; x += non_tile_overlap_x) {
int dx = 0;
if (!circular_x && x + tile_size_x >= small_width) {
if (!circular_x && x + tile_size_w >= small_width) {
int original_x = x;
x = small_width - tile_size_x;
x = small_width - tile_size_w;
dx = original_x - x;
if (decode) {
dx *= scale;
@@ -235,12 +235,12 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
int overlap_y_out = decode ? tile_overlap_y * scale : tile_overlap_y;
int64_t t1 = ggml_time_ms();
auto input_tile = sd_tensor_split_2d(input, input_tile_size_x, input_tile_size_y, x_in, y_in);
auto input_tile = sd_tensor_split_2d(input, input_tile_size_w, input_tile_size_h, x_in, y_in);
auto output_tile = on_processing(input_tile);
if (output_tile.empty()) {
return {};
}
GGML_ASSERT(output_tile.shape()[0] == output_tile_size_x && output_tile.shape()[1] == output_tile_size_y);
GGML_ASSERT(output_tile.shape()[0] == output_tile_size_w && output_tile.shape()[1] == output_tile_size_h);
if (output.empty()) {
std::vector<int64_t> output_shape = output_tile.shape();
output_shape[0] = output_width;
+2 -2
View File
@@ -11,8 +11,8 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
int output_width,
int output_height,
int scale,
int p_tile_size_x,
int p_tile_size_y,
int p_tile_size_w,
int p_tile_size_h,
float tile_overlap_factor,
bool circular_x,
bool circular_y,