wgblas 1.2.1 → 2.1.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 +54 -72
- package/dist/wgblas.browser.js +1637 -858
- package/index.d.mts +51 -6
- package/index.mjs +8 -0
- package/package.json +56 -2
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +6 -6
- package/src/dasum/dasum.mjs +25 -21
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +69 -0
- package/src/idamax/idamax.mjs +130 -0
- package/src/init.mjs +115 -49
- package/src/isamax/isamax.d.mts +21 -3
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +3 -3
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +3 -3
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +3 -3
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +3 -3
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +208 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +204 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +42 -26
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +12 -8
- package/src/shaders/block_transfer.wgsl +42 -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/index.mjs +164 -14
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +206 -0
- package/src/shaders/sgemm_small.wgsl +212 -0
- package/src/shaders/sgemmtr_large.wgsl +120 -0
- package/src/shaders/sgemmtr_small.wgsl +113 -0
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +3 -3
- package/src/snrm2/snrm2.mjs +33 -21
- package/src/srot/srot.d.mts +3 -5
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +3 -5
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +4 -4
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +3 -3
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +218 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +12 -8
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +11 -7
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +12 -8
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +202 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +177 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +226 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +12 -8
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +360 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +18 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +116 -20
- package/src/util/compute.mjs +12 -12
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/f64.mjs +3 -3
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -55
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Performs the triangular matrix-matrix operation
|
|
5
|
+
* B := alpha * op(A) * B (`side='left'`) or
|
|
6
|
+
* B := alpha * B * op(A) (`side='right'`) — `A` is triangular, only its
|
|
7
|
+
* `uplo` triangle stored; `B` is a general m×n matrix, overwritten in place.
|
|
8
|
+
*
|
|
9
|
+
* - `side='left'`: `A` is m×m — `A` premultiplies `B`
|
|
10
|
+
* - `side='right'`: `A` is n×n — `A` postmultiplies `B`
|
|
11
|
+
*
|
|
12
|
+
* No dedicated fused kernel — a `triangularize` pass materializes a dense
|
|
13
|
+
* copy of `op(A)` (zero-filling the unstored triangle, substituting the
|
|
14
|
+
* implicit diagonal when `diag='unit'`), then a plain `sgemm` pass
|
|
15
|
+
* (`sgemm_small.wgsl`/`sgemm_large.wgsl`, unmodified) does the actual
|
|
16
|
+
* multiply, both on one command encoder.
|
|
17
|
+
*
|
|
18
|
+
* {@includeCode ../../examples/strmm/strmm.js}
|
|
19
|
+
*
|
|
20
|
+
* **Browser (standalone HTML):**
|
|
21
|
+
* {@includeCode ../../examples/strmm/web/strmm.html}
|
|
22
|
+
*
|
|
23
|
+
* @param device - GPUDevice from `init()`
|
|
24
|
+
* @param side - `'left'` for op(A)*B, `'right'` for B*op(A)
|
|
25
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
26
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
27
|
+
* @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
28
|
+
* @param m - rows of B
|
|
29
|
+
* @param n - columns of B
|
|
30
|
+
* @param alpha - scalar multiplier for the matrix product
|
|
31
|
+
* @param A - Float32Array, triangular, row-major or column-major (see `layout`)
|
|
32
|
+
* @param lda - leading dimension of A as stored
|
|
33
|
+
* @param B - Float32Array input/output matrix, overwritten with the result, row-major or column-major
|
|
34
|
+
* @param ldb - leading dimension of B as stored
|
|
35
|
+
* @param layout - storage layout shared by A/B when they're Float32Array
|
|
36
|
+
* (default: `'row-major'`) — column-major A is a genuine transpose (A isn't
|
|
37
|
+
* symmetric like ssymm's), so both `transA` and `uplo` are adjusted
|
|
38
|
+
* internally to compensate
|
|
39
|
+
* @returns updated B as a Float32Array
|
|
40
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmm/strmm.mjs#L25">Source code: strmm.mjs (L25)</a>
|
|
41
|
+
* @category BLAS Level 3
|
|
42
|
+
*/
|
|
43
|
+
export declare function strmm(
|
|
44
|
+
device: GPUDevice,
|
|
45
|
+
side: 'left' | 'right',
|
|
46
|
+
uplo: 'lower' | 'upper',
|
|
47
|
+
transA: 'no-transpose' | 'transpose',
|
|
48
|
+
diag: 'unit' | 'non-unit',
|
|
49
|
+
m: number,
|
|
50
|
+
n: number,
|
|
51
|
+
alpha: number,
|
|
52
|
+
A: Float32Array,
|
|
53
|
+
lda: number,
|
|
54
|
+
B: Float32Array,
|
|
55
|
+
ldb: number,
|
|
56
|
+
layout?: 'row-major' | 'column-major',
|
|
57
|
+
): Promise<{ B: Float32Array; gpuTimeMs?: number }>;
|
|
58
|
+
|
|
59
|
+
/**
|
|
60
|
+
* Performs the triangular matrix-matrix operation
|
|
61
|
+
* B := alpha * op(A) * B (`side='left'`) or
|
|
62
|
+
* B := alpha * B * op(A) (`side='right'`)
|
|
63
|
+
*
|
|
64
|
+
* A and B are both kept GPU-resident; B is mutated in place. Each matrix's
|
|
65
|
+
* own `layout` (set at `GpuMatrix.from` time) determines the operation —
|
|
66
|
+
* there is no separate `layout` argument here. A and B must both be
|
|
67
|
+
* GpuMatrix or both be Float32Array — mixing is not supported.
|
|
68
|
+
*
|
|
69
|
+
* {@includeCode ../../examples/strmm/gpu.strmm.js}
|
|
70
|
+
*
|
|
71
|
+
* @param device - GPUDevice from `init()`
|
|
72
|
+
* @param side - `'left'` for op(A)*B, `'right'` for B*op(A)
|
|
73
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
74
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
75
|
+
* @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
76
|
+
* @param m - rows of B
|
|
77
|
+
* @param n - columns of B
|
|
78
|
+
* @param alpha - scalar multiplier for the matrix product
|
|
79
|
+
* @param A - GpuMatrix, triangular
|
|
80
|
+
* @param lda - leading dimension of A (must equal A.lda)
|
|
81
|
+
* @param B - GpuMatrix (mutated in place)
|
|
82
|
+
* @param ldb - leading dimension of B (must equal B.lda)
|
|
83
|
+
* @returns no B — it stays GPU-resident; call `B.read()` yourself for a CPU readback (see the example)
|
|
84
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmm/strmm.mjs#L25">Source code: strmm.mjs (L25)</a>
|
|
85
|
+
* @category BLAS Level 3
|
|
86
|
+
*/
|
|
87
|
+
export declare function strmm(
|
|
88
|
+
device: GPUDevice,
|
|
89
|
+
side: 'left' | 'right',
|
|
90
|
+
uplo: 'lower' | 'upper',
|
|
91
|
+
transA: 'no-transpose' | 'transpose',
|
|
92
|
+
diag: 'unit' | 'non-unit',
|
|
93
|
+
m: number,
|
|
94
|
+
n: number,
|
|
95
|
+
alpha: number,
|
|
96
|
+
A: GpuMatrix,
|
|
97
|
+
lda: number,
|
|
98
|
+
B: GpuMatrix,
|
|
99
|
+
ldb: number,
|
|
100
|
+
): Promise<{ gpuTimeMs?: number }>;
|
|
@@ -0,0 +1,226 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
createStorageBuffer,
|
|
5
|
+
stageReadback,
|
|
6
|
+
destroyBuffers,
|
|
7
|
+
vec4ViewBinding,
|
|
8
|
+
} from "../util/buffer.mjs";
|
|
9
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
10
|
+
import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
|
|
11
|
+
import { extractResult } from "../util/result.mjs";
|
|
12
|
+
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
13
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
16
|
+
import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
|
|
17
|
+
import { TILE_WG_2D } from "../util/constants.mjs";
|
|
18
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
// strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
|
|
22
|
+
// triangular. Triangularize then sgemm, one command encoder. B is both
|
|
23
|
+
// input and output, so gemm writes to a fresh buffer (no aliasing race),
|
|
24
|
+
// copied back into B (GpuMatrix) or read back directly (Float32Array).
|
|
25
|
+
export async function strmm(
|
|
26
|
+
device, side, uplo, transA, diag, m, n, alpha, A, lda, B, ldb, layout = "row-major",
|
|
27
|
+
) {
|
|
28
|
+
const AIsGpu = A instanceof GpuMatrix;
|
|
29
|
+
const BIsGpu = B instanceof GpuMatrix;
|
|
30
|
+
const isUnit = diag === "unit";
|
|
31
|
+
|
|
32
|
+
if (!(device instanceof GPUDevice))
|
|
33
|
+
throw new Error("device must be a GPUDevice.");
|
|
34
|
+
requireSameDevice(device, "strmm", { A, B });
|
|
35
|
+
if (side !== "left" && side !== "right")
|
|
36
|
+
throw new Error("side must be 'left' or 'right'.");
|
|
37
|
+
if (uplo !== "lower" && uplo !== "upper")
|
|
38
|
+
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
39
|
+
if (transA !== "no-transpose" && transA !== "transpose")
|
|
40
|
+
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
41
|
+
if (!isUnit && diag !== "non-unit")
|
|
42
|
+
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
43
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
44
|
+
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
45
|
+
if (typeof alpha !== "number")
|
|
46
|
+
throw new Error("alpha must be a number.");
|
|
47
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
48
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
49
|
+
if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
|
|
50
|
+
throw new Error("m, n, lda, and ldb must be integers.");
|
|
51
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
52
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
53
|
+
if (!BIsGpu && !(B instanceof Float32Array))
|
|
54
|
+
throw new Error("B must be a Float32Array or GpuMatrix.");
|
|
55
|
+
if (AIsGpu !== BIsGpu)
|
|
56
|
+
throw new Error("A and B must both be GpuMatrix or both be Float32Array.");
|
|
57
|
+
if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
|
|
58
|
+
if (m === 0 || n === 0) return BIsGpu ? {} : { B };
|
|
59
|
+
|
|
60
|
+
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
61
|
+
const effLayoutB = BIsGpu ? B.layout : layout;
|
|
62
|
+
|
|
63
|
+
// A: triangular, order = m (side='left') or n (side='right').
|
|
64
|
+
const aOrder = side === "left" ? m : n;
|
|
65
|
+
if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
66
|
+
if (AIsGpu) {
|
|
67
|
+
if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
68
|
+
if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
|
|
69
|
+
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
70
|
+
throw new Error("A does not have enough elements for the given dimensions and lda.");
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
// B: always m x n, overwritten in place with the same ldb.
|
|
74
|
+
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
75
|
+
const bInner = effLayoutB === "column-major" ? m : n;
|
|
76
|
+
if (ldb < bInner)
|
|
77
|
+
throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
|
|
78
|
+
if (BIsGpu) {
|
|
79
|
+
if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
80
|
+
if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
|
|
81
|
+
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
82
|
+
throw new Error("B does not have enough elements for the given dimensions and ldb.");
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
86
|
+
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
87
|
+
const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
|
|
88
|
+
const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
|
|
89
|
+
|
|
90
|
+
const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
|
|
91
|
+
const transDense = "no-transpose"; // Adense already embodies op(A)
|
|
92
|
+
|
|
93
|
+
// X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
|
|
94
|
+
// output: compute (B_out)^T instead — sgemm's own trick, same as ssymm's.
|
|
95
|
+
let mg = m, ng = n;
|
|
96
|
+
const kg = aOrder;
|
|
97
|
+
let transX = side === "left" ? transDense : transB;
|
|
98
|
+
let transY = side === "left" ? transB : transDense;
|
|
99
|
+
const flip = (t) => (t === "no-transpose" ? "transpose" : "no-transpose");
|
|
100
|
+
let swapXY = side === "right";
|
|
101
|
+
if (effLayoutB === "column-major") {
|
|
102
|
+
[transX, transY] = [flip(transY), flip(transX)];
|
|
103
|
+
swapXY = !swapXY;
|
|
104
|
+
[mg, ng] = [ng, mg];
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
const ldDense = aOrder; // Adense is tightly packed, row-major
|
|
108
|
+
const largeWgX = Math.ceil(ng / BN_LARGE);
|
|
109
|
+
const largeWgY = Math.ceil(mg / BM_LARGE);
|
|
110
|
+
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
111
|
+
const gemmPipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
|
|
112
|
+
const triPipeline = await getPipeline(device, "triangularize");
|
|
113
|
+
const gemmWgCount = useLargeTile
|
|
114
|
+
? {
|
|
115
|
+
x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
|
|
116
|
+
y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
|
|
117
|
+
}
|
|
118
|
+
: {
|
|
119
|
+
x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strmm", "x"),
|
|
120
|
+
y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strmm", "y"),
|
|
121
|
+
};
|
|
122
|
+
|
|
123
|
+
// Null-init here and allocate inside the try below, so a throw partway
|
|
124
|
+
// through the sequence still reaches finally with every handle visible
|
|
125
|
+
// (strsv.mjs is the reference for this pattern).
|
|
126
|
+
let ABuffer = null, BBuffer = null;
|
|
127
|
+
let AdenseBuffer = null, outBuffer = null;
|
|
128
|
+
let triParams = null, gemmParams = null;
|
|
129
|
+
let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
|
|
130
|
+
|
|
131
|
+
try {
|
|
132
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
|
|
133
|
+
// readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
|
|
134
|
+
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
|
|
135
|
+
AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "strmm-Adense");
|
|
136
|
+
// COPY_DST: seeded from B's own content before gemm runs, so stride-padding
|
|
137
|
+
// gaps (never written by gemm's tight m x n loop) keep B's original bytes
|
|
138
|
+
// instead of reading back as zero. COPY_SRC: read back / adopted by B after.
|
|
139
|
+
outBuffer = createStorageBuffer(device,
|
|
140
|
+
bOuter * ldb * 4, "strmm-out", GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
|
|
141
|
+
);
|
|
142
|
+
|
|
143
|
+
triParams = createParamsBuffer(device,
|
|
144
|
+
[
|
|
145
|
+
{ value: aOrder, type: "u32" },
|
|
146
|
+
{ value: lda, type: "u32" },
|
|
147
|
+
{ value: ldDense, type: "u32" },
|
|
148
|
+
{ value: uploEffA === "upper" ? 1 : 0, type: "u32" },
|
|
149
|
+
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
150
|
+
{ value: isUnit ? 1 : 0, type: "u32" },
|
|
151
|
+
],
|
|
152
|
+
"strmm-tri-params",
|
|
153
|
+
);
|
|
154
|
+
const triBindGroup = createBindGroup(device, triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
|
|
155
|
+
|
|
156
|
+
// X/Y buffers and their own ld, matching swapXY above.
|
|
157
|
+
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
158
|
+
const ldX = swapXY ? ldb : ldDense;
|
|
159
|
+
const YBuffer = swapXY ? AdenseBuffer : BBuffer;
|
|
160
|
+
const ldY = swapXY ? ldDense : ldb;
|
|
161
|
+
|
|
162
|
+
gemmParams = createParamsBuffer(device,
|
|
163
|
+
[
|
|
164
|
+
{ value: mg, type: "u32" },
|
|
165
|
+
{ value: ng, type: "u32" },
|
|
166
|
+
{ value: kg, type: "u32" },
|
|
167
|
+
{ value: alpha, type: "f32" },
|
|
168
|
+
{ value: 0.0, type: "f32" }, // beta — strmm has no C accumulation term
|
|
169
|
+
{ value: ldX, type: "u32" },
|
|
170
|
+
{ value: ldY, type: "u32" },
|
|
171
|
+
{ value: ldb, type: "u32" },
|
|
172
|
+
{ value: transX === "transpose" ? 1 : 0, type: "u32" },
|
|
173
|
+
{ value: transY === "transpose" ? 1 : 0, type: "u32" },
|
|
174
|
+
],
|
|
175
|
+
"strmm-gemm-params",
|
|
176
|
+
);
|
|
177
|
+
const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
|
|
178
|
+
XBuffer,
|
|
179
|
+
vec4ViewBinding(device, XBuffer),
|
|
180
|
+
YBuffer,
|
|
181
|
+
vec4ViewBinding(device, YBuffer),
|
|
182
|
+
outBuffer,
|
|
183
|
+
gemmParams,
|
|
184
|
+
]);
|
|
185
|
+
|
|
186
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
187
|
+
// Seed outBuffer with B's own bytes first, so gemm's tight m x n write
|
|
188
|
+
// leaves stride-padding gaps holding B's original content, not zero.
|
|
189
|
+
// BBuffer may be larger than outBuffer (e.g. a validation-test baseline
|
|
190
|
+
// over-provisioned for a bigger ldb it might later be substituted with).
|
|
191
|
+
commandEncoder.copyBufferToBuffer(BBuffer, 0, outBuffer, 0, Math.min(BBuffer.size, outBuffer.size));
|
|
192
|
+
const triDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
193
|
+
const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
194
|
+
encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, triDesc);
|
|
195
|
+
encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
|
|
196
|
+
|
|
197
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
198
|
+
const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, outBuffer);
|
|
199
|
+
|
|
200
|
+
submit(device, commandEncoder);
|
|
201
|
+
|
|
202
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
203
|
+
|
|
204
|
+
if (BIsGpu) {
|
|
205
|
+
// Adopt outBuffer as B's own backing buffer instead of copying into the
|
|
206
|
+
// old one (B._buf has no COPY_DST usage) — cheaper and avoids needing
|
|
207
|
+
// an extra buffer-usage flag on every GpuMatrix for this one routine.
|
|
208
|
+
destroyBuffers(B._buf);
|
|
209
|
+
B._buf = outBuffer;
|
|
210
|
+
outBufferAdopted = true;
|
|
211
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
212
|
+
return {};
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
216
|
+
if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
|
|
217
|
+
return { B: result };
|
|
218
|
+
} finally {
|
|
219
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
220
|
+
if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
|
|
221
|
+
if (AdenseBuffer) destroyBuffers(AdenseBuffer);
|
|
222
|
+
if (outBuffer && !outBufferAdopted) destroyBuffers(outBuffer);
|
|
223
|
+
if (triParams) destroyBuffers(triParams);
|
|
224
|
+
if (gemmParams) destroyBuffers(gemmParams);
|
|
225
|
+
}
|
|
226
|
+
}
|
package/src/strmv/strmv.d.mts
CHANGED
|
@@ -44,41 +44,6 @@ export declare function strmv(
|
|
|
44
44
|
layout?: 'row-major' | 'column-major',
|
|
45
45
|
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
46
46
|
|
|
47
|
-
/**
|
|
48
|
-
* Performs the triangular matrix-vector operation y = op(A) * x
|
|
49
|
-
*
|
|
50
|
-
* A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
|
|
51
|
-
* (set at `GpuMatrix.from` time) determines the operation — there is no
|
|
52
|
-
* separate `layout` argument here.
|
|
53
|
-
*
|
|
54
|
-
* @param device - GPUDevice from `init()`
|
|
55
|
-
* @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
|
|
56
|
-
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
|
57
|
-
* @param diag - `'unit'` to treat the diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
58
|
-
* @param n - order of the matrix A
|
|
59
|
-
* @param A - GpuMatrix, GPU-resident
|
|
60
|
-
* @param lda - leading dimension of A (must equal A.lda)
|
|
61
|
-
* @param x - Float32Array input vector
|
|
62
|
-
* @param incx - stride for x (must be a positive integer)
|
|
63
|
-
* @param y - Float32Array output vector
|
|
64
|
-
* @param incy - stride for y (must be a positive integer)
|
|
65
|
-
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmv/strmv.mjs#L15">Source code: strmv.mjs (L15)</a>
|
|
66
|
-
* @category BLAS Level 2
|
|
67
|
-
*/
|
|
68
|
-
export declare function strmv(
|
|
69
|
-
device: GPUDevice,
|
|
70
|
-
uplo: 'lower' | 'upper',
|
|
71
|
-
trans: 'no-transpose' | 'transpose',
|
|
72
|
-
diag: 'unit' | 'non-unit',
|
|
73
|
-
n: number,
|
|
74
|
-
A: GpuMatrix,
|
|
75
|
-
lda: number,
|
|
76
|
-
x: Float32Array,
|
|
77
|
-
incx: number,
|
|
78
|
-
y: Float32Array,
|
|
79
|
-
incy: number,
|
|
80
|
-
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
81
|
-
|
|
82
47
|
/**
|
|
83
48
|
* Performs the triangular matrix-vector operation y = op(A) * x
|
|
84
49
|
*
|
|
@@ -86,7 +51,7 @@ export declare function strmv(
|
|
|
86
51
|
* `layout` (set at `GpuMatrix.from` time) determines the operation — there is
|
|
87
52
|
* no separate `layout` argument here.
|
|
88
53
|
*
|
|
89
|
-
* {@includeCode ../../examples/strmv/
|
|
54
|
+
* {@includeCode ../../examples/strmv/gpu.strmv.js}
|
|
90
55
|
*
|
|
91
56
|
* @param device - GPUDevice from `init()`
|
|
92
57
|
* @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
|
package/src/strmv/strmv.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, incy, layout = "row-major") {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -20,6 +21,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
20
21
|
|
|
21
22
|
if (!(device instanceof GPUDevice))
|
|
22
23
|
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
requireSameDevice(device, "strmv", { A, x, y });
|
|
23
25
|
if (uplo !== "lower" && uplo !== "upper")
|
|
24
26
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
25
27
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
@@ -54,6 +56,8 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
54
56
|
);
|
|
55
57
|
if (xIsGpu && !AIsGpu)
|
|
56
58
|
throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
|
|
59
|
+
if (AIsGpu && !xIsGpu)
|
|
60
|
+
throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
|
|
57
61
|
if (AIsGpu && yIsGpu && A._buf === y._buf)
|
|
58
62
|
throw new Error("A and y must not reference the same GPU buffer.");
|
|
59
63
|
if (AIsGpu && lda !== A.lda)
|
|
@@ -90,10 +94,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
90
94
|
let paramsBuffer = null;
|
|
91
95
|
|
|
92
96
|
try {
|
|
93
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strmv-A", false);
|
|
94
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strmv-x", false);
|
|
95
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "strmv-y", true);
|
|
96
|
-
paramsBuffer = createParamsBuffer(
|
|
97
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
|
|
98
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
|
|
99
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
|
|
100
|
+
paramsBuffer = createParamsBuffer(device,
|
|
97
101
|
[
|
|
98
102
|
{ value: n, type: "u32" },
|
|
99
103
|
{ value: incx, type: "u32" },
|
|
@@ -106,7 +110,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
106
110
|
"strmv-params",
|
|
107
111
|
);
|
|
108
112
|
|
|
109
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
113
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
110
114
|
ABuffer,
|
|
111
115
|
xBuffer,
|
|
112
116
|
yBuffer,
|
|
@@ -114,10 +118,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
114
118
|
]);
|
|
115
119
|
|
|
116
120
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
117
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
118
|
-
const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
121
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
122
|
+
const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
119
123
|
|
|
120
|
-
submit(commandEncoder);
|
|
124
|
+
submit(device, commandEncoder);
|
|
121
125
|
|
|
122
126
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
123
127
|
|
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Solves the triangular matrix equation
|
|
5
|
+
* op(A) * X = alpha * B (`side='left'`) or
|
|
6
|
+
* X * op(A) = alpha * B (`side='right'`), overwriting `B` with the
|
|
7
|
+
* solution `X` — `A` is triangular, only its `uplo` triangle stored; `B` is
|
|
8
|
+
* a general m×n matrix.
|
|
9
|
+
*
|
|
10
|
+
* - `side='left'`: `A` is m×m — solves `op(A)*X = alpha*B`
|
|
11
|
+
* - `side='right'`: `A` is n×n — solves `X*op(A) = alpha*B`
|
|
12
|
+
*
|
|
13
|
+
* Blocked substitution (strsv's own technique, generalized to a matrix
|
|
14
|
+
* RHS): strsv_invert_block + sgemm, both unmodified. A near-zero diagonal
|
|
15
|
+
* entry amplifies error, same as any triangular solve.
|
|
16
|
+
*
|
|
17
|
+
* {@includeCode ../../examples/strsm/strsm.js}
|
|
18
|
+
*
|
|
19
|
+
* **Browser (standalone HTML):**
|
|
20
|
+
* {@includeCode ../../examples/strsm/web/strsm.html}
|
|
21
|
+
*
|
|
22
|
+
* @param device - GPUDevice from `init()`
|
|
23
|
+
* @param side - `'left'` to solve op(A)*X=alpha*B, `'right'` to solve X*op(A)=alpha*B
|
|
24
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
25
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
26
|
+
* @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
27
|
+
* @param m - rows of B
|
|
28
|
+
* @param n - columns of B
|
|
29
|
+
* @param alpha - scalar multiplier for B
|
|
30
|
+
* @param A - Float32Array, triangular, row-major or column-major (see `layout`)
|
|
31
|
+
* @param lda - leading dimension of A as stored
|
|
32
|
+
* @param B - Float32Array input/output matrix, overwritten with the solution, row-major or column-major
|
|
33
|
+
* @param ldb - leading dimension of B as stored
|
|
34
|
+
* @param layout - storage layout shared by A/B when they're Float32Array
|
|
35
|
+
* (default: `'row-major'`) — column-major A is a genuine transpose (A isn't
|
|
36
|
+
* symmetric like ssymm's), so both `transA` and `uplo` are adjusted
|
|
37
|
+
* internally to compensate
|
|
38
|
+
* @returns the solution X, written into B, as a Float32Array
|
|
39
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsm/strsm.mjs#L25">Source code: strsm.mjs (L25)</a>
|
|
40
|
+
* @category BLAS Level 3
|
|
41
|
+
*/
|
|
42
|
+
export declare function strsm(
|
|
43
|
+
device: GPUDevice,
|
|
44
|
+
side: 'left' | 'right',
|
|
45
|
+
uplo: 'lower' | 'upper',
|
|
46
|
+
transA: 'no-transpose' | 'transpose',
|
|
47
|
+
diag: 'unit' | 'non-unit',
|
|
48
|
+
m: number,
|
|
49
|
+
n: number,
|
|
50
|
+
alpha: number,
|
|
51
|
+
A: Float32Array,
|
|
52
|
+
lda: number,
|
|
53
|
+
B: Float32Array,
|
|
54
|
+
ldb: number,
|
|
55
|
+
layout?: 'row-major' | 'column-major',
|
|
56
|
+
): Promise<{ B: Float32Array; gpuTimeMs?: number }>;
|
|
57
|
+
|
|
58
|
+
/**
|
|
59
|
+
* Solves the triangular matrix equation
|
|
60
|
+
* op(A) * X = alpha * B (`side='left'`) or
|
|
61
|
+
* X * op(A) = alpha * B (`side='right'`), overwriting `B` in place with `X`.
|
|
62
|
+
*
|
|
63
|
+
* A and B are both kept GPU-resident. Each matrix's own `layout` (set at
|
|
64
|
+
* `GpuMatrix.from` time) determines the operation — there is no separate
|
|
65
|
+
* `layout` argument here. A and B must both be GpuMatrix or both be
|
|
66
|
+
* Float32Array — mixing is not supported.
|
|
67
|
+
*
|
|
68
|
+
* {@includeCode ../../examples/strsm/gpu.strsm.js}
|
|
69
|
+
*
|
|
70
|
+
* @param device - GPUDevice from `init()`
|
|
71
|
+
* @param side - `'left'` to solve op(A)*X=alpha*B, `'right'` to solve X*op(A)=alpha*B
|
|
72
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
73
|
+
* @param transA - `'no-transpose'` for A, `'transpose'` for A^T
|
|
74
|
+
* @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
75
|
+
* @param m - rows of B
|
|
76
|
+
* @param n - columns of B
|
|
77
|
+
* @param alpha - scalar multiplier for B
|
|
78
|
+
* @param A - GpuMatrix, triangular
|
|
79
|
+
* @param lda - leading dimension of A (must equal A.lda)
|
|
80
|
+
* @param B - GpuMatrix (overwritten in place with the solution)
|
|
81
|
+
* @param ldb - leading dimension of B (must equal B.lda)
|
|
82
|
+
* @returns no B — it stays GPU-resident; call `B.read()` yourself for a CPU readback (see the example)
|
|
83
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsm/strsm.mjs#L25">Source code: strsm.mjs (L25)</a>
|
|
84
|
+
* @category BLAS Level 3
|
|
85
|
+
*/
|
|
86
|
+
export declare function strsm(
|
|
87
|
+
device: GPUDevice,
|
|
88
|
+
side: 'left' | 'right',
|
|
89
|
+
uplo: 'lower' | 'upper',
|
|
90
|
+
transA: 'no-transpose' | 'transpose',
|
|
91
|
+
diag: 'unit' | 'non-unit',
|
|
92
|
+
m: number,
|
|
93
|
+
n: number,
|
|
94
|
+
alpha: number,
|
|
95
|
+
A: GpuMatrix,
|
|
96
|
+
lda: number,
|
|
97
|
+
B: GpuMatrix,
|
|
98
|
+
ldb: number,
|
|
99
|
+
): Promise<{ gpuTimeMs?: number }>;
|