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