@driftengine/texture 4.0.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 (99) hide show
  1. package/LICENSE +202 -0
  2. package/NOTICE +29 -0
  3. package/README.md +106 -0
  4. package/dist/decodeCpu.d.ts +59 -0
  5. package/dist/decodeCpu.js +234 -0
  6. package/dist/decodeGraph.d.ts +105 -0
  7. package/dist/decodeGraph.js +180 -0
  8. package/dist/half.d.ts +24 -0
  9. package/dist/half.js +86 -0
  10. package/dist/index.d.ts +66 -0
  11. package/dist/index.js +55 -0
  12. package/dist/inference.d.ts +53 -0
  13. package/dist/inference.js +243 -0
  14. package/dist/materialArray.d.ts +38 -0
  15. package/dist/materialArray.js +40 -0
  16. package/dist/mipNdf.d.ts +29 -0
  17. package/dist/mipNdf.js +53 -0
  18. package/dist/overlay/journal.d.ts +78 -0
  19. package/dist/overlay/journal.js +171 -0
  20. package/dist/overlay/sparse.d.ts +68 -0
  21. package/dist/overlay/sparse.js +212 -0
  22. package/dist/progressive.d.ts +30 -0
  23. package/dist/progressive.js +56 -0
  24. package/dist/residency/pageCache.d.ts +103 -0
  25. package/dist/residency/pageCache.js +184 -0
  26. package/dist/residency/predict.d.ts +55 -0
  27. package/dist/residency/predict.js +51 -0
  28. package/dist/residency/predictor.d.ts +16 -0
  29. package/dist/residency/predictor.js +44 -0
  30. package/dist/residency/queue.d.ts +26 -0
  31. package/dist/residency/queue.js +52 -0
  32. package/dist/residency/stream.d.ts +66 -0
  33. package/dist/residency/stream.js +142 -0
  34. package/dist/residency/table.d.ts +36 -0
  35. package/dist/residency/table.js +72 -0
  36. package/dist/residency/viewTiles.d.ts +108 -0
  37. package/dist/residency/viewTiles.js +419 -0
  38. package/dist/semantics.d.ts +52 -0
  39. package/dist/semantics.js +76 -0
  40. package/dist/tensor/architecture.d.ts +53 -0
  41. package/dist/tensor/architecture.js +96 -0
  42. package/dist/tensor/attention.d.ts +5 -0
  43. package/dist/tensor/attention.js +62 -0
  44. package/dist/tensor/denseOperators.d.ts +2 -0
  45. package/dist/tensor/denseOperators.js +136 -0
  46. package/dist/tensor/graph.d.ts +83 -0
  47. package/dist/tensor/graph.js +175 -0
  48. package/dist/tensor/linear.d.ts +49 -0
  49. package/dist/tensor/linear.js +136 -0
  50. package/dist/tensor/operatorKit.d.ts +27 -0
  51. package/dist/tensor/operatorKit.js +45 -0
  52. package/dist/tensor/operators.d.ts +3 -0
  53. package/dist/tensor/operators.js +24 -0
  54. package/dist/tensor/resize.d.ts +6 -0
  55. package/dist/tensor/resize.js +107 -0
  56. package/dist/tensor/reuse.d.ts +33 -0
  57. package/dist/tensor/reuse.js +59 -0
  58. package/dist/tensor/shapeOperators.d.ts +3 -0
  59. package/dist/tensor/shapeOperators.js +173 -0
  60. package/dist/tensor/spatial.d.ts +34 -0
  61. package/dist/tensor/spatial.js +131 -0
  62. package/dist/tensor/spatialOperators.d.ts +2 -0
  63. package/dist/tensor/spatialOperators.js +138 -0
  64. package/dist/tileHash.d.ts +29 -0
  65. package/dist/tileHash.js +50 -0
  66. package/dist/timeNodes.d.ts +26 -0
  67. package/dist/timeNodes.js +48 -0
  68. package/package.json +59 -0
  69. package/src/decodeCpu.ts +308 -0
  70. package/src/decodeGraph.ts +214 -0
  71. package/src/half.ts +86 -0
  72. package/src/index.ts +175 -0
  73. package/src/inference.ts +278 -0
  74. package/src/materialArray.ts +67 -0
  75. package/src/mipNdf.ts +63 -0
  76. package/src/overlay/journal.ts +218 -0
  77. package/src/overlay/sparse.ts +275 -0
  78. package/src/progressive.ts +60 -0
  79. package/src/residency/pageCache.ts +233 -0
  80. package/src/residency/predict.ts +74 -0
  81. package/src/residency/predictor.ts +62 -0
  82. package/src/residency/queue.ts +78 -0
  83. package/src/residency/stream.ts +194 -0
  84. package/src/residency/table.ts +89 -0
  85. package/src/residency/viewTiles.ts +553 -0
  86. package/src/semantics.ts +114 -0
  87. package/src/tensor/architecture.ts +140 -0
  88. package/src/tensor/attention.ts +75 -0
  89. package/src/tensor/denseOperators.ts +153 -0
  90. package/src/tensor/graph.ts +244 -0
  91. package/src/tensor/linear.ts +153 -0
  92. package/src/tensor/operatorKit.ts +76 -0
  93. package/src/tensor/operators.ts +28 -0
  94. package/src/tensor/resize.ts +140 -0
  95. package/src/tensor/shapeOperators.ts +173 -0
  96. package/src/tensor/spatial.ts +178 -0
  97. package/src/tensor/spatialOperators.ts +182 -0
  98. package/src/tileHash.ts +60 -0
  99. package/src/timeNodes.ts +60 -0
@@ -0,0 +1,114 @@
1
+ /**
2
+ * What a channel *is*, declared rather than conventional.
3
+ *
4
+ * **The family of defects this removes is the one nobody attributes correctly.** A normal map
5
+ * upside down in one asset and not in another; an albedo that is linear here and sRGB there; a
6
+ * gloss map read as roughness, which inverts every highlight in the scene. Each is a convention
7
+ * held in somebody's head, and each survives review because the file looks fine on its own.
8
+ *
9
+ * Declared, the format knows. A consumer never sees a convention at all, because `normaliseSample`
10
+ * has already applied it — a Y-down normal comes back Y-up, an sRGB value comes back linear, and
11
+ * gloss comes back as the roughness the engine shades with.
12
+ */
13
+ export type ChannelSemantic =
14
+ | 'albedo-srgb'
15
+ | 'albedo-linear'
16
+ | 'normal-tangent-yup'
17
+ | 'normal-tangent-ydown'
18
+ | 'roughness-linear'
19
+ | 'gloss-linear'
20
+ | 'metallic-linear'
21
+ | 'occlusion-linear'
22
+ | 'height-linear'
23
+ | 'emissive-srgb'
24
+ | 'mask-linear';
25
+
26
+ /**
27
+ * Every semantic, in the order a file stores them by.
28
+ *
29
+ * **A name crosses a package boundary as a number**, because `@driftengine/drft` is the container
30
+ * and knows nothing about what a channel means — a `DTEX` chunk stores `(semanticIndex << 4) |
31
+ * component` and would have to carry strings otherwise. The index is therefore part of the format:
32
+ * **a semantic may be appended and none may be reordered or removed**, or every file written before
33
+ * the change decodes its channels as something else. `semantics.test.ts` holds the order.
34
+ */
35
+ export const CHANNEL_SEMANTICS: readonly ChannelSemantic[] = [
36
+ 'albedo-srgb',
37
+ 'albedo-linear',
38
+ 'normal-tangent-yup',
39
+ 'normal-tangent-ydown',
40
+ 'roughness-linear',
41
+ 'gloss-linear',
42
+ 'metallic-linear',
43
+ 'occlusion-linear',
44
+ 'height-linear',
45
+ 'emissive-srgb',
46
+ 'mask-linear',
47
+ ];
48
+
49
+ /** Where a semantic sits in that order, or −1 for one this build does not know. */
50
+ export function semanticIndex(semantic: ChannelSemantic): number {
51
+ return CHANNEL_SEMANTICS.indexOf(semantic);
52
+ }
53
+
54
+ /** The semantic an index names, or null where a file names one from a later version. */
55
+ export function semanticAt(index: number): ChannelSemantic | null {
56
+ return CHANNEL_SEMANTICS[index] ?? null;
57
+ }
58
+
59
+ export interface ChannelSpec {
60
+ readonly semantic: ChannelSemantic;
61
+ /** Which component of the decoded vector this channel occupies. */
62
+ readonly component: number;
63
+ }
64
+
65
+ /** Whether this channel carries colour, and therefore whether a transfer curve applies to it. */
66
+ export function isColour(semantic: ChannelSemantic): boolean {
67
+ return semantic === 'albedo-srgb' || semantic === 'albedo-linear' || semantic === 'emissive-srgb';
68
+ }
69
+
70
+ /** Whether mipping this channel must preserve the normal distribution. See `mipNdf.ts`. */
71
+ export function needsVarianceMips(semantic: ChannelSemantic): boolean {
72
+ return semantic === 'normal-tangent-yup' || semantic === 'normal-tangent-ydown';
73
+ }
74
+
75
+ /**
76
+ * The exact sRGB transfer function, piecewise, not the 2.2 power approximation.
77
+ *
78
+ * The approximation is within about two percent almost everywhere and wrong by more than that in
79
+ * the dark end, which in a texture pipeline is a colour shift nobody attributes to the right cause
80
+ * for weeks. The midpoint is the value that tells them apart: 0.5 linearises to 0.2140 exactly and
81
+ * to 0.2176 approximately.
82
+ */
83
+ export function srgbToLinear(value: number): number {
84
+ return value <= 0.04045 ? value / 12.92 : Math.pow((value + 0.055) / 1.055, 2.4);
85
+ }
86
+
87
+ export function linearToSrgb(value: number): number {
88
+ return value <= 0.0031308 ? value * 12.92 : 1.055 * Math.pow(value, 1 / 2.4) - 0.055;
89
+ }
90
+
91
+ /**
92
+ * Write one raw channel value into its component of `out`, with its convention already applied.
93
+ *
94
+ * A consumer of this never asks which way a normal points or which curve an albedo carries.
95
+ */
96
+ export function normaliseSample(out: Float32Array, spec: ChannelSpec, raw: number): void {
97
+ const at = spec.component;
98
+ switch (spec.semantic) {
99
+ case 'albedo-srgb':
100
+ case 'emissive-srgb':
101
+ out[at] = srgbToLinear(raw);
102
+ return;
103
+ case 'normal-tangent-ydown':
104
+ /* Flipped about the midpoint, because a tangent-space normal is stored biased into 0..1. */
105
+ out[at] = 1 - raw;
106
+ return;
107
+ case 'gloss-linear':
108
+ /* The engine shades with roughness. One of the two has to win, and it is not this one. */
109
+ out[at] = 1 - raw;
110
+ return;
111
+ default:
112
+ out[at] = raw;
113
+ }
114
+ }
@@ -0,0 +1,140 @@
1
+ /**
2
+ * A network's definition, written as a function of its weights, built into a graph the runtime runs.
3
+ *
4
+ * **One definition, run over two sources.** At conversion it reads a checkpoint; afterwards it reads
5
+ * the converted file, whenever the runtime needs the graph at a new size — a vision transformer's
6
+ * patch grid follows the image's aspect ratio, and every device kernel bakes its shapes, so a model
7
+ * is one graph per size and not one graph. The two runs agree because of three rules:
8
+ *
9
+ * - **A weight is read by name, at the shape the definition expects**, and a name the source lacks
10
+ * or a shape it disagrees with is refused naming both.
11
+ * - **A weight derived from others** — a batch norm folded into its convolution — is computed once,
12
+ * at conversion, and kept under its own name; a source that already holds that name answers with
13
+ * it and the computation never runs, which is why a converted file needs none of the parts.
14
+ * - **A constant** — a table computed from the shapes, such as rotary angles — carries a name
15
+ * beginning `@`, is rebuilt on every run, and is never taken for a weight.
16
+ *
17
+ * **And every weight the source holds is read, derived from, or set aside by name.** A forgotten
18
+ * bias or layer scale is a network that runs and answers wrongly — the one failure nothing later can
19
+ * see — so a weight left over is refused. What that costs is an explicit `ignore` for every head a
20
+ * definition does not run, which is the point.
21
+ */
22
+ import { validateGraph, type GraphTensor, type GraphValue, type NetworkGraph } from './graph.ts';
23
+ import type { Attributes } from './operators.ts';
24
+
25
+ /** Where a definition's weights come from: a checkpoint being converted, or a converted file. */
26
+ export interface WeightSource {
27
+ get(name: string): GraphTensor | undefined;
28
+ names(): Iterable<string>;
29
+ }
30
+
31
+ /** What a constant's name begins with: a table of the shapes, never a weight. */
32
+ export const CONSTANT_PREFIX = '@';
33
+
34
+ export interface Weights {
35
+ /** The weight `name` as a value of the graph; refused if absent or held at another shape. */
36
+ read(name: string, shape?: readonly number[]): string;
37
+ /**
38
+ * A weight computed once from others, kept under `name`. `compute` reads its parts through
39
+ * `values`, which counts each as read; a source already holding `name` answers with it instead.
40
+ */
41
+ derive(
42
+ name: string,
43
+ shape: readonly number[],
44
+ compute: (values: (part: string, shape?: readonly number[]) => Float32Array) => Float32Array,
45
+ ): string;
46
+ /** A table computed from the graph's shapes, rebuilt on every run and never a weight. */
47
+ constant(name: string, shape: readonly number[], data: Float32Array): string;
48
+ /** Weights the definition deliberately does not read: a name, or a prefix ending in a dot. */
49
+ ignore(nameOrPrefix: string): void;
50
+ }
51
+
52
+ export interface GraphBuilder {
53
+ /** A node; its output is named `output` when given, and otherwise by the builder. */
54
+ node(op: string, inputs: readonly string[], attributes?: Attributes, output?: string): string;
55
+ }
56
+
57
+ /** A definition: it reads its weights, adds its nodes, and says what goes in and comes out. */
58
+ export type Architecture = (
59
+ weights: Weights,
60
+ graph: GraphBuilder,
61
+ ) => { readonly inputs: readonly GraphValue[]; readonly outputs: readonly string[] };
62
+
63
+ const shapeText = (shape: readonly number[]): string => `[${shape.join(', ')}]`;
64
+
65
+ export function graphFromWeights(source: WeightSource, architecture: Architecture): NetworkGraph {
66
+ const tensors = new Map<string, GraphTensor>();
67
+ const read = new Set<string>();
68
+ const ignored: string[] = [];
69
+ const nodes: NetworkGraph['nodes'][number][] = [];
70
+ const take = (name: string, shape?: readonly number[]): GraphTensor => {
71
+ const tensor = source.get(name);
72
+ if (tensor === undefined) {
73
+ throw new Error(`the source has no weight "${name}", which the definition reads`);
74
+ }
75
+ if (shape !== undefined && shapeText(shape) !== shapeText(tensor.shape)) {
76
+ throw new Error(
77
+ `"${name}" is ${shapeText(tensor.shape)} in the source and the definition reads it as ` +
78
+ shapeText(shape),
79
+ );
80
+ }
81
+ read.add(name);
82
+ return tensor;
83
+ };
84
+ const weights: Weights = {
85
+ read(name, shape) {
86
+ if (!tensors.has(name)) tensors.set(name, take(name, shape));
87
+ else take(name, shape);
88
+ return name;
89
+ },
90
+ derive(name, shape, compute) {
91
+ if (tensors.has(name)) return name;
92
+ const held = source.get(name);
93
+ const data =
94
+ held === undefined
95
+ ? compute((part, partShape) => take(part, partShape).data)
96
+ : take(name, shape).data;
97
+ if (data.length !== shape.reduce((total, d) => total * d, 1)) {
98
+ throw new Error(
99
+ `the derived weight "${name}" has ${data.length} values for ${shapeText(shape)}`,
100
+ );
101
+ }
102
+ tensors.set(name, { shape, data });
103
+ return name;
104
+ },
105
+ constant(name, shape, data) {
106
+ const key = `${CONSTANT_PREFIX}${name}`;
107
+ tensors.set(key, { shape, data });
108
+ return key;
109
+ },
110
+ ignore(nameOrPrefix) {
111
+ ignored.push(nameOrPrefix);
112
+ },
113
+ };
114
+ let count = 0;
115
+ const graph: GraphBuilder = {
116
+ node(op, inputs, attributes = {}, output = `v${(count += 1)}`) {
117
+ nodes.push({ op, inputs, output, attributes });
118
+ return output;
119
+ },
120
+ };
121
+ const { inputs, outputs } = architecture(weights, graph);
122
+
123
+ const aside = (name: string): boolean =>
124
+ name.startsWith(CONSTANT_PREFIX) ||
125
+ ignored.some((rule) => (rule.endsWith('.') ? name.startsWith(rule) : name === rule));
126
+ const unread = [...source.names()].filter((name) => !read.has(name) && !aside(name));
127
+ if (unread.length > 0) {
128
+ const named = unread.slice(0, 5).map((name) => `"${name}"`);
129
+ throw new Error(
130
+ `the definition never reads ${named.join(', ')}` +
131
+ (unread.length > 5 ? ` and ${unread.length - 5} more` : '') +
132
+ ': a weight left behind is a network that runs and answers wrongly, so read it or set it ' +
133
+ 'aside by name',
134
+ );
135
+ }
136
+ const network: NetworkGraph = { inputs, outputs, nodes, tensors };
137
+ const problem = validateGraph(network);
138
+ if (problem !== null) throw new Error(`the definition's graph cannot run: ${problem}`);
139
+ return network;
140
+ }
@@ -0,0 +1,75 @@
1
+ /**
2
+ * Multi-head attention, as a composition of the dense operators.
3
+ *
4
+ * **Heads are contiguous runs of channels, in the upstream order**: head `h` of `heads` owns
5
+ * channels `h·d` to `(h+1)·d − 1`, where `d` is the channel count over the head count. A port that
6
+ * interleaves them instead evaluates a network nobody trained.
7
+ *
8
+ * The queries, keys and values arrive already projected — the projections are matrix multiplies a
9
+ * graph states on its own — so this is the part that is attention's alone: scores scaled by
10
+ * `1/√d`, a stable softmax over each query's row, and the weighted sum of values.
11
+ *
12
+ * **A bias is added to the scaled scores**, `[heads][queries][keys]`, as a relative-position table
13
+ * or a mask is upstream; and **a batch is independent windows**, each attending over its own keys
14
+ * with the one bias. Queries and keys may differ in number, as they do where one set of tokens
15
+ * reads another.
16
+ */
17
+ import { softmax } from './linear.ts';
18
+
19
+ /**
20
+ * `out[batch][queries][channels]` from `q` of that shape and `k` and `v` of
21
+ * `[batch][keys][channels]`. `scratch` holds at least `queries × keys` values and is overwritten.
22
+ */
23
+ export function attention(
24
+ out: Float32Array,
25
+ q: Float32Array,
26
+ k: Float32Array,
27
+ v: Float32Array,
28
+ queries: number,
29
+ keys: number,
30
+ channels: number,
31
+ heads: number,
32
+ scratch: Float32Array,
33
+ bias: Float32Array | null = null,
34
+ batch = 1,
35
+ ): void {
36
+ const d = channels / heads;
37
+ if (!Number.isInteger(d)) {
38
+ throw new RangeError(`attention: ${channels} channels do not split into ${heads} heads`);
39
+ }
40
+ const scale = 1 / Math.sqrt(d);
41
+ for (let b = 0; b < batch; b += 1) {
42
+ const qAt = b * queries * channels;
43
+ const kAt = b * keys * channels;
44
+ for (let head = 0; head < heads; head += 1) {
45
+ const first = head * d;
46
+ for (let query = 0; query < queries; query += 1) {
47
+ for (let key = 0; key < keys; key += 1) {
48
+ let dot = 0;
49
+ for (let c = 0; c < d; c += 1) {
50
+ dot +=
51
+ (q[qAt + query * channels + first + c] as number) *
52
+ (k[kAt + key * channels + first + c] as number);
53
+ }
54
+ const at = query * keys + key;
55
+ scratch[at] = dot * scale;
56
+ if (bias !== null) {
57
+ scratch[at] = (scratch[at] as number) + (bias[head * queries * keys + at] as number);
58
+ }
59
+ }
60
+ }
61
+ softmax(scratch, queries, keys, scratch);
62
+ for (let query = 0; query < queries; query += 1) {
63
+ for (let c = 0; c < d; c += 1) {
64
+ let sum = 0;
65
+ for (let key = 0; key < keys; key += 1) {
66
+ sum +=
67
+ (scratch[query * keys + key] as number) *
68
+ (v[kAt + key * channels + first + c] as number);
69
+ }
70
+ out[qAt + query * channels + first + c] = sum;
71
+ }
72
+ }
73
+ }
74
+ }
75
+ }
@@ -0,0 +1,153 @@
1
+ /** The dense operators' rows: projections, activations, normalisation, attention. */
2
+ import { attention } from './attention.ts';
3
+ import { gelu, layerNorm, sigmoid, softmax } from './linear.ts';
4
+ import { type Operator, elementwise, num, product, same, type Shape } from './operatorKit.ts';
5
+
6
+ export const DENSE_OPERATORS: readonly (readonly [string, Operator])[] = [
7
+ [
8
+ 'linear',
9
+ {
10
+ ranks: [2, 2, 1],
11
+ arity: [2, 3],
12
+ shape: ([x, w, b]) => {
13
+ const [rows, width] = x as Shape;
14
+ const [outs, ins] = w as Shape;
15
+ if (width !== ins) return `x has ${width} columns and the weight takes ${ins}`;
16
+ if (b !== undefined && b[0] !== outs) return `the bias has ${b[0]} and the weight ${outs}`;
17
+ return [rows as number, outs as number];
18
+ },
19
+ evaluate: ([x, w, b], [xs, ws], _attributes, out) => {
20
+ const [rows, ins] = xs as Shape as [number, number];
21
+ const outs = (ws as Shape)[0] as number;
22
+ for (let row = 0; row < rows; row += 1) {
23
+ for (let o = 0; o < outs; o += 1) {
24
+ let dot = 0;
25
+ for (let i = 0; i < ins; i += 1) {
26
+ dot +=
27
+ ((x as Float32Array)[row * ins + i] as number) *
28
+ ((w as Float32Array)[o * ins + i] as number);
29
+ }
30
+ /* The bias joins the double-precision sum, and the result is rounded once. */
31
+ out[row * outs + o] = dot + (b === undefined ? 0 : (b[o] as number));
32
+ }
33
+ }
34
+ },
35
+ },
36
+ ],
37
+ [
38
+ 'relu',
39
+ {
40
+ arity: [1, 1],
41
+ shape: ([x]) => [...(x as Shape)],
42
+ evaluate: ([x], _shapes, _attributes, out) => {
43
+ for (let i = 0; i < out.length; i += 1)
44
+ out[i] = Math.max(0, (x as Float32Array)[i] as number);
45
+ },
46
+ },
47
+ ],
48
+ [
49
+ 'gelu',
50
+ {
51
+ arity: [1, 1],
52
+ shape: ([x]) => [...(x as Shape)],
53
+ evaluate: ([x], _shapes, _attributes, out) => gelu(x as Float32Array, out),
54
+ },
55
+ ],
56
+ [
57
+ 'sigmoid',
58
+ {
59
+ arity: [1, 1],
60
+ shape: ([x]) => [...(x as Shape)],
61
+ evaluate: ([x], _shapes, _attributes, out) => sigmoid(x as Float32Array, out),
62
+ },
63
+ ],
64
+ ['add', elementwise((a, b) => a + b)],
65
+ ['mul', elementwise((a, b) => a * b)],
66
+ [
67
+ 'layerNorm',
68
+ {
69
+ ranks: [2, 1, 1],
70
+ arity: [3, 3],
71
+ shape: ([x, gamma]) => {
72
+ const cols = (x as Shape)[1];
73
+ return gamma?.[0] === cols ? [...(x as Shape)] : `gamma has ${gamma?.[0]} and x ${cols}`;
74
+ },
75
+ evaluate: ([x, gamma, beta], [xs], attributes, out) => {
76
+ const [rows, cols] = xs as Shape as [number, number];
77
+ const epsilon = num(attributes, 'epsilon', 1e-6);
78
+ layerNorm(
79
+ x as Float32Array,
80
+ rows,
81
+ cols,
82
+ gamma as Float32Array,
83
+ beta as Float32Array,
84
+ epsilon,
85
+ out,
86
+ );
87
+ },
88
+ },
89
+ ],
90
+ [
91
+ 'softmax',
92
+ {
93
+ arity: [1, 1],
94
+ shape: ([x]) => [...(x as Shape)],
95
+ evaluate: ([x], [xs], _attributes, out) => {
96
+ const shape = xs as Shape;
97
+ const cols = shape[shape.length - 1] as number;
98
+ softmax(x as Float32Array, product(shape) / cols, cols, out);
99
+ },
100
+ },
101
+ ],
102
+ [
103
+ 'attention',
104
+ {
105
+ /*
106
+ * Rank 2, or 3 for a batch of windows; checked here rather than by `ranks`, which states one.
107
+ * The bias, where there is one, is a score for each head, query and key.
108
+ */
109
+ arity: [3, 4],
110
+ shape: ([q, k, v, b], attributes) => {
111
+ const qs = q as Shape;
112
+ const ks = k as Shape;
113
+ if (qs.length !== 2 && qs.length !== 3)
114
+ return `the queries have rank ${qs.length}, not 2 or 3`;
115
+ if (ks.length !== qs.length)
116
+ return `the keys have rank ${ks.length} and the queries ${qs.length}`;
117
+ if (!same(ks, v as Shape)) return 'the keys and values differ in shape';
118
+ if (qs.length === 3 && qs[0] !== ks[0]) {
119
+ return `a batch of ${qs[0]} queries and of ${ks[0]} keys`;
120
+ }
121
+ const [queries, channels] = qs.slice(-2) as [number, number];
122
+ const [keys, keyChannels] = ks.slice(-2) as [number, number];
123
+ if (channels !== keyChannels) {
124
+ return `the queries have ${channels} channels and the keys ${keyChannels}`;
125
+ }
126
+ const heads = num(attributes, 'heads', 1);
127
+ if (channels % heads !== 0) return `${channels} channels do not split into ${heads} heads`;
128
+ if (b !== undefined && !same(b, [heads, queries, keys])) {
129
+ return `the bias is [${b.join(', ')}], not [${heads}, ${queries}, ${keys}]`;
130
+ }
131
+ return [...qs];
132
+ },
133
+ scratch: ([q, k]) => ((q as Shape).at(-2) as number) * ((k as Shape).at(-2) as number),
134
+ evaluate: ([q, k, v, b], [qs, ks], attributes, out, scratch) => {
135
+ const shape = qs as Shape;
136
+ const [queries, channels] = shape.slice(-2) as [number, number];
137
+ attention(
138
+ out,
139
+ q as Float32Array,
140
+ k as Float32Array,
141
+ v as Float32Array,
142
+ queries,
143
+ (ks as Shape).at(-2) as number,
144
+ channels,
145
+ num(attributes, 'heads', 1),
146
+ scratch,
147
+ b ?? null,
148
+ shape.length === 3 ? (shape[0] as number) : 1,
149
+ );
150
+ },
151
+ },
152
+ ],
153
+ ];