feat: add native CUDA Sol-Attn support

This commit is contained in:
leejet
2026-09-21 21:15:00 +08:00
parent c678dfe704
commit 0657e6cdfe
13 changed files with 258 additions and 10 deletions
+21 -3
View File
@@ -623,7 +623,9 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
bool skip_reshape,
bool flash_attn,
float kv_scale,
bool sage_attn) { // avoid overflow
bool sage_attn,
bool sol_attn,
float sol_attn_tau) { // avoid overflow
int64_t L_q;
int64_t L_k;
int64_t C;
@@ -715,7 +717,23 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
};
#ifndef SD_USE_UPSTREAM_GGML
if (sage_attn && mask == nullptr && d_head > 0 && d_head <= 128) {
if (sol_attn && mask == nullptr && d_head == 128 && L_q == L_k && n_head == n_kv_head) {
auto q_in = ggml_reshape_4d(ctx, ggml_ext_cont(ctx, q->type == GGML_TYPE_F32 ? q : ggml_cast(ctx, q, GGML_TYPE_F32)), d_head, L_q, n_head, N);
auto k_in = ggml_reshape_4d(ctx, ggml_ext_cont(ctx, k->type == GGML_TYPE_F32 ? k : ggml_cast(ctx, k, GGML_TYPE_F32)), d_head, L_k, n_kv_head, N);
auto v_in = ggml_ext_cont(ctx, ggml_permute(ctx, v, 0, 2, 1, 3));
if (v_in->type != GGML_TYPE_F32) {
v_in = ggml_cast(ctx, v_in, GGML_TYPE_F32);
}
if (kv_scale != 1.0f) {
k_in = ggml_ext_scale(ctx, k_in, kv_scale);
v_in = ggml_ext_scale(ctx, v_in, kv_scale);
}
auto out = ggml_sol_attn(ctx, q_in, k_in, v_in, scale / kv_scale, sol_attn_tau);
if (ggml_backend_supports_op(backend, out)) {
kqv = kv_scale != 1.0f ? ggml_ext_scale(ctx, out, 1.0f / kv_scale) : out;
}
}
if (kqv == nullptr && sage_attn && mask == nullptr && d_head > 0 && d_head <= 128) {
auto q_in = ggml_reshape_4d(ctx, ggml_ext_cont(ctx, q->type == GGML_TYPE_F32 ? q : ggml_cast(ctx, q, GGML_TYPE_F32)), d_head, L_q, n_head, N);
auto k_in = ggml_reshape_4d(ctx, ggml_ext_cont(ctx, k->type == GGML_TYPE_F32 ? k : ggml_cast(ctx, k, GGML_TYPE_F32)), d_head, L_k, n_kv_head, N);
auto v_in = ggml_ext_cont(ctx, ggml_permute(ctx, v, 0, 2, 1, 3));
@@ -744,7 +762,7 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
}
#endif
if (kqv == nullptr && (flash_attn || sage_attn)) {
if (kqv == nullptr && (flash_attn || sage_attn || sol_attn)) {
// LOG_VERBOSE("attention_ext L_q:%d L_k:%d n_head:%d C:%d d_head:%d N:%d", L_q, L_k, n_head, C, d_head, N);
bool can_use_flash_attn = true;
if (mask != nullptr) {
+7 -5
View File
@@ -217,11 +217,13 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
ggml_tensor* k,
ggml_tensor* v,
int64_t n_head,
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.0f,
bool sage_attn = false);
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.0f,
bool sage_attn = false,
bool sol_attn = false,
float sol_attn_tau = 1.0f);
ggml_tensor* ggml_ext_layer_norm(ggml_context* ctx,
ggml_tensor* x,
+19 -1
View File
@@ -25,7 +25,7 @@ ggml_tensor* ggml_ext_attention_ext(GGMLRunnerContext* ctx,
if (ctx->attn_scale > 0.f) {
kv_scale = ctx->attn_scale;
}
return ggml_ext_attention_ext(ctx->ggml_ctx, ctx->backend, q, k, v, n_head, mask, skip_reshape, flash_attn, kv_scale, ctx->sage_attn_enabled);
return ggml_ext_attention_ext(ctx->ggml_ctx, ctx->backend, q, k, v, n_head, mask, skip_reshape, flash_attn, kv_scale, ctx->sage_attn_enabled, ctx->sol_attn_enabled, ctx->sol_attn_tau);
}
void GGMLRunner::alloc_params_ctx() {
@@ -164,6 +164,22 @@ ggml_cgraph* GGMLRunner::get_compute_graph(get_graph_cb_t get_graph) {
}
}
prepare_build_in_tensor_after(gf);
#ifndef SD_USE_UPSTREAM_GGML
if (sol_attn_enabled && !sol_attn_graph_logged) {
int sol_nodes = 0;
int flash_nodes = 0;
for (int i = 0; i < ggml_graph_n_nodes(gf); ++i) {
const auto op = ggml_graph_node(gf, i)->op;
sol_nodes += op == GGML_OP_SOL_ATTN;
flash_nodes += op == GGML_OP_FLASH_ATTN_EXT;
}
LOG_INFO("Sol-Attn graph: %d Sol-Attn nodes, %d FlashAttention nodes", sol_nodes, flash_nodes);
if (sol_nodes == 0) {
LOG_WARN("This graph has no attention operations supported by Sol-Attn");
}
sol_attn_graph_logged = true;
}
#endif
return gf;
}
@@ -521,6 +537,8 @@ GGMLRunnerContext GGMLRunner::get_context() {
runner_ctx.backend = runtime_backend;
runner_ctx.flash_attn_enabled = flash_attn_enabled;
runner_ctx.sage_attn_enabled = sage_attn_enabled;
runner_ctx.sol_attn_enabled = sol_attn_enabled;
runner_ctx.sol_attn_tau = sol_attn_tau;
runner_ctx.linear_scale = linear_scale;
runner_ctx.attn_scale = attn_scale;
runner_ctx.conv2d_direct_enabled = conv2d_direct_enabled;
+15
View File
@@ -69,6 +69,8 @@ struct GGMLRunnerContext {
ggml_context* ggml_ctx = nullptr;
bool flash_attn_enabled = false;
bool sage_attn_enabled = false;
bool sol_attn_enabled = false;
float sol_attn_tau = 1.f;
float linear_scale = 0.f;
float attn_scale = 0.f;
bool conv2d_direct_enabled = false;
@@ -178,6 +180,9 @@ protected:
bool flash_attn_enabled = false;
bool sage_attn_enabled = false;
bool sol_attn_enabled = false;
float sol_attn_tau = 1.f;
bool sol_attn_graph_logged = false;
float linear_scale = 0.f;
float attn_scale = 0.f;
bool conv2d_direct_enabled = false;
@@ -347,6 +352,16 @@ public:
}
}
void set_sol_attention_enabled(bool enabled, float tau) {
if (sol_attn_enabled != enabled || sol_attn_tau != tau) {
free_cache_ctx_and_buffer();
graph_cut_plan_cache_.graph_cut_plans.clear();
sol_attn_enabled = enabled;
sol_attn_tau = tau;
sol_attn_graph_logged = false;
}
}
void set_scale_overrides(float linear_scale, float attn_scale) {
this->linear_scale = linear_scale;
this->attn_scale = attn_scale;
+56
View File
@@ -899,7 +899,60 @@ bool StableDiffusionGGML::set_sage_attention_enabled(bool enabled) {
return true;
}
bool StableDiffusionGGML::set_sol_attention_enabled(bool enabled, float tau) {
if (!diffusion_model || !std::isfinite(tau)) {
LOG_ERROR("Sol-Attn requires a diffusion model and finite tau");
return false;
}
if (enabled) {
if (config_->params.sage_attn) {
LOG_ERROR("Sol-Attn and SageAttention cannot be enabled together");
return false;
}
#ifndef SD_USE_UPSTREAM_GGML
auto* ctx = ggml_init({4 * ggml_tensor_overhead(), nullptr, true});
if (!ctx) {
return false;
}
auto* q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 128, 128, 1, 1);
auto* k = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 128, 128, 1, 1);
auto* v = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 128, 128, 1, 1);
auto* op = ggml_sol_attn(ctx, q, k, v, 1.f / sqrtf(128.f), tau);
bool supported = true;
for (auto backend : backend_manager.runtime_backends(SDBackendModule::DIFFUSION)) {
if (!ggml_backend_supports_op(backend, op)) {
LOG_ERROR("Sol-Attn is unavailable on %s; it requires patched GGML, CUDA 12.0 or newer, and SM80 or newer kernels", ggml_backend_name(backend));
supported = false;
}
}
ggml_free(ctx);
if (!supported) {
return false;
}
#else
LOG_ERROR("Sol-Attn requires -DSD_USE_UPSTREAM_GGML=OFF and a CUDA backend");
return false;
#endif
}
diffusion_model->set_sol_attention_enabled(enabled, tau);
if (high_noise_diffusion_model) {
high_noise_diffusion_model->set_sol_attention_enabled(enabled, tau);
}
if (enabled) {
LOG_INFO("Using Sol-Attn (tau=%g, diagonal threshold) in diffusion; unsupported attention uses flash/default attention", tau);
}
return true;
}
bool StableDiffusionGGML::init(const sd_ctx_params_t* sd_ctx_params) {
if (sd_ctx_params->sol_attn && sd_ctx_params->sage_attn) {
LOG_ERROR("Sol-Attn and SageAttention cannot be enabled together");
return false;
}
if (!std::isfinite(sd_ctx_params->sol_attn_tau)) {
LOG_ERROR("Sol-Attn tau must be finite");
return false;
}
#ifdef SD_USE_UPSTREAM_GGML
LOG_WARN(
"Using upstream GGML: INT8 tensorwise/convrot is disabled and FP8 weights are "
@@ -1180,6 +1233,9 @@ bool StableDiffusionGGML::validate_and_load_runners() {
if (sd_ctx_params->sage_attn && !set_sage_attention_enabled(true)) {
return false;
}
if (sd_ctx_params->sol_attn && !set_sol_attention_enabled(true, sd_ctx_params->sol_attn_tau)) {
return false;
}
LOG_VERBOSE("validating model metadata");
std::set<std::string> ignore_tensors;
+1
View File
@@ -313,6 +313,7 @@ public:
bool init(const sd_ctx_params_t* sd_ctx_params);
bool set_sage_attention_enabled(bool enabled);
bool set_sol_attention_enabled(bool enabled, float tau);
bool uses_tae() const;
+6
View File
@@ -338,6 +338,8 @@ void sd_ctx_params_init(sd_ctx_params_t* sd_ctx_params) {
sd_ctx_params->enable_mmap = false;
sd_ctx_params->diffusion_flash_attn = false;
sd_ctx_params->sage_attn = false;
sd_ctx_params->sol_attn = false;
sd_ctx_params->sol_attn_tau = 1.f;
sd_ctx_params->linear_scale = 0.f;
sd_ctx_params->attn_scale = 0.f;
sd_ctx_params->vae_format = SD_VAE_FORMAT_AUTO;
@@ -394,6 +396,8 @@ char* sd_ctx_params_to_str(const sd_ctx_params_t* sd_ctx_params) {
"flash_attn: %s\n"
"diffusion_flash_attn: %s\n"
"sage_attn: %s\n"
"sol_attn: %s\n"
"sol_attn_tau: %g\n"
"linear_scale: %g\n"
"attn_scale: %g\n"
"vae_format: %s\n",
@@ -434,6 +438,8 @@ char* sd_ctx_params_to_str(const sd_ctx_params_t* sd_ctx_params) {
BOOL_STR(sd_ctx_params->flash_attn),
BOOL_STR(sd_ctx_params->diffusion_flash_attn),
BOOL_STR(sd_ctx_params->sage_attn),
BOOL_STR(sd_ctx_params->sol_attn),
sd_ctx_params->sol_attn_tau,
sd_ctx_params->linear_scale,
sd_ctx_params->attn_scale,
sd_vae_format_name(sd_ctx_params->vae_format));