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
@@ -0,0 +1,170 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ } from "../util/buffer.mjs";
7
+ import { createBindGroup } from "../util/bindgroup.mjs";
8
+ import { runComputePass, submit } from "../util/compute.mjs";
9
+ import { extractResult } from "../util/result.mjs";
10
+ import { extractTimestamp } from "../util/benchmark.mjs";
11
+ import { getPipeline } from "../util/pipeline.mjs";
12
+ import { calcWorkgroups } from "../util/workgroup.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
16
+
17
+ // drot: x := c*x + s*y, y := -s*x + c*y, double-double (Dekker) f64
18
+ // emulation of srot — x, y, c, and s are each split into an f32 (hi, lo)
19
+ // pair; WGSL has no f64 type.
20
+ export async function drot(device, n, x, incx, y, incy, c, s) {
21
+ const xIsGpu = x instanceof GpuVector;
22
+ const yIsGpu = y instanceof GpuVector;
23
+
24
+ requireGpuDevice(device);
25
+ requireSameDevice(device, "drot", { x, y });
26
+ if (
27
+ !Number.isInteger(n) ||
28
+ !Number.isInteger(incx) ||
29
+ !Number.isInteger(incy)
30
+ )
31
+ throw new Error("n, incx, and incy must be integers.");
32
+ if (typeof c !== "number") throw new Error("c must be a number.");
33
+ if (typeof s !== "number") throw new Error("s must be a number.");
34
+ if (Number.isNaN(c) || Number.isNaN(s))
35
+ throw new Error("c and s must not be NaN.");
36
+ if (!Number.isFinite(c)) throw new Error("c must be finite.");
37
+ if (!Number.isFinite(s)) throw new Error("s must be finite.");
38
+ if (incx <= 0 || incy <= 0)
39
+ throw new Error("incx and incy must be positive.");
40
+ if (!(x instanceof Float64Array) && !xIsGpu)
41
+ throw new Error("x must be a Float64Array or GpuVector.");
42
+ if (!(y instanceof Float64Array) && !yIsGpu)
43
+ throw new Error("y must be a Float64Array or GpuVector.");
44
+ if (xIsGpu && x.dtype !== Float64Array)
45
+ throw new Error("x must be a Float64Array-backed GpuVector.");
46
+ if (yIsGpu && y.dtype !== Float64Array)
47
+ throw new Error("y must be a Float64Array-backed GpuVector.");
48
+ if (xIsGpu !== yIsGpu)
49
+ throw new Error(
50
+ "x and y must be the same type (both Float64Array or both GpuVector).",
51
+ );
52
+ if (n <= 0) return xIsGpu ? {} : { x, y };
53
+ if (x.length < (n - 1) * incx + 1)
54
+ throw new Error(
55
+ "x does not have enough elements for the given n and incx.",
56
+ );
57
+ if (y.length < (n - 1) * incy + 1)
58
+ throw new Error(
59
+ "y does not have enough elements for the given n and incy.",
60
+ );
61
+
62
+ // Concatenated with f64/dekker.wgsl (DD struct), f64/utils/add.wgsl
63
+ // (fsub/negf/fastTwoSumProtected/ddAddProtected), and f64/utils/multiply.wgsl
64
+ // (ddMulProtected) — WGSL has no #include.
65
+ const f64Deps = ["f64/dekker", "f64/utils/add", "f64/utils/multiply"];
66
+ const pipeline = await getPipeline(device, [...f64Deps, "drot"]);
67
+
68
+ const { hi: cHi, lo: cLo } = splitDoubleDouble(new Float64Array([c]));
69
+ const { hi: sHi, lo: sLo } = splitDoubleDouble(new Float64Array([s]));
70
+
71
+ let xHiBuffer = null;
72
+ let xLoBuffer = null;
73
+ let yHiBuffer = null;
74
+ let yLoBuffer = null;
75
+ let paramsBuffer = null;
76
+ let xReadHiBuffer = null;
77
+ let xReadLoBuffer = null;
78
+ let yReadHiBuffer = null;
79
+ let yReadLoBuffer = null;
80
+
81
+ try {
82
+ if (xIsGpu) {
83
+ xHiBuffer = x._buf;
84
+ xLoBuffer = x._loBuf;
85
+ yHiBuffer = y._buf;
86
+ yLoBuffer = y._loBuf;
87
+ } else {
88
+ const xSplit = splitDoubleDouble(x);
89
+ const ySplit = splitDoubleDouble(y);
90
+ xHiBuffer = uploadBuffer(device, xSplit.hi, "drot-xHi", true);
91
+ xLoBuffer = uploadBuffer(device, xSplit.lo, "drot-xLo", true);
92
+ yHiBuffer = uploadBuffer(device, ySplit.hi, "drot-yHi", true);
93
+ yLoBuffer = uploadBuffer(device, ySplit.lo, "drot-yLo", true);
94
+ }
95
+ paramsBuffer = createParamsBuffer(
96
+ device,
97
+ [
98
+ { value: n, type: "u32" },
99
+ { value: cHi[0], type: "f32" },
100
+ { value: cLo[0], type: "f32" },
101
+ { value: sHi[0], type: "f32" },
102
+ { value: sLo[0], type: "f32" },
103
+ { value: incx, type: "u32" },
104
+ { value: incy, type: "u32" },
105
+ ],
106
+ "drot-params",
107
+ );
108
+
109
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
110
+ xHiBuffer,
111
+ xLoBuffer,
112
+ yHiBuffer,
113
+ yLoBuffer,
114
+ paramsBuffer,
115
+ ]);
116
+ const { commandEncoder, ts } = runComputePass(
117
+ device,
118
+ pipeline,
119
+ bindGroup,
120
+ calcWorkgroups(device, n),
121
+ );
122
+ xReadHiBuffer = xIsGpu
123
+ ? null
124
+ : stageReadback(device, commandEncoder, xHiBuffer);
125
+ xReadLoBuffer = xIsGpu
126
+ ? null
127
+ : stageReadback(device, commandEncoder, xLoBuffer);
128
+ yReadHiBuffer = yIsGpu
129
+ ? null
130
+ : stageReadback(device, commandEncoder, yHiBuffer);
131
+ yReadLoBuffer = yIsGpu
132
+ ? null
133
+ : stageReadback(device, commandEncoder, yLoBuffer);
134
+
135
+ submit(device, commandEncoder);
136
+
137
+ const gpuTimeMs = await extractTimestamp(ts);
138
+
139
+ if (xIsGpu) {
140
+ // xIsGpu === yIsGpu, enforced above
141
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
142
+ return {};
143
+ }
144
+
145
+ const xHi = await extractResult(xReadHiBuffer, Float32Array);
146
+ xReadHiBuffer = null; // extractResult already destroyed it
147
+ const xLo = await extractResult(xReadLoBuffer, Float32Array);
148
+ xReadLoBuffer = null;
149
+ const yHi = await extractResult(yReadHiBuffer, Float32Array);
150
+ yReadHiBuffer = null;
151
+ const yLo = await extractResult(yReadLoBuffer, Float32Array);
152
+ yReadLoBuffer = null;
153
+ const resultX = mergeDoubleDouble(xHi, xLo);
154
+ const resultY = mergeDoubleDouble(yHi, yLo);
155
+ if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
156
+ return { x: resultX, y: resultY };
157
+ } finally {
158
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
159
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
160
+ if (!yIsGpu && yHiBuffer) destroyBuffers(yHiBuffer);
161
+ if (!yIsGpu && yLoBuffer) destroyBuffers(yLoBuffer);
162
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
163
+ // Only reached if extractTimestamp or extractResult threw before
164
+ // clearing these — on the success path they're already null.
165
+ if (xReadHiBuffer) destroyBuffers(xReadHiBuffer);
166
+ if (xReadLoBuffer) destroyBuffers(xReadLoBuffer);
167
+ if (yReadHiBuffer) destroyBuffers(yReadHiBuffer);
168
+ if (yReadLoBuffer) destroyBuffers(yReadLoBuffer);
169
+ }
170
+ }
@@ -0,0 +1,67 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Applies a modified Givens plane rotation H to double-precision vectors x
5
+ * and y:
6
+ * $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
7
+ * — double-double (Dekker) f64 emulation of {@link srotm}, since WGSL has
8
+ * no native f64 type.
9
+ *
10
+ * {@includeCode ../../examples/drotm/drotm.js}
11
+ *
12
+ * **Browser (standalone HTML):**
13
+ * {@includeCode ../../examples/drotm/web/drotm.html}
14
+ *
15
+ * @param device - GPUDevice from `init()`
16
+ * @param n - number of elements (must be a positive integer)
17
+ * @param x - Float64Array input/output vector
18
+ * @param incx - stride for x (must be a positive integer)
19
+ * @param y - Float64Array input/output vector
20
+ * @param incy - stride for y (must be a positive integer)
21
+ * @param param - 5-element Float64Array: [flag, h11, h21, h12, h22]
22
+ * flag = -2: identity (no-op), -1: full H, 0: unit diagonal, 1: unit off-diagonal
23
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/drotm/drotm.mjs#L18">Source code: drotm.mjs (L18)</a>
24
+ * @category BLAS Level 1
25
+ */
26
+ export declare function drotm(
27
+ device: GPUDevice,
28
+ n: number,
29
+ x: Float64Array,
30
+ incx: number,
31
+ y: Float64Array,
32
+ incy: number,
33
+ param: Float64Array,
34
+ ): Promise<
35
+ | { x: Float64Array; y: Float64Array }
36
+ | { x: Float64Array; y: Float64Array; gpuTimeMs: number }
37
+ >;
38
+
39
+ /**
40
+ * Applies a modified Givens plane rotation H to double-precision vectors x
41
+ * and y:
42
+ * $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
43
+ * — GPU-resident overload; see the Float64Array overload above for the
44
+ * routine itself.
45
+ *
46
+ * {@includeCode ../../examples/drotm/gpu.drotm.js}
47
+ *
48
+ * @param device - GPUDevice from `init()`
49
+ * @param n - number of elements (must be a positive integer)
50
+ * @param x - GpuVector input/output vector (must be Float64Array-backed, mutated in place)
51
+ * @param incx - stride for x (must be a positive integer)
52
+ * @param y - GpuVector input/output vector (must be Float64Array-backed, mutated in place)
53
+ * @param incy - stride for y (must be a positive integer)
54
+ * @param param - 5-element Float64Array: [flag, h11, h21, h12, h22]
55
+ * flag = -2: identity (no-op), -1: full H, 0: unit diagonal, 1: unit off-diagonal
56
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/drotm/drotm.mjs#L18">Source code: drotm.mjs (L18)</a>
57
+ * @category BLAS Level 1
58
+ */
59
+ export declare function drotm(
60
+ device: GPUDevice,
61
+ n: number,
62
+ x: GpuVector,
63
+ incx: number,
64
+ y: GpuVector,
65
+ incy: number,
66
+ param: Float64Array,
67
+ ): Promise<{} | { gpuTimeMs: number }>;
@@ -0,0 +1,171 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ } from "../util/buffer.mjs";
7
+ import { createBindGroup } from "../util/bindgroup.mjs";
8
+ import { runComputePass, submit } from "../util/compute.mjs";
9
+ import { extractTimestamp } from "../util/benchmark.mjs";
10
+ import { extractResult } from "../util/result.mjs";
11
+ import { getPipeline } from "../util/pipeline.mjs";
12
+ import { calcWorkgroups } from "../util/workgroup.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
16
+
17
+ // drotm: applies a modified Givens rotation H to vectors x and y —
18
+ // double-double (Dekker) f64 emulation of srotm. x, y, and every entry of
19
+ // param are each split into an f32 (hi, lo) pair; WGSL has no f64 type.
20
+ export async function drotm(device, n, x, incx, y, incy, param) {
21
+ const xIsGpu = x instanceof GpuVector;
22
+ const yIsGpu = y instanceof GpuVector;
23
+
24
+ requireGpuDevice(device);
25
+ requireSameDevice(device, "drotm", { x, y });
26
+ if (
27
+ !Number.isInteger(n) ||
28
+ !Number.isInteger(incx) ||
29
+ !Number.isInteger(incy)
30
+ )
31
+ throw new Error("n, incx, and incy must be integers.");
32
+ if (!(param instanceof Float64Array) || param.length !== 5)
33
+ throw new Error("param must be a Float64Array of length 5.");
34
+ if (param[0] !== -2 && param[0] !== -1 && param[0] !== 0 && param[0] !== 1)
35
+ throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
36
+ if (incx <= 0 || incy <= 0)
37
+ throw new Error("incx and incy must be positive.");
38
+ if (!(x instanceof Float64Array) && !xIsGpu)
39
+ throw new Error("x must be a Float64Array or GpuVector.");
40
+ if (!(y instanceof Float64Array) && !yIsGpu)
41
+ throw new Error("y must be a Float64Array or GpuVector.");
42
+ if (xIsGpu && x.dtype !== Float64Array)
43
+ throw new Error("x must be a Float64Array-backed GpuVector.");
44
+ if (yIsGpu && y.dtype !== Float64Array)
45
+ throw new Error("y must be a Float64Array-backed GpuVector.");
46
+ if (xIsGpu !== yIsGpu)
47
+ throw new Error(
48
+ "x and y must be the same type (both Float64Array or both GpuVector).",
49
+ );
50
+ if (n <= 0 || param[0] === -2.0) return xIsGpu ? {} : { x, y };
51
+ if (x.length < (n - 1) * incx + 1)
52
+ throw new Error(
53
+ "x does not have enough elements for the given n and incx.",
54
+ );
55
+ if (y.length < (n - 1) * incy + 1)
56
+ throw new Error(
57
+ "y does not have enough elements for the given n and incy.",
58
+ );
59
+
60
+ // Concatenated with f64/dekker.wgsl (DD struct), f64/utils/add.wgsl
61
+ // (fsub/negf/fastTwoSumProtected/ddAddProtected), and f64/utils/multiply.wgsl
62
+ // (ddMulProtected) — WGSL has no #include.
63
+ const f64Deps = ["f64/dekker", "f64/utils/add", "f64/utils/multiply"];
64
+ const pipeline = await getPipeline(device, [...f64Deps, "drotm"]);
65
+
66
+ const { hi: paramHi, lo: paramLo } = splitDoubleDouble(param);
67
+
68
+ let xHiBuffer = null;
69
+ let xLoBuffer = null;
70
+ let yHiBuffer = null;
71
+ let yLoBuffer = null;
72
+ let paramHiBuffer = null;
73
+ let paramLoBuffer = null;
74
+ let paramsBuffer = null;
75
+ let xReadHiBuffer = null;
76
+ let xReadLoBuffer = null;
77
+ let yReadHiBuffer = null;
78
+ let yReadLoBuffer = null;
79
+
80
+ try {
81
+ if (xIsGpu) {
82
+ xHiBuffer = x._buf;
83
+ xLoBuffer = x._loBuf;
84
+ yHiBuffer = y._buf;
85
+ yLoBuffer = y._loBuf;
86
+ } else {
87
+ const xSplit = splitDoubleDouble(x);
88
+ const ySplit = splitDoubleDouble(y);
89
+ xHiBuffer = uploadBuffer(device, xSplit.hi, "drotm-xHi", true);
90
+ xLoBuffer = uploadBuffer(device, xSplit.lo, "drotm-xLo", true);
91
+ yHiBuffer = uploadBuffer(device, ySplit.hi, "drotm-yHi", true);
92
+ yLoBuffer = uploadBuffer(device, ySplit.lo, "drotm-yLo", true);
93
+ }
94
+ paramHiBuffer = uploadBuffer(device, paramHi, "drotm-paramHi", false);
95
+ paramLoBuffer = uploadBuffer(device, paramLo, "drotm-paramLo", false);
96
+ paramsBuffer = createParamsBuffer(
97
+ device,
98
+ [
99
+ { value: n, type: "u32" },
100
+ { value: incx, type: "u32" },
101
+ { value: incy, type: "u32" },
102
+ ],
103
+ "drotm-params",
104
+ );
105
+
106
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
107
+ xHiBuffer,
108
+ xLoBuffer,
109
+ yHiBuffer,
110
+ yLoBuffer,
111
+ paramHiBuffer,
112
+ paramLoBuffer,
113
+ paramsBuffer,
114
+ ]);
115
+ const { commandEncoder, ts } = runComputePass(
116
+ device,
117
+ pipeline,
118
+ bindGroup,
119
+ calcWorkgroups(device, n),
120
+ );
121
+ xReadHiBuffer = xIsGpu
122
+ ? null
123
+ : stageReadback(device, commandEncoder, xHiBuffer);
124
+ xReadLoBuffer = xIsGpu
125
+ ? null
126
+ : stageReadback(device, commandEncoder, xLoBuffer);
127
+ yReadHiBuffer = yIsGpu
128
+ ? null
129
+ : stageReadback(device, commandEncoder, yHiBuffer);
130
+ yReadLoBuffer = yIsGpu
131
+ ? null
132
+ : stageReadback(device, commandEncoder, yLoBuffer);
133
+
134
+ submit(device, commandEncoder);
135
+
136
+ const gpuTimeMs = await extractTimestamp(ts);
137
+
138
+ if (xIsGpu) {
139
+ // xIsGpu === yIsGpu, enforced above
140
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
141
+ return {};
142
+ }
143
+
144
+ const xHi = await extractResult(xReadHiBuffer, Float32Array);
145
+ xReadHiBuffer = null; // extractResult already destroyed it
146
+ const xLo = await extractResult(xReadLoBuffer, Float32Array);
147
+ xReadLoBuffer = null;
148
+ const yHi = await extractResult(yReadHiBuffer, Float32Array);
149
+ yReadHiBuffer = null;
150
+ const yLo = await extractResult(yReadLoBuffer, Float32Array);
151
+ yReadLoBuffer = null;
152
+ const resultX = mergeDoubleDouble(xHi, xLo);
153
+ const resultY = mergeDoubleDouble(yHi, yLo);
154
+ if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
155
+ return { x: resultX, y: resultY };
156
+ } finally {
157
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
158
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
159
+ if (!yIsGpu && yHiBuffer) destroyBuffers(yHiBuffer);
160
+ if (!yIsGpu && yLoBuffer) destroyBuffers(yLoBuffer);
161
+ if (paramHiBuffer) destroyBuffers(paramHiBuffer);
162
+ if (paramLoBuffer) destroyBuffers(paramLoBuffer);
163
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
164
+ // Only reached if extractTimestamp or extractResult threw before
165
+ // clearing these — on the success path they're already null.
166
+ if (xReadHiBuffer) destroyBuffers(xReadHiBuffer);
167
+ if (xReadLoBuffer) destroyBuffers(xReadLoBuffer);
168
+ if (yReadHiBuffer) destroyBuffers(yReadHiBuffer);
169
+ if (yReadLoBuffer) destroyBuffers(yReadLoBuffer);
170
+ }
171
+ }
@@ -0,0 +1,52 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Scales a double-precision vector by a constant: $$x \leftarrow \alpha x$$
5
+ *
6
+ * `x` and `alpha` are each split into a (hi, lo) double-double f32 pair (see
7
+ * `splitDoubleDouble`/`f64.mjs`) since WGSL has no f64 type; the multiply
8
+ * uses Dekker's double-double algorithm (see `shaders/f64/`), giving ~48
9
+ * bits of mantissa — more than a single f32 (24 bits) but less than true
10
+ * f64 (52 bits), so results are not bit-exact with a CPU double.
11
+ *
12
+ * {@includeCode ../../examples/dscal/dscal.js}
13
+ *
14
+ * **Browser (standalone HTML):**
15
+ * {@includeCode ../../examples/dscal/web/dscal.html}
16
+ *
17
+ * @param device - GPUDevice from `init()`
18
+ * @param n - number of elements to scale (must be a positive integer)
19
+ * @param alpha - scalar multiplier
20
+ * @param x - Float64Array input/output vector
21
+ * @param incx - stride for x (must be a positive integer)
22
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dscal/dscal.mjs">Source code: dscal.mjs</a>
23
+ * @category BLAS Level 1
24
+ */
25
+ export declare function dscal(
26
+ device: GPUDevice,
27
+ n: number,
28
+ alpha: number,
29
+ x: Float64Array,
30
+ incx: number,
31
+ ): Promise<{ x: Float64Array } | { x: Float64Array; gpuTimeMs: number }>;
32
+
33
+ /**
34
+ * Scales a double-precision vector by a constant: $$x \leftarrow \alpha x$$
35
+ *
36
+ * {@includeCode ../../examples/dscal/gpu.dscal.js}
37
+ *
38
+ * @param device - GPUDevice from `init()`
39
+ * @param n - number of elements to scale (must be a positive integer)
40
+ * @param alpha - scalar multiplier
41
+ * @param x - Float64Array-backed GpuVector input/output vector (mutated in place)
42
+ * @param incx - stride for x (must be a positive integer)
43
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dscal/dscal.mjs">Source code: dscal.mjs</a>
44
+ * @category BLAS Level 1
45
+ */
46
+ export declare function dscal(
47
+ device: GPUDevice,
48
+ n: number,
49
+ alpha: number,
50
+ x: GpuVector,
51
+ incx: number,
52
+ ): Promise<{} | { gpuTimeMs: number }>;
@@ -0,0 +1,119 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ } from "../util/buffer.mjs";
7
+ import { createBindGroup } from "../util/bindgroup.mjs";
8
+ import { runComputePass, submit } from "../util/compute.mjs";
9
+ import { extractResult } from "../util/result.mjs";
10
+ import { extractTimestamp } from "../util/benchmark.mjs";
11
+ import { getPipeline } from "../util/pipeline.mjs";
12
+ import { calcWorkgroups } from "../util/workgroup.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
16
+
17
+ // dscal: x := alpha * x, double-double (Dekker) f64 emulation of sscal — x
18
+ // and alpha are each split into an f32 (hi, lo) pair; WGSL has no f64 type.
19
+ export async function dscal(device, n, alpha, x, incx) {
20
+ const xIsGpu = x instanceof GpuVector;
21
+
22
+ requireGpuDevice(device);
23
+ if (!Number.isInteger(n) || !Number.isInteger(incx))
24
+ throw new Error("n and incx must be integers.");
25
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
26
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
27
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
28
+ if (!(x instanceof Float64Array) && !xIsGpu)
29
+ throw new Error("x must be a Float64Array or GpuVector.");
30
+ if (xIsGpu && x.dtype !== Float64Array)
31
+ throw new Error("x must be a Float64Array-backed GpuVector.");
32
+ if (incx <= 0) throw new Error("incx must be positive.");
33
+ requireSameDevice(device, "dscal", { x });
34
+ if (n <= 0) return xIsGpu ? {} : { x };
35
+ if (x.length < (n - 1) * incx + 1)
36
+ throw new Error(
37
+ "x does not have enough elements for the given n and incx.",
38
+ );
39
+
40
+ // Concatenated with f64/dekker.wgsl (DD struct), f64/utils/add.wgsl
41
+ // (fsub/negf/fastTwoSumProtected), and f64/utils/multiply.wgsl
42
+ // (ddMulProtected) — WGSL has no #include.
43
+ const f64Deps = ["f64/dekker", "f64/utils/add", "f64/utils/multiply"];
44
+ const pipeline = await getPipeline(device, [...f64Deps, "dscal"]);
45
+
46
+ const { hi: alphaHi, lo: alphaLo } = splitDoubleDouble(
47
+ new Float64Array([alpha]),
48
+ );
49
+
50
+ let xHiBuffer = null;
51
+ let xLoBuffer = null;
52
+ let paramsBuffer = null;
53
+ let readHiBuffer = null;
54
+ let readLoBuffer = null;
55
+
56
+ try {
57
+ if (xIsGpu) {
58
+ xHiBuffer = x._buf;
59
+ xLoBuffer = x._loBuf;
60
+ } else {
61
+ const { hi, lo } = splitDoubleDouble(x);
62
+ xHiBuffer = uploadBuffer(device, hi, "dscal-xHi", true);
63
+ xLoBuffer = uploadBuffer(device, lo, "dscal-xLo", true);
64
+ }
65
+ paramsBuffer = createParamsBuffer(
66
+ device,
67
+ [
68
+ { value: n, type: "u32" },
69
+ { value: alphaHi[0], type: "f32" },
70
+ { value: alphaLo[0], type: "f32" },
71
+ { value: incx, type: "u32" },
72
+ ],
73
+ "dscal-params",
74
+ );
75
+
76
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
77
+ xHiBuffer,
78
+ xLoBuffer,
79
+ paramsBuffer,
80
+ ]);
81
+ const { commandEncoder, ts } = runComputePass(
82
+ device,
83
+ pipeline,
84
+ bindGroup,
85
+ calcWorkgroups(device, n),
86
+ );
87
+ readHiBuffer = xIsGpu
88
+ ? null
89
+ : stageReadback(device, commandEncoder, xHiBuffer);
90
+ readLoBuffer = xIsGpu
91
+ ? null
92
+ : stageReadback(device, commandEncoder, xLoBuffer);
93
+
94
+ submit(device, commandEncoder);
95
+
96
+ const gpuTimeMs = await extractTimestamp(ts);
97
+
98
+ if (xIsGpu) {
99
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
100
+ return {};
101
+ }
102
+
103
+ const hi = await extractResult(readHiBuffer, Float32Array);
104
+ readHiBuffer = null; // extractResult already destroyed it
105
+ const lo = await extractResult(readLoBuffer, Float32Array);
106
+ readLoBuffer = null;
107
+ const result = mergeDoubleDouble(hi, lo);
108
+ if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
109
+ return { x: result };
110
+ } finally {
111
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
112
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
113
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
114
+ // Only reached if extractTimestamp or extractResult threw before
115
+ // clearing these — on the success path they're already null.
116
+ if (readHiBuffer) destroyBuffers(readHiBuffer);
117
+ if (readLoBuffer) destroyBuffers(readLoBuffer);
118
+ }
119
+ }
@@ -0,0 +1,57 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Swaps the elements of two double-precision vectors: $$x \leftrightarrow y$$
5
+ * — double-double (Dekker) f64 emulation of {@link sswap}, since WGSL has
6
+ * no native f64 type.
7
+ *
8
+ * {@includeCode ../../examples/dswap/dswap.js}
9
+ *
10
+ * **Browser (standalone HTML):**
11
+ * {@includeCode ../../examples/dswap/web/dswap.html}
12
+ *
13
+ * @param device - GPUDevice from `init()`
14
+ * @param n - number of elements to swap (must be a positive integer)
15
+ * @param x - Float64Array first input/output vector
16
+ * @param incx - stride for x (must be a positive integer)
17
+ * @param y - Float64Array second input/output vector
18
+ * @param incy - stride for y (must be a positive integer)
19
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dswap/dswap.mjs#L22">Source code: dswap.mjs (L22)</a>
20
+ * @category BLAS Level 1
21
+ */
22
+ export declare function dswap(
23
+ device: GPUDevice,
24
+ n: number,
25
+ x: Float64Array,
26
+ incx: number,
27
+ y: Float64Array,
28
+ incy: number,
29
+ ): Promise<
30
+ | { x: Float64Array; y: Float64Array }
31
+ | { x: Float64Array; y: Float64Array; gpuTimeMs: number }
32
+ >;
33
+
34
+ /**
35
+ * Swaps the elements of two double-precision vectors: $$x \leftrightarrow y$$
36
+ * — GPU-resident overload; see the Float64Array overload above for the
37
+ * routine itself.
38
+ *
39
+ * {@includeCode ../../examples/dswap/gpu.dswap.js}
40
+ *
41
+ * @param device - GPUDevice from `init()`
42
+ * @param n - number of elements to swap (must be a positive integer)
43
+ * @param x - GpuVector first input/output vector (must be Float64Array-backed, mutated in place)
44
+ * @param incx - stride for x (must be a positive integer)
45
+ * @param y - GpuVector second input/output vector (must be Float64Array-backed, mutated in place)
46
+ * @param incy - stride for y (must be a positive integer)
47
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dswap/dswap.mjs#L22">Source code: dswap.mjs (L22)</a>
48
+ * @category BLAS Level 1
49
+ */
50
+ export declare function dswap(
51
+ device: GPUDevice,
52
+ n: number,
53
+ x: GpuVector,
54
+ incx: number,
55
+ y: GpuVector,
56
+ incy: number,
57
+ ): Promise<{} | { gpuTimeMs: number }>;