wgblas 2.0.0 → 2.2.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 (124) hide show
  1. package/README.md +20 -18
  2. package/dist/wgblas.browser.js +2172 -1174
  3. package/index.d.mts +49 -44
  4. package/index.mjs +11 -0
  5. package/package.json +133 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +126 -17
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +66 -11
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.d.mts +4 -4
  17. package/src/dasum/dasum.mjs +38 -20
  18. package/src/daxpy/daxpy.d.mts +56 -0
  19. package/src/daxpy/daxpy.mjs +150 -0
  20. package/src/dcopy/dcopy.d.mts +52 -0
  21. package/src/dcopy/dcopy.mjs +140 -0
  22. package/src/ddot/ddot.d.mts +62 -0
  23. package/src/ddot/ddot.mjs +184 -0
  24. package/src/devdocs.mjs +13 -0
  25. package/src/dnrm2/dnrm2.d.mts +50 -0
  26. package/src/dnrm2/dnrm2.mjs +189 -0
  27. package/src/drot/drot.d.mts +67 -0
  28. package/src/drot/drot.mjs +170 -0
  29. package/src/drotm/drotm.d.mts +67 -0
  30. package/src/drotm/drotm.mjs +171 -0
  31. package/src/dscal/dscal.d.mts +52 -0
  32. package/src/dscal/dscal.mjs +119 -0
  33. package/src/dswap/dswap.d.mts +57 -0
  34. package/src/dswap/dswap.mjs +155 -0
  35. package/src/idamax/idamax.d.mts +20 -2
  36. package/src/idamax/idamax.mjs +56 -24
  37. package/src/init.mjs +117 -56
  38. package/src/isamax/isamax.d.mts +20 -2
  39. package/src/isamax/isamax.mjs +21 -16
  40. package/src/random/random.d.mts +37 -39
  41. package/src/random/random.mjs +39 -7
  42. package/src/sasum/sasum.d.mts +2 -2
  43. package/src/sasum/sasum.mjs +20 -16
  44. package/src/saxpy/saxpy.d.mts +2 -2
  45. package/src/saxpy/saxpy.mjs +14 -11
  46. package/src/scopy/scopy.d.mts +2 -2
  47. package/src/scopy/scopy.mjs +13 -9
  48. package/src/sdot/sdot.d.mts +2 -2
  49. package/src/sdot/sdot.mjs +21 -17
  50. package/src/sgemm/sgemm.d.mts +2 -2
  51. package/src/sgemm/sgemm.mjs +109 -40
  52. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  53. package/src/sgemmtr/sgemmtr.mjs +98 -40
  54. package/src/sgemv/sgemv.d.mts +2 -2
  55. package/src/sgemv/sgemv.mjs +69 -41
  56. package/src/sger/sger.d.mts +2 -2
  57. package/src/sger/sger.mjs +43 -19
  58. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  59. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  60. package/src/shaders/cscal.wgsl +33 -0
  61. package/src/shaders/daxpy.wgsl +66 -0
  62. package/src/shaders/dcopy.wgsl +34 -0
  63. package/src/shaders/ddot.wgsl +106 -0
  64. package/src/shaders/dnrm2.wgsl +167 -0
  65. package/src/shaders/drot.wgsl +81 -0
  66. package/src/shaders/drotm.wgsl +99 -0
  67. package/src/shaders/dscal.wgsl +60 -0
  68. package/src/shaders/dswap.wgsl +38 -0
  69. package/src/shaders/f64/utils/add.wgsl +6 -0
  70. package/src/shaders/f64/utils/divide.wgsl +45 -0
  71. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  72. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  73. package/src/shaders/index.mjs +233 -14
  74. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  75. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  76. package/src/shaders/sgemm_large.wgsl +107 -18
  77. package/src/shaders/sgemm_small.wgsl +115 -15
  78. package/src/shaders/sgemmtr_large.wgsl +4 -1
  79. package/src/shaders/sgemmtr_small.wgsl +4 -1
  80. package/src/shaders/sgemv_n.wgsl +3 -1
  81. package/src/shaders/sgemv_t.wgsl +3 -1
  82. package/src/shaders/snrm2.wgsl +72 -23
  83. package/src/shaders/ssymv.wgsl +3 -1
  84. package/src/snrm2/snrm2.d.mts +2 -2
  85. package/src/snrm2/snrm2.mjs +41 -23
  86. package/src/srot/srot.d.mts +2 -4
  87. package/src/srot/srot.mjs +16 -11
  88. package/src/srotm/srotm.d.mts +2 -4
  89. package/src/srotm/srotm.mjs +17 -11
  90. package/src/sscal/sscal.d.mts +3 -3
  91. package/src/sscal/sscal.mjs +14 -12
  92. package/src/sswap/sswap.d.mts +2 -2
  93. package/src/sswap/sswap.mjs +18 -10
  94. package/src/ssymm/ssymm.d.mts +5 -4
  95. package/src/ssymm/ssymm.mjs +150 -54
  96. package/src/ssymv/ssymv.d.mts +2 -2
  97. package/src/ssymv/ssymv.mjs +47 -26
  98. package/src/ssyr/ssyr.d.mts +2 -2
  99. package/src/ssyr/ssyr.mjs +38 -17
  100. package/src/ssyr2/ssyr2.d.mts +2 -2
  101. package/src/ssyr2/ssyr2.mjs +48 -21
  102. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  103. package/src/ssyr2k/ssyr2k.mjs +140 -62
  104. package/src/ssyrk/ssyrk.d.mts +3 -2
  105. package/src/ssyrk/ssyrk.mjs +91 -39
  106. package/src/strmm/strmm.d.mts +5 -4
  107. package/src/strmm/strmm.mjs +174 -60
  108. package/src/strmv/strmv.d.mts +2 -2
  109. package/src/strmv/strmv.mjs +42 -20
  110. package/src/strsm/strsm.d.mts +6 -4
  111. package/src/strsm/strsm.mjs +438 -174
  112. package/src/strsv/strsv.d.mts +5 -3
  113. package/src/strsv/strsv.mjs +89 -34
  114. package/src/util/benchmark.mjs +9 -9
  115. package/src/util/bindgroup.mjs +1 -3
  116. package/src/util/buffer.mjs +139 -24
  117. package/src/util/complex.mjs +87 -0
  118. package/src/util/compute.mjs +19 -16
  119. package/src/util/constants.mjs +57 -0
  120. package/src/util/device.mjs +49 -0
  121. package/src/util/pipeline.mjs +44 -10
  122. package/src/util/workgroup.mjs +72 -7
  123. package/src/shaders/browser-shaders.mjs +0 -81
  124. package/src/shaders/f64add.wgsl +0 -281
@@ -11,13 +11,14 @@ import { extractResult } from "../util/result.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 { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function srotm(device, n, x, incx, y, incy, param) {
16
17
  const xIsGpu = x instanceof GpuVector;
17
18
  const yIsGpu = y instanceof GpuVector;
18
19
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
21
+ requireSameDevice(device, "srotm", { x, y });
21
22
  if (
22
23
  !Number.isInteger(n) ||
23
24
  !Number.isInteger(incx) ||
@@ -26,6 +27,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
26
27
  throw new Error("n, incx, and incy must be integers.");
27
28
  if (!(param instanceof Float32Array) || param.length !== 5)
28
29
  throw new Error("param must be a Float32Array of length 5.");
30
+ if (param[0] !== -2 && param[0] !== -1 && param[0] !== 0 && param[0] !== 1)
31
+ throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
29
32
  if (incx <= 0 || incy <= 0)
30
33
  throw new Error("incx and incy must be positive.");
31
34
  if (!xIsGpu && !(x instanceof Float32Array))
@@ -56,10 +59,11 @@ export async function srotm(device, n, x, incx, y, incy, param) {
56
59
  let readY = null;
57
60
 
58
61
  try {
59
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
60
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
61
- paramBuffer = uploadBuffer(param, "srotm-param", false);
62
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
63
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
64
+ paramBuffer = uploadBuffer(device, param, "srotm-param", false);
62
65
  paramsBuffer = createParamsBuffer(
66
+ device,
63
67
  [
64
68
  { value: n, type: "u32" },
65
69
  { value: incx, type: "u32" },
@@ -68,24 +72,26 @@ export async function srotm(device, n, x, incx, y, incy, param) {
68
72
  "srotm-params",
69
73
  );
70
74
 
71
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
75
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
72
76
  xBuffer,
73
77
  yBuffer,
74
78
  paramBuffer,
75
79
  paramsBuffer,
76
80
  ]);
77
81
  const { commandEncoder, ts } = runComputePass(
82
+ device,
78
83
  pipeline,
79
84
  bindGroup,
80
- calcWorkgroups(n),
85
+ calcWorkgroups(device, n),
81
86
  );
82
- readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
83
- readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
84
- submit(commandEncoder);
87
+ readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
88
+ readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
89
+ submit(device, commandEncoder);
85
90
 
86
91
  const gpuTimeMs = await extractTimestamp(ts);
87
92
 
88
- if (xIsGpu && yIsGpu) {
93
+ if (xIsGpu) {
94
+ // xIsGpu === yIsGpu, enforced above
89
95
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
90
96
  return {};
91
97
  }
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Scales a single-precision vector by a constant: x = alpha * x
4
+ * Scales a single-precision vector by a constant: $$x \leftarrow \alpha x$$
5
5
  *
6
6
  * {@includeCode ../../examples/sscal/sscal.js}
7
7
  *
@@ -22,10 +22,10 @@ export declare function sscal(
22
22
  alpha: number,
23
23
  x: Float32Array,
24
24
  incx: number,
25
- ): Promise<Float32Array | { result: Float32Array; gpuTimeMs: number }>;
25
+ ): Promise<{ x: Float32Array } | { x: Float32Array; gpuTimeMs: number }>;
26
26
 
27
27
  /**
28
- * Scales a single-precision vector by a constant: x = alpha * x
28
+ * Scales a single-precision vector by a constant: $$x \leftarrow \alpha x$$
29
29
  *
30
30
  * {@includeCode ../../examples/sscal/gpu.sscal.js}
31
31
  *
@@ -11,22 +11,22 @@ 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 { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function sscal(device, n, alpha, x, incx) {
16
17
  const xIsGpu = x instanceof GpuVector;
17
-
18
- if (!(device instanceof GPUDevice))
19
- throw new Error("device must be a GPUDevice.");
18
+
19
+ requireGpuDevice(device);
20
+ requireSameDevice(device, "sscal", { x });
20
21
  if (!Number.isInteger(n) || !Number.isInteger(incx))
21
22
  throw new Error("n and incx must be integers.");
22
- if (typeof alpha !== "number")
23
- throw new Error("alpha must be a number.");
23
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
24
24
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
25
25
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
26
26
  if (incx <= 0) throw new Error("incx must be positive.");
27
27
  if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
28
28
  throw new Error("x must be a Float32Array or GpuVector.");
29
- if (n <= 0) return xIsGpu ? {} : x;
29
+ if (n <= 0) return xIsGpu ? {} : { x };
30
30
  if (x.length < (n - 1) * incx + 1)
31
31
  throw new Error(
32
32
  "x does not have enough elements for the given n and incx.",
@@ -39,8 +39,9 @@ export async function sscal(device, n, alpha, x, incx) {
39
39
  let readBuffer = null;
40
40
 
41
41
  try {
42
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
42
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
43
43
  paramsBuffer = createParamsBuffer(
44
+ device,
44
45
  [
45
46
  { value: n, type: "u32" },
46
47
  { value: alpha, type: "f32" },
@@ -49,18 +50,19 @@ export async function sscal(device, n, alpha, x, incx) {
49
50
  "sscal-params",
50
51
  );
51
52
 
52
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
53
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
53
54
  xBuffer,
54
55
  paramsBuffer,
55
56
  ]);
56
57
  const { commandEncoder, ts } = runComputePass(
58
+ device,
57
59
  pipeline,
58
60
  bindGroup,
59
- calcWorkgroups(n),
61
+ calcWorkgroups(device, n),
60
62
  );
61
- readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
63
+ readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
62
64
 
63
- submit(commandEncoder);
65
+ submit(device, commandEncoder);
64
66
 
65
67
  const gpuTimeMs = await extractTimestamp(ts);
66
68
 
@@ -72,7 +74,7 @@ export async function sscal(device, n, alpha, x, incx) {
72
74
  const result = await extractResult(readBuffer, Float32Array);
73
75
  readBuffer = null; // extractResult already destroyed it
74
76
  if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
75
- return result;
77
+ return { x: result };
76
78
  } finally {
77
79
  if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
78
80
  if (paramsBuffer) destroyBuffers(paramsBuffer);
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Swaps the elements of two single-precision vectors: x <-> y
4
+ * Swaps the elements of two single-precision vectors: $$x \leftrightarrow y$$
5
5
  *
6
6
  * {@includeCode ../../examples/sswap/sswap.js}
7
7
  *
@@ -27,7 +27,7 @@ export declare function sswap(
27
27
  ): Promise<{ x: Float32Array; y: Float32Array } | { x: Float32Array; y: Float32Array; gpuTimeMs: number }>;
28
28
 
29
29
  /**
30
- * Swaps the elements of two single-precision vectors: x <-> y
30
+ * Swaps the elements of two single-precision vectors: $$x \leftrightarrow y$$
31
31
  *
32
32
  * {@includeCode ../../examples/sswap/gpu.sswap.js}
33
33
  *
@@ -11,13 +11,14 @@ 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 { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function sswap(device, n, x, incx, y, incy) {
16
17
  const xIsGpu = x instanceof GpuVector;
17
18
  const yIsGpu = y instanceof GpuVector;
18
19
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
21
+ requireSameDevice(device, "sswap", { x, y });
21
22
  if (
22
23
  !Number.isInteger(n) ||
23
24
  !Number.isInteger(incx) ||
@@ -53,9 +54,10 @@ export async function sswap(device, n, x, incx, y, incy) {
53
54
  let yReadBuffer = null;
54
55
 
55
56
  try {
56
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
57
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
57
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sswap-x", true);
58
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sswap-y", true);
58
59
  paramsBuffer = createParamsBuffer(
60
+ device,
59
61
  [
60
62
  { value: n, type: "u32" },
61
63
  { value: incx, type: "u32" },
@@ -64,24 +66,30 @@ export async function sswap(device, n, x, incx, y, incy) {
64
66
  "sswap-params",
65
67
  );
66
68
 
67
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
69
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
68
70
  xBuffer,
69
71
  yBuffer,
70
72
  paramsBuffer,
71
73
  ]);
72
74
  const { commandEncoder, ts } = runComputePass(
75
+ device,
73
76
  pipeline,
74
77
  bindGroup,
75
- calcWorkgroups(n),
78
+ calcWorkgroups(device, n),
76
79
  );
77
- xReadBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
78
- yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
80
+ xReadBuffer = xIsGpu
81
+ ? null
82
+ : stageReadback(device, commandEncoder, xBuffer);
83
+ yReadBuffer = yIsGpu
84
+ ? null
85
+ : stageReadback(device, commandEncoder, yBuffer);
79
86
 
80
- submit(commandEncoder);
87
+ submit(device, commandEncoder);
81
88
 
82
89
  const gpuTimeMs = await extractTimestamp(ts);
83
90
 
84
- if (xIsGpu && yIsGpu) {
91
+ if (xIsGpu) {
92
+ // xIsGpu === yIsGpu, enforced above (x.constructor !== y.constructor throws)
85
93
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
86
94
  return {};
87
95
  }
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
4
  * Performs the symmetric matrix-matrix operation
5
- * C := alpha * A * B + beta * C (`side='left'`) or
6
- * C := alpha * B * A + beta * C (`side='right'`) — `A` is symmetric, only
5
+ * $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
6
+ * $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
7
+ * `A` is symmetric, only
7
8
  * its `uplo` triangle stored; `B` and `C` are general m×n matrices.
8
9
  *
9
10
  * - `side='left'`: `A` is m×m — `A` premultiplies `B`
@@ -59,8 +60,8 @@ export declare function ssymm(
59
60
 
60
61
  /**
61
62
  * Performs the symmetric matrix-matrix operation
62
- * C := alpha * A * B + beta * C (`side='left'`) or
63
- * C := alpha * B * A + beta * C (`side='right'`)
63
+ * $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
64
+ * $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
64
65
  *
65
66
  * A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
66
67
  * `GpuMatrix.from` time) determines the operation — there is no separate
@@ -4,6 +4,7 @@ import {
4
4
  createStorageBuffer,
5
5
  stageReadback,
6
6
  destroyBuffers,
7
+ vec4ViewBinding,
7
8
  } from "../util/buffer.mjs";
8
9
  import { createBindGroup } from "../util/bindgroup.mjs";
9
10
  import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
@@ -11,41 +12,60 @@ import { extractResult } from "../util/result.mjs";
11
12
  import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
13
  import { getPipeline } from "../util/pipeline.mjs";
13
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
-
15
- const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
16
- const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
17
- const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr/ssyrk use
18
- const SYM_WG = 8; // symmetrize.wgsl's @workgroup_size(8, 8)
15
+ import { requireWorkgroupCount } from "../util/workgroup.mjs";
16
+ import {
17
+ BM_SMALL,
18
+ BN_SMALL,
19
+ BM_LARGE,
20
+ BN_LARGE,
21
+ LARGE_TILE_WORKGROUP_THRESHOLD,
22
+ } from "../util/constants.mjs";
23
+ import { TILE_WG_2D } from "../util/constants.mjs";
24
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
19
25
 
20
26
  // ssymm: C := alpha*A*B + beta*C (side='left') or alpha*B*A + beta*C
21
27
  // (side='right'), A symmetric. No fused kernel — symmetrize then sgemm,
22
28
  // both on one command encoder. See symmetrize.wgsl.
23
29
  export async function ssymm(
24
- device, side, uplo, m, n, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
30
+ device,
31
+ side,
32
+ uplo,
33
+ m,
34
+ n,
35
+ alpha,
36
+ A,
37
+ lda,
38
+ B,
39
+ ldb,
40
+ beta,
41
+ C,
42
+ ldc,
43
+ layout = "row-major",
25
44
  ) {
26
45
  const AIsGpu = A instanceof GpuMatrix;
27
46
  const BIsGpu = B instanceof GpuMatrix;
28
47
  const CIsGpu = C instanceof GpuMatrix;
29
48
 
30
- if (!(device instanceof GPUDevice))
31
- throw new Error("device must be a GPUDevice.");
49
+ requireGpuDevice(device);
50
+ requireSameDevice(device, "ssymm", { A, B, C });
32
51
  if (side !== "left" && side !== "right")
33
52
  throw new Error("side must be 'left' or 'right'.");
34
53
  if (uplo !== "lower" && uplo !== "upper")
35
54
  throw new Error("uplo must be 'lower' or 'upper'.");
36
55
  if (layout !== "row-major" && layout !== "column-major")
37
56
  throw new Error("layout must be 'row-major' or 'column-major'.");
38
- if (typeof alpha !== "number")
39
- throw new Error("alpha must be a number.");
57
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
40
58
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
41
59
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
42
- if (typeof beta !== "number")
43
- throw new Error("beta must be a number.");
60
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
44
61
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
45
62
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
46
63
  if (
47
- !Number.isInteger(m) || !Number.isInteger(n) ||
48
- !Number.isInteger(lda) || !Number.isInteger(ldb) || !Number.isInteger(ldc)
64
+ !Number.isInteger(m) ||
65
+ !Number.isInteger(n) ||
66
+ !Number.isInteger(lda) ||
67
+ !Number.isInteger(ldb) ||
68
+ !Number.isInteger(ldc)
49
69
  )
50
70
  throw new Error("m, n, lda, ldb, and ldc must be integers.");
51
71
  if (!AIsGpu && !(A instanceof Float32Array))
@@ -67,48 +87,71 @@ export async function ssymm(
67
87
 
68
88
  // A: symmetric, order = m (side='left') or n (side='right').
69
89
  const aOrder = side === "left" ? m : n;
70
- if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
90
+ if (lda < aOrder)
91
+ throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
71
92
  if (AIsGpu) {
72
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
73
- if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
93
+ if (lda !== A.lda)
94
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
95
+ if (A.rows < aOrder || A.cols < aOrder)
96
+ throw new Error("A is too small for the given m/n and side.");
74
97
  } else if (A.length < (aOrder - 1) * lda + aOrder) {
75
- throw new Error("A does not have enough elements for the given dimensions and lda.");
98
+ throw new Error(
99
+ "A does not have enough elements for the given dimensions and lda.",
100
+ );
76
101
  }
77
102
 
78
103
  // B: always m x n, no trans flag — same shape rule as sgemm's C.
79
104
  const bOuter = effLayoutB === "column-major" ? n : m;
80
105
  const bInner = effLayoutB === "column-major" ? m : n;
81
106
  if (ldb < bInner)
82
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
107
+ throw new Error(
108
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
109
+ );
83
110
  if (BIsGpu) {
84
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
85
- if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
111
+ if (ldb !== B.lda)
112
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
113
+ if (B.rows < m || B.cols < n)
114
+ throw new Error("B is too small for the given m and n.");
86
115
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
87
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
116
+ throw new Error(
117
+ "B does not have enough elements for the given dimensions and ldb.",
118
+ );
88
119
  }
89
120
 
90
121
  // C: always m x n.
91
122
  const cOuter = effLayoutC === "column-major" ? n : m;
92
123
  const cInner = effLayoutC === "column-major" ? m : n;
93
124
  if (ldc < cInner)
94
- throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
125
+ throw new Error(
126
+ `ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
127
+ );
95
128
  if (CIsGpu) {
96
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
97
- if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
129
+ if (ldc !== C.lda)
130
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
131
+ if (C.rows < m || C.cols < n)
132
+ throw new Error("C is too small for the given m and n.");
98
133
  } else if (C.length < (cOuter - 1) * ldc + cInner) {
99
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
134
+ throw new Error(
135
+ "C does not have enough elements for the given dimensions and ldc.",
136
+ );
100
137
  }
101
138
 
102
139
  // A = A^T, so column-major storage is still A, but the populated
103
140
  // triangle swaps — uplo flips (same reasoning ssyr/ssyrk use).
104
- const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
141
+ const uploEffA =
142
+ effLayoutA === "column-major"
143
+ ? uplo === "lower"
144
+ ? "upper"
145
+ : "lower"
146
+ : uplo;
105
147
 
106
148
  const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
107
149
  const transDense = "no-transpose"; // Adense is always row-major
108
150
 
109
151
  // X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
110
152
  // C: compute C^T instead (swap X/Y, flip trans, swap m_g/n_g) — sgemm's own trick.
111
- let mg = m, ng = n;
153
+ let mg = m,
154
+ ng = n;
112
155
  const kg = aOrder;
113
156
  let transX = side === "left" ? transDense : transB;
114
157
  let transY = side === "left" ? transB : transDense;
@@ -124,26 +167,45 @@ export async function ssymm(
124
167
  const largeWgX = Math.ceil(ng / BN_LARGE);
125
168
  const largeWgY = Math.ceil(mg / BM_LARGE);
126
169
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
127
- const gemmPipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
170
+ const gemmPipeline = await getPipeline(
171
+ device,
172
+ useLargeTile ? "sgemm_large" : "sgemm_small",
173
+ );
128
174
  const symPipeline = await getPipeline(device, "symmetrize");
129
175
  const gemmWgCount = useLargeTile
130
176
  ? {
131
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
132
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
133
- }
177
+ x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
178
+ y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
179
+ }
134
180
  : {
135
- x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
136
- y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
137
- };
181
+ x: requireWorkgroupCount(
182
+ device,
183
+ Math.ceil(ng / BN_SMALL),
184
+ "ssymm",
185
+ "x",
186
+ ),
187
+ y: requireWorkgroupCount(
188
+ device,
189
+ Math.ceil(mg / BM_SMALL),
190
+ "ssymm",
191
+ "y",
192
+ ),
193
+ };
138
194
 
139
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssymm-A", false);
140
- const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "ssymm-B", false);
141
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssymm-C", true);
142
- const AdenseBuffer = createStorageBuffer(aOrder * ldDense * 4, "ssymm-Adense");
143
- let symParams = null, gemmParams = null;
195
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymm-A", false);
196
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssymm-B", false);
197
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssymm-C", true);
198
+ const AdenseBuffer = createStorageBuffer(
199
+ device,
200
+ aOrder * ldDense * 4,
201
+ "ssymm-Adense",
202
+ );
203
+ let symParams = null,
204
+ gemmParams = null;
144
205
 
145
206
  try {
146
207
  symParams = createParamsBuffer(
208
+ device,
147
209
  [
148
210
  { value: aOrder, type: "u32" },
149
211
  { value: lda, type: "u32" },
@@ -152,7 +214,11 @@ export async function ssymm(
152
214
  ],
153
215
  "ssymm-sym-params",
154
216
  );
155
- const symBindGroup = createBindGroup(symPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, symParams]);
217
+ const symBindGroup = createBindGroup(
218
+ device,
219
+ symPipeline.getBindGroupLayout(0),
220
+ [ABuffer, AdenseBuffer, symParams],
221
+ );
156
222
 
157
223
  // X/Y buffers and their own ld, matching swapXY above.
158
224
  const XBuffer = swapXY ? BBuffer : AdenseBuffer;
@@ -161,12 +227,13 @@ export async function ssymm(
161
227
  const ldY = swapXY ? ldDense : ldb;
162
228
 
163
229
  gemmParams = createParamsBuffer(
230
+ device,
164
231
  [
165
- { value: mg, type: "u32" },
166
- { value: ng, type: "u32" },
167
- { value: kg, type: "u32" },
232
+ { value: mg, type: "u32" },
233
+ { value: ng, type: "u32" },
234
+ { value: kg, type: "u32" },
168
235
  { value: alpha, type: "f32" },
169
- { value: beta, type: "f32" },
236
+ { value: beta, type: "f32" },
170
237
  { value: ldX, type: "u32" },
171
238
  { value: ldY, type: "u32" },
172
239
  { value: ldc, type: "u32" },
@@ -175,18 +242,47 @@ export async function ssymm(
175
242
  ],
176
243
  "ssymm-gemm-params",
177
244
  );
178
- const gemmBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [XBuffer, YBuffer, CBuffer, gemmParams]);
245
+ const gemmBindGroup = createBindGroup(
246
+ device,
247
+ gemmPipeline.getBindGroupLayout(0),
248
+ [
249
+ XBuffer,
250
+ vec4ViewBinding(device, XBuffer),
251
+ YBuffer,
252
+ vec4ViewBinding(device, YBuffer),
253
+ CBuffer,
254
+ gemmParams,
255
+ ],
256
+ );
179
257
 
180
- const { commandEncoder, querySet } = beginTimedEncoder();
181
- const symDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
182
- const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
183
- encodePass(commandEncoder, symPipeline, symBindGroup, { x: Math.ceil(aOrder / SYM_WG), y: Math.ceil(aOrder / SYM_WG) }, symDesc);
184
- encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
258
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
259
+ const symDesc = querySet
260
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
261
+ : undefined;
262
+ const gemmDesc = querySet
263
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
264
+ : undefined;
265
+ encodePass(
266
+ commandEncoder,
267
+ symPipeline,
268
+ symBindGroup,
269
+ { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) },
270
+ symDesc,
271
+ );
272
+ encodePass(
273
+ commandEncoder,
274
+ gemmPipeline,
275
+ gemmBindGroup,
276
+ gemmWgCount,
277
+ gemmDesc,
278
+ );
185
279
 
186
- const ts = resolveTimestamp(commandEncoder, querySet);
187
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
280
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
281
+ const readBuffer = CIsGpu
282
+ ? null
283
+ : stageReadback(device, commandEncoder, CBuffer);
188
284
 
189
- submit(commandEncoder);
285
+ submit(device, commandEncoder);
190
286
 
191
287
  const gpuTimeMs = await extractTimestamp(ts);
192
288
 
@@ -2,7 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
3
 
4
4
  /**
5
- * Performs the symmetric matrix-vector operation y = alpha * A * x + beta * y
5
+ * Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
6
6
  *
7
7
  * A is an n×n symmetric matrix stored in row-major order. Only the triangle
8
8
  * specified by `uplo` is referenced; the other triangle is inferred by symmetry.
@@ -45,7 +45,7 @@ export declare function ssymv(
45
45
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
46
46
 
47
47
  /**
48
- * Performs the symmetric matrix-vector operation y = alpha * A * x + beta * y
48
+ * Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
49
49
  *
50
50
  * x and y are kept resident on the GPU. A must be a GpuMatrix; its own
51
51
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is