fix: restrict VAE tiling retries to allocation failures (#2019)

This commit is contained in:
leejet 2026-09-21 23:40:05 +08:00 committed by GitHub
parent 78557f88d9
commit 97d932b8f8
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 39 additions and 10 deletions

View File

@ -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",

View File

@ -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;
} }

View File

@ -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

View File

@ -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.

View File

@ -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;
} }

View File

@ -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;