@dxo/train 0.0.8 → 0.0.10
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/README.md +8 -8
- package/dist/index.d.ts +91 -0
- package/dist/index.d.ts.map +1 -0
- package/dist/index.js +142 -0
- package/dist/index.js.map +1 -0
- package/package.json +19 -13
package/README.md
CHANGED
|
@@ -1,4 +1,4 @@
|
|
|
1
|
-
# @dxo/train
|
|
1
|
+
# 🏋️ @dxo/train
|
|
2
2
|
|
|
3
3
|
**Developer preview — API unstable (`0.0.x`).**
|
|
4
4
|
|
|
@@ -6,13 +6,13 @@ Async training loop over `@dxo/core` / `@dxo/nn` / `@dxo/optimizer` / `@dxo/data
|
|
|
6
6
|
|
|
7
7
|
## Contract (0.0.7 / G5)
|
|
8
8
|
|
|
9
|
-
| API
|
|
10
|
-
|
|
11
|
-
| `Trainer`
|
|
12
|
-
| `fitIter` / `run` | `AsyncGenerator<TrainEvent>`; respects `AbortSignal`
|
|
13
|
-
| `fit`
|
|
14
|
-
| Checkpoint
|
|
15
|
-
| Device
|
|
9
|
+
| API | Behavior |
|
|
10
|
+
|-------------------|-------------------------------------------------------------------|
|
|
11
|
+
| `Trainer` | Owns model + optimizer + batch factory + epoch count |
|
|
12
|
+
| `fitIter` / `run` | `AsyncGenerator<TrainEvent>`; respects `AbortSignal` |
|
|
13
|
+
| `fit` | Consumes the event stream; returns a short summary |
|
|
14
|
+
| Checkpoint | Emits `encodeLinearState(model.state())` documents (no FS I/O) |
|
|
15
|
+
| Device | **CPU-only** in this slice (GPU spike not required for this gate) |
|
|
16
16
|
|
|
17
17
|
```typescript
|
|
18
18
|
import { batch, dataset } from '@dxo/data';
|
package/dist/index.d.ts
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
import type { Tensor } from '@dxo/core';
|
|
2
|
+
import type { Batch } from '@dxo/data';
|
|
3
|
+
import type { Linear, LinearState } from '@dxo/nn';
|
|
4
|
+
import type { Optimizer } from '@dxo/optimizer';
|
|
5
|
+
import { type StateDocument } from '@dxo/serialize';
|
|
6
|
+
/** Events yielded by {@link Trainer.fitIter} / {@link Trainer.run}. */
|
|
7
|
+
export type TrainEvent = {
|
|
8
|
+
type: 'epoch_start';
|
|
9
|
+
epoch: number;
|
|
10
|
+
epochs: number;
|
|
11
|
+
} | {
|
|
12
|
+
type: 'batch';
|
|
13
|
+
epoch: number;
|
|
14
|
+
step: number;
|
|
15
|
+
loss: number;
|
|
16
|
+
} | {
|
|
17
|
+
type: 'epoch_end';
|
|
18
|
+
epoch: number;
|
|
19
|
+
meanLoss: number;
|
|
20
|
+
steps: number;
|
|
21
|
+
} | {
|
|
22
|
+
type: 'checkpoint';
|
|
23
|
+
epoch: number;
|
|
24
|
+
document: StateDocument;
|
|
25
|
+
state: LinearState;
|
|
26
|
+
} | {
|
|
27
|
+
type: 'aborted';
|
|
28
|
+
reason: 'signal';
|
|
29
|
+
epoch: number;
|
|
30
|
+
step: number;
|
|
31
|
+
} | {
|
|
32
|
+
type: 'done';
|
|
33
|
+
epochs: number;
|
|
34
|
+
steps: number;
|
|
35
|
+
finalMeanLoss?: number;
|
|
36
|
+
};
|
|
37
|
+
export interface FitSummary {
|
|
38
|
+
epochs: number;
|
|
39
|
+
steps: number;
|
|
40
|
+
aborted: boolean;
|
|
41
|
+
finalMeanLoss?: number;
|
|
42
|
+
/** Last checkpoint document, if any was emitted. */
|
|
43
|
+
lastCheckpoint?: StateDocument;
|
|
44
|
+
}
|
|
45
|
+
export type BatchSource = Iterable<Batch> | AsyncIterable<Batch>;
|
|
46
|
+
export interface TrainerOptions {
|
|
47
|
+
model: Linear;
|
|
48
|
+
optimizer: Optimizer;
|
|
49
|
+
/** Called each epoch so iterators restart (sync or async batches). */
|
|
50
|
+
batches: () => BatchSource;
|
|
51
|
+
epochs: number;
|
|
52
|
+
/**
|
|
53
|
+
* Scalar loss from prediction and target. Default: mean squared error.
|
|
54
|
+
* Must return a tensor with `numel === 1`.
|
|
55
|
+
*/
|
|
56
|
+
loss?: (pred: Tensor, y: Tensor) => Tensor;
|
|
57
|
+
/** Emit a checkpoint every N epochs (and always after the last completed epoch). Default: every epoch. */
|
|
58
|
+
checkpointEvery?: number;
|
|
59
|
+
}
|
|
60
|
+
/** Default MSE: `mean((pred - y)^2)`. */
|
|
61
|
+
export declare function mseLoss(pred: Tensor, y: Tensor): Tensor;
|
|
62
|
+
/**
|
|
63
|
+
* Minimal CPU training loop (G5 / 0.0.7).
|
|
64
|
+
*
|
|
65
|
+
* GPU / multi-device is out of scope; declare CPU-only when GPU is deferred.
|
|
66
|
+
*/
|
|
67
|
+
export declare class Trainer {
|
|
68
|
+
readonly model: Linear;
|
|
69
|
+
readonly optimizer: Optimizer;
|
|
70
|
+
readonly epochs: number;
|
|
71
|
+
readonly checkpointEvery: number;
|
|
72
|
+
private readonly batches;
|
|
73
|
+
private readonly lossFn;
|
|
74
|
+
constructor(options: TrainerOptions);
|
|
75
|
+
/** Alias for {@link fitIter} (Living API name). */
|
|
76
|
+
run(opts?: {
|
|
77
|
+
signal?: AbortSignal;
|
|
78
|
+
}): AsyncGenerator<TrainEvent, void, undefined>;
|
|
79
|
+
/**
|
|
80
|
+
* Yield training events. Honors `signal.aborted` between batches.
|
|
81
|
+
* Checkpoint payloads are in-memory `dxo-state` documents — callers persist them.
|
|
82
|
+
*/
|
|
83
|
+
fitIter(opts?: {
|
|
84
|
+
signal?: AbortSignal;
|
|
85
|
+
}): AsyncGenerator<TrainEvent, void, undefined>;
|
|
86
|
+
/** Consume {@link fitIter} and return a summary (including last checkpoint document). */
|
|
87
|
+
fit(opts?: {
|
|
88
|
+
signal?: AbortSignal;
|
|
89
|
+
}): Promise<FitSummary>;
|
|
90
|
+
}
|
|
91
|
+
//# sourceMappingURL=index.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.d.ts","sourceRoot":"","sources":["../src/index.ts"],"names":[],"mappings":"AAAA,OAAO,KAAK,EAAE,MAAM,EAAE,MAAM,WAAW,CAAC;AACxC,OAAO,KAAK,EAAE,KAAK,EAAE,MAAM,WAAW,CAAC;AACvC,OAAO,KAAK,EAAE,MAAM,EAAE,WAAW,EAAE,MAAM,SAAS,CAAC;AACnD,OAAO,KAAK,EAAE,SAAS,EAAE,MAAM,gBAAgB,CAAC;AAChD,OAAO,EAAqB,KAAK,aAAa,EAAE,MAAM,gBAAgB,CAAC;AAEvE,uEAAuE;AACvE,MAAM,MAAM,UAAU,GAChB;IAAE,IAAI,EAAE,aAAa,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,MAAM,EAAE,MAAM,CAAA;CAAE,GACtD;IAAE,IAAI,EAAE,OAAO,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,IAAI,EAAE,MAAM,CAAC;IAAC,IAAI,EAAE,MAAM,CAAA;CAAE,GAC5D;IAAE,IAAI,EAAE,WAAW,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,QAAQ,EAAE,MAAM,CAAC;IAAC,KAAK,EAAE,MAAM,CAAA;CAAE,GACrE;IAAE,IAAI,EAAE,YAAY,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,QAAQ,EAAE,aAAa,CAAC;IAAC,KAAK,EAAE,WAAW,CAAA;CAAE,GAClF;IAAE,IAAI,EAAE,SAAS,CAAC;IAAC,MAAM,EAAE,QAAQ,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,IAAI,EAAE,MAAM,CAAA;CAAE,GAClE;IAAE,IAAI,EAAE,MAAM,CAAC;IAAC,MAAM,EAAE,MAAM,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,aAAa,CAAC,EAAE,MAAM,CAAA;CAAE,CAAC;AAE9E,MAAM,WAAW,UAAU;IACvB,MAAM,EAAE,MAAM,CAAC;IACf,KAAK,EAAE,MAAM,CAAC;IACd,OAAO,EAAE,OAAO,CAAC;IACjB,aAAa,CAAC,EAAE,MAAM,CAAC;IACvB,oDAAoD;IACpD,cAAc,CAAC,EAAE,aAAa,CAAC;CAClC;AAED,MAAM,MAAM,WAAW,GAAG,QAAQ,CAAC,KAAK,CAAC,GAAG,aAAa,CAAC,KAAK,CAAC,CAAC;AAEjE,MAAM,WAAW,cAAc;IAC3B,KAAK,EAAE,MAAM,CAAC;IACd,SAAS,EAAE,SAAS,CAAC;IACrB,sEAAsE;IACtE,OAAO,EAAE,MAAM,WAAW,CAAC;IAC3B,MAAM,EAAE,MAAM,CAAC;IACf;;;OAGG;IACH,IAAI,CAAC,EAAE,CAAC,IAAI,EAAE,MAAM,EAAE,CAAC,EAAE,MAAM,KAAK,MAAM,CAAC;IAC3C,0GAA0G;IAC1G,eAAe,CAAC,EAAE,MAAM,CAAC;CAC5B;AAcD,yCAAyC;AACzC,wBAAgB,OAAO,CAAC,IAAI,EAAE,MAAM,EAAE,CAAC,EAAE,MAAM,GAAG,MAAM,CAGvD;AAED;;;;GAIG;AACH,qBAAa,OAAO;IAChB,QAAQ,CAAC,KAAK,EAAE,MAAM,CAAC;IACvB,QAAQ,CAAC,SAAS,EAAE,SAAS,CAAC;IAC9B,QAAQ,CAAC,MAAM,EAAE,MAAM,CAAC;IACxB,QAAQ,CAAC,eAAe,EAAE,MAAM,CAAC;IACjC,OAAO,CAAC,QAAQ,CAAC,OAAO,CAAoB;IAC5C,OAAO,CAAC,QAAQ,CAAC,MAAM,CAAsC;gBAEjD,OAAO,EAAE,cAAc;IAgBnC,mDAAmD;IACnD,GAAG,CAAC,IAAI,CAAC,EAAE;QAAE,MAAM,CAAC,EAAE,WAAW,CAAA;KAAE,GAAG,cAAc,CAAC,UAAU,EAAE,IAAI,EAAE,SAAS,CAAC;IAIjF;;;OAGG;IACI,OAAO,CAAC,IAAI,CAAC,EAAE;QAAE,MAAM,CAAC,EAAE,WAAW,CAAA;KAAE,GAAG,cAAc,CAAC,UAAU,EAAE,IAAI,EAAE,SAAS,CAAC;IA+D5F,yFAAyF;IACnF,GAAG,CAAC,IAAI,CAAC,EAAE;QAAE,MAAM,CAAC,EAAE,WAAW,CAAA;KAAE,GAAG,OAAO,CAAC,UAAU,CAAC;CAoClE"}
|
package/dist/index.js
ADDED
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
import { encodeLinearState } from '@dxo/serialize';
|
|
2
|
+
function isAsyncIterable(source) {
|
|
3
|
+
return typeof source[Symbol.asyncIterator] === 'function';
|
|
4
|
+
}
|
|
5
|
+
async function* iterateBatches(source) {
|
|
6
|
+
if (isAsyncIterable(source)) {
|
|
7
|
+
for await (const b of source)
|
|
8
|
+
yield b;
|
|
9
|
+
return;
|
|
10
|
+
}
|
|
11
|
+
for (const b of source)
|
|
12
|
+
yield b;
|
|
13
|
+
}
|
|
14
|
+
/** Default MSE: `mean((pred - y)^2)`. */
|
|
15
|
+
export function mseLoss(pred, y) {
|
|
16
|
+
const diff = pred.sub(y);
|
|
17
|
+
return diff.mul(diff).mean();
|
|
18
|
+
}
|
|
19
|
+
/**
|
|
20
|
+
* Minimal CPU training loop (G5 / 0.0.7).
|
|
21
|
+
*
|
|
22
|
+
* GPU / multi-device is out of scope; declare CPU-only when GPU is deferred.
|
|
23
|
+
*/
|
|
24
|
+
export class Trainer {
|
|
25
|
+
model;
|
|
26
|
+
optimizer;
|
|
27
|
+
epochs;
|
|
28
|
+
checkpointEvery;
|
|
29
|
+
batches;
|
|
30
|
+
lossFn;
|
|
31
|
+
constructor(options) {
|
|
32
|
+
if (!(options.epochs > 0) || !Number.isInteger(options.epochs)) {
|
|
33
|
+
throw new Error('epochs must be a positive integer');
|
|
34
|
+
}
|
|
35
|
+
const every = options.checkpointEvery ?? 1;
|
|
36
|
+
if (!(every > 0) || !Number.isInteger(every)) {
|
|
37
|
+
throw new Error('checkpointEvery must be a positive integer');
|
|
38
|
+
}
|
|
39
|
+
this.model = options.model;
|
|
40
|
+
this.optimizer = options.optimizer;
|
|
41
|
+
this.epochs = options.epochs;
|
|
42
|
+
this.checkpointEvery = every;
|
|
43
|
+
this.batches = options.batches;
|
|
44
|
+
this.lossFn = options.loss ?? mseLoss;
|
|
45
|
+
}
|
|
46
|
+
/** Alias for {@link fitIter} (Living API name). */
|
|
47
|
+
run(opts) {
|
|
48
|
+
return this.fitIter(opts);
|
|
49
|
+
}
|
|
50
|
+
/**
|
|
51
|
+
* Yield training events. Honors `signal.aborted` between batches.
|
|
52
|
+
* Checkpoint payloads are in-memory `dxo-state` documents — callers persist them.
|
|
53
|
+
*/
|
|
54
|
+
async *fitIter(opts) {
|
|
55
|
+
const signal = opts?.signal;
|
|
56
|
+
let globalStep = 0;
|
|
57
|
+
let lastMean;
|
|
58
|
+
for (let epoch = 1; epoch <= this.epochs; epoch++) {
|
|
59
|
+
if (signal?.aborted) {
|
|
60
|
+
yield { type: 'aborted', reason: 'signal', epoch, step: globalStep };
|
|
61
|
+
return;
|
|
62
|
+
}
|
|
63
|
+
yield { type: 'epoch_start', epoch, epochs: this.epochs };
|
|
64
|
+
let epochLossSum = 0;
|
|
65
|
+
let epochSteps = 0;
|
|
66
|
+
for await (const batch of iterateBatches(this.batches())) {
|
|
67
|
+
if (signal?.aborted) {
|
|
68
|
+
yield { type: 'aborted', reason: 'signal', epoch, step: globalStep };
|
|
69
|
+
return;
|
|
70
|
+
}
|
|
71
|
+
if (!batch.y) {
|
|
72
|
+
throw new Error('Trainer requires batch.y targets');
|
|
73
|
+
}
|
|
74
|
+
this.model.zeroGrad();
|
|
75
|
+
const pred = this.model.forward(batch.x);
|
|
76
|
+
const loss = this.lossFn(pred, batch.y);
|
|
77
|
+
const value = await loss.item();
|
|
78
|
+
if (!Number.isFinite(value)) {
|
|
79
|
+
throw new Error(`non-finite loss at epoch ${epoch} step ${globalStep + 1}: ${value}`);
|
|
80
|
+
}
|
|
81
|
+
loss.backward();
|
|
82
|
+
this.model.loadParameters(await this.optimizer.step(this.model.parameters()));
|
|
83
|
+
globalStep += 1;
|
|
84
|
+
epochSteps += 1;
|
|
85
|
+
epochLossSum += value;
|
|
86
|
+
yield { type: 'batch', epoch, step: globalStep, loss: value };
|
|
87
|
+
}
|
|
88
|
+
if (epochSteps === 0) {
|
|
89
|
+
throw new Error(`epoch ${epoch} produced zero batches`);
|
|
90
|
+
}
|
|
91
|
+
lastMean = epochLossSum / epochSteps;
|
|
92
|
+
yield { type: 'epoch_end', epoch, meanLoss: lastMean, steps: epochSteps };
|
|
93
|
+
const shouldCheckpoint = epoch % this.checkpointEvery === 0 || epoch === this.epochs;
|
|
94
|
+
if (shouldCheckpoint) {
|
|
95
|
+
const state = await this.model.state();
|
|
96
|
+
yield {
|
|
97
|
+
type: 'checkpoint',
|
|
98
|
+
epoch,
|
|
99
|
+
state,
|
|
100
|
+
document: encodeLinearState(state),
|
|
101
|
+
};
|
|
102
|
+
}
|
|
103
|
+
}
|
|
104
|
+
yield { type: 'done', epochs: this.epochs, steps: globalStep, finalMeanLoss: lastMean };
|
|
105
|
+
}
|
|
106
|
+
/** Consume {@link fitIter} and return a summary (including last checkpoint document). */
|
|
107
|
+
async fit(opts) {
|
|
108
|
+
let epochs = 0;
|
|
109
|
+
let steps = 0;
|
|
110
|
+
let aborted = false;
|
|
111
|
+
let finalMeanLoss;
|
|
112
|
+
let lastCheckpoint;
|
|
113
|
+
for await (const event of this.fitIter(opts)) {
|
|
114
|
+
switch (event.type) {
|
|
115
|
+
case 'batch':
|
|
116
|
+
steps = event.step;
|
|
117
|
+
break;
|
|
118
|
+
case 'epoch_end':
|
|
119
|
+
epochs = event.epoch;
|
|
120
|
+
finalMeanLoss = event.meanLoss;
|
|
121
|
+
break;
|
|
122
|
+
case 'checkpoint':
|
|
123
|
+
lastCheckpoint = event.document;
|
|
124
|
+
break;
|
|
125
|
+
case 'aborted':
|
|
126
|
+
aborted = true;
|
|
127
|
+
epochs = event.epoch;
|
|
128
|
+
steps = event.step;
|
|
129
|
+
break;
|
|
130
|
+
case 'done':
|
|
131
|
+
epochs = event.epochs;
|
|
132
|
+
steps = event.steps;
|
|
133
|
+
finalMeanLoss = event.finalMeanLoss;
|
|
134
|
+
break;
|
|
135
|
+
default:
|
|
136
|
+
break;
|
|
137
|
+
}
|
|
138
|
+
}
|
|
139
|
+
return { epochs, steps, aborted, finalMeanLoss, lastCheckpoint };
|
|
140
|
+
}
|
|
141
|
+
}
|
|
142
|
+
//# sourceMappingURL=index.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.js","sourceRoot":"","sources":["../src/index.ts"],"names":[],"mappings":"AAIA,OAAO,EAAE,iBAAiB,EAAsB,MAAM,gBAAgB,CAAC;AAqCvE,SAAS,eAAe,CAAC,MAAmB;IACxC,OAAO,OAAQ,MAA+B,CAAC,MAAM,CAAC,aAAa,CAAC,KAAK,UAAU,CAAC;AACxF,CAAC;AAED,KAAK,SAAS,CAAC,CAAC,cAAc,CAAC,MAAmB;IAC9C,IAAI,eAAe,CAAC,MAAM,CAAC,EAAE,CAAC;QAC1B,IAAI,KAAK,EAAE,MAAM,CAAC,IAAI,MAAM;YAAE,MAAM,CAAC,CAAC;QACtC,OAAO;IACX,CAAC;IACD,KAAK,MAAM,CAAC,IAAI,MAAM;QAAE,MAAM,CAAC,CAAC;AACpC,CAAC;AAED,yCAAyC;AACzC,MAAM,UAAU,OAAO,CAAC,IAAY,EAAE,CAAS;IAC3C,MAAM,IAAI,GAAG,IAAI,CAAC,GAAG,CAAC,CAAC,CAAC,CAAC;IACzB,OAAO,IAAI,CAAC,GAAG,CAAC,IAAI,CAAC,CAAC,IAAI,EAAE,CAAC;AACjC,CAAC;AAED;;;;GAIG;AACH,MAAM,OAAO,OAAO;IACP,KAAK,CAAS;IACd,SAAS,CAAY;IACrB,MAAM,CAAS;IACf,eAAe,CAAS;IAChB,OAAO,CAAoB;IAC3B,MAAM,CAAsC;IAE7D,YAAY,OAAuB;QAC/B,IAAI,CAAC,CAAC,OAAO,CAAC,MAAM,GAAG,CAAC,CAAC,IAAI,CAAC,MAAM,CAAC,SAAS,CAAC,OAAO,CAAC,MAAM,CAAC,EAAE,CAAC;YAC7D,MAAM,IAAI,KAAK,CAAC,mCAAmC,CAAC,CAAC;QACzD,CAAC;QACD,MAAM,KAAK,GAAG,OAAO,CAAC,eAAe,IAAI,CAAC,CAAC;QAC3C,IAAI,CAAC,CAAC,KAAK,GAAG,CAAC,CAAC,IAAI,CAAC,MAAM,CAAC,SAAS,CAAC,KAAK,CAAC,EAAE,CAAC;YAC3C,MAAM,IAAI,KAAK,CAAC,4CAA4C,CAAC,CAAC;QAClE,CAAC;QACD,IAAI,CAAC,KAAK,GAAG,OAAO,CAAC,KAAK,CAAC;QAC3B,IAAI,CAAC,SAAS,GAAG,OAAO,CAAC,SAAS,CAAC;QACnC,IAAI,CAAC,MAAM,GAAG,OAAO,CAAC,MAAM,CAAC;QAC7B,IAAI,CAAC,eAAe,GAAG,KAAK,CAAC;QAC7B,IAAI,CAAC,OAAO,GAAG,OAAO,CAAC,OAAO,CAAC;QAC/B,IAAI,CAAC,MAAM,GAAG,OAAO,CAAC,IAAI,IAAI,OAAO,CAAC;IAC1C,CAAC;IAED,mDAAmD;IACnD,GAAG,CAAC,IAA+B;QAC/B,OAAO,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC,CAAC;IAC9B,CAAC;IAED;;;OAGG;IACH,KAAK,CAAC,CAAC,OAAO,CAAC,IAA+B;QAC1C,MAAM,MAAM,GAAG,IAAI,EAAE,MAAM,CAAC;QAC5B,IAAI,UAAU,GAAG,CAAC,CAAC;QACnB,IAAI,QAA4B,CAAC;QAEjC,KAAK,IAAI,KAAK,GAAG,CAAC,EAAE,KAAK,IAAI,IAAI,CAAC,MAAM,EAAE,KAAK,EAAE,EAAE,CAAC;YAChD,IAAI,MAAM,EAAE,OAAO,EAAE,CAAC;gBAClB,MAAM,EAAE,IAAI,EAAE,SAAS,EAAE,MAAM,EAAE,QAAQ,EAAE,KAAK,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;gBACrE,OAAO;YACX,CAAC;YAED,MAAM,EAAE,IAAI,EAAE,aAAa,EAAE,KAAK,EAAE,MAAM,EAAE,IAAI,CAAC,MAAM,EAAE,CAAC;YAE1D,IAAI,YAAY,GAAG,CAAC,CAAC;YACrB,IAAI,UAAU,GAAG,CAAC,CAAC;YAEnB,IAAI,KAAK,EAAE,MAAM,KAAK,IAAI,cAAc,CAAC,IAAI,CAAC,OAAO,EAAE,CAAC,EAAE,CAAC;gBACvD,IAAI,MAAM,EAAE,OAAO,EAAE,CAAC;oBAClB,MAAM,EAAE,IAAI,EAAE,SAAS,EAAE,MAAM,EAAE,QAAQ,EAAE,KAAK,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;oBACrE,OAAO;gBACX,CAAC;gBACD,IAAI,CAAC,KAAK,CAAC,CAAC,EAAE,CAAC;oBACX,MAAM,IAAI,KAAK,CAAC,kCAAkC,CAAC,CAAC;gBACxD,CAAC;gBAED,IAAI,CAAC,KAAK,CAAC,QAAQ,EAAE,CAAC;gBACtB,MAAM,IAAI,GAAG,IAAI,CAAC,KAAK,CAAC,OAAO,CAAC,KAAK,CAAC,CAAC,CAAC,CAAC;gBACzC,MAAM,IAAI,GAAG,IAAI,CAAC,MAAM,CAAC,IAAI,EAAE,KAAK,CAAC,CAAC,CAAC,CAAC;gBACxC,MAAM,KAAK,GAAG,MAAM,IAAI,CAAC,IAAI,EAAE,CAAC;gBAChC,IAAI,CAAC,MAAM,CAAC,QAAQ,CAAC,KAAK,CAAC,EAAE,CAAC;oBAC1B,MAAM,IAAI,KAAK,CAAC,4BAA4B,KAAK,SAAS,UAAU,GAAG,CAAC,KAAK,KAAK,EAAE,CAAC,CAAC;gBAC1F,CAAC;gBACD,IAAI,CAAC,QAAQ,EAAE,CAAC;gBAChB,IAAI,CAAC,KAAK,CAAC,cAAc,CAAC,MAAM,IAAI,CAAC,SAAS,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,CAAC,UAAU,EAAE,CAAC,CAAC,CAAC;gBAE9E,UAAU,IAAI,CAAC,CAAC;gBAChB,UAAU,IAAI,CAAC,CAAC;gBAChB,YAAY,IAAI,KAAK,CAAC;gBACtB,MAAM,EAAE,IAAI,EAAE,OAAO,EAAE,KAAK,EAAE,IAAI,EAAE,UAAU,EAAE,IAAI,EAAE,KAAK,EAAE,CAAC;YAClE,CAAC;YAED,IAAI,UAAU,KAAK,CAAC,EAAE,CAAC;gBACnB,MAAM,IAAI,KAAK,CAAC,SAAS,KAAK,wBAAwB,CAAC,CAAC;YAC5D,CAAC;YAED,QAAQ,GAAG,YAAY,GAAG,UAAU,CAAC;YACrC,MAAM,EAAE,IAAI,EAAE,WAAW,EAAE,KAAK,EAAE,QAAQ,EAAE,QAAQ,EAAE,KAAK,EAAE,UAAU,EAAE,CAAC;YAE1E,MAAM,gBAAgB,GAAG,KAAK,GAAG,IAAI,CAAC,eAAe,KAAK,CAAC,IAAI,KAAK,KAAK,IAAI,CAAC,MAAM,CAAC;YACrF,IAAI,gBAAgB,EAAE,CAAC;gBACnB,MAAM,KAAK,GAAG,MAAM,IAAI,CAAC,KAAK,CAAC,KAAK,EAAE,CAAC;gBACvC,MAAM;oBACF,IAAI,EAAE,YAAY;oBAClB,KAAK;oBACL,KAAK;oBACL,QAAQ,EAAE,iBAAiB,CAAC,KAAK,CAAC;iBACrC,CAAC;YACN,CAAC;QACL,CAAC;QAED,MAAM,EAAE,IAAI,EAAE,MAAM,EAAE,MAAM,EAAE,IAAI,CAAC,MAAM,EAAE,KAAK,EAAE,UAAU,EAAE,aAAa,EAAE,QAAQ,EAAE,CAAC;IAC5F,CAAC;IAED,yFAAyF;IACzF,KAAK,CAAC,GAAG,CAAC,IAA+B;QACrC,IAAI,MAAM,GAAG,CAAC,CAAC;QACf,IAAI,KAAK,GAAG,CAAC,CAAC;QACd,IAAI,OAAO,GAAG,KAAK,CAAC;QACpB,IAAI,aAAiC,CAAC;QACtC,IAAI,cAAyC,CAAC;QAE9C,IAAI,KAAK,EAAE,MAAM,KAAK,IAAI,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC,EAAE,CAAC;YAC3C,QAAQ,KAAK,CAAC,IAAI,EAAE,CAAC;gBACjB,KAAK,OAAO;oBACR,KAAK,GAAG,KAAK,CAAC,IAAI,CAAC;oBACnB,MAAM;gBACV,KAAK,WAAW;oBACZ,MAAM,GAAG,KAAK,CAAC,KAAK,CAAC;oBACrB,aAAa,GAAG,KAAK,CAAC,QAAQ,CAAC;oBAC/B,MAAM;gBACV,KAAK,YAAY;oBACb,cAAc,GAAG,KAAK,CAAC,QAAQ,CAAC;oBAChC,MAAM;gBACV,KAAK,SAAS;oBACV,OAAO,GAAG,IAAI,CAAC;oBACf,MAAM,GAAG,KAAK,CAAC,KAAK,CAAC;oBACrB,KAAK,GAAG,KAAK,CAAC,IAAI,CAAC;oBACnB,MAAM;gBACV,KAAK,MAAM;oBACP,MAAM,GAAG,KAAK,CAAC,MAAM,CAAC;oBACtB,KAAK,GAAG,KAAK,CAAC,KAAK,CAAC;oBACpB,aAAa,GAAG,KAAK,CAAC,aAAa,CAAC;oBACpC,MAAM;gBACV;oBACI,MAAM;YACd,CAAC;QACL,CAAC;QAED,OAAO,EAAE,MAAM,EAAE,KAAK,EAAE,OAAO,EAAE,aAAa,EAAE,cAAc,EAAE,CAAC;IACrE,CAAC;CACJ"}
|
package/package.json
CHANGED
|
@@ -1,9 +1,9 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@dxo/train",
|
|
3
|
-
"version": "0.0.
|
|
3
|
+
"version": "0.0.10",
|
|
4
4
|
"type": "module",
|
|
5
|
-
"description": "DXO training loop
|
|
6
|
-
"license": "
|
|
5
|
+
"description": "DXO training loop - fitIter / TrainEvent / checkpoint (developer preview, API unstable)",
|
|
6
|
+
"license": "Apache-2.0",
|
|
7
7
|
"main": "./dist/index.js",
|
|
8
8
|
"types": "./dist/index.d.ts",
|
|
9
9
|
"exports": {
|
|
@@ -20,23 +20,29 @@
|
|
|
20
20
|
"build": "tsc -p tsconfig.json"
|
|
21
21
|
},
|
|
22
22
|
"dependencies": {
|
|
23
|
-
"@dxo/core": "0.0.
|
|
24
|
-
"@dxo/data": "0.0.
|
|
25
|
-
"@dxo/nn": "0.0.
|
|
26
|
-
"@dxo/optimizer": "0.0.
|
|
27
|
-
"@dxo/serialize": "0.0.
|
|
23
|
+
"@dxo/core": "0.0.10",
|
|
24
|
+
"@dxo/data": "0.0.10",
|
|
25
|
+
"@dxo/nn": "0.0.10",
|
|
26
|
+
"@dxo/optimizer": "0.0.10",
|
|
27
|
+
"@dxo/serialize": "0.0.10"
|
|
28
28
|
},
|
|
29
29
|
"keywords": [
|
|
30
30
|
"dxo",
|
|
31
31
|
"trainer",
|
|
32
32
|
"training",
|
|
33
|
-
"typescript"
|
|
33
|
+
"typescript",
|
|
34
|
+
"deep-learning"
|
|
34
35
|
],
|
|
35
|
-
"
|
|
36
|
-
"access": "public"
|
|
37
|
-
},
|
|
36
|
+
"homepage": "https://github.com/ai4waifu/dxo-framework#readme",
|
|
38
37
|
"repository": {
|
|
39
38
|
"type": "git",
|
|
40
|
-
"url": "
|
|
39
|
+
"url": "https://github.com/ai4waifu/dxo-framework.git",
|
|
40
|
+
"directory": "projects/runtimes/dxo-train"
|
|
41
|
+
},
|
|
42
|
+
"bugs": {
|
|
43
|
+
"url": "https://github.com/ai4waifu/dxo-framework/issues"
|
|
44
|
+
},
|
|
45
|
+
"publishConfig": {
|
|
46
|
+
"access": "public"
|
|
41
47
|
}
|
|
42
48
|
}
|