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.
Files changed (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { calcWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function saxpy(device, n, alpha, x, incx, y, incy) {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -18,6 +19,7 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
18
19
 
19
20
  if (!(device instanceof GPUDevice))
20
21
  throw new Error("device must be a GPUDevice.");
22
+ requireSameDevice(device, "saxpy", { x, y });
21
23
  if (
22
24
  !Number.isInteger(n) ||
23
25
  !Number.isInteger(incx) ||
@@ -56,9 +58,9 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
56
58
  let readBuffer = null;
57
59
 
58
60
  try {
59
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
60
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
61
- paramsBuffer = createParamsBuffer(
61
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "saxpy-x", false);
62
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "saxpy-y", true);
63
+ paramsBuffer = createParamsBuffer(device,
62
64
  [
63
65
  { value: n, type: "u32" },
64
66
  { value: alpha, type: "f32" },
@@ -68,19 +70,19 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
68
70
  "saxpy-params",
69
71
  );
70
72
 
71
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
73
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
72
74
  xBuffer,
73
75
  yBuffer,
74
76
  paramsBuffer,
75
77
  ]);
76
- const { commandEncoder, ts } = runComputePass(
78
+ const { commandEncoder, ts } = runComputePass(device,
77
79
  pipeline,
78
80
  bindGroup,
79
- calcWorkgroups(n),
81
+ calcWorkgroups(device, n),
80
82
  );
81
- readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
83
+ readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
82
84
 
83
- submit(commandEncoder);
85
+ submit(device, commandEncoder);
84
86
 
85
87
  const gpuTimeMs = await extractTimestamp(ts);
86
88
 
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Performs the operation y = x
4
+ * Performs the operation $$y \leftarrow x$$
5
5
  *
6
6
  * {@includeCode ../../examples/scopy/scopy.js}
7
7
  *
@@ -27,9 +27,9 @@ export declare function scopy(
27
27
  ): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
28
28
 
29
29
  /**
30
- * Performs the operation y = x
30
+ * Performs the operation $$y \leftarrow x$$
31
31
  *
32
- * {@includeCode ../../examples/scopy/gpuvec.scopy.js}
32
+ * {@includeCode ../../examples/scopy/gpu.scopy.js}
33
33
  *
34
34
  * @param device - GPUDevice from `init()`
35
35
  * @param n - number of elements (must be a positive integer)
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { calcWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function scopy(device, n, x, incx, y, incy) {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -18,6 +19,7 @@ export async function scopy(device, n, x, incx, y, incy) {
18
19
 
19
20
  if (!(device instanceof GPUDevice))
20
21
  throw new Error("device must be a GPUDevice.");
22
+ requireSameDevice(device, "scopy", { x, y });
21
23
  if (
22
24
  !Number.isInteger(n) ||
23
25
  !Number.isInteger(incx) ||
@@ -52,9 +54,9 @@ export async function scopy(device, n, x, incx, y, incy) {
52
54
  let readBuffer = null;
53
55
 
54
56
  try {
55
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
56
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
57
- paramsBuffer = createParamsBuffer(
57
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "scopy-x", false);
58
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "scopy-y", true);
59
+ paramsBuffer = createParamsBuffer(device,
58
60
  [
59
61
  { value: n, type: "u32" },
60
62
  { value: incx, type: "u32" },
@@ -63,19 +65,19 @@ export async function scopy(device, n, x, incx, y, incy) {
63
65
  "scopy-params",
64
66
  );
65
67
 
66
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
68
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
67
69
  xBuffer,
68
70
  yBuffer,
69
71
  paramsBuffer,
70
72
  ]);
71
- const { commandEncoder, ts } = runComputePass(
73
+ const { commandEncoder, ts } = runComputePass(device,
72
74
  pipeline,
73
75
  bindGroup,
74
- calcWorkgroups(n),
76
+ calcWorkgroups(device, n),
75
77
  );
76
- readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
78
+ readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
77
79
 
78
- submit(commandEncoder);
80
+ submit(device, commandEncoder);
79
81
 
80
82
  const gpuTimeMs = await extractTimestamp(ts);
81
83
 
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Computes the dot product of two vectors: result = sum(x[i] * y[i])
4
+ * Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
5
5
  *
6
6
  * {@includeCode ../../examples/sdot/sdot.js}
7
7
  *
@@ -28,9 +28,9 @@ export declare function sdot(
28
28
  ): Promise<{ dot: number } | { dot: number; gpuTimeMs: number }>;
29
29
 
30
30
  /**
31
- * Computes the dot product of two vectors: result = sum(x[i] * y[i])
31
+ * Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
32
32
  *
33
- * {@includeCode ../../examples/sdot/gpuvec.sdot.js}
33
+ * {@includeCode ../../examples/sdot/gpu.sdot.js}
34
34
  *
35
35
  * @param device - GPUDevice from `init()`
36
36
  * @param n - number of elements (must be a positive integer)
package/src/sdot/sdot.mjs CHANGED
@@ -12,8 +12,9 @@ import { extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
+ import { WGS } from "../util/constants.mjs";
16
+ import { requireSameDevice } from "../util/device.mjs";
15
17
 
16
- const WGS = 64; //workgroup size
17
18
 
18
19
  export async function sdot(device, n, x, incx, y, incy) {
19
20
  const xIsGpu = x instanceof GpuVector;
@@ -21,6 +22,7 @@ export async function sdot(device, n, x, incx, y, incy) {
21
22
 
22
23
  if (!(device instanceof GPUDevice))
23
24
  throw new Error("device must be a GPUDevice.");
25
+ requireSameDevice(device, "sdot", { x, y });
24
26
  if (
25
27
  !Number.isInteger(n) ||
26
28
  !Number.isInteger(incx) ||
@@ -58,11 +60,11 @@ export async function sdot(device, n, x, incx, y, incy) {
58
60
  let readBuffer = null;
59
61
 
60
62
  try {
61
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
62
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
63
- partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
64
- resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
65
- paramsBuffer = createParamsBuffer(
63
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sdot-x", false);
64
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sdot-y", false);
65
+ partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
66
+ resultBuffer = createResultBuffer(device, 4, "sdot-result"); //to hold the final float32 dot product
67
+ paramsBuffer = createParamsBuffer(device,
66
68
  [
67
69
  { value: n, type: "u32" },
68
70
  { value: incx, type: "u32" },
@@ -71,32 +73,32 @@ export async function sdot(device, n, x, incx, y, incy) {
71
73
  "sdot-params",
72
74
  );
73
75
 
74
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
76
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
75
77
  xBuffer,
76
78
  yBuffer,
77
79
  partialsBuffer,
78
80
  paramsBuffer,
79
81
  ]);
80
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
82
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
81
83
  pipelineMain,
82
84
  bgMain,
83
85
  2 * WGS,
84
86
  ); //dispatch 2*WGS workgroups
85
87
 
86
- submit(enc1);
88
+ submit(device, enc1);
87
89
 
88
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
90
+ const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
89
91
  partialsBuffer,
90
92
  resultBuffer,
91
93
  ]);
92
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
94
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
93
95
  pipelineReduce,
94
96
  bgReduce,
95
97
  1,
96
98
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
97
- readBuffer = stageReadback(enc2, resultBuffer);
99
+ readBuffer = stageReadback(device, enc2, resultBuffer);
98
100
 
99
- submit(enc2);
101
+ submit(device, enc2);
100
102
 
101
103
  const resultPromise = extractResult(readBuffer, Float32Array);
102
104
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -117,7 +119,7 @@ export async function sdot(device, n, x, incx, y, incy) {
117
119
  if (partialsBuffer) destroyBuffers(partialsBuffer);
118
120
  if (resultBuffer) destroyBuffers(resultBuffer);
119
121
  if (paramsBuffer) destroyBuffers(paramsBuffer);
120
- // Only reached if submit(enc2) threw before ownership was transferred above.
122
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
121
123
  if (readBuffer) destroyBuffers(readBuffer);
122
124
  }
123
125
  }
@@ -0,0 +1,102 @@
1
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
+
3
+ /**
4
+ * Performs the matrix-matrix operation C = alpha * op(A) * op(B) + beta * C
5
+ *
6
+ * - `transA/transB='no-transpose'`: op(A) = A (m×k), op(B) = B (k×n)
7
+ * - `transA/transB='transpose'`: op(A) = A^T, op(B) = B^T
8
+ *
9
+ * A, B, C are row-major or column-major (see `layout`) — backed by one of
10
+ * two shared-memory-tiled, register-blocked kernels chosen by shape:
11
+ * `sgemm_small.wgsl` (BM=BN=32) below a 6x6 workgroup grid, `sgemm_large.wgsl`
12
+ * (BM=BN=64) above it.
13
+ *
14
+ * {@includeCode ../../examples/sgemm/sgemm.js}
15
+ *
16
+ * **Browser (standalone HTML):**
17
+ * {@includeCode ../../examples/sgemm/web/sgemm.html}
18
+ *
19
+ * @param device - GPUDevice from `init()`
20
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
21
+ * @param transB - `'no-transpose'` for B, `'transpose'` for B^T
22
+ * @param m - rows of op(A) and C
23
+ * @param n - columns of op(B) and C
24
+ * @param k - columns of op(A), rows of op(B)
25
+ * @param alpha - scalar multiplier for op(A)*op(B)
26
+ * @param A - Float32Array, row-major or column-major (see `layout`)
27
+ * @param lda - leading dimension of A as stored
28
+ * @param B - Float32Array, row-major or column-major (see `layout`)
29
+ * @param ldb - leading dimension of B as stored
30
+ * @param beta - scalar multiplier for C
31
+ * @param C - Float32Array input/output matrix, row-major or column-major
32
+ * @param ldc - leading dimension of C as stored
33
+ * @param layout - storage layout shared by A/B/C when they're Float32Array
34
+ * (default: `'row-major'`); column-major A/B flips the respective trans
35
+ * flag internally, column-major C computes C^T = op(B)^T*op(A)^T instead
36
+ * (same underlying bytes) — op(A)*op(B) stays what you asked for either way
37
+ * @returns updated C as a Float32Array
38
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemm/sgemm.mjs#L17">Source code: sgemm.mjs (L17)</a>
39
+ * @category BLAS Level 3
40
+ */
41
+ export declare function sgemm(
42
+ device: GPUDevice,
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 = 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/sgemm/gpu.sgemm.js}
68
+ *
69
+ * @param device - GPUDevice from `init()`
70
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
71
+ * @param transB - `'no-transpose'` for B, `'transpose'` for B^T
72
+ * @param m - rows of op(A) and C
73
+ * @param n - columns of op(B) and C
74
+ * @param k - columns of op(A), rows of op(B)
75
+ * @param alpha - scalar multiplier for op(A)*op(B)
76
+ * @param A - GpuMatrix
77
+ * @param lda - leading dimension of A (must equal A.lda)
78
+ * @param B - GpuMatrix
79
+ * @param ldb - leading dimension of B (must equal B.lda)
80
+ * @param beta - scalar multiplier for C
81
+ * @param C - GpuMatrix (mutated in place)
82
+ * @param ldc - leading dimension of C (must equal C.lda)
83
+ * @returns no C — it stays GPU-resident; call `C.read()` yourself for a CPU readback (see the example)
84
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemm/sgemm.mjs#L17">Source code: sgemm.mjs (L17)</a>
85
+ * @category BLAS Level 3
86
+ */
87
+ export declare function sgemm(
88
+ device: GPUDevice,
89
+ transA: 'no-transpose' | 'transpose',
90
+ transB: 'no-transpose' | 'transpose',
91
+ m: number,
92
+ n: number,
93
+ k: number,
94
+ alpha: number,
95
+ A: GpuMatrix,
96
+ lda: number,
97
+ B: GpuMatrix,
98
+ ldb: number,
99
+ beta: number,
100
+ C: GpuMatrix,
101
+ ldc: number,
102
+ ): Promise<{ gpuTimeMs?: number }>;
@@ -0,0 +1,208 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ vec4ViewBinding,
7
+ vec4Usable,
8
+ } from "../util/buffer.mjs";
9
+ import { createBindGroup } from "../util/bindgroup.mjs";
10
+ import { runComputePass, submit } from "../util/compute.mjs";
11
+ import { extractResult } from "../util/result.mjs";
12
+ import { 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 { requireSameDevice } from "../util/device.mjs";
18
+
19
+
20
+ export async function sgemm(
21
+ device, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
22
+ ) {
23
+ let AIsGpu = A instanceof GpuMatrix;
24
+ let BIsGpu = B instanceof GpuMatrix;
25
+ const CIsGpu = C instanceof GpuMatrix;
26
+
27
+ if (!(device instanceof GPUDevice))
28
+ throw new Error("device must be a GPUDevice.");
29
+ requireSameDevice(device, "sgemm", { A, B, C });
30
+ if (transA !== "no-transpose" && transA !== "transpose")
31
+ throw new Error("transA must be 'no-transpose' or 'transpose'.");
32
+ if (transB !== "no-transpose" && transB !== "transpose")
33
+ throw new Error("transB must be 'no-transpose' or 'transpose'.");
34
+ if (layout !== "row-major" && layout !== "column-major")
35
+ throw new Error("layout must be 'row-major' or 'column-major'.");
36
+ if (typeof alpha !== "number")
37
+ throw new Error("alpha must be a number.");
38
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
40
+ if (typeof beta !== "number")
41
+ throw new Error("beta must be a number.");
42
+ if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
+ if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
+ if (
45
+ !Number.isInteger(m) ||
46
+ !Number.isInteger(n) ||
47
+ !Number.isInteger(k) ||
48
+ !Number.isInteger(lda) ||
49
+ !Number.isInteger(ldb) ||
50
+ !Number.isInteger(ldc)
51
+ )
52
+ throw new Error("m, n, k, lda, ldb, and ldc must be integers.");
53
+ if (!AIsGpu && !(A instanceof Float32Array))
54
+ throw new Error("A must be a Float32Array or GpuMatrix.");
55
+ if (!BIsGpu && !(B instanceof Float32Array))
56
+ throw new Error("B must be a Float32Array or GpuMatrix.");
57
+ if (!CIsGpu && !(C instanceof Float32Array))
58
+ throw new Error("C must be a Float32Array or GpuMatrix.");
59
+ if ((AIsGpu || BIsGpu) && !CIsGpu)
60
+ throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
61
+ if (CIsGpu && (!AIsGpu || !BIsGpu))
62
+ throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
63
+ if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
64
+ if (m === 0 || n === 0) return CIsGpu ? {} : { C };
65
+
66
+ // GpuMatrix's own .layout wins over the shared `layout` argument.
67
+ const effLayoutA = AIsGpu ? A.layout : layout;
68
+ const effLayoutB = BIsGpu ? B.layout : layout;
69
+ const effLayoutC = CIsGpu ? C.layout : layout;
70
+
71
+ // Shape validation, before any layout-driven swapping below.
72
+ // A: op(A) is m x k; A itself is m x k or k x m depending on transA.
73
+ const aRows = effLayoutA === "column-major" ? k : m;
74
+ const aCols = effLayoutA === "column-major" ? m : k;
75
+ const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
76
+ const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
77
+ if (lda < aInner)
78
+ throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
79
+ if (AIsGpu) {
80
+ if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
81
+ const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
82
+ if (A.rows < aLogRows || A.cols < aLogCols)
83
+ throw new Error("A is too small for the given m, k, and transA.");
84
+ } else if (A.length < (aOuter - 1) * lda + aInner) {
85
+ throw new Error("A does not have enough elements for the given dimensions and lda.");
86
+ }
87
+
88
+ // B: same reasoning as A, with op(B) = k x n.
89
+ const bRows = effLayoutB === "column-major" ? n : k;
90
+ const bCols = effLayoutB === "column-major" ? k : n;
91
+ const bOuter = transB === "no-transpose" ? bRows : bCols;
92
+ const bInner = transB === "no-transpose" ? bCols : bRows;
93
+ if (ldb < bInner)
94
+ throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
95
+ if (BIsGpu) {
96
+ if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
97
+ const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
98
+ if (B.rows < bLogRows || B.cols < bLogCols)
99
+ throw new Error("B is too small for the given n, k, and transB.");
100
+ } else if (B.length < (bOuter - 1) * ldb + bInner) {
101
+ throw new Error("B does not have enough elements for the given dimensions and ldb.");
102
+ }
103
+
104
+ // C: always m x n (no trans flag) — layout only affects lda/storage order.
105
+ const cOuter = effLayoutC === "column-major" ? n : m;
106
+ const cInner = effLayoutC === "column-major" ? m : n;
107
+ if (ldc < cInner)
108
+ throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
109
+ if (CIsGpu) {
110
+ if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
111
+ if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
112
+ } else if (C.length < (cOuter - 1) * ldc + cInner) {
113
+ throw new Error("C does not have enough elements for the given dimensions and ldc.");
114
+ }
115
+
116
+ // Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
117
+ if (effLayoutA === "column-major")
118
+ transA = transA === "no-transpose" ? "transpose" : "no-transpose";
119
+ if (effLayoutB === "column-major")
120
+ transB = transB === "no-transpose" ? "transpose" : "no-transpose";
121
+
122
+ // Column-major C: compute C^T = op(B)^T * op(A)^T instead (swap A/B, flip trans, swap m<->n).
123
+ if (effLayoutC === "column-major") {
124
+ [A, B] = [B, A];
125
+ [AIsGpu, BIsGpu] = [BIsGpu, AIsGpu];
126
+ [lda, ldb] = [ldb, lda];
127
+ [transA, transB] = [
128
+ transB === "no-transpose" ? "transpose" : "no-transpose",
129
+ transA === "no-transpose" ? "transpose" : "no-transpose",
130
+ ];
131
+ [m, n] = [n, m];
132
+ }
133
+
134
+ // Shape-based auto-select — see sgemm_small.wgsl/sgemm_large.wgsl.
135
+ const largeWgX = Math.ceil(n / BN_LARGE);
136
+ const largeWgY = Math.ceil(m / BM_LARGE);
137
+ const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
138
+
139
+ const pipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
140
+
141
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemm-A", false);
142
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemm-B", false);
143
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemm-C", true);
144
+ // Vectorized-load enablement — kernel-side view after the column-major swap.
145
+ // op(A) is m×k (no-trans) or k×m (trans); op(B) is k×n or n×k.
146
+ const aNot = transA === "no-transpose";
147
+ const bNot = transB === "no-transpose";
148
+ const useVecA = aNot && vec4Usable(ABuffer, lda, m, k); // transposed-A vec stores bank-conflict smem — measured slower than scalar
149
+ const useVecB = vec4Usable(BBuffer, ldb, bNot ? k : n, bNot ? n : k);
150
+ const paramsBuffer = createParamsBuffer(device,
151
+ [
152
+ { value: m, type: "u32" },
153
+ { value: n, type: "u32" },
154
+ { value: k, type: "u32" },
155
+ { value: alpha, type: "f32" },
156
+ { value: beta, type: "f32" },
157
+ { value: lda, type: "u32" },
158
+ { value: ldb, type: "u32" },
159
+ { value: ldc, type: "u32" },
160
+ { value: transA === "transpose" ? 1 : 0, type: "u32" },
161
+ { value: transB === "transpose" ? 1 : 0, type: "u32" },
162
+ { value: useVecA ? 1 : 0, type: "u32" },
163
+ { value: useVecB ? 1 : 0, type: "u32" },
164
+ ],
165
+ "sgemm-params",
166
+ );
167
+
168
+ try {
169
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
170
+ ABuffer,
171
+ vec4ViewBinding(device, ABuffer),
172
+ BBuffer,
173
+ vec4ViewBinding(device, BBuffer),
174
+ CBuffer,
175
+ paramsBuffer,
176
+ ]);
177
+
178
+ const wgCount = useLargeTile
179
+ ? {
180
+ x: requireWorkgroupCount(device, largeWgX, "sgemm", "x"),
181
+ y: requireWorkgroupCount(device, largeWgY, "sgemm", "y"),
182
+ }
183
+ : {
184
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemm", "x"),
185
+ y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemm", "y"),
186
+ };
187
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
188
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
189
+
190
+ submit(device, commandEncoder);
191
+
192
+ const gpuTimeMs = await extractTimestamp(ts);
193
+
194
+ if (CIsGpu) {
195
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
196
+ return {};
197
+ }
198
+
199
+ const result = await extractResult(readBuffer, Float32Array);
200
+ if (gpuTimeMs !== undefined) return { C: result, gpuTimeMs };
201
+ return { C: result };
202
+ } finally {
203
+ if (!AIsGpu) destroyBuffers(ABuffer);
204
+ if (!BIsGpu) destroyBuffers(BBuffer);
205
+ if (!CIsGpu) destroyBuffers(CBuffer);
206
+ destroyBuffers(paramsBuffer);
207
+ }
208
+ }
@@ -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 }>;