mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-09-24 20:20:37 +00:00
fix: restrict VAE tiling retries to allocation failures (#2019)
This commit is contained in:
parent
78557f88d9
commit
97d932b8f8
@ -1754,7 +1754,7 @@ ArgOptions SDGenerationParams::get_options() {
|
|||||||
on_scm_policy_arg},
|
on_scm_policy_arg},
|
||||||
{"",
|
{"",
|
||||||
"--vae-tile-size",
|
"--vae-tile-size",
|
||||||
"tile size for vae tiling, format [X]x[Y] (default: 32x32)",
|
"tile size for vae tiling in latent units, not image pixels, format [X]x[Y] (default: 32x32)",
|
||||||
on_tile_size_arg},
|
on_tile_size_arg},
|
||||||
{"",
|
{"",
|
||||||
"--vae-relative-tile-size",
|
"--vae-relative-tile-size",
|
||||||
|
|||||||
@ -478,7 +478,11 @@ namespace sd::backend_fit {
|
|||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params, bool prefer_temporal_tiling) {
|
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params, bool prefer_temporal_tiling, ggml_status status) {
|
||||||
|
// Execution failures can leave the device unusable; tiling only helps with allocation failures.
|
||||||
|
if (status != GGML_STATUS_ALLOC_FAILED) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
const char* retry_mode = nullptr;
|
const char* retry_mode = nullptr;
|
||||||
if (prefer_temporal_tiling && !tiling_params.temporal_tiling) {
|
if (prefer_temporal_tiling && !tiling_params.temporal_tiling) {
|
||||||
tiling_params.temporal_tiling = true;
|
tiling_params.temporal_tiling = true;
|
||||||
@ -498,7 +502,7 @@ namespace sd::backend_fit {
|
|||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
LOG_WARN("VAE decode failed (likely out of memory); retrying with %s tiling",
|
LOG_WARN("VAE decode ran out of memory; retrying with %s tiling",
|
||||||
retry_mode);
|
retry_mode);
|
||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
|||||||
@ -16,7 +16,8 @@ namespace sd::backend_fit {
|
|||||||
std::string& params_spec);
|
std::string& params_spec);
|
||||||
|
|
||||||
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params,
|
bool prepare_vae_decode_retry_tiling(sd_tiling_params_t& tiling_params,
|
||||||
bool prefer_temporal_tiling);
|
bool prefer_temporal_tiling,
|
||||||
|
ggml_status status);
|
||||||
|
|
||||||
} // namespace sd::backend_fit
|
} // namespace sd::backend_fit
|
||||||
|
|
||||||
|
|||||||
@ -590,6 +590,7 @@ std::optional<sd::Tensor<float>> GGMLRunner::compute(get_graph_cb_t get_graph,
|
|||||||
bool auto_runner_end,
|
bool auto_runner_end,
|
||||||
bool no_return,
|
bool no_return,
|
||||||
const std::function<bool()>& read_outputs) {
|
const std::function<bool()>& read_outputs) {
|
||||||
|
last_compute_status_ = GGML_STATUS_FAILED;
|
||||||
if (graph_active_) {
|
if (graph_active_) {
|
||||||
LOG_ERROR("%s does not support reentrant graph execution", get_desc().c_str());
|
LOG_ERROR("%s does not support reentrant graph execution", get_desc().c_str());
|
||||||
return std::nullopt;
|
return std::nullopt;
|
||||||
@ -613,7 +614,9 @@ std::optional<sd::Tensor<float>> GGMLRunner::compute(get_graph_cb_t get_graph,
|
|||||||
GGMLRunner& runner;
|
GGMLRunner& runner;
|
||||||
const bool& success;
|
const bool& success;
|
||||||
~GraphEndGuard() {
|
~GraphEndGuard() {
|
||||||
runner.workspace_.segment_end();
|
if (!runner.workspace_.segment_end()) {
|
||||||
|
runner.last_compute_status_ = GGML_STATUS_FAILED;
|
||||||
|
}
|
||||||
runner.cache_.graph_end(false);
|
runner.cache_.graph_end(false);
|
||||||
runner.cut_cache_.clear();
|
runner.cut_cache_.clear();
|
||||||
runner.free_compute_ctx();
|
runner.free_compute_ctx();
|
||||||
@ -642,6 +645,7 @@ std::optional<sd::Tensor<float>> GGMLRunner::compute(get_graph_cb_t get_graph,
|
|||||||
try {
|
try {
|
||||||
output = execute_graph(graph, n_threads, no_return, read_outputs);
|
output = execute_graph(graph, n_threads, no_return, read_outputs);
|
||||||
} catch (const std::exception& error) {
|
} catch (const std::exception& error) {
|
||||||
|
last_compute_status_ = GGML_STATUS_FAILED;
|
||||||
LOG_ERROR("%s graph execution failed on %s: %s", get_desc().c_str(),
|
LOG_ERROR("%s graph execution failed on %s: %s", get_desc().c_str(),
|
||||||
ggml_backend_name(runtime_backend), error.what());
|
ggml_backend_name(runtime_backend), error.what());
|
||||||
return std::nullopt;
|
return std::nullopt;
|
||||||
@ -649,6 +653,7 @@ std::optional<sd::Tensor<float>> GGMLRunner::compute(get_graph_cb_t get_graph,
|
|||||||
success = output.has_value();
|
success = output.has_value();
|
||||||
if (success) {
|
if (success) {
|
||||||
cache_.graph_end(true);
|
cache_.graph_end(true);
|
||||||
|
last_compute_status_ = GGML_STATUS_SUCCESS;
|
||||||
}
|
}
|
||||||
return output;
|
return output;
|
||||||
}
|
}
|
||||||
@ -766,6 +771,7 @@ bool GGMLRunner::execute_segment(ggml_cgraph* graph, int n_threads) {
|
|||||||
}
|
}
|
||||||
workspace_.synchronize();
|
workspace_.synchronize();
|
||||||
if (status != GGML_STATUS_SUCCESS) {
|
if (status != GGML_STATUS_SUCCESS) {
|
||||||
|
last_compute_status_ = status;
|
||||||
LOG_ERROR("%s compute failed: %s", get_desc().c_str(), ggml_status_to_string(status));
|
LOG_ERROR("%s compute failed: %s", get_desc().c_str(), ggml_status_to_string(status));
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
@ -818,6 +824,7 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
const auto& cached_plan = resolve_graph_cut_plan(graph);
|
const auto& cached_plan = resolve_graph_cut_plan(graph);
|
||||||
const auto full_measurement = measure(graph, cached_plan.compute_buffer_size);
|
const auto full_measurement = measure(graph, cached_plan.compute_buffer_size);
|
||||||
if (full_measurement.buffers.empty()) {
|
if (full_measurement.buffers.empty()) {
|
||||||
|
last_compute_status_ = GGML_STATUS_ALLOC_FAILED;
|
||||||
return std::nullopt;
|
return std::nullopt;
|
||||||
}
|
}
|
||||||
auto manager = residency_manager.lock();
|
auto manager = residency_manager.lock();
|
||||||
@ -888,7 +895,9 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
SegmentGraphBindings& bindings;
|
SegmentGraphBindings& bindings;
|
||||||
ggml_context* context;
|
ggml_context* context;
|
||||||
~SegmentCleanup() {
|
~SegmentCleanup() {
|
||||||
runner.workspace_.segment_end();
|
if (!runner.workspace_.segment_end()) {
|
||||||
|
runner.last_compute_status_ = GGML_STATUS_FAILED;
|
||||||
|
}
|
||||||
bindings.restore();
|
bindings.restore();
|
||||||
weights.segment_end();
|
weights.segment_end();
|
||||||
ggml_free(context);
|
ggml_free(context);
|
||||||
@ -898,6 +907,7 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
|
|
||||||
auto measurement = segmented ? measure(segment_graph, segment.compute_buffer_size) : full_measurement;
|
auto measurement = segmented ? measure(segment_graph, segment.compute_buffer_size) : full_measurement;
|
||||||
if (!workspace_.prepare(measurement)) {
|
if (!workspace_.prepare(measurement)) {
|
||||||
|
last_compute_status_ = GGML_STATUS_ALLOC_FAILED;
|
||||||
return fail_segment("workspace preparation");
|
return fail_segment("workspace preparation");
|
||||||
}
|
}
|
||||||
const size_t cut_bytes = last ? 0 : cut_cache_.estimate_output_bytes(graph, segment);
|
const size_t cut_bytes = last ? 0 : cut_cache_.estimate_output_bytes(graph, segment);
|
||||||
@ -912,7 +922,11 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
sync_runtime_residency();
|
sync_runtime_residency();
|
||||||
requests = memory_requests(measurement.buffers, new_cache_bytes);
|
requests = memory_requests(measurement.buffers, new_cache_bytes);
|
||||||
}
|
}
|
||||||
return weights.ensure_segment_capacity(index, requests);
|
const bool ready = weights.ensure_segment_capacity(index, requests);
|
||||||
|
if (!ready && manager != nullptr) {
|
||||||
|
last_compute_status_ = GGML_STATUS_ALLOC_FAILED;
|
||||||
|
}
|
||||||
|
return ready;
|
||||||
};
|
};
|
||||||
if (!weights.segment_start(index, ensure_capacity)) {
|
if (!weights.segment_start(index, ensure_capacity)) {
|
||||||
return fail_segment("weight preparation");
|
return fail_segment("weight preparation");
|
||||||
@ -921,12 +935,17 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
if (!workspace_.measurement_matches(segment_graph, measurement)) {
|
if (!workspace_.measurement_matches(segment_graph, measurement)) {
|
||||||
measurement = measure(segment_graph, segment.compute_buffer_size);
|
measurement = measure(segment_graph, segment.compute_buffer_size);
|
||||||
}
|
}
|
||||||
if (!workspace_.prepare(measurement) || !ensure_capacity()) {
|
if (!workspace_.prepare(measurement)) {
|
||||||
|
last_compute_status_ = GGML_STATUS_ALLOC_FAILED;
|
||||||
|
return fail_segment("workspace preparation");
|
||||||
|
}
|
||||||
|
if (!ensure_capacity()) {
|
||||||
return fail_segment("workspace capacity check");
|
return fail_segment("workspace capacity check");
|
||||||
}
|
}
|
||||||
if (!workspace_.allocate(segment_graph, [&](ggml_backend_sched_t scheduler, ggml_cgraph* current) {
|
if (!workspace_.allocate(segment_graph, [&](ggml_backend_sched_t scheduler, ggml_cgraph* current) {
|
||||||
pin_multi_device_nodes(scheduler, current);
|
pin_multi_device_nodes(scheduler, current);
|
||||||
})) {
|
})) {
|
||||||
|
last_compute_status_ = GGML_STATUS_ALLOC_FAILED;
|
||||||
return fail_segment("workspace allocation");
|
return fail_segment("workspace allocation");
|
||||||
}
|
}
|
||||||
for (const auto& size : measurement.buffers) {
|
for (const auto& size : measurement.buffers) {
|
||||||
@ -964,6 +983,7 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if (!workspace_.segment_end()) {
|
if (!workspace_.segment_end()) {
|
||||||
|
last_compute_status_ = GGML_STATUS_FAILED;
|
||||||
return fail_segment("workspace synchronization");
|
return fail_segment("workspace synchronization");
|
||||||
}
|
}
|
||||||
// Final outputs and their callbacks may still be views of consumed cuts.
|
// Final outputs and their callbacks may still be views of consumed cuts.
|
||||||
|
|||||||
@ -131,6 +131,7 @@ struct GGMLRunner {
|
|||||||
private:
|
private:
|
||||||
std::map<ggml_backend_t, size_t> logged_compute_bytes_;
|
std::map<ggml_backend_t, size_t> logged_compute_bytes_;
|
||||||
size_t logged_segment_count_ = 0;
|
size_t logged_segment_count_ = 0;
|
||||||
|
ggml_status last_compute_status_ = GGML_STATUS_SUCCESS;
|
||||||
|
|
||||||
sd::ComputeWorkspace::Measurement measure(ggml_cgraph* graph, size_t direct_bytes);
|
sd::ComputeWorkspace::Measurement measure(ggml_cgraph* graph, size_t direct_bytes);
|
||||||
std::vector<DeviceMemoryRequest> memory_requests(const std::vector<sd::BackendBufferSize>& sizes,
|
std::vector<DeviceMemoryRequest> memory_requests(const std::vector<sd::BackendBufferSize>& sizes,
|
||||||
@ -335,6 +336,8 @@ public:
|
|||||||
bool no_return = false,
|
bool no_return = false,
|
||||||
const std::function<bool()>& read_outputs = {});
|
const std::function<bool()>& read_outputs = {});
|
||||||
|
|
||||||
|
ggml_status last_compute_status() const { return last_compute_status_; }
|
||||||
|
|
||||||
void set_flash_attention_enabled(bool enabled) {
|
void set_flash_attention_enabled(bool enabled) {
|
||||||
flash_attn_enabled = enabled;
|
flash_attn_enabled = enabled;
|
||||||
}
|
}
|
||||||
|
|||||||
@ -2766,7 +2766,8 @@ sd::Tensor<float> StableDiffusionGGML::decode_first_stage(const sd::Tensor<float
|
|||||||
auto decoded = first_stage_model->decode(n_threads, latents, vae_tiling_params, decode_video, circular_x, circular_y);
|
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();
|
const bool prefer_temporal_tiling = decode_video && first_stage_model->can_temporal_tile_decode();
|
||||||
while (decoded.empty() &&
|
while (decoded.empty() &&
|
||||||
sd::backend_fit::prepare_vae_decode_retry_tiling(vae_tiling_params, prefer_temporal_tiling)) {
|
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);
|
decoded = first_stage_model->decode(n_threads, latents, vae_tiling_params, decode_video, circular_x, circular_y);
|
||||||
}
|
}
|
||||||
return decoded;
|
return decoded;
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user