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
@@ -1,17 +1,12 @@
1
1
  /** @module devdocs/utility-functions/compute */
2
- import { getDevice } from "../init.mjs";
3
2
  import { beginTimestamp, resolveTimestamp } from "./benchmark.mjs";
4
3
 
5
- // Anchors the pass encoder to its command encoder to prevent premature GC.
6
- const _passEncoders = new WeakMap();
7
-
8
4
  /**
9
5
  * Finalises `commandEncoder` into a command buffer and submits it to the GPU queue.
10
6
  * @param {GPUCommandEncoder} commandEncoder
11
7
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUQueue/submit GPUQueue.submit()}
12
8
  */
13
- export function submit(commandEncoder) {
14
- const device = getDevice();
9
+ export function submit(device, commandEncoder) {
15
10
  device.queue.submit([commandEncoder.finish()]);
16
11
  }
17
12
 
@@ -23,9 +18,8 @@ export function submit(commandEncoder) {
23
18
  * @returns {{ commandEncoder: GPUCommandEncoder, querySet: GPUQuerySet|null, passDescriptor: object|undefined }}
24
19
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
25
20
  */
26
- export function beginTimedEncoder() {
27
- const device = getDevice();
28
- const { querySet, passDescriptor } = beginTimestamp();
21
+ export function beginTimedEncoder(device) {
22
+ const { querySet, passDescriptor } = beginTimestamp(device);
29
23
  const commandEncoder = device.createCommandEncoder();
30
24
  return { commandEncoder, querySet, passDescriptor };
31
25
  }
@@ -43,7 +37,13 @@ export function beginTimedEncoder() {
43
37
  * @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
44
38
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
45
39
  */
46
- export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor) {
40
+ export function encodePass(
41
+ commandEncoder,
42
+ pipeline,
43
+ bindGroup,
44
+ workgroups,
45
+ passDescriptor,
46
+ ) {
47
47
  const passEncoder = commandEncoder.beginComputePass(passDescriptor);
48
48
 
49
49
  passEncoder.setPipeline(pipeline);
@@ -53,12 +53,14 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
53
53
  passEncoder.dispatchWorkgroups(workgroups);
54
54
  } else {
55
55
  // `?? 1` is load-bearing — confirmed passing undefined (from an {x,y}-only caller) crashes the process, not just no-ops.
56
- passEncoder.dispatchWorkgroups(workgroups.x, workgroups.y, workgroups.z ?? 1);
56
+ passEncoder.dispatchWorkgroups(
57
+ workgroups.x,
58
+ workgroups.y,
59
+ workgroups.z ?? 1,
60
+ );
57
61
  }
58
62
 
59
63
  passEncoder.end();
60
-
61
- _passEncoders.set(commandEncoder, passEncoder);
62
64
  }
63
65
 
64
66
  /**
@@ -71,11 +73,12 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
71
73
  * @returns {{ commandEncoder: GPUCommandEncoder, ts: any }} encoded commands and timestamp handle
72
74
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
73
75
  */
74
- export function runComputePass(pipeline, bindGroup, workgroups) {
75
- const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder();
76
+ export function runComputePass(device, pipeline, bindGroup, workgroups) {
77
+ const { commandEncoder, querySet, passDescriptor } =
78
+ beginTimedEncoder(device);
76
79
  encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor);
77
80
 
78
- const ts = resolveTimestamp(commandEncoder, querySet);
81
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
79
82
 
80
83
  return { commandEncoder, ts };
81
84
  }
@@ -0,0 +1,57 @@
1
+ /** @module devdocs/utility-functions/constants */
2
+
3
+ /**
4
+ * Single source of truth for every constant the JS dispatch side shares with a
5
+ * shader.
6
+ *
7
+ * WGSL has no import mechanism, so each shader necessarily declares its own
8
+ * copy of these values. That makes them a silent-corruption hazard: change
9
+ * `BM` in `sgemm_large.wgsl` without changing `BM_LARGE` here and the host
10
+ * dispatches the wrong grid — too few workgroups computes part of the matrix
11
+ * and reports success. Hoisting them here removes the JS-to-JS duplication;
12
+ * `tests/utils/test.constants.js` closes the remaining JS-to-WGSL gap
13
+ * by parsing the shader sources and asserting they still agree.
14
+ *
15
+ * Every export below names the shader declaration it mirrors. Changing one
16
+ * means changing both, and the test will tell you if you forget.
17
+ */
18
+
19
+ // --- gemm block tiles ------------------------------------------------------
20
+
21
+ /** `BM` in sgemm_small.wgsl / sgemmtr_small.wgsl. */
22
+ export const BM_SMALL = 32;
23
+ /** `BN` in sgemm_small.wgsl / sgemmtr_small.wgsl. */
24
+ export const BN_SMALL = 32;
25
+ /** `BM` in sgemm_large.wgsl / sgemmtr_large.wgsl. */
26
+ export const BM_LARGE = 64;
27
+ /** `BN` in sgemm_large.wgsl / sgemmtr_large.wgsl. */
28
+ export const BN_LARGE = 64;
29
+
30
+ /**
31
+ * The large tile only pays for its bigger workgroups once the problem needs at
32
+ * least a 6x6 grid of them; below that the small tile wins. JS-only (no shader
33
+ * counterpart) — it selects *which* shader runs.
34
+ */
35
+ export const LARGE_TILE_WORKGROUP_THRESHOLD = 36;
36
+
37
+ // --- 1D / reduction kernels ------------------------------------------------
38
+
39
+ /**
40
+ * `const WGS: u32 = 64` — declared by every 1D and reduction shader, and the
41
+ * workgroup size `calcWorkgroups` divides by for a 1D dispatch.
42
+ */
43
+ export const WGS = 64;
44
+
45
+ // --- 2D helper kernels -----------------------------------------------------
46
+
47
+ /**
48
+ * `@workgroup_size(8, 8)` in symmetrize.wgsl, triangularize.wgsl and
49
+ * block_transfer.wgsl, and the size `calcWorkgroups` divides by per dimension
50
+ * for a 2D dispatch.
51
+ */
52
+ export const TILE_WG_2D = 8;
53
+
54
+ // --- triangular solve ------------------------------------------------------
55
+
56
+ /** `BLOCK_SIZE` in strsv_invert_block.wgsl — the diagonal block order. */
57
+ export const BLOCK_SIZE = 64;
@@ -0,0 +1,49 @@
1
+ /** @module devdocs/utility-functions/device */
2
+ import { GpuVector } from "../classes/GpuVector.mjs";
3
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
4
+
5
+ /**
6
+ * Throws if `device` is not a `GPUDevice`.
7
+ *
8
+ * Every routine's first guard, extracted since the check and message are
9
+ * identical across all of them.
10
+ *
11
+ * @param {GPUDevice} device - the value to check
12
+ * @throws {Error} if `device` is not a `GPUDevice`
13
+ */
14
+ export function requireGpuDevice(device) {
15
+ if (!(device instanceof GPUDevice))
16
+ throw new Error("device must be a GPUDevice.");
17
+ }
18
+
19
+ /**
20
+ * Throws if any GPU-resident operand belongs to a device other than the one
21
+ * the routine was called with.
22
+ *
23
+ * A `GPUBuffer` is bound to the device that created it, and WebGPU has no way
24
+ * to share one across devices. Handing a routine a `GpuMatrix` from device A
25
+ * while passing device B fails deep inside bind-group creation as a
26
+ * `GPUValidationError` with no indication that two devices are involved —
27
+ * this turns it into a named, actionable error at the call boundary.
28
+ *
29
+ * Scalars, plain typed arrays and `undefined` entries are ignored, so callers
30
+ * can pass their whole operand set without filtering.
31
+ *
32
+ * @param {GPUDevice} device - the device the routine will dispatch on
33
+ * @param {string} routine - routine name, for the error message
34
+ * @param {Record<string, unknown>} operands - operand name -> value
35
+ * @throws {Error} if an operand is GPU-resident on a different device
36
+ */
37
+ export function requireSameDevice(device, routine, operands) {
38
+ for (const [name, value] of Object.entries(operands)) {
39
+ if (!(value instanceof GpuVector) && !(value instanceof GpuMatrix))
40
+ continue;
41
+ if (value.device !== device) {
42
+ throw new Error(
43
+ `${routine}: ${name} belongs to a different GPUDevice than the one passed in. ` +
44
+ "GPU buffers cannot be shared across devices — recreate the operand on this " +
45
+ "device, or call the routine with the device that owns it.",
46
+ );
47
+ }
48
+ }
49
+ }
@@ -1,5 +1,4 @@
1
1
  /** @module devdocs/utility-functions/pipeline */
2
- import { getDevice } from "../init.mjs";
3
2
 
4
3
  // WeakMap keyed by GPUDevice so pipelines are released automatically when the device is destroyed.
5
4
  const _pipelines = new WeakMap();
@@ -24,23 +23,31 @@ export async function getPipeline(device, shaderName, entryPoint = "main") {
24
23
  const names = Array.isArray(shaderName) ? shaderName : [shaderName];
25
24
  const key = `${names.join("+")}::${entryPoint}`;
26
25
  if (!byName.has(key)) {
27
- byName.set(key, await loadShader(names, entryPoint));
26
+ // Cache the in-flight promise (not its resolved value) so a concurrent
27
+ // call awaits the same compile instead of starting a duplicate one;
28
+ // drop the entry on failure so a later call can retry.
29
+ const pending = loadShader(device, names, entryPoint).catch((err) => {
30
+ byName.delete(key);
31
+ throw err;
32
+ });
33
+ byName.set(key, pending);
28
34
  }
29
35
  return byName.get(key);
30
36
  }
31
37
 
32
38
  /**
33
39
  * Loads WGSL source for `shaderName`. In the browser, reads from the inline bundle
34
- * (`browser-shaders.mjs`); in Node.js, reads the `.wgsl` file directly from disk.
40
+ * (`shaders/index.mjs`'s `shaderSources`); in Node.js, reads the `.wgsl` file directly from disk.
35
41
  * @param {string} shaderName
36
42
  * @returns {Promise<string>}
37
43
  */
38
44
  async function loadCode(shaderName) {
39
45
  // Check for Node.js explicitly — `window` is undefined in Web Workers too, so it's not a reliable signal.
40
46
  if (typeof process === "undefined" || !process.versions?.node) {
41
- const { shaderSources } = await import("../shaders/browser-shaders.mjs");
47
+ const { shaderSources } = await import("../shaders/index.mjs");
42
48
  const src = shaderSources[shaderName];
43
- if (!src) throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
49
+ if (!src)
50
+ throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
44
51
  return src;
45
52
  } else {
46
53
  const { readFileSync } = await import("fs");
@@ -56,6 +63,7 @@ async function loadCode(shaderName) {
56
63
  * `GPUComputePipeline`. Throws with line-level detail if compilation fails, rather than
57
64
  * surfacing a raw GPU error. Uses `layout: "auto"` so the pipeline derives its bind group
58
65
  * layout from the shader — no manual layout definition needed.
66
+ * @param {GPUDevice} device
59
67
  * @param {string[]} shaderNames - filenames without `.wgsl`, concatenated in array order
60
68
  * @param {string} [entryPoint="main"] - which `@compute` function in the combined module to run
61
69
  * @returns {Promise<GPUComputePipeline>}
@@ -64,11 +72,34 @@ async function loadCode(shaderName) {
64
72
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUShaderModule/getCompilationInfo GPUShaderModule.getCompilationInfo()}
65
73
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createComputePipeline GPUDevice.createComputePipeline()}
66
74
  */
67
- export async function loadShader(shaderNames, entryPoint = "main") {
68
- const device = getDevice();
75
+ export async function loadShader(device, shaderNames, entryPoint = "main") {
69
76
  const label = shaderNames.join("+");
70
- const code = (await Promise.all(shaderNames.map(loadCode))).join("\n");
77
+ const codes = await Promise.all(shaderNames.map(loadCode));
78
+
79
+ // Per-file line ranges within the concatenated module, so a compile error
80
+ // reports as "<file>.wgsl:<local line>" instead of a whole-module line
81
+ // number — confusing for multi-file pipelines like dasum's.
82
+ let offset = 0;
83
+ const ranges = codes.map((c, i) => {
84
+ const lineCount = c.split("\n").length;
85
+ const range = {
86
+ name: shaderNames[i],
87
+ startLine: offset + 1,
88
+ endLine: offset + lineCount,
89
+ };
90
+ offset += lineCount;
91
+ return range;
92
+ });
93
+ const locate = (lineNum) => {
94
+ const range =
95
+ lineNum &&
96
+ ranges.find((r) => lineNum >= r.startLine && lineNum <= r.endLine);
97
+ return range
98
+ ? `${range.name}.wgsl:${lineNum - range.startLine + 1}`
99
+ : `line ${lineNum}`;
100
+ };
71
101
 
102
+ const code = codes.join("\n");
72
103
  const shaderModule = device.createShaderModule({ label, code });
73
104
 
74
105
  const info = await shaderModule.getCompilationInfo();
@@ -76,7 +107,7 @@ export async function loadShader(shaderNames, entryPoint = "main") {
76
107
  const errors = info.messages.filter((m) => m.type === "error");
77
108
  if (errors.length > 0) {
78
109
  throw new Error(
79
- `Shader "${label}" compilation failed:\n${errors.map((m) => ` line ${m.lineNum}: ${m.message}`).join("\n")}`,
110
+ `Shader "${label}" compilation failed:\n${errors.map((m) => ` ${locate(m.lineNum)}: ${m.message}`).join("\n")}`,
80
111
  );
81
112
  }
82
113
 
@@ -84,7 +115,10 @@ export async function loadShader(shaderNames, entryPoint = "main") {
84
115
  // this project's WebGPU backend is unstable (intermittent multi-minute hangs and wrong
85
116
  // results, confirmed by bisection) when entryPoint is set explicitly, even to the shader's
86
117
  // only/correct entry point. Auto-detecting the single entry point is the stable path.
87
- const compute = entryPoint === "main" ? { module: shaderModule } : { module: shaderModule, entryPoint };
118
+ const compute =
119
+ entryPoint === "main"
120
+ ? { module: shaderModule }
121
+ : { module: shaderModule, entryPoint };
88
122
  const pipeline = device.createComputePipeline({
89
123
  label,
90
124
  layout: "auto",
@@ -1,14 +1,23 @@
1
1
  /** @module devdocs/utility-functions/workgroup */
2
- import { getDevice } from "../init.mjs";
3
-
4
- // Fixed sizes match the shader declarations (WGS = 64 for 1D, 8×8 = 64 threads for 2D).
5
- const WORKGROUP_SIZE_1D = 64;
6
- const WORKGROUP_SIZE_2D = 8;
2
+ // Fixed sizes match the shader declarations (WGS = 64 for 1D, 8×8 = 64 threads
3
+ // for 2D) — see constants.mjs, which is where both values are defined and
4
+ // where the WGSL cross-check hangs off.
5
+ import {
6
+ WGS as WORKGROUP_SIZE_1D,
7
+ TILE_WG_2D as WORKGROUP_SIZE_2D,
8
+ } from "./constants.mjs";
7
9
 
8
10
  /**
9
11
  * Calculates the number of workgroups to dispatch, clamped to the device's
10
12
  * `maxComputeWorkgroupsPerDimension` limit (default 65535 across most devices).
11
13
  *
14
+ * ONLY for shaders whose kernel is a grid-stride loop driven by
15
+ * `num_workgroups` — those re-walk the whole domain regardless of how many
16
+ * workgroups actually launch, so clamping costs a little parallelism and
17
+ * nothing else. A shader that indexes straight off `workgroup_id` or
18
+ * `global_invocation_id` silently drops every row past the clamp; those must
19
+ * use {@link requireWorkgroups} / {@link requireWorkgroupCount} instead.
20
+ *
12
21
  * - 1D (pass only `rows`): returns a single count for `dispatchWorkgroups(n)`.
13
22
  * - 2D (pass both `rows` and `cols`): returns `{ x, y }` for `dispatchWorkgroups(x, y)`.
14
23
  * `rows` maps to the y dimension and `cols` maps to the x dimension.
@@ -18,8 +27,8 @@ const WORKGROUP_SIZE_2D = 8;
18
27
  * @returns {number | { x: number, y: number }}
19
28
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUSupportedLimits GPUSupportedLimits} (`maxComputeWorkgroupsPerDimension`)
20
29
  */
21
- export function calcWorkgroups(rows, cols) {
22
- const max = getDevice().limits.maxComputeWorkgroupsPerDimension;
30
+ export function calcWorkgroups(device, rows, cols) {
31
+ const max = device.limits.maxComputeWorkgroupsPerDimension;
23
32
  if (cols === undefined) {
24
33
  return Math.min(Math.ceil(rows / WORKGROUP_SIZE_1D), max);
25
34
  } else {
@@ -29,3 +38,59 @@ export function calcWorkgroups(rows, cols) {
29
38
  };
30
39
  }
31
40
  }
41
+
42
+ /**
43
+ * Returns `count` unchanged if it fits the device's dispatch limit, and throws
44
+ * otherwise. The counterpart to {@link calcWorkgroups} for shaders with no
45
+ * grid-stride fallback: silently clamping those computes only part of the
46
+ * result and reports success, so a refusal is the safer failure.
47
+ *
48
+ * @param {number} count - workgroups this dispatch requires in one dimension
49
+ * @param {string} routine - routine name, for the error message
50
+ * @param {string} [dim] - dimension label ("x"/"y"), for the error message
51
+ * @returns {number} `count`
52
+ * @throws {Error} when `count` exceeds `maxComputeWorkgroupsPerDimension`
53
+ */
54
+ export function requireWorkgroupCount(device, count, routine, dim = "x") {
55
+ const max = device.limits.maxComputeWorkgroupsPerDimension;
56
+ if (count > max)
57
+ throw new Error(
58
+ `${routine}: this problem needs ${count} workgroups in ${dim}, but the device allows ` +
59
+ `${max} (maxComputeWorkgroupsPerDimension). The operands are too large for this device — ` +
60
+ `split the operation into smaller blocks.`,
61
+ );
62
+ return count;
63
+ }
64
+
65
+ /**
66
+ * {@link calcWorkgroups} with the clamp replaced by a throw — same arguments
67
+ * and same return shape, for shaders without a grid-stride fallback.
68
+ *
69
+ * @param {string} routine - routine name, for the error message
70
+ * @param {number} rows - row count (1D: element count)
71
+ * @param {number} [cols] - column count; omit for a 1D dispatch
72
+ * @returns {number | { x: number, y: number }}
73
+ * @throws {Error} when either dimension exceeds `maxComputeWorkgroupsPerDimension`
74
+ */
75
+ export function requireWorkgroups(device, routine, rows, cols) {
76
+ if (cols === undefined)
77
+ return requireWorkgroupCount(
78
+ device,
79
+ Math.ceil(rows / WORKGROUP_SIZE_1D),
80
+ routine,
81
+ );
82
+ return {
83
+ x: requireWorkgroupCount(
84
+ device,
85
+ Math.ceil(cols / WORKGROUP_SIZE_2D),
86
+ routine,
87
+ "x",
88
+ ),
89
+ y: requireWorkgroupCount(
90
+ device,
91
+ Math.ceil(rows / WORKGROUP_SIZE_2D),
92
+ routine,
93
+ "y",
94
+ ),
95
+ };
96
+ }
@@ -1,81 +0,0 @@
1
- import argmax from "./reduction/argmax.wgsl";
2
- import argmaxF64 from "./reduction/argmaxF64.wgsl";
3
- import sum from "./reduction/sum.wgsl";
4
- import sumF64 from "./reduction/sumF64.wgsl";
5
- import sscal from "./sscal.wgsl";
6
- import sswap from "./sswap.wgsl";
7
- import saxpy from "./saxpy.wgsl";
8
- import scopy from "./scopy.wgsl";
9
- import sdot from "./sdot.wgsl";
10
- import sasum from "./sasum.wgsl";
11
- import snrm2 from "./snrm2.wgsl";
12
- import srot from "./srot.wgsl";
13
- import srotm from "./srotm.wgsl";
14
- import isamax from "./isamax.wgsl";
15
- import sgemv_n from "./sgemv_n.wgsl";
16
- import sgemv_t from "./sgemv_t.wgsl";
17
- import ssymv from "./ssymv.wgsl";
18
- import strmv from "./strmv.wgsl";
19
- import sger from "./sger.wgsl";
20
- import ssyr from "./ssyr.wgsl";
21
- import ssyr2 from "./ssyr2.wgsl";
22
- import f64add from "./f64add.wgsl";
23
- import dekker from "./f64/dekker.wgsl";
24
- import ddAbs from "./f64/utils/abs.wgsl";
25
- import ddAddUtil from "./f64/utils/add.wgsl";
26
- import ddGreater from "./f64/utils/greater.wgsl";
27
- import ddEqual from "./f64/utils/equal.wgsl";
28
- import dasum from "./dasum.wgsl";
29
- import idamax from "./idamax.wgsl";
30
- import strsv_invert_block from "./strsv_invert_block.wgsl";
31
- import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
32
- import strsv_update from "./strsv_update.wgsl";
33
- import sgemm_small from "./sgemm_small.wgsl";
34
- import sgemm_large from "./sgemm_large.wgsl";
35
- import sgemmtr_small from "./sgemmtr_small.wgsl";
36
- import sgemmtr_large from "./sgemmtr_large.wgsl";
37
- import symmetrize from "./symmetrize.wgsl";
38
- import triangularize from "./triangularize.wgsl";
39
- import blockTransfer from "./block_transfer.wgsl";
40
-
41
- export const shaderSources = {
42
- "reduction/argmax": argmax,
43
- "reduction/argmaxF64": argmaxF64,
44
- "reduction/sum": sum,
45
- "reduction/sumF64": sumF64,
46
- sscal,
47
- sswap,
48
- saxpy,
49
- scopy,
50
- sdot,
51
- sasum,
52
- snrm2,
53
- srot,
54
- srotm,
55
- isamax,
56
- sgemv_n,
57
- sgemv_t,
58
- ssymv,
59
- strmv,
60
- sger,
61
- ssyr,
62
- ssyr2,
63
- f64add,
64
- "f64/dekker": dekker,
65
- "f64/utils/abs": ddAbs,
66
- "f64/utils/add": ddAddUtil,
67
- "f64/utils/greater": ddGreater,
68
- "f64/utils/equal": ddEqual,
69
- dasum,
70
- idamax,
71
- strsv_invert_block,
72
- strsv_apply_inverse,
73
- strsv_update,
74
- sgemm_small,
75
- sgemm_large,
76
- sgemmtr_small,
77
- sgemmtr_large,
78
- symmetrize,
79
- triangularize,
80
- block_transfer: blockTransfer,
81
- };