@dxo/nn 0.0.7 → 0.0.9
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.
- package/dist/index.d.ts +3 -0
- package/dist/index.d.ts.map +1 -0
- package/dist/index.js +2 -0
- package/dist/index.js.map +1 -0
- package/dist/module.d.ts +44 -0
- package/dist/module.d.ts.map +1 -0
- package/dist/module.js +82 -0
- package/dist/module.js.map +1 -0
- package/package.json +2 -2
package/dist/index.d.ts
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.d.ts","sourceRoot":"","sources":["../src/index.ts"],"names":[],"mappings":"AAAA,YAAY,EAAE,WAAW,EAAE,gBAAgB,EAAE,MAAM,aAAa,CAAC;AACjE,OAAO,EAAE,MAAM,EAAE,MAAM,EAAE,IAAI,EAAE,IAAI,EAAE,UAAU,EAAE,MAAM,aAAa,CAAC"}
|
package/dist/index.js
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.js","sourceRoot":"","sources":["../src/index.ts"],"names":[],"mappings":"AACA,OAAO,EAAE,MAAM,EAAE,MAAM,EAAE,IAAI,EAAE,IAAI,EAAE,UAAU,EAAE,MAAM,aAAa,CAAC"}
|
package/dist/module.d.ts
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
import { Tensor } from '@dxo/core';
|
|
2
|
+
export interface TensorStateSlice {
|
|
3
|
+
shape: number[];
|
|
4
|
+
data: number[];
|
|
5
|
+
}
|
|
6
|
+
export interface LinearState {
|
|
7
|
+
weight: TensorStateSlice;
|
|
8
|
+
bias: TensorStateSlice;
|
|
9
|
+
}
|
|
10
|
+
export declare abstract class Module {
|
|
11
|
+
abstract forward(x: Tensor): Tensor;
|
|
12
|
+
parameters(): Tensor[];
|
|
13
|
+
zeroGrad(): void;
|
|
14
|
+
}
|
|
15
|
+
export declare function relu(x: Tensor): Tensor;
|
|
16
|
+
/** Element-wise ReLU module. */
|
|
17
|
+
export declare class Relu extends Module {
|
|
18
|
+
forward(x: Tensor): Tensor;
|
|
19
|
+
}
|
|
20
|
+
/** Fully-connected affine map: `y = x @ weight + bias` (no activation).
|
|
21
|
+
*
|
|
22
|
+
* Weight layout: `[inFeatures, outFeatures]`. Default leaves use `requiresGrad: true`.
|
|
23
|
+
* After `optimizer.step(parameters())`, call `loadParameters` to install new leaves.
|
|
24
|
+
*/
|
|
25
|
+
export declare class Linear extends Module {
|
|
26
|
+
readonly inFeatures: number;
|
|
27
|
+
readonly outFeatures: number;
|
|
28
|
+
weight: Tensor;
|
|
29
|
+
bias: Tensor;
|
|
30
|
+
constructor(inFeatures: number, outFeatures: number, opts?: {
|
|
31
|
+
requiresGrad?: boolean;
|
|
32
|
+
});
|
|
33
|
+
forward(x: Tensor): Tensor;
|
|
34
|
+
/** Replace parameter leaves after an optimizer step. */
|
|
35
|
+
loadParameters(params: Tensor[]): void;
|
|
36
|
+
state(): Promise<LinearState>;
|
|
37
|
+
loadState(saved: LinearState): void;
|
|
38
|
+
}
|
|
39
|
+
export declare class Sequential extends Module {
|
|
40
|
+
readonly layers: Module[];
|
|
41
|
+
constructor(layers: Module[]);
|
|
42
|
+
forward(x: Tensor): Tensor;
|
|
43
|
+
}
|
|
44
|
+
//# sourceMappingURL=module.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"module.d.ts","sourceRoot":"","sources":["../src/module.ts"],"names":[],"mappings":"AAAA,OAAO,EAAe,MAAM,EAAiB,MAAM,WAAW,CAAC;AAE/D,MAAM,WAAW,gBAAgB;IAC7B,KAAK,EAAE,MAAM,EAAE,CAAC;IAChB,IAAI,EAAE,MAAM,EAAE,CAAC;CAClB;AAED,MAAM,WAAW,WAAW;IACxB,MAAM,EAAE,gBAAgB,CAAC;IACzB,IAAI,EAAE,gBAAgB,CAAC;CAC1B;AAED,8BAAsB,MAAM;IACxB,QAAQ,CAAC,OAAO,CAAC,CAAC,EAAE,MAAM,GAAG,MAAM;IAEnC,UAAU,IAAI,MAAM,EAAE;IAStB,QAAQ,IAAI,IAAI;CAGnB;AAED,wBAAgB,IAAI,CAAC,CAAC,EAAE,MAAM,GAAG,MAAM,CAEtC;AAED,gCAAgC;AAChC,qBAAa,IAAK,SAAQ,MAAM;IAC5B,OAAO,CAAC,CAAC,EAAE,MAAM,GAAG,MAAM;CAG7B;AAED;;;;GAIG;AACH,qBAAa,MAAO,SAAQ,MAAM;IAK1B,QAAQ,CAAC,UAAU,EAAE,MAAM;IAC3B,QAAQ,CAAC,WAAW,EAAE,MAAM;IALhC,MAAM,EAAE,MAAM,CAAC;IACf,IAAI,EAAE,MAAM,CAAC;gBAGA,UAAU,EAAE,MAAM,EAClB,WAAW,EAAE,MAAM,EAC5B,IAAI,GAAE;QAAE,YAAY,CAAC,EAAE,OAAO,CAAA;KAAO;IAUzC,OAAO,CAAC,CAAC,EAAE,MAAM,GAAG,MAAM;IAI1B,wDAAwD;IACxD,cAAc,CAAC,MAAM,EAAE,MAAM,EAAE,GAAG,IAAI;IAMhC,KAAK,IAAI,OAAO,CAAC,WAAW,CAAC;IAOnC,SAAS,CAAC,KAAK,EAAE,WAAW,GAAG,IAAI;CAItC;AAED,qBAAa,UAAW,SAAQ,MAAM;IACtB,QAAQ,CAAC,MAAM,EAAE,MAAM,EAAE;gBAAhB,MAAM,EAAE,MAAM,EAAE;IAIrC,OAAO,CAAC,CAAC,EAAE,MAAM,GAAG,MAAM;CAO7B"}
|
package/dist/module.js
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import { randnValues, Tensor, tensor, zeros } from '@dxo/core';
|
|
2
|
+
export class Module {
|
|
3
|
+
parameters() {
|
|
4
|
+
const out = [];
|
|
5
|
+
for (const value of Object.values(this)) {
|
|
6
|
+
if (value instanceof Tensor)
|
|
7
|
+
out.push(value);
|
|
8
|
+
if (value instanceof Module)
|
|
9
|
+
out.push(...value.parameters());
|
|
10
|
+
}
|
|
11
|
+
return out;
|
|
12
|
+
}
|
|
13
|
+
zeroGrad() {
|
|
14
|
+
for (const p of this.parameters())
|
|
15
|
+
p.zeroGrad();
|
|
16
|
+
}
|
|
17
|
+
}
|
|
18
|
+
export function relu(x) {
|
|
19
|
+
return x.relu();
|
|
20
|
+
}
|
|
21
|
+
/** Element-wise ReLU module. */
|
|
22
|
+
export class Relu extends Module {
|
|
23
|
+
forward(x) {
|
|
24
|
+
return x.relu();
|
|
25
|
+
}
|
|
26
|
+
}
|
|
27
|
+
/** Fully-connected affine map: `y = x @ weight + bias` (no activation).
|
|
28
|
+
*
|
|
29
|
+
* Weight layout: `[inFeatures, outFeatures]`. Default leaves use `requiresGrad: true`.
|
|
30
|
+
* After `optimizer.step(parameters())`, call `loadParameters` to install new leaves.
|
|
31
|
+
*/
|
|
32
|
+
export class Linear extends Module {
|
|
33
|
+
inFeatures;
|
|
34
|
+
outFeatures;
|
|
35
|
+
weight;
|
|
36
|
+
bias;
|
|
37
|
+
constructor(inFeatures, outFeatures, opts = {}) {
|
|
38
|
+
super();
|
|
39
|
+
this.inFeatures = inFeatures;
|
|
40
|
+
this.outFeatures = outFeatures;
|
|
41
|
+
const rg = opts.requiresGrad ?? true;
|
|
42
|
+
const scale = Math.sqrt(2 / (inFeatures + outFeatures));
|
|
43
|
+
const raw = randnValues([inFeatures, outFeatures]).map((v) => v * scale);
|
|
44
|
+
this.weight = tensor(raw, [inFeatures, outFeatures], { requiresGrad: rg });
|
|
45
|
+
this.bias = zeros([outFeatures], { requiresGrad: rg });
|
|
46
|
+
}
|
|
47
|
+
forward(x) {
|
|
48
|
+
return x.matmul(this.weight).add(this.bias);
|
|
49
|
+
}
|
|
50
|
+
/** Replace parameter leaves after an optimizer step. */
|
|
51
|
+
loadParameters(params) {
|
|
52
|
+
if (params.length < 2)
|
|
53
|
+
throw new Error('Linear expects [weight, bias]');
|
|
54
|
+
this.weight = params[0];
|
|
55
|
+
this.bias = params[1];
|
|
56
|
+
}
|
|
57
|
+
async state() {
|
|
58
|
+
return {
|
|
59
|
+
weight: { shape: [...this.weight.shape], data: await this.weight.toArray() },
|
|
60
|
+
bias: { shape: [...this.bias.shape], data: await this.bias.toArray() },
|
|
61
|
+
};
|
|
62
|
+
}
|
|
63
|
+
loadState(saved) {
|
|
64
|
+
this.weight = tensor(saved.weight.data, saved.weight.shape, { requiresGrad: true });
|
|
65
|
+
this.bias = tensor(saved.bias.data, saved.bias.shape, { requiresGrad: true });
|
|
66
|
+
}
|
|
67
|
+
}
|
|
68
|
+
export class Sequential extends Module {
|
|
69
|
+
layers;
|
|
70
|
+
constructor(layers) {
|
|
71
|
+
super();
|
|
72
|
+
this.layers = layers;
|
|
73
|
+
}
|
|
74
|
+
forward(x) {
|
|
75
|
+
let out = x;
|
|
76
|
+
for (const layer of this.layers) {
|
|
77
|
+
out = layer.forward(out);
|
|
78
|
+
}
|
|
79
|
+
return out;
|
|
80
|
+
}
|
|
81
|
+
}
|
|
82
|
+
//# sourceMappingURL=module.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"module.js","sourceRoot":"","sources":["../src/module.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,WAAW,EAAE,MAAM,EAAE,MAAM,EAAE,KAAK,EAAE,MAAM,WAAW,CAAC;AAY/D,MAAM,OAAgB,MAAM;IAGxB,UAAU;QACN,MAAM,GAAG,GAAa,EAAE,CAAC;QACzB,KAAK,MAAM,KAAK,IAAI,MAAM,CAAC,MAAM,CAAC,IAAI,CAAc,EAAE,CAAC;YACnD,IAAI,KAAK,YAAY,MAAM;gBAAE,GAAG,CAAC,IAAI,CAAC,KAAK,CAAC,CAAC;YAC7C,IAAI,KAAK,YAAY,MAAM;gBAAE,GAAG,CAAC,IAAI,CAAC,GAAG,KAAK,CAAC,UAAU,EAAE,CAAC,CAAC;QACjE,CAAC;QACD,OAAO,GAAG,CAAC;IACf,CAAC;IAED,QAAQ;QACJ,KAAK,MAAM,CAAC,IAAI,IAAI,CAAC,UAAU,EAAE;YAAE,CAAC,CAAC,QAAQ,EAAE,CAAC;IACpD,CAAC;CACJ;AAED,MAAM,UAAU,IAAI,CAAC,CAAS;IAC1B,OAAO,CAAC,CAAC,IAAI,EAAE,CAAC;AACpB,CAAC;AAED,gCAAgC;AAChC,MAAM,OAAO,IAAK,SAAQ,MAAM;IAC5B,OAAO,CAAC,CAAS;QACb,OAAO,CAAC,CAAC,IAAI,EAAE,CAAC;IACpB,CAAC;CACJ;AAED;;;;GAIG;AACH,MAAM,OAAO,MAAO,SAAQ,MAAM;IAKjB;IACA;IALb,MAAM,CAAS;IACf,IAAI,CAAS;IAEb,YACa,UAAkB,EAClB,WAAmB,EAC5B,OAAmC,EAAE;QAErC,KAAK,EAAE,CAAC;QAJC,eAAU,GAAV,UAAU,CAAQ;QAClB,gBAAW,GAAX,WAAW,CAAQ;QAI5B,MAAM,EAAE,GAAG,IAAI,CAAC,YAAY,IAAI,IAAI,CAAC;QACrC,MAAM,KAAK,GAAG,IAAI,CAAC,IAAI,CAAC,CAAC,GAAG,CAAC,UAAU,GAAG,WAAW,CAAC,CAAC,CAAC;QACxD,MAAM,GAAG,GAAG,WAAW,CAAC,CAAC,UAAU,EAAE,WAAW,CAAC,CAAC,CAAC,GAAG,CAAC,CAAC,CAAC,EAAE,EAAE,CAAC,CAAC,GAAG,KAAK,CAAC,CAAC;QACzE,IAAI,CAAC,MAAM,GAAG,MAAM,CAAC,GAAG,EAAE,CAAC,UAAU,EAAE,WAAW,CAAC,EAAE,EAAE,YAAY,EAAE,EAAE,EAAE,CAAC,CAAC;QAC3E,IAAI,CAAC,IAAI,GAAG,KAAK,CAAC,CAAC,WAAW,CAAC,EAAE,EAAE,YAAY,EAAE,EAAE,EAAE,CAAC,CAAC;IAC3D,CAAC;IAED,OAAO,CAAC,CAAS;QACb,OAAO,CAAC,CAAC,MAAM,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC,GAAG,CAAC,IAAI,CAAC,IAAI,CAAC,CAAC;IAChD,CAAC;IAED,wDAAwD;IACxD,cAAc,CAAC,MAAgB;QAC3B,IAAI,MAAM,CAAC,MAAM,GAAG,CAAC;YAAE,MAAM,IAAI,KAAK,CAAC,+BAA+B,CAAC,CAAC;QACxE,IAAI,CAAC,MAAM,GAAG,MAAM,CAAC,CAAC,CAAE,CAAC;QACzB,IAAI,CAAC,IAAI,GAAG,MAAM,CAAC,CAAC,CAAE,CAAC;IAC3B,CAAC;IAED,KAAK,CAAC,KAAK;QACP,OAAO;YACH,MAAM,EAAE,EAAE,KAAK,EAAE,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,KAAK,CAAC,EAAE,IAAI,EAAE,MAAM,IAAI,CAAC,MAAM,CAAC,OAAO,EAAE,EAAE;YAC5E,IAAI,EAAE,EAAE,KAAK,EAAE,CAAC,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,CAAC,EAAE,IAAI,EAAE,MAAM,IAAI,CAAC,IAAI,CAAC,OAAO,EAAE,EAAE;SACzE,CAAC;IACN,CAAC;IAED,SAAS,CAAC,KAAkB;QACxB,IAAI,CAAC,MAAM,GAAG,MAAM,CAAC,KAAK,CAAC,MAAM,CAAC,IAAI,EAAE,KAAK,CAAC,MAAM,CAAC,KAAK,EAAE,EAAE,YAAY,EAAE,IAAI,EAAE,CAAC,CAAC;QACpF,IAAI,CAAC,IAAI,GAAG,MAAM,CAAC,KAAK,CAAC,IAAI,CAAC,IAAI,EAAE,KAAK,CAAC,IAAI,CAAC,KAAK,EAAE,EAAE,YAAY,EAAE,IAAI,EAAE,CAAC,CAAC;IAClF,CAAC;CACJ;AAED,MAAM,OAAO,UAAW,SAAQ,MAAM;IACb;IAArB,YAAqB,MAAgB;QACjC,KAAK,EAAE,CAAC;QADS,WAAM,GAAN,MAAM,CAAU;IAErC,CAAC;IAED,OAAO,CAAC,CAAS;QACb,IAAI,GAAG,GAAG,CAAC,CAAC;QACZ,KAAK,MAAM,KAAK,IAAI,IAAI,CAAC,MAAM,EAAE,CAAC;YAC9B,GAAG,GAAG,KAAK,CAAC,OAAO,CAAC,GAAG,CAAC,CAAC;QAC7B,CAAC;QACD,OAAO,GAAG,CAAC;IACf,CAAC;CACJ"}
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@dxo/nn",
|
|
3
|
-
"version": "0.0.
|
|
3
|
+
"version": "0.0.9",
|
|
4
4
|
"type": "module",
|
|
5
5
|
"description": "DXO neural modules over @dxo/core (developer preview, API unstable)",
|
|
6
6
|
"license": "MIT OR Apache-2.0",
|
|
@@ -21,7 +21,7 @@
|
|
|
21
21
|
"test:forward": "tsx ../../../scripts/test/nn-forward.ts"
|
|
22
22
|
},
|
|
23
23
|
"dependencies": {
|
|
24
|
-
"@dxo/core": "0.0.
|
|
24
|
+
"@dxo/core": "0.0.9"
|
|
25
25
|
},
|
|
26
26
|
"keywords": [
|
|
27
27
|
"dxo",
|