Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
297 changes: 205 additions & 92 deletions src/runtime/tiling.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,32 @@
#include "core/util.h"
#include "ggml.h"

static int sd_tiling_calc_num_tiles(int dimension, int tile_size, float target_overlap_factor) {
if (dimension <= tile_size) {
return 1;
} else if (dimension < 2 * tile_size) {
return 2;
} else if (dimension == 2 * tile_size) {
return 3;
} else {
float target_num_tiles = 1.0f + (dimension - tile_size) / ((1.0f - target_overlap_factor) * tile_size);
int num_tiles_lower = static_cast<int>(std::floor(target_num_tiles));
int num_tiles_upper = static_cast<int>(std::ceil(target_num_tiles));
int num_tiles_min = 1 + (dimension - 2) / (tile_size - 1); // positive adjacent overlap
int num_tiles_max = 2 * dimension / tile_size - 1; // no triple overlap under Bresenham placement
num_tiles_lower = std::clamp(num_tiles_lower, num_tiles_min, num_tiles_max);
num_tiles_upper = std::clamp(num_tiles_upper, num_tiles_min, num_tiles_max);
auto overlap_error = [target_num_tiles](int num_tiles) -> float {
return std::abs(1.0f / (num_tiles - 1.0f) - 1.0f / (target_num_tiles - 1.0f));
}; // (dimension - tile_size) / tile_size factors out
return (overlap_error(num_tiles_upper) < overlap_error(num_tiles_lower)) ? num_tiles_upper : num_tiles_lower; // use lower if tie
}
}

static float sd_tiling_calc_average_stride_factor(int dimension, int tile_size, int num_tiles) {
return static_cast<float>(dimension - tile_size) / static_cast<float>(tile_size * (num_tiles - 1));
}

static void sd_tiling_calc_tiles(int& num_tiles_dim,
float& tile_overlap_factor_dim,
int small_dim,
Expand All @@ -32,29 +58,9 @@ static void sd_tiling_calc_tiles(int& num_tiles_dim,
num_tiles_dim++;
tile_overlap_factor_dim = 0.5;
}

return;
}
// else, non-circular means the last and first tile are not overlapping

num_tiles_dim = (small_dim - tile_overlap) / non_tile_overlap;
int overshoot_dim = ((num_tiles_dim + 1) * non_tile_overlap + tile_overlap) % small_dim;

if ((overshoot_dim != non_tile_overlap) && (overshoot_dim <= num_tiles_dim * (tile_size / 2 - tile_overlap))) {
// if tiles don't fit perfectly using the desired overlap
// and there is enough room to squeeze an extra tile without overlap becoming >0.5
num_tiles_dim++;
}

tile_overlap_factor_dim = (float)(tile_size * num_tiles_dim - small_dim) / (float)(tile_size * (num_tiles_dim - 1));
if (num_tiles_dim <= 2) {
if (small_dim <= tile_size) {
num_tiles_dim = 1;
tile_overlap_factor_dim = 0;
} else {
num_tiles_dim = 2;
tile_overlap_factor_dim = (2 * tile_size - small_dim) / (float)tile_size;
}
} else {
num_tiles_dim = sd_tiling_calc_num_tiles(small_dim, tile_size, tile_overlap_factor);
tile_overlap_factor_dim = (num_tiles_dim == 1) ? 0 : (1.0f - sd_tiling_calc_average_stride_factor(small_dim, tile_size, num_tiles_dim));
}
}

Expand Down Expand Up @@ -138,6 +144,59 @@ static void sd_tensor_merge_2d(const sd::Tensor<float>& input,
}
}

static void sd_tensor_merge_2d_non_circular(const sd::Tensor<float>& input,
sd::Tensor<float>* output,
int x,
int y,
int overlap_left,
int overlap_right,
int overlap_top,
int overlap_bottom) {
GGML_ASSERT(output != nullptr);

int64_t in_width = input.shape()[0];
int64_t in_height = input.shape()[1];
int64_t out_width = output->shape()[0];
int64_t out_height = output->shape()[1];
int64_t in_size = sd_tensor_plane_size(input);
int64_t out_size = sd_tensor_plane_size(*output);
int64_t plane_count = input.numel() / in_size;

GGML_ASSERT(output->numel() == plane_count * out_size);
GGML_ASSERT(x >= 0 && y >= 0);
GGML_ASSERT(x + in_width <= out_width);
GGML_ASSERT(y + in_height <= out_height);
GGML_ASSERT(overlap_left >= 0 && overlap_right >= 0);
GGML_ASSERT(overlap_top >= 0 && overlap_bottom >= 0);

auto smootherstep_f32 = [](const float x) -> float {
return x * x * x * (x * (6.0f * x - 15.0f) + 10.0f);
};
for (int64_t plane = 0; plane < plane_count; ++plane) {
for (int iy = 0; iy < in_height; ++iy) {
float y_f = 1.0f;
if (iy < overlap_top) {
y_f = static_cast<float>(iy) / overlap_top;
}
if (iy >= in_height - overlap_bottom) {
y_f = static_cast<float>(in_height - iy) / overlap_bottom;
}
const float y_weight = smootherstep_f32(std::clamp(y_f, 0.0f, 1.0f));
for (int ix = 0; ix < in_width; ++ix) {
float x_f = 1.0f;
if (ix < overlap_left) {
x_f = static_cast<float>(ix) / overlap_left;
}
if (ix >= in_width - overlap_right) {
x_f = static_cast<float>(in_width - ix) / overlap_right;
}
float x_weight = smootherstep_f32(std::clamp(x_f, 0.0f, 1.0f));
(*output)[plane * out_size + out_width * (y + iy) + (x + ix)] += x_weight * y_weight * input[plane * in_size + in_width * iy + ix];
}
}
}
}

sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
int output_width,
int output_height,
Expand All @@ -158,13 +217,11 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
GGML_ASSERT(((input_width / output_width) == scale) ||
((output_width / input_width) == scale));

int small_width = output_width;
int small_height = output_height;
bool decode = output_width > input_width;
if (decode) {
small_width = input_width;
small_height = input_height;
}
bool decode = output_width > input_width; // scale up
int small_width = decode ? input_width : output_width;
int small_height = decode ? input_height : output_height;
int scale_in = decode ? 1 : scale;
int scale_out = decode ? scale : 1;

int num_tiles_x;
float tile_overlap_factor_x;
Expand All @@ -174,89 +231,145 @@ sd::Tensor<float> process_tiles_2d(const sd::Tensor<float>& input,
float tile_overlap_factor_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_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_w *= scale;
output_tile_size_h *= scale;
} else {
input_tile_size_w *= scale;
input_tile_size_h *= scale;
}
int tile_width = std::min(p_tile_size_w, small_width);
int tile_height = std::min(p_tile_size_h, small_height);
int input_tile_width = tile_width * scale_in;
int input_tile_height = tile_height * scale_in;
int output_tile_width = tile_width * scale_out;
int output_tile_height = tile_height * scale_out;

int num_tiles = num_tiles_x * num_tiles_y;
int tile_count = 1;
bool last_y = false;
bool last_x = false;
float last_time = 0.0f;
if (!silent) {
LOG_VERBOSE("num tiles : %d, %d ", num_tiles_x, num_tiles_y);
LOG_VERBOSE("optimal overlap : %f, %f (targeting %f)", tile_overlap_factor_x, tile_overlap_factor_y, tile_overlap_factor);
LOG_VERBOSE("processing %i tiles", num_tiles);
pretty_progress(0, num_tiles, 0.0f);
}
for (int y = 0; y < small_height && !last_y; y += non_tile_overlap_y) {
int dy = 0;
if (!circular_y && y + tile_size_h >= small_height) {
int original_y = y;
y = small_height - tile_size_h;
dy = original_y - y;
if (decode) {
dy *= scale;
}
last_y = true;
}
for (int x = 0; x < small_width && !last_x; x += non_tile_overlap_x) {
int dx = 0;
if (!circular_x && x + tile_size_w >= small_width) {
int original_x = x;
x = small_width - tile_size_w;
dx = original_x - x;
if (circular_x || circular_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;

bool last_y = false;
bool last_x = false;

for (int y = 0; y < small_height && !last_y; y += non_tile_overlap_y) {
int dy = 0;
if (!circular_y && y + tile_height >= small_height) {
int original_y = y;
y = small_height - tile_height;
dy = original_y - y;
if (decode) {
dx *= scale;
dy *= scale;
}
last_x = true;
last_y = true;
}
for (int x = 0; x < small_width && !last_x; x += non_tile_overlap_x) {
int dx = 0;
if (!circular_x && x + tile_width >= small_width) {
int original_x = x;
x = small_width - tile_width;
dx = original_x - x;
if (decode) {
dx *= scale;
}
last_x = true;
}

int x_in = decode ? x : scale * x;
int y_in = decode ? y : scale * y;
int x_out = decode ? x * scale : x;
int y_out = decode ? y * scale : y;
int x_in = decode ? x : scale * x;
int y_in = decode ? y : scale * y;
int x_out = decode ? x * scale : x;
int y_out = decode ? y * scale : y;

int overlap_x_out = decode ? tile_overlap_x * scale : tile_overlap_x;
int overlap_y_out = decode ? tile_overlap_y * scale : tile_overlap_y;
int overlap_x_out = decode ? tile_overlap_x * scale : tile_overlap_x;
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_w, input_tile_size_h, x_in, y_in);
auto output_tile = on_processing(input_tile);
if (output_tile.empty()) {
return {};
int64_t t1 = ggml_time_ms();
auto input_tile = sd_tensor_split_2d(input, input_tile_width, input_tile_height, x_in, y_in);
auto output_tile = on_processing(input_tile);
if (output_tile.empty()) {
return {};
}
GGML_ASSERT(output_tile.shape()[0] == output_tile_width && output_tile.shape()[1] == output_tile_height);
if (output.empty()) {
std::vector<int64_t> output_shape = output_tile.shape();
output_shape[0] = output_width;
output_shape[1] = output_height;
output = sd::Tensor<float>::zeros(std::move(output_shape));
}
sd_tensor_merge_2d(output_tile, &output, x_out, y_out, overlap_x_out, overlap_y_out, circular_x, circular_y, dx, dy);

if (!silent) {
int64_t t2 = ggml_time_ms();
last_time = (t2 - t1) / 1000.0f;
pretty_progress(tile_count, num_tiles, last_time);
}
tile_count++;
}
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;
output_shape[1] = output_height;
output = sd::Tensor<float>::zeros(std::move(output_shape));
last_x = false;
}
} else {
for (int j = 0; j < num_tiles_y; ++j) {
int y = 0;
int overlap_top = 0;
int overlap_bottom = 0;
if (num_tiles_y > 1) {
y = j * (small_height - tile_height) / (num_tiles_y - 1);
if (j > 0) {
int y_prev = (j - 1) * (small_height - tile_height) / (num_tiles_y - 1);
overlap_top = y_prev + tile_height - y;
}
if (j < num_tiles_y - 1) {
int y_next = (j + 1) * (small_height - tile_height) / (num_tiles_y - 1);
overlap_bottom = y + tile_height - y_next;
}
}
sd_tensor_merge_2d(output_tile, &output, x_out, y_out, overlap_x_out, overlap_y_out, circular_x, circular_y, dx, dy);
for (int i = 0; i < num_tiles_x; ++i) {
int x = 0;
int overlap_left = 0;
int overlap_right = 0;
if (num_tiles_x > 1) {
x = i * (small_width - tile_width) / (num_tiles_x - 1);
if (i > 0) {
int x_prev = (i - 1) * (small_width - tile_width) / (num_tiles_x - 1);
overlap_left = x_prev + tile_width - x;
}
if (i < num_tiles_x - 1) {
int x_next = (i + 1) * (small_width - tile_width) / (num_tiles_x - 1);
overlap_right = x + tile_width - x_next;
}
}

if (!silent) {
int64_t t2 = ggml_time_ms();
last_time = (t2 - t1) / 1000.0f;
pretty_progress(tile_count, num_tiles, last_time);
int64_t t1 = ggml_time_ms();
auto input_tile = sd_tensor_split_2d(input, input_tile_width, input_tile_height, x * scale_in, y * scale_in);
auto output_tile = on_processing(input_tile);
if (output_tile.empty()) {
return {};
}
GGML_ASSERT(output_tile.shape()[0] == output_tile_width && output_tile.shape()[1] == output_tile_height);
if (output.empty()) {
std::vector<int64_t> output_shape = output_tile.shape();
output_shape[0] = output_width;
output_shape[1] = output_height;
output = sd::Tensor<float>::zeros(std::move(output_shape));
}
sd_tensor_merge_2d_non_circular(
output_tile, &output,
x * scale_out, y * scale_out,
overlap_left * scale_out,
overlap_right * scale_out,
overlap_top * scale_out,
overlap_bottom * scale_out);
if (!silent) {
last_time = (ggml_time_ms() - t1) / 1000.0f;
pretty_progress(tile_count, num_tiles, last_time);
}
tile_count++;
}
tile_count++;
}
last_x = false;
}
if (!silent && tile_count < num_tiles) {
pretty_progress(num_tiles, num_tiles, last_time);
Expand Down