wgblas 1.2.1 → 2.0.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +1 -1
- package/README.md +36 -54
- package/dist/wgblas.browser.js +779 -34
- package/index.d.mts +13 -0
- package/index.mjs +8 -0
- package/package.json +55 -2
- package/src/dasum/dasum.d.mts +2 -2
- package/src/dasum/dasum.mjs +6 -4
- package/src/idamax/idamax.d.mts +51 -0
- package/src/idamax/idamax.mjs +128 -0
- package/src/init.mjs +9 -1
- package/src/isamax/isamax.d.mts +1 -1
- package/src/sasum/sasum.d.mts +1 -1
- package/src/saxpy/saxpy.d.mts +1 -1
- package/src/scopy/scopy.d.mts +1 -1
- package/src/sdot/sdot.d.mts +1 -1
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +195 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +203 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +4 -0
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +2 -0
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/browser-shaders.mjs +26 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +117 -0
- package/src/shaders/sgemm_small.wgsl +112 -0
- package/src/shaders/sgemmtr_large.wgsl +117 -0
- package/src/shaders/sgemmtr_small.wgsl +110 -0
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +1 -1
- package/src/srot/srot.d.mts +1 -1
- package/src/srotm/srotm.d.mts +1 -1
- package/src/sscal/sscal.d.mts +1 -1
- package/src/sswap/sswap.d.mts +1 -1
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +209 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +2 -0
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +2 -0
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +2 -0
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +201 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +176 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +211 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +2 -0
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +342 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +2 -0
- package/src/util/buffer.mjs +4 -2
- package/src/util/compute.mjs +6 -3
- package/src/util/f64.mjs +3 -3
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
stageReadback,
|
|
5
|
+
destroyBuffers,
|
|
6
|
+
} from "../util/buffer.mjs";
|
|
7
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
8
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
9
|
+
import { extractResult } from "../util/result.mjs";
|
|
10
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
13
|
+
|
|
14
|
+
const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
|
|
15
|
+
const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
|
|
16
|
+
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // large tile needs >= a 6x6 grid of its own tiles to beat the small tile
|
|
17
|
+
|
|
18
|
+
export async function sgemm(
|
|
19
|
+
device, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
|
|
20
|
+
) {
|
|
21
|
+
let AIsGpu = A instanceof GpuMatrix;
|
|
22
|
+
let BIsGpu = B instanceof GpuMatrix;
|
|
23
|
+
const CIsGpu = C instanceof GpuMatrix;
|
|
24
|
+
|
|
25
|
+
if (!(device instanceof GPUDevice))
|
|
26
|
+
throw new Error("device must be a GPUDevice.");
|
|
27
|
+
if (transA !== "no-transpose" && transA !== "transpose")
|
|
28
|
+
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
29
|
+
if (transB !== "no-transpose" && transB !== "transpose")
|
|
30
|
+
throw new Error("transB must be 'no-transpose' or 'transpose'.");
|
|
31
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
32
|
+
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
33
|
+
if (typeof alpha !== "number")
|
|
34
|
+
throw new Error("alpha must be a number.");
|
|
35
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
36
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
37
|
+
if (typeof beta !== "number")
|
|
38
|
+
throw new Error("beta must be a number.");
|
|
39
|
+
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
40
|
+
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
41
|
+
if (
|
|
42
|
+
!Number.isInteger(m) ||
|
|
43
|
+
!Number.isInteger(n) ||
|
|
44
|
+
!Number.isInteger(k) ||
|
|
45
|
+
!Number.isInteger(lda) ||
|
|
46
|
+
!Number.isInteger(ldb) ||
|
|
47
|
+
!Number.isInteger(ldc)
|
|
48
|
+
)
|
|
49
|
+
throw new Error("m, n, k, lda, ldb, and ldc must be integers.");
|
|
50
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
51
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
52
|
+
if (!BIsGpu && !(B instanceof Float32Array))
|
|
53
|
+
throw new Error("B must be a Float32Array or GpuMatrix.");
|
|
54
|
+
if (!CIsGpu && !(C instanceof Float32Array))
|
|
55
|
+
throw new Error("C must be a Float32Array or GpuMatrix.");
|
|
56
|
+
if ((AIsGpu || BIsGpu) && !CIsGpu)
|
|
57
|
+
throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
|
|
58
|
+
if (CIsGpu && (!AIsGpu || !BIsGpu))
|
|
59
|
+
throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
|
|
60
|
+
if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
|
|
61
|
+
if (m === 0 || n === 0) return CIsGpu ? {} : { C };
|
|
62
|
+
|
|
63
|
+
// GpuMatrix's own .layout wins over the shared `layout` argument.
|
|
64
|
+
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
65
|
+
const effLayoutB = BIsGpu ? B.layout : layout;
|
|
66
|
+
const effLayoutC = CIsGpu ? C.layout : layout;
|
|
67
|
+
|
|
68
|
+
// Shape validation, before any layout-driven swapping below.
|
|
69
|
+
// A: op(A) is m x k; A itself is m x k or k x m depending on transA.
|
|
70
|
+
const aRows = effLayoutA === "column-major" ? k : m;
|
|
71
|
+
const aCols = effLayoutA === "column-major" ? m : k;
|
|
72
|
+
const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
|
|
73
|
+
const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
|
|
74
|
+
if (lda < aInner)
|
|
75
|
+
throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
|
|
76
|
+
if (AIsGpu) {
|
|
77
|
+
if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
78
|
+
const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
|
|
79
|
+
if (A.rows < aLogRows || A.cols < aLogCols)
|
|
80
|
+
throw new Error("A is too small for the given m, k, and transA.");
|
|
81
|
+
} else if (A.length < (aOuter - 1) * lda + aInner) {
|
|
82
|
+
throw new Error("A does not have enough elements for the given dimensions and lda.");
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
// B: same reasoning as A, with op(B) = k x n.
|
|
86
|
+
const bRows = effLayoutB === "column-major" ? n : k;
|
|
87
|
+
const bCols = effLayoutB === "column-major" ? k : n;
|
|
88
|
+
const bOuter = transB === "no-transpose" ? bRows : bCols;
|
|
89
|
+
const bInner = transB === "no-transpose" ? bCols : bRows;
|
|
90
|
+
if (ldb < bInner)
|
|
91
|
+
throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
|
|
92
|
+
if (BIsGpu) {
|
|
93
|
+
if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
94
|
+
const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
|
|
95
|
+
if (B.rows < bLogRows || B.cols < bLogCols)
|
|
96
|
+
throw new Error("B is too small for the given n, k, and transB.");
|
|
97
|
+
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
98
|
+
throw new Error("B does not have enough elements for the given dimensions and ldb.");
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
// C: always m x n (no trans flag) — layout only affects lda/storage order.
|
|
102
|
+
const cOuter = effLayoutC === "column-major" ? n : m;
|
|
103
|
+
const cInner = effLayoutC === "column-major" ? m : n;
|
|
104
|
+
if (ldc < cInner)
|
|
105
|
+
throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
|
|
106
|
+
if (CIsGpu) {
|
|
107
|
+
if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
|
|
108
|
+
if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
|
|
109
|
+
} else if (C.length < (cOuter - 1) * ldc + cInner) {
|
|
110
|
+
throw new Error("C does not have enough elements for the given dimensions and ldc.");
|
|
111
|
+
}
|
|
112
|
+
|
|
113
|
+
// Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
|
|
114
|
+
if (effLayoutA === "column-major")
|
|
115
|
+
transA = transA === "no-transpose" ? "transpose" : "no-transpose";
|
|
116
|
+
if (effLayoutB === "column-major")
|
|
117
|
+
transB = transB === "no-transpose" ? "transpose" : "no-transpose";
|
|
118
|
+
|
|
119
|
+
// Column-major C: compute C^T = op(B)^T * op(A)^T instead (swap A/B, flip trans, swap m<->n).
|
|
120
|
+
if (effLayoutC === "column-major") {
|
|
121
|
+
[A, B] = [B, A];
|
|
122
|
+
[AIsGpu, BIsGpu] = [BIsGpu, AIsGpu];
|
|
123
|
+
[lda, ldb] = [ldb, lda];
|
|
124
|
+
[transA, transB] = [
|
|
125
|
+
transB === "no-transpose" ? "transpose" : "no-transpose",
|
|
126
|
+
transA === "no-transpose" ? "transpose" : "no-transpose",
|
|
127
|
+
];
|
|
128
|
+
[m, n] = [n, m];
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
// Shape-based auto-select — see sgemm_small.wgsl/sgemm_large.wgsl.
|
|
132
|
+
const largeWgX = Math.ceil(n / BN_LARGE);
|
|
133
|
+
const largeWgY = Math.ceil(m / BM_LARGE);
|
|
134
|
+
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
135
|
+
|
|
136
|
+
const pipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
|
|
137
|
+
|
|
138
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemm-A", false);
|
|
139
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemm-B", false);
|
|
140
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemm-C", true);
|
|
141
|
+
const paramsBuffer = createParamsBuffer(
|
|
142
|
+
[
|
|
143
|
+
{ value: m, type: "u32" },
|
|
144
|
+
{ value: n, type: "u32" },
|
|
145
|
+
{ value: k, type: "u32" },
|
|
146
|
+
{ value: alpha, type: "f32" },
|
|
147
|
+
{ value: beta, type: "f32" },
|
|
148
|
+
{ value: lda, type: "u32" },
|
|
149
|
+
{ value: ldb, type: "u32" },
|
|
150
|
+
{ value: ldc, type: "u32" },
|
|
151
|
+
{ value: transA === "transpose" ? 1 : 0, type: "u32" },
|
|
152
|
+
{ value: transB === "transpose" ? 1 : 0, type: "u32" },
|
|
153
|
+
],
|
|
154
|
+
"sgemm-params",
|
|
155
|
+
);
|
|
156
|
+
|
|
157
|
+
try {
|
|
158
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
159
|
+
ABuffer,
|
|
160
|
+
BBuffer,
|
|
161
|
+
CBuffer,
|
|
162
|
+
paramsBuffer,
|
|
163
|
+
]);
|
|
164
|
+
|
|
165
|
+
const wgCount = useLargeTile
|
|
166
|
+
? {
|
|
167
|
+
x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
|
|
168
|
+
y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
|
|
169
|
+
}
|
|
170
|
+
: {
|
|
171
|
+
x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
|
|
172
|
+
y: Math.min(Math.ceil(m / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
|
|
173
|
+
};
|
|
174
|
+
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
175
|
+
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
176
|
+
|
|
177
|
+
submit(commandEncoder);
|
|
178
|
+
|
|
179
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
180
|
+
|
|
181
|
+
if (CIsGpu) {
|
|
182
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
183
|
+
return {};
|
|
184
|
+
}
|
|
185
|
+
|
|
186
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
187
|
+
if (gpuTimeMs !== undefined) return { C: result, gpuTimeMs };
|
|
188
|
+
return { C: result };
|
|
189
|
+
} finally {
|
|
190
|
+
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
191
|
+
if (!BIsGpu) destroyBuffers(BBuffer);
|
|
192
|
+
if (!CIsGpu) destroyBuffers(CBuffer);
|
|
193
|
+
destroyBuffers(paramsBuffer);
|
|
194
|
+
}
|
|
195
|
+
}
|
|
@@ -0,0 +1,104 @@
|
|
|
1
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Performs the matrix-matrix operation C := uplo(alpha * op(A) * op(B) + beta * C) —
|
|
5
|
+
* `sgemm`'s operation, but only the triangle of C named by `uplo` is read or
|
|
6
|
+
* written (`'lower'`: `col <= row`, `'upper'`: `col >= row`; C need not be
|
|
7
|
+
* square — the test applies over the full m×n grid).
|
|
8
|
+
*
|
|
9
|
+
* Same kernels as `sgemm` (`sgemmtr_small.wgsl`/`sgemmtr_large.wgsl`,
|
|
10
|
+
* identical tiling), with the final output write masked to one triangle.
|
|
11
|
+
*
|
|
12
|
+
* {@includeCode ../../examples/sgemmtr/sgemmtr.js}
|
|
13
|
+
*
|
|
14
|
+
* **Browser (standalone HTML):**
|
|
15
|
+
* {@includeCode ../../examples/sgemmtr/web/sgemmtr.html}
|
|
16
|
+
*
|
|
17
|
+
* @param device - GPUDevice from `init()`
|
|
18
|
+
* @param uplo - `'lower'` to update only `col <= row`, `'upper'` for `col >= row`
|
|
19
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
20
|
+
* @param transB - `'no-transpose'` for B, `'transpose'` for B^T
|
|
21
|
+
* @param m - rows of op(A) and C
|
|
22
|
+
* @param n - columns of op(B) and C
|
|
23
|
+
* @param k - columns of op(A), rows of op(B)
|
|
24
|
+
* @param alpha - scalar multiplier for op(A)*op(B)
|
|
25
|
+
* @param A - Float32Array, row-major or column-major (see `layout`)
|
|
26
|
+
* @param lda - leading dimension of A as stored
|
|
27
|
+
* @param B - Float32Array, row-major or column-major (see `layout`)
|
|
28
|
+
* @param ldb - leading dimension of B as stored
|
|
29
|
+
* @param beta - scalar multiplier for C
|
|
30
|
+
* @param C - Float32Array input/output matrix, row-major or column-major
|
|
31
|
+
* @param ldc - leading dimension of C as stored
|
|
32
|
+
* @param layout - storage layout shared by A/B/C when they're Float32Array
|
|
33
|
+
* (default: `'row-major'`) — same handling as `sgemm`, plus `uplo` is
|
|
34
|
+
* flipped internally for column-major C so it still names the triangle
|
|
35
|
+
* you asked for
|
|
36
|
+
* @returns updated C as a Float32Array (only the requested triangle changed)
|
|
37
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemmtr/sgemmtr.mjs#L19">Source code: sgemmtr.mjs (L19)</a>
|
|
38
|
+
* @category BLAS Level 3
|
|
39
|
+
*/
|
|
40
|
+
export declare function sgemmtr(
|
|
41
|
+
device: GPUDevice,
|
|
42
|
+
uplo: 'lower' | 'upper',
|
|
43
|
+
transA: 'no-transpose' | 'transpose',
|
|
44
|
+
transB: 'no-transpose' | 'transpose',
|
|
45
|
+
m: number,
|
|
46
|
+
n: number,
|
|
47
|
+
k: number,
|
|
48
|
+
alpha: number,
|
|
49
|
+
A: Float32Array,
|
|
50
|
+
lda: number,
|
|
51
|
+
B: Float32Array,
|
|
52
|
+
ldb: number,
|
|
53
|
+
beta: number,
|
|
54
|
+
C: Float32Array,
|
|
55
|
+
ldc: number,
|
|
56
|
+
layout?: 'row-major' | 'column-major',
|
|
57
|
+
): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
|
|
58
|
+
|
|
59
|
+
/**
|
|
60
|
+
* Performs the matrix-matrix operation C := uplo(alpha * op(A) * op(B) + beta * C)
|
|
61
|
+
*
|
|
62
|
+
* A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
|
|
63
|
+
* `GpuMatrix.from` time) determines the operation — there is no separate
|
|
64
|
+
* `layout` argument here. A and B must be GpuMatrix whenever C is, and vice
|
|
65
|
+
* versa — mixing a GpuMatrix with a plain Float32Array is not supported.
|
|
66
|
+
*
|
|
67
|
+
* {@includeCode ../../examples/sgemmtr/gpu.sgemmtr.js}
|
|
68
|
+
*
|
|
69
|
+
* @param device - GPUDevice from `init()`
|
|
70
|
+
* @param uplo - `'lower'` to update only `col <= row`, `'upper'` for `col >= row`
|
|
71
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
72
|
+
* @param transB - `'no-transpose'` for B, `'transpose'` for B^T
|
|
73
|
+
* @param m - rows of op(A) and C
|
|
74
|
+
* @param n - columns of op(B) and C
|
|
75
|
+
* @param k - columns of op(A), rows of op(B)
|
|
76
|
+
* @param alpha - scalar multiplier for op(A)*op(B)
|
|
77
|
+
* @param A - GpuMatrix
|
|
78
|
+
* @param lda - leading dimension of A (must equal A.lda)
|
|
79
|
+
* @param B - GpuMatrix
|
|
80
|
+
* @param ldb - leading dimension of B (must equal B.lda)
|
|
81
|
+
* @param beta - scalar multiplier for C
|
|
82
|
+
* @param C - GpuMatrix (mutated in place; only the requested triangle changes)
|
|
83
|
+
* @param ldc - leading dimension of C (must equal C.lda)
|
|
84
|
+
* @returns no C — it stays GPU-resident; call `C.read()` yourself for a CPU readback (see the example)
|
|
85
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemmtr/sgemmtr.mjs#L19">Source code: sgemmtr.mjs (L19)</a>
|
|
86
|
+
* @category BLAS Level 3
|
|
87
|
+
*/
|
|
88
|
+
export declare function sgemmtr(
|
|
89
|
+
device: GPUDevice,
|
|
90
|
+
uplo: 'lower' | 'upper',
|
|
91
|
+
transA: 'no-transpose' | 'transpose',
|
|
92
|
+
transB: 'no-transpose' | 'transpose',
|
|
93
|
+
m: number,
|
|
94
|
+
n: number,
|
|
95
|
+
k: number,
|
|
96
|
+
alpha: number,
|
|
97
|
+
A: GpuMatrix,
|
|
98
|
+
lda: number,
|
|
99
|
+
B: GpuMatrix,
|
|
100
|
+
ldb: number,
|
|
101
|
+
beta: number,
|
|
102
|
+
C: GpuMatrix,
|
|
103
|
+
ldc: number,
|
|
104
|
+
): Promise<{ gpuTimeMs?: number }>;
|
|
@@ -0,0 +1,203 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
stageReadback,
|
|
5
|
+
destroyBuffers,
|
|
6
|
+
} from "../util/buffer.mjs";
|
|
7
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
8
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
9
|
+
import { extractResult } from "../util/result.mjs";
|
|
10
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
13
|
+
|
|
14
|
+
const BM_SMALL = 32, BN_SMALL = 32; // sgemmtr_small.wgsl's block tile
|
|
15
|
+
const BM_LARGE = 64, BN_LARGE = 64; // sgemmtr_large.wgsl's block tile
|
|
16
|
+
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm uses — see sgemm.mjs
|
|
17
|
+
|
|
18
|
+
export async function sgemmtr(
|
|
19
|
+
device, uplo, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
|
|
20
|
+
) {
|
|
21
|
+
let AIsGpu = A instanceof GpuMatrix;
|
|
22
|
+
let BIsGpu = B instanceof GpuMatrix;
|
|
23
|
+
const CIsGpu = C instanceof GpuMatrix;
|
|
24
|
+
|
|
25
|
+
if (!(device instanceof GPUDevice))
|
|
26
|
+
throw new Error("device must be a GPUDevice.");
|
|
27
|
+
if (uplo !== "lower" && uplo !== "upper")
|
|
28
|
+
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
29
|
+
if (transA !== "no-transpose" && transA !== "transpose")
|
|
30
|
+
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
31
|
+
if (transB !== "no-transpose" && transB !== "transpose")
|
|
32
|
+
throw new Error("transB must be 'no-transpose' or 'transpose'.");
|
|
33
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
34
|
+
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
35
|
+
if (typeof alpha !== "number")
|
|
36
|
+
throw new Error("alpha must be a number.");
|
|
37
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
38
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
39
|
+
if (typeof beta !== "number")
|
|
40
|
+
throw new Error("beta must be a number.");
|
|
41
|
+
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
42
|
+
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
43
|
+
if (
|
|
44
|
+
!Number.isInteger(m) ||
|
|
45
|
+
!Number.isInteger(n) ||
|
|
46
|
+
!Number.isInteger(k) ||
|
|
47
|
+
!Number.isInteger(lda) ||
|
|
48
|
+
!Number.isInteger(ldb) ||
|
|
49
|
+
!Number.isInteger(ldc)
|
|
50
|
+
)
|
|
51
|
+
throw new Error("m, n, k, lda, ldb, and ldc must be integers.");
|
|
52
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
53
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
54
|
+
if (!BIsGpu && !(B instanceof Float32Array))
|
|
55
|
+
throw new Error("B must be a Float32Array or GpuMatrix.");
|
|
56
|
+
if (!CIsGpu && !(C instanceof Float32Array))
|
|
57
|
+
throw new Error("C must be a Float32Array or GpuMatrix.");
|
|
58
|
+
if ((AIsGpu || BIsGpu) && !CIsGpu)
|
|
59
|
+
throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
|
|
60
|
+
if (CIsGpu && (!AIsGpu || !BIsGpu))
|
|
61
|
+
throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
|
|
62
|
+
if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
|
|
63
|
+
if (m === 0 || n === 0) return CIsGpu ? {} : { C };
|
|
64
|
+
|
|
65
|
+
// GpuMatrix's own .layout wins over the shared `layout` argument.
|
|
66
|
+
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
67
|
+
const effLayoutB = BIsGpu ? B.layout : layout;
|
|
68
|
+
const effLayoutC = CIsGpu ? C.layout : layout;
|
|
69
|
+
|
|
70
|
+
// Shape validation, before any layout-driven swapping below.
|
|
71
|
+
// A: op(A) is m x k; A itself is m x k or k x m depending on transA.
|
|
72
|
+
const aRows = effLayoutA === "column-major" ? k : m;
|
|
73
|
+
const aCols = effLayoutA === "column-major" ? m : k;
|
|
74
|
+
const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
|
|
75
|
+
const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
|
|
76
|
+
if (lda < aInner)
|
|
77
|
+
throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
|
|
78
|
+
if (AIsGpu) {
|
|
79
|
+
if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
80
|
+
const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
|
|
81
|
+
if (A.rows < aLogRows || A.cols < aLogCols)
|
|
82
|
+
throw new Error("A is too small for the given m, k, and transA.");
|
|
83
|
+
} else if (A.length < (aOuter - 1) * lda + aInner) {
|
|
84
|
+
throw new Error("A does not have enough elements for the given dimensions and lda.");
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
// B: same reasoning as A, with op(B) = k x n.
|
|
88
|
+
const bRows = effLayoutB === "column-major" ? n : k;
|
|
89
|
+
const bCols = effLayoutB === "column-major" ? k : n;
|
|
90
|
+
const bOuter = transB === "no-transpose" ? bRows : bCols;
|
|
91
|
+
const bInner = transB === "no-transpose" ? bCols : bRows;
|
|
92
|
+
if (ldb < bInner)
|
|
93
|
+
throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
|
|
94
|
+
if (BIsGpu) {
|
|
95
|
+
if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
96
|
+
const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
|
|
97
|
+
if (B.rows < bLogRows || B.cols < bLogCols)
|
|
98
|
+
throw new Error("B is too small for the given n, k, and transB.");
|
|
99
|
+
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
100
|
+
throw new Error("B does not have enough elements for the given dimensions and ldb.");
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
// C: always m x n (no trans flag) — layout only affects lda/storage order.
|
|
104
|
+
const cOuter = effLayoutC === "column-major" ? n : m;
|
|
105
|
+
const cInner = effLayoutC === "column-major" ? m : n;
|
|
106
|
+
if (ldc < cInner)
|
|
107
|
+
throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
|
|
108
|
+
if (CIsGpu) {
|
|
109
|
+
if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
|
|
110
|
+
if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
|
|
111
|
+
} else if (C.length < (cOuter - 1) * ldc + cInner) {
|
|
112
|
+
throw new Error("C does not have enough elements for the given dimensions and ldc.");
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
// Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
|
|
116
|
+
if (effLayoutA === "column-major")
|
|
117
|
+
transA = transA === "no-transpose" ? "transpose" : "no-transpose";
|
|
118
|
+
if (effLayoutB === "column-major")
|
|
119
|
+
transB = transB === "no-transpose" ? "transpose" : "no-transpose";
|
|
120
|
+
|
|
121
|
+
// Column-major C: compute C^T = op(B)^T * op(A)^T instead (swap A/B, flip trans, swap m<->n)
|
|
122
|
+
// — same trick sgemm uses. uplo(C) in row/col terms becomes uplo(C^T) with row/col swapped,
|
|
123
|
+
// i.e. the opposite triangle, so uplo must flip here too (same reasoning ssyr's isLower flip
|
|
124
|
+
// uses for column-major A) — nothing else that's uplo-specific needs to change, since the
|
|
125
|
+
// shader's row/col test is applied to whatever (m, n) grid it's actually given.
|
|
126
|
+
if (effLayoutC === "column-major") {
|
|
127
|
+
[A, B] = [B, A];
|
|
128
|
+
[AIsGpu, BIsGpu] = [BIsGpu, AIsGpu];
|
|
129
|
+
[lda, ldb] = [ldb, lda];
|
|
130
|
+
[transA, transB] = [
|
|
131
|
+
transB === "no-transpose" ? "transpose" : "no-transpose",
|
|
132
|
+
transA === "no-transpose" ? "transpose" : "no-transpose",
|
|
133
|
+
];
|
|
134
|
+
[m, n] = [n, m];
|
|
135
|
+
uplo = uplo === "lower" ? "upper" : "lower";
|
|
136
|
+
}
|
|
137
|
+
|
|
138
|
+
// Shape-based auto-select — see sgemmtr_small.wgsl/sgemmtr_large.wgsl.
|
|
139
|
+
const largeWgX = Math.ceil(n / BN_LARGE);
|
|
140
|
+
const largeWgY = Math.ceil(m / BM_LARGE);
|
|
141
|
+
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
142
|
+
|
|
143
|
+
const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
|
|
144
|
+
|
|
145
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemmtr-A", false);
|
|
146
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemmtr-B", false);
|
|
147
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemmtr-C", true);
|
|
148
|
+
const paramsBuffer = createParamsBuffer(
|
|
149
|
+
[
|
|
150
|
+
{ value: m, type: "u32" },
|
|
151
|
+
{ value: n, type: "u32" },
|
|
152
|
+
{ value: k, type: "u32" },
|
|
153
|
+
{ value: alpha, type: "f32" },
|
|
154
|
+
{ value: beta, type: "f32" },
|
|
155
|
+
{ value: lda, type: "u32" },
|
|
156
|
+
{ value: ldb, type: "u32" },
|
|
157
|
+
{ value: ldc, type: "u32" },
|
|
158
|
+
{ value: transA === "transpose" ? 1 : 0, type: "u32" },
|
|
159
|
+
{ value: transB === "transpose" ? 1 : 0, type: "u32" },
|
|
160
|
+
{ value: uplo === "upper" ? 1 : 0, type: "u32" },
|
|
161
|
+
],
|
|
162
|
+
"sgemmtr-params",
|
|
163
|
+
);
|
|
164
|
+
|
|
165
|
+
try {
|
|
166
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
167
|
+
ABuffer,
|
|
168
|
+
BBuffer,
|
|
169
|
+
CBuffer,
|
|
170
|
+
paramsBuffer,
|
|
171
|
+
]);
|
|
172
|
+
|
|
173
|
+
const wgCount = useLargeTile
|
|
174
|
+
? {
|
|
175
|
+
x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
|
|
176
|
+
y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
|
|
177
|
+
}
|
|
178
|
+
: {
|
|
179
|
+
x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
|
|
180
|
+
y: Math.min(Math.ceil(m / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
|
|
181
|
+
};
|
|
182
|
+
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
183
|
+
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
184
|
+
|
|
185
|
+
submit(commandEncoder);
|
|
186
|
+
|
|
187
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
188
|
+
|
|
189
|
+
if (CIsGpu) {
|
|
190
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
191
|
+
return {};
|
|
192
|
+
}
|
|
193
|
+
|
|
194
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
195
|
+
if (gpuTimeMs !== undefined) return { C: result, gpuTimeMs };
|
|
196
|
+
return { C: result };
|
|
197
|
+
} finally {
|
|
198
|
+
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
199
|
+
if (!BIsGpu) destroyBuffers(BBuffer);
|
|
200
|
+
if (!CIsGpu) destroyBuffers(CBuffer);
|
|
201
|
+
destroyBuffers(paramsBuffer);
|
|
202
|
+
}
|
|
203
|
+
}
|
package/src/sgemv/sgemv.d.mts
CHANGED
|
@@ -50,43 +50,6 @@ export declare function sgemv(
|
|
|
50
50
|
layout?: 'row-major' | 'column-major',
|
|
51
51
|
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
52
52
|
|
|
53
|
-
/**
|
|
54
|
-
* Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
|
|
55
|
-
*
|
|
56
|
-
* A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
|
|
57
|
-
* (set at `GpuMatrix.from` time) determines the operation — there is no
|
|
58
|
-
* separate `layout` argument here.
|
|
59
|
-
*
|
|
60
|
-
* @param device - GPUDevice from `init()`
|
|
61
|
-
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
|
62
|
-
* @param m - number of rows in A
|
|
63
|
-
* @param n - number of columns in A
|
|
64
|
-
* @param alpha - scalar multiplier for op(A)*x
|
|
65
|
-
* @param A - GpuMatrix, GPU-resident
|
|
66
|
-
* @param lda - leading dimension of A (must equal A.lda)
|
|
67
|
-
* @param x - Float32Array input vector
|
|
68
|
-
* @param incx - stride for x (must be a positive integer)
|
|
69
|
-
* @param beta - scalar multiplier for y
|
|
70
|
-
* @param y - Float32Array input/output vector
|
|
71
|
-
* @param incy - stride for y (must be a positive integer)
|
|
72
|
-
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
|
|
73
|
-
* @category BLAS Level 2
|
|
74
|
-
*/
|
|
75
|
-
export declare function sgemv(
|
|
76
|
-
device: GPUDevice,
|
|
77
|
-
trans: 'no-transpose' | 'transpose',
|
|
78
|
-
m: number,
|
|
79
|
-
n: number,
|
|
80
|
-
alpha: number,
|
|
81
|
-
A: GpuMatrix,
|
|
82
|
-
lda: number,
|
|
83
|
-
x: Float32Array,
|
|
84
|
-
incx: number,
|
|
85
|
-
beta: number,
|
|
86
|
-
y: Float32Array,
|
|
87
|
-
incy: number,
|
|
88
|
-
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
89
|
-
|
|
90
53
|
/**
|
|
91
54
|
* Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
|
|
92
55
|
*
|
|
@@ -94,7 +57,7 @@ export declare function sgemv(
|
|
|
94
57
|
* `layout` (set at `GpuMatrix.from` time) determines the operation — there is
|
|
95
58
|
* no separate `layout` argument here.
|
|
96
59
|
*
|
|
97
|
-
* {@includeCode ../../examples/sgemv/
|
|
60
|
+
* {@includeCode ../../examples/sgemv/gpu.sgemv.js}
|
|
98
61
|
*
|
|
99
62
|
* @param device - GPUDevice from `init()`
|
|
100
63
|
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
package/src/sgemv/sgemv.mjs
CHANGED
|
@@ -56,6 +56,10 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
56
56
|
throw new Error(
|
|
57
57
|
"A must be a GpuMatrix when x and y are GpuVectors.",
|
|
58
58
|
);
|
|
59
|
+
if (AIsGpu && !xIsGpu)
|
|
60
|
+
throw new Error(
|
|
61
|
+
"x and y must be GpuVectors when A is a GpuMatrix.",
|
|
62
|
+
);
|
|
59
63
|
if (xIsGpu && x._buf === y._buf)
|
|
60
64
|
throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
|
|
61
65
|
if (AIsGpu && lda !== A.lda)
|
package/src/sger/sger.d.mts
CHANGED
|
@@ -42,39 +42,6 @@ export declare function sger(
|
|
|
42
42
|
layout?: 'row-major' | 'column-major',
|
|
43
43
|
): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
|
|
44
44
|
|
|
45
|
-
/**
|
|
46
|
-
* Performs the rank-1 update A = alpha * x * y^T + A
|
|
47
|
-
*
|
|
48
|
-
* A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
|
|
49
|
-
* (set at `GpuMatrix.from` time) determines the operation — there is no
|
|
50
|
-
* separate `layout` argument here.
|
|
51
|
-
*
|
|
52
|
-
* @param device - GPUDevice from `init()`
|
|
53
|
-
* @param m - number of rows in A
|
|
54
|
-
* @param n - number of columns in A
|
|
55
|
-
* @param alpha - scalar multiplier for x*y^T
|
|
56
|
-
* @param x - Float32Array input vector
|
|
57
|
-
* @param incx - stride for x (must be a positive integer)
|
|
58
|
-
* @param y - Float32Array input vector
|
|
59
|
-
* @param incy - stride for y (must be a positive integer)
|
|
60
|
-
* @param A - GpuMatrix, GPU-resident
|
|
61
|
-
* @param lda - leading dimension of A (must equal A.lda)
|
|
62
|
-
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sger/sger.mjs#L15">Source code: sger.mjs (L15)</a>
|
|
63
|
-
* @category BLAS Level 2
|
|
64
|
-
*/
|
|
65
|
-
export declare function sger(
|
|
66
|
-
device: GPUDevice,
|
|
67
|
-
m: number,
|
|
68
|
-
n: number,
|
|
69
|
-
alpha: number,
|
|
70
|
-
x: Float32Array,
|
|
71
|
-
incx: number,
|
|
72
|
-
y: Float32Array,
|
|
73
|
-
incy: number,
|
|
74
|
-
A: GpuMatrix,
|
|
75
|
-
lda: number,
|
|
76
|
-
): Promise<{ gpuTimeMs?: number }>;
|
|
77
|
-
|
|
78
45
|
/**
|
|
79
46
|
* Performs the rank-1 update A = alpha * x * y^T + A
|
|
80
47
|
*
|
|
@@ -82,7 +49,7 @@ export declare function sger(
|
|
|
82
49
|
* `GpuMatrix.from` time) determines the operation — there is no separate
|
|
83
50
|
* `layout` argument here.
|
|
84
51
|
*
|
|
85
|
-
* {@includeCode ../../examples/sger/
|
|
52
|
+
* {@includeCode ../../examples/sger/gpu.sger.js}
|
|
86
53
|
*
|
|
87
54
|
* @param device - GPUDevice from `init()`
|
|
88
55
|
* @param m - number of rows in A
|
package/src/sger/sger.mjs
CHANGED
|
@@ -63,6 +63,8 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
63
63
|
);
|
|
64
64
|
if (xIsGpu && !AIsGpu)
|
|
65
65
|
throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
|
|
66
|
+
if (AIsGpu && !xIsGpu)
|
|
67
|
+
throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
|
|
66
68
|
if (AIsGpu && xIsGpu && A._buf === x._buf)
|
|
67
69
|
throw new Error("A and x must not reference the same GPU buffer.");
|
|
68
70
|
if (AIsGpu && yIsGpu && A._buf === y._buf)
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
// block_transfer: gather/scatter/scatter-subtract between a tight (blockLen
|
|
2
|
+
// x otherLen) block and a sub-range of a strided (any ld, row/col-major)
|
|
3
|
+
// buffer — needed since block offsets aren't 256-byte-aligned and block
|
|
4
|
+
// rows/cols aren't always one contiguous range for copyBufferToBuffer.
|
|
5
|
+
|
|
6
|
+
@group(0) @binding(0) var<storage, read_write> block: array<f32>; // blockLen x otherLen, block[i*otherLen+j]
|
|
7
|
+
@group(0) @binding(1) var<storage, read_write> strided: array<f32>; // B's or A's own buffer
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
blockStart: u32,
|
|
11
|
+
blockLen: u32,
|
|
12
|
+
otherStart: u32,
|
|
13
|
+
otherLen: u32,
|
|
14
|
+
ld: u32,
|
|
15
|
+
isColMajor: u32, // 0 = row-major addressing, 1 = column-major (row/col swapped)
|
|
16
|
+
blockIsRow: u32, // 1 = blockStart indexes strided's rows, 0 = its columns
|
|
17
|
+
mode: u32, // 0 = scatter (strided := block), 1 = scatter_sub (strided -= block), 2 = gather (block := strided)
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
21
|
+
|
|
22
|
+
@compute @workgroup_size(8, 8)
|
|
23
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
24
|
+
let i = gid.y; // index along the blocked axis, within the block
|
|
25
|
+
let j = gid.x; // index along the other axis, within the block
|
|
26
|
+
if (i >= params.blockLen || j >= params.otherLen) {
|
|
27
|
+
return;
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
let row = select(params.otherStart + j, params.blockStart + i, params.blockIsRow == 1u);
|
|
31
|
+
let col = select(params.blockStart + i, params.otherStart + j, params.blockIsRow == 1u);
|
|
32
|
+
let stridedIdx = select(row * params.ld + col, col * params.ld + row, params.isColMajor == 1u);
|
|
33
|
+
let blockIdx = i * params.otherLen + j;
|
|
34
|
+
|
|
35
|
+
if (params.mode == 2u) {
|
|
36
|
+
block[blockIdx] = strided[stridedIdx];
|
|
37
|
+
} else if (params.mode == 1u) {
|
|
38
|
+
strided[stridedIdx] -= block[blockIdx];
|
|
39
|
+
} else {
|
|
40
|
+
strided[stridedIdx] = block[blockIdx];
|
|
41
|
+
}
|
|
42
|
+
}
|