fix: extend f32 matmul precision to ROCm for Qwen-Image, Krea2 and Boogu (#1772)

This commit is contained in:
Piotr Wilkin (ilintar) 2026-07-10 17:23:01 +02:00 committed by GitHub
parent 9beb6aca69
commit ead6bf521b
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 6 additions and 6 deletions

View File

@ -294,7 +294,7 @@ public:
auto net_0 = std::dynamic_pointer_cast<UnaryBlock>(blocks["net.0"]); auto net_0 = std::dynamic_pointer_cast<UnaryBlock>(blocks["net.0"]);
auto net_2 = std::dynamic_pointer_cast<Linear>(blocks["net.2"]); auto net_2 = std::dynamic_pointer_cast<Linear>(blocks["net.2"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
net_2->set_force_prec_f32(true); net_2->set_force_prec_f32(true);
} }

View File

@ -199,7 +199,7 @@ namespace Boogu {
auto linear_2 = std::dynamic_pointer_cast<Linear>(blocks["linear_2"]); auto linear_2 = std::dynamic_pointer_cast<Linear>(blocks["linear_2"]);
auto linear_3 = std::dynamic_pointer_cast<Linear>(blocks["linear_3"]); auto linear_3 = std::dynamic_pointer_cast<Linear>(blocks["linear_3"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
linear_2->set_force_prec_f32(true); linear_2->set_force_prec_f32(true);
} }
@ -259,7 +259,7 @@ namespace Boogu {
auto norm_k = std::dynamic_pointer_cast<RMSNorm>(blocks["norm_k"]); auto norm_k = std::dynamic_pointer_cast<RMSNorm>(blocks["norm_k"]);
auto to_out_0 = std::dynamic_pointer_cast<Linear>(blocks["to_out.0"]); auto to_out_0 = std::dynamic_pointer_cast<Linear>(blocks["to_out.0"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
to_out_0->set_force_prec_f32(true); to_out_0->set_force_prec_f32(true);
} }
@ -383,7 +383,7 @@ namespace Boogu {
auto instruct_out = std::dynamic_pointer_cast<Linear>(blocks["processor.instruct_out"]); auto instruct_out = std::dynamic_pointer_cast<Linear>(blocks["processor.instruct_out"]);
auto img_out = std::dynamic_pointer_cast<Linear>(blocks["processor.img_out"]); auto img_out = std::dynamic_pointer_cast<Linear>(blocks["processor.img_out"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
to_out_0->set_force_prec_f32(true); to_out_0->set_force_prec_f32(true);
} }

View File

@ -267,7 +267,7 @@ namespace Krea2 {
auto knorm = std::dynamic_pointer_cast<KreaRMSNorm>(blocks["qknorm.knorm"]); auto knorm = std::dynamic_pointer_cast<KreaRMSNorm>(blocks["qknorm.knorm"]);
auto wo = std::dynamic_pointer_cast<Linear>(blocks["wo"]); auto wo = std::dynamic_pointer_cast<Linear>(blocks["wo"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
wo->set_force_prec_f32(true); wo->set_force_prec_f32(true);
} }

View File

@ -183,7 +183,7 @@ namespace Qwen {
auto to_v = std::dynamic_pointer_cast<Linear>(blocks["to_v"]); auto to_v = std::dynamic_pointer_cast<Linear>(blocks["to_v"]);
auto to_out_0 = std::dynamic_pointer_cast<Linear>(blocks["to_out.0"]); auto to_out_0 = std::dynamic_pointer_cast<Linear>(blocks["to_out.0"]);
if (sd_backend_is(ctx->backend, "Vulkan")) { if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
to_out_0->set_force_prec_f32(true); to_out_0->set_force_prec_f32(true);
} }