Skip to content
Open
Show file tree
Hide file tree
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
9 changes: 8 additions & 1 deletion src/abot_world.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -974,9 +974,16 @@ class AbotWalkSession {
// return false to fall back to the RNG.
std::function<bool(int block, int step, float* dst, size_t n)> noise_override;

std::shared_ptr<ModelManager> model_manager;
std::unique_ptr<AbotWorldRunner> runner;
std::shared_ptr<AbotTinyVideoAutoEncoder> tae;
// Declared after the runners on purpose (mirrors StableDiffusionGGML):
// ~ModelManager force-frees the param storage blocks and writes through
// the registered ggml tensors (state->tensor->buffer = nullptr), which
// live in the runner/tae contexts above. Members destroy in reverse
// declaration order and the runners hold only weak_ptr refs to the
// manager, so this ordering runs ~ModelManager first, while every
// registered tensor is still alive.
std::shared_ptr<ModelManager> model_manager;

// finalized walk state (per-frame ggml {W,H,C} latents, torch [C,H,W] flat)
std::vector<std::vector<float>> history;
Expand Down
109 changes: 78 additions & 31 deletions src/core/ggml_extend.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -2071,7 +2071,12 @@ struct GGMLRunner {
if (cpu_fallback_backend == nullptr && !sd_backend_is_cpu(runtime_backend)) {
cpu_fallback_backend = sd_backend_cpu_init();
}
if (cpu_fallback_backend != nullptr) {
// The internal CPU fallback backend is only needed as the
// scheduler's trailing fallback while the primary runtime is
// non-CPU. When a VAE graph is retried on its explicit CPU backend,
// adding a second CPU backend creates an ambiguous scheduler route.
if (cpu_fallback_backend != nullptr &&
!sd_backend_is_cpu(runtime_backend)) {
backends.push_back(cpu_fallback_backend);
}

Expand Down Expand Up @@ -2227,6 +2232,36 @@ struct GGMLRunner {
return bytes;
}

bool assign_graph_params_compute_backend(const std::vector<ggml_tensor*>& graph_params,
ggml_backend_t backend,
const char* action) {
GGML_ASSERT(backend != nullptr);
auto manager = weight_manager.lock();
if (manager == nullptr) {
if (graph_params.empty()) {
return true;
}
LOG_ERROR("%s VAE CPU fallback cannot %s graph params without a weight manager",
get_desc().c_str(),
action);
return false;
}
if (!manager->assign_compute_backend(graph_params, backend)) {
LOG_ERROR("%s VAE CPU fallback failed to %s %zu graph params to %s",
get_desc().c_str(),
action,
graph_params.size(),
ggml_backend_name(backend));
return false;
}
LOG_DEBUG("%s VAE CPU fallback %s %zu graph params to %s",
get_desc().c_str(),
action,
graph_params.size(),
ggml_backend_name(backend));
return true;
}

void switch_runtime_backend(ggml_backend_t backend) {
GGML_ASSERT(backend != nullptr);
free_compute_buffer();
Expand Down Expand Up @@ -3233,37 +3268,59 @@ struct GGMLRunner {
}
const bool has_stateful_cache =
!cache_tensor_map.empty() || cache_ctx != nullptr;
auto retry_stateless_on_cpu =
[&](const char* failure) -> std::optional<sd::Tensor<T>> {
if (!vae_auto_cpu_fallback_enabled ||
sd_backend_is_cpu(runtime_backend) ||
vae_fallback_backend == nullptr ||
!sd_backend_is_cpu(vae_fallback_backend) ||
has_stateful_cache) {
auto run_stateless_on_cpu =
[&](const std::vector<ggml_tensor*>& graph_params) -> std::optional<sd::Tensor<T>> {
ggml_backend_t previous_backend = runtime_backend;
const std::string previous_backend_name = ggml_backend_name(previous_backend);
if (!assign_graph_params_compute_backend(graph_params,
vae_fallback_backend,
"assign")) {
return std::nullopt;
}
ggml_backend_t previous_backend = runtime_backend;
const std::string previous_backend_name =
ggml_backend_name(previous_backend);
LOG_WARN("%s VAE %s on %s; retrying stateless graph on CPU",
get_desc().c_str(),
failure,
previous_backend_name.c_str());
switch_runtime_backend(vae_fallback_backend);

auto restore_previous_backend = [&]() {
const bool restored = assign_graph_params_compute_backend(graph_params,
previous_backend,
"restore");
switch_runtime_backend(previous_backend);
return restored;
};

try {
// Fallback staging cannot remain active after runtime restoration.
// Always release CPU compute params before assigning the graph back.
auto output = compute<T>(get_graph,
n_threads,
false,
free_compute_buffer,
free_compute_params,
true,
no_return);
switch_runtime_backend(previous_backend);
if (!restore_previous_backend()) {
return std::nullopt;
}
return output;
} catch (...) {
switch_runtime_backend(previous_backend);
restore_previous_backend();
throw;
}
};
auto retry_stateless_on_cpu =
[&](const char* failure) -> std::optional<sd::Tensor<T>> {
if (!vae_auto_cpu_fallback_enabled ||
sd_backend_is_cpu(runtime_backend) ||
vae_fallback_backend == nullptr ||
!sd_backend_is_cpu(vae_fallback_backend) ||
has_stateful_cache) {
return std::nullopt;
}
const std::string previous_backend_name = ggml_backend_name(runtime_backend);
LOG_WARN("%s VAE %s on %s; retrying stateless graph on CPU",
get_desc().c_str(),
failure,
previous_backend_name.c_str());
return run_stateless_on_cpu(collect_used_param_tensors(gf));
};

if (vae_auto_cpu_fallback_enabled &&
!sd_backend_is_cpu(runtime_backend) &&
Expand Down Expand Up @@ -3315,23 +3372,13 @@ struct GGMLRunner {
previous_backend_name.c_str(),
cpu_backend_name.c_str());

switch_runtime_backend(vae_fallback_backend);
try {
auto output = compute<T>(get_graph,
n_threads,
false,
free_compute_buffer,
free_compute_params,
no_return);
switch_runtime_backend(previous_backend);
auto output = run_stateless_on_cpu(collect_used_param_tensors(gf));
if (output.has_value()) {
LOG_INFO("%s VAE CPU fallback complete; restored runtime backend %s",
get_desc().c_str(),
previous_backend_name.c_str());
return output;
} catch (...) {
switch_runtime_backend(previous_backend);
throw;
}
return output;
}

if (decision.reason == sd::VaeGraphRouteReason::STATEFUL_GRAPH) {
Expand Down
19 changes: 13 additions & 6 deletions src/stable-diffusion.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7517,16 +7517,23 @@ bool sd_abot_scene_create(const sd_abot_scene_params_t* p) {
LOG_ERROR("sd_abot_scene_create: backend init failed: %s", error.c_str()); return false;
}
const int threads = p->n_threads > 0 ? p->n_threads : sd_get_num_physical_cores();
// The runners below own the ggml tensors that ~ModelManager's
// storage-block teardown writes through (state->tensor->buffer =
// nullptr), and they hold only weak_ptr refs to the manager. Declare
// them BEFORE the manager so reverse local destruction runs
// ~ModelManager first, while every registered tensor is still alive
// (mirrors the member ordering in StableDiffusionGGML).
std::shared_ptr<AbotT5Runner> t5;
std::shared_ptr<AbotWanVAE> vae;
auto model_manager = std::make_shared<ModelManager>();
model_manager->set_n_threads(threads);
ModelLoader& model_loader = model_manager->loader();
if (!model_loader.init_from_file(p->t5_path, "text_encoders.t5xxl.transformer.")) return false;
auto t5 = std::make_shared<AbotT5Runner>(backends.runtime_backend(SDBackendModule::TE),
model_loader.get_tensor_storage_map(),
"text_encoders.t5xxl.transformer",
true,
model_manager);
std::shared_ptr<AbotWanVAE> vae;
t5 = std::make_shared<AbotT5Runner>(backends.runtime_backend(SDBackendModule::TE),
model_loader.get_tensor_storage_map(),
"text_encoders.t5xxl.transformer",
true,
model_manager);
if (has_image) {
if (!model_loader.init_from_file(p->vae_path, "first_stage_model.")) return false;
vae = std::make_shared<AbotWanVAE>(backends.runtime_backend(SDBackendModule::VAE),
Expand Down
65 changes: 65 additions & 0 deletions tests/test-vae-routing.cpp
Original file line number Diff line number Diff line change
@@ -1,12 +1,53 @@
#include "core/ggml_extend.hpp"
#include "core/ggml_extend_backend.h"
#include "vae_fallback.hpp"

#include <cstddef>
#include <limits>
#include <memory>
#include <string>
#include <vector>

namespace {
constexpr size_t MIB = 1024ull * 1024ull;

class RecordingWeightManager : public RunnerWeightManager {
public:
std::vector<ggml_backend_t> assignments;
size_t fail_on_assignment = 0;

bool assign_compute_backend(const std::vector<ggml_tensor*>& tensors,
ggml_backend_t compute_backend) override {
GGML_ASSERT(!tensors.empty());
assignments.push_back(compute_backend);
return fail_on_assignment == 0 || assignments.size() != fail_on_assignment;
}

bool prepare_params(const std::vector<ggml_tensor*>&) override { return true; }
void release_compute_backend_params(const std::vector<ggml_tensor*>&) override {}
void release_params_backend_params(const std::vector<ggml_tensor*>&) override {}
};

class RoutingTestRunner : public GGMLRunner {
public:
RoutingTestRunner(ggml_backend_t backend,
const std::shared_ptr<RunnerWeightManager>& manager)
: GGMLRunner(backend, manager) {}

std::string get_desc() override { return "routing_test"; }

ggml_tensor* make_param() {
ggml_tensor* tensor = ggml_new_tensor_1d(params_ctx, GGML_TYPE_F32, 1);
ggml_set_name(tensor, "routing_test.weight");
return tensor;
}

bool assign_graph_params(const std::vector<ggml_tensor*>& tensors,
ggml_backend_t backend,
const char* action) {
return assign_graph_params_compute_backend(tensors, backend, action);
}
};
}

int main() {
Expand All @@ -21,6 +62,30 @@ int main() {
&compatibility_error));
compatibility_manager.reset();

ggml_backend_t original_backend = sd_backend_cpu_init();
ggml_backend_t fallback_backend = sd_backend_cpu_init();
GGML_ASSERT(original_backend != nullptr);
GGML_ASSERT(fallback_backend != nullptr);
{
auto manager = std::make_shared<RecordingWeightManager>();
RoutingTestRunner runner(original_backend, manager);
std::vector<ggml_tensor*> graph_params = {runner.make_param()};

GGML_ASSERT(runner.assign_graph_params(graph_params, fallback_backend, "assign"));
GGML_ASSERT(runner.assign_graph_params(graph_params, original_backend, "restore"));
GGML_ASSERT(manager->assignments.size() == 2);
GGML_ASSERT(manager->assignments[0] == fallback_backend);
GGML_ASSERT(manager->assignments[1] == original_backend);

manager->assignments.clear();
manager->fail_on_assignment = 2;
GGML_ASSERT(runner.assign_graph_params(graph_params, fallback_backend, "assign"));
GGML_ASSERT(!runner.assign_graph_params(graph_params, original_backend, "restore"));
GGML_ASSERT(manager->assignments.size() == 2);
}
ggml_backend_free(fallback_backend);
ggml_backend_free(original_backend);

SDBackendAssignment runtime;
runtime.set_default("vulkan0");
runtime.set_module(SDBackendModule::VAE, "metal");
Expand Down