aboutsummaryrefslogtreecommitdiffstats
path: root/stable-diffusion.cpp-vulkan/patches/lokr-nd-input.patch
blob: 6dd2d01759d2184a8410b9ba964125ec1b97caac (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
diff --git a/src/model/adapter/lora_ops.cpp b/src/model/adapter/lora_ops.cpp
index 7dacfd1..686b8d3 100644
--- a/src/model/adapter/lora_ops.cpp
+++ b/src/model/adapter/lora_ops.cpp
@@ -65,6 +65,13 @@ ggml_tensor* ggml_ext_lokr_forward(
     ggml_tensor* hb;
 
     if (!is_conv) {
+        // Linear inputs may carry extra dims ([q, tokens, batch, ...], e.g. Krea2):
+        // flatten them into one batch dim and restore the shape on output.
+        ggml_tensor* h_in = h;
+        if (!ggml_is_contiguous(h)) {
+            h = ggml_cont(ctx, h);
+        }
+        h                  = ggml_reshape_2d(ctx, h, h->ne[0], ggml_nelements(h) / h->ne[0]);
         int batch          = (int)h->ne[1];
         int merge_batch_uq = batch;
         int merge_batch_vp = batch;
@@ -118,7 +125,7 @@ ggml_tensor* ggml_ext_lokr_forward(
         }
 
         ggml_tensor* hc  = ggml_transpose(ctx, hc_t);
-        ggml_tensor* out = ggml_reshape_2d(ctx, ggml_cont(ctx, hc), up * vp, batch);
+        ggml_tensor* out = ggml_reshape_4d(ctx, ggml_cont(ctx, hc), up * vp, h_in->ne[1], h_in->ne[2], h_in->ne[3]);
         return ggml_ext_scale(ctx, out, scale);
     } else {
         int batch = (int)h->ne[3];