@genai-fi/nanogpt 0.24.0 → 1.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.
- package/README.md +78 -281
- package/dist/{RealDiv-CNsvC4AU.js → RealDiv-CSnvtN2E.js} +20 -20
- package/dist/{Reshape-dnm9bO3B.js → Reshape-BlylqwWy.js} +12 -12
- package/dist/TeachableLLM.d.ts +10 -14
- package/dist/TeachableLLM.js +201 -2
- package/dist/api/responses.d.ts +81 -0
- package/dist/api/responses.js +169 -0
- package/dist/api/training.d.ts +70 -0
- package/dist/api/training.js +205 -0
- package/dist/data/docx.js +9 -3036
- package/dist/data/stream.d.ts +2 -2
- package/dist/data/textLoader.d.ts +1 -1
- package/dist/data/textLoader.js +1 -1
- package/dist/data.d.ts +3 -0
- package/dist/data.js +12 -0
- package/dist/{dist-BqAU9-yi.js → dist-CwK5S7Ls.js} +2168 -2168
- package/dist/{gpgpu_math-DBYEAAdI.js → gpgpu_math-20tPK8LM.js} +458 -458
- package/dist/{Generator.d.ts → inference/Generator.d.ts} +16 -44
- package/dist/inference/Generator.js +271 -0
- package/dist/inference/tokenisePrompt.d.ts +4 -0
- package/dist/inference/tokenisePrompt.js +13 -0
- package/dist/inference/types.d.ts +44 -8
- package/dist/inference/utilities.d.ts +9 -0
- package/dist/inference/utilities.js +20 -0
- package/dist/jszip.min-DKa1Rjyn.js +3033 -0
- package/dist/{kernel_funcs_utils-D-mATnGR.js → kernel_funcs_utils-ql8Y8qPn.js} +96 -93
- package/dist/layers/MLP.d.ts +1 -1
- package/dist/layers/PositionEmbedding.d.ts +2 -1
- package/dist/layers/PositionEmbedding.js +1 -1
- package/dist/layers/RMSNorm.d.ts +1 -1
- package/dist/layers/TiedEmbedding.js +1 -1
- package/dist/layers.d.ts +4 -0
- package/dist/layers.js +14 -0
- package/dist/loader/load.js +58 -2
- package/dist/loader/loadHF.d.ts +1 -1
- package/dist/loader/loadHF.js +17 -2
- package/dist/loader/loadTransformers.js +46 -2
- package/dist/loader/newZipLoad.js +25 -2
- package/dist/loader/oldZipLoad.d.ts +1 -1
- package/dist/loader/oldZipLoad.js +37 -2
- package/dist/loader/save.js +75 -2
- package/dist/loader/types.d.ts +3 -3
- package/dist/main.d.ts +34 -35
- package/dist/main.js +12327 -20
- package/dist/{matMulGelu-BAIgQaRx.js → matMulGelu-CBoqTZM7.js} +2 -2
- package/dist/models/NanoGPTV1.js +95 -2
- package/dist/models/NanoGPTV2.js +86 -2
- package/dist/models/factory.js +13 -2
- package/dist/models/model.js +76 -2
- package/dist/models.d.ts +4 -0
- package/dist/models.js +14 -0
- package/dist/ops/dot16.js +1 -1
- package/dist/ops/matMulGelu.js +1 -1
- package/dist/ops/webgl/adamAdjust.js +1 -1
- package/dist/ops/webgl/fusedSoftmax.js +2 -2
- package/dist/ops/webgl/gelu.js +2 -2
- package/dist/ops/webgl/log.js +5 -5
- package/dist/ops/webgl/matMulGelu.js +1 -1
- package/dist/ops/webgl/matMulMul.js +1 -1
- package/dist/{tfjs_backend-CydPRQTc.js → tfjs_backend-h5weiy1O.js} +36 -36
- package/dist/tokenise.d.ts +4 -0
- package/dist/tokenise.js +15 -0
- package/dist/training/BasicTrainer.d.ts +5 -10
- package/dist/training/BasicTrainer.js +80 -88
- package/dist/training/configure.d.ts +3 -0
- package/dist/training/configure.js +32 -0
- package/dist/training/factory.d.ts +6 -0
- package/dist/training/factory.js +8 -0
- package/dist/training/prepareData.d.ts +22 -0
- package/dist/training/prepareData.js +49 -0
- package/dist/training/tasks/tokenStream.d.ts +2 -1
- package/dist/training/types.d.ts +14 -1
- package/dist/training/validateOptions.d.ts +2 -0
- package/dist/training/validateOptions.js +19 -0
- package/dist/utilities/arrayShape.d.ts +1 -0
- package/dist/utilities/arrayShape.js +8 -0
- package/dist/utilities/waitForModel.d.ts +1 -1
- package/dist/v4-BK7K-jy_.js +30 -0
- package/package.json +8 -2
- package/dist/Generator.js +0 -2
- package/dist/Trainer-Cr7csbTD.js +0 -228
- package/dist/Trainer.d.ts +0 -45
- package/dist/Trainer.js +0 -2
- package/dist/main-Dz72vadm.js +0 -13267
package/dist/TeachableLLM.js
CHANGED
|
@@ -1,2 +1,201 @@
|
|
|
1
|
-
import {
|
|
2
|
-
|
|
1
|
+
import { t as e } from "./eventemitter3-D_qV3Lof.js";
|
|
2
|
+
import t from "./tokeniser/CharTokeniser.js";
|
|
3
|
+
import n from "./tokeniser/bpe.js";
|
|
4
|
+
import { validateConfig as r } from "./models/config.js";
|
|
5
|
+
import { dummyPassTrainAsync as i } from "./utilities/dummy.js";
|
|
6
|
+
import a from "./models/factory.js";
|
|
7
|
+
import { loadModel as o } from "./loader/load.js";
|
|
8
|
+
import { saveModel as s } from "./loader/save.js";
|
|
9
|
+
import c from "./utilities/profile.js";
|
|
10
|
+
import l from "./api/responses.js";
|
|
11
|
+
import u from "./api/training.js";
|
|
12
|
+
import { selectBackend as d } from "./backend.js";
|
|
13
|
+
//#region lib/TeachableLLM.ts
|
|
14
|
+
var f = class f {
|
|
15
|
+
ee = new e();
|
|
16
|
+
_config;
|
|
17
|
+
_model;
|
|
18
|
+
_tokeniser;
|
|
19
|
+
_status = "loading";
|
|
20
|
+
_memoryRequirements;
|
|
21
|
+
_responses = null;
|
|
22
|
+
_training = null;
|
|
23
|
+
meta = {
|
|
24
|
+
version: 2,
|
|
25
|
+
application: "@genai-fi/nanogpt"
|
|
26
|
+
};
|
|
27
|
+
static selectBackend(e, t) {
|
|
28
|
+
return d(e, t);
|
|
29
|
+
}
|
|
30
|
+
constructor(e, t) {
|
|
31
|
+
this._config = t?.config, this._tokeniser = e, this._model = t, t?.metaData && (this.meta = t.metaData);
|
|
32
|
+
}
|
|
33
|
+
get vocab() {
|
|
34
|
+
return this._tokeniser?.getVocab() || [];
|
|
35
|
+
}
|
|
36
|
+
get mode() {
|
|
37
|
+
return this._model?.metaData?.mode ?? "untrained";
|
|
38
|
+
}
|
|
39
|
+
set mode(e) {
|
|
40
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
41
|
+
this._model.metaData.mode === "conversational" && e === "completion" || e !== "untrained" && (this._model.metaData.mode = e, this.ee.emit("mode", e));
|
|
42
|
+
}
|
|
43
|
+
get loaded() {
|
|
44
|
+
return !!this._model && !!this._tokeniser && !!this._config;
|
|
45
|
+
}
|
|
46
|
+
get config() {
|
|
47
|
+
if (!this._config) throw Error("configuration_not_initialized.");
|
|
48
|
+
return this._config;
|
|
49
|
+
}
|
|
50
|
+
get model() {
|
|
51
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
52
|
+
return this._model;
|
|
53
|
+
}
|
|
54
|
+
get tokeniser() {
|
|
55
|
+
if (!this._tokeniser) throw Error("tokeniser_not_initialized.");
|
|
56
|
+
return this._tokeniser;
|
|
57
|
+
}
|
|
58
|
+
get status() {
|
|
59
|
+
return this._status;
|
|
60
|
+
}
|
|
61
|
+
get ready() {
|
|
62
|
+
return this._status === "ready" && !!this._model && !!this._tokeniser;
|
|
63
|
+
}
|
|
64
|
+
get busy() {
|
|
65
|
+
return this._status === "busy" || this._status === "training";
|
|
66
|
+
}
|
|
67
|
+
createLoRA(e, t) {
|
|
68
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
69
|
+
this._model.createLoRA(e, t), this.ee.emit("changeLoRA");
|
|
70
|
+
}
|
|
71
|
+
deleteLoRA(e) {
|
|
72
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
73
|
+
this._model.deleteLoRA(e), this.ee.emit("changeLoRA");
|
|
74
|
+
}
|
|
75
|
+
renameLoRA(e, t) {
|
|
76
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
77
|
+
this._model.renameLoRA(e, t), this.ee.emit("changeLoRA");
|
|
78
|
+
}
|
|
79
|
+
attachLoRA(e) {
|
|
80
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
81
|
+
this.model.lora?.name !== e && (this._model.attachLoRA(e), this.ee.emit("changeLoRA"));
|
|
82
|
+
}
|
|
83
|
+
detachLoRA() {
|
|
84
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
85
|
+
this._model.detachLoRA(), this.ee.emit("changeLoRA");
|
|
86
|
+
}
|
|
87
|
+
hasLoRA(e) {
|
|
88
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
89
|
+
return this._model.hasLoRA(e);
|
|
90
|
+
}
|
|
91
|
+
listLoRAs() {
|
|
92
|
+
if (!this._model) throw Error("model_not_initialized.");
|
|
93
|
+
return this._model.listLoRAs();
|
|
94
|
+
}
|
|
95
|
+
estimateTrainingMemoryUsage(e) {
|
|
96
|
+
let t = this._memoryRequirements ?? {
|
|
97
|
+
perBatch: 0,
|
|
98
|
+
tapeSize: 0,
|
|
99
|
+
gradients: 0
|
|
100
|
+
}, n = t.perBatch * e, r = t.gradients;
|
|
101
|
+
return n * .66 + r * 4;
|
|
102
|
+
}
|
|
103
|
+
setStatus(e) {
|
|
104
|
+
this._status !== e && (this._status = e, this.ee.emit("status", e));
|
|
105
|
+
}
|
|
106
|
+
saveModel(e) {
|
|
107
|
+
if (!this._model || !this._tokeniser) throw Error("model_or_tokeniser_not_initialized.");
|
|
108
|
+
let t = (e?.includeOptimizer && this._training?.getPretrainingJob()) ?? null;
|
|
109
|
+
return s(this._model, this._tokeniser, {
|
|
110
|
+
...e,
|
|
111
|
+
name: e?.name || this.meta.name
|
|
112
|
+
}, t ? {
|
|
113
|
+
optimizer: t.trainer.optimizer,
|
|
114
|
+
trainingLog: t.trainer.log
|
|
115
|
+
} : void 0);
|
|
116
|
+
}
|
|
117
|
+
static loadModel(e, t) {
|
|
118
|
+
let n = new f();
|
|
119
|
+
return o(e, t).then(({ model: e, tokeniser: t, metaData: a, optimizer: o, log: s }) => {
|
|
120
|
+
r(e.config), n._model = e, n._tokeniser = t, n._config = e.config, a && (n.meta = a), n.setStatus("warmup"), i(e).then((t) => {
|
|
121
|
+
n._memoryRequirements = t, o && e.metaData.pretrainingSettings && e.metaData.pretrainingData && n.training.restore(e.metaData.pretrainingSettings, s || [], o, e.metaData.pretrainingData), n.setStatus("ready"), n.ee.emit("loaded"), n.ee.emit("mode", n.mode);
|
|
122
|
+
}).catch((e) => {
|
|
123
|
+
n.setStatus("error"), n.ee.emit("error", e), console.error("Error during warmup:", e);
|
|
124
|
+
});
|
|
125
|
+
}).catch((e) => {
|
|
126
|
+
n.setStatus("error"), n.ee.emit("error", e), console.error("Error loading model:", e);
|
|
127
|
+
}), n;
|
|
128
|
+
}
|
|
129
|
+
static create(e, o) {
|
|
130
|
+
r(o);
|
|
131
|
+
let s = o, c = e === "char" ? new t(s.vocabSize) : e === "bpe" ? new n(s.vocabSize) : e, l = a(s), u = new f(c, l);
|
|
132
|
+
return u.setStatus("warmup"), i(l).then((e) => {
|
|
133
|
+
u._memoryRequirements = e, u.tokeniser.trained ? (u.setStatus("ready"), u.ee.emit("loaded"), u.ee.emit("mode", u.mode)) : (u.setStatus("awaitingTokens"), u.ee.emit("loaded"), u.ee.emit("mode", u.mode), u.tokeniser.once("trainStatus", (e) => {
|
|
134
|
+
e === "trained" && u.setStatus("ready");
|
|
135
|
+
}));
|
|
136
|
+
}).catch((e) => {
|
|
137
|
+
u.setStatus("error"), u.ee.emit("error", e), console.error("Error during warmup:", e);
|
|
138
|
+
}), u;
|
|
139
|
+
}
|
|
140
|
+
getProfiler() {
|
|
141
|
+
return this._model?.getProfiler();
|
|
142
|
+
}
|
|
143
|
+
get enableProfiler() {
|
|
144
|
+
return !!this._model?.getProfiler();
|
|
145
|
+
}
|
|
146
|
+
set enableProfiler(e) {
|
|
147
|
+
if (e) {
|
|
148
|
+
if (!this._config) return;
|
|
149
|
+
this.model.getProfiler() || this.model.setProfiler(new c());
|
|
150
|
+
} else this.model.getProfiler() && this.model.setProfiler(null);
|
|
151
|
+
}
|
|
152
|
+
getNumParams() {
|
|
153
|
+
return this._model ? this._model.getNumParams() : 0;
|
|
154
|
+
}
|
|
155
|
+
async trainTokeniser(e) {
|
|
156
|
+
if (!this._tokeniser) throw Error("tokeniser_not_initialized.");
|
|
157
|
+
let t = await this._tokeniser.train(e);
|
|
158
|
+
return this._status === "awaitingTokens" && this.setStatus("ready"), t;
|
|
159
|
+
}
|
|
160
|
+
get responses() {
|
|
161
|
+
if (!this._responses) {
|
|
162
|
+
if (!this._model || !this._tokeniser) throw Error("model_or_tokeniser_not_initialized.");
|
|
163
|
+
this._responses = new l(this._model, this._tokeniser), this._responses.on("error", (e) => {
|
|
164
|
+
this.setStatus("error"), this.ee.emit("error", e);
|
|
165
|
+
}), this._responses.on("status", (e) => {
|
|
166
|
+
e === "busy" ? this.setStatus("busy") : e === "ready" && this.setStatus("ready");
|
|
167
|
+
});
|
|
168
|
+
}
|
|
169
|
+
return this._responses;
|
|
170
|
+
}
|
|
171
|
+
get training() {
|
|
172
|
+
if (!this._training) {
|
|
173
|
+
if (!this._model || !this._tokeniser) throw Error("model_or_tokeniser_not_initialized.");
|
|
174
|
+
this._training = new u(this._model, this._tokeniser), this._training.on("running", () => {
|
|
175
|
+
this.setStatus("busy");
|
|
176
|
+
}), this._training.on("completed", () => {
|
|
177
|
+
this._training?.training || this.setStatus("ready");
|
|
178
|
+
}), this._training.on("cancelled", () => {
|
|
179
|
+
this._training?.training || this.setStatus("ready");
|
|
180
|
+
}), this._training.on("error", (e) => {
|
|
181
|
+
this.setStatus("error"), this.ee.emit("error", e);
|
|
182
|
+
});
|
|
183
|
+
}
|
|
184
|
+
return this._training;
|
|
185
|
+
}
|
|
186
|
+
dispose() {
|
|
187
|
+
this._responses &&= (this._responses.dispose(), null), this._training &&= (this._training.dispose(), null), this._model?.dispose(), this.ee.removeAllListeners();
|
|
188
|
+
}
|
|
189
|
+
on(e, t) {
|
|
190
|
+
if (e === "loaded" && this.loaded) {
|
|
191
|
+
setTimeout(() => t(), 0);
|
|
192
|
+
return;
|
|
193
|
+
}
|
|
194
|
+
this.ee.on(e, t);
|
|
195
|
+
}
|
|
196
|
+
off(e, t) {
|
|
197
|
+
this.ee.off(e, t);
|
|
198
|
+
}
|
|
199
|
+
};
|
|
200
|
+
//#endregion
|
|
201
|
+
export { f as default };
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import { IGenerateOptions, IGeneratorResponse } from '../../inference/types';
|
|
2
|
+
import { ITokeniser } from '../../tokeniser/type';
|
|
3
|
+
import { default as Model, ModelForwardAttributes } from '../../models/model';
|
|
4
|
+
import { GPTConfig } from '../../models/config';
|
|
5
|
+
interface ResponseEvents {
|
|
6
|
+
status: (status: 'ready' | 'busy') => void;
|
|
7
|
+
error: (error: Error) => void;
|
|
8
|
+
done: (id: string) => void;
|
|
9
|
+
hook: (id: string) => void;
|
|
10
|
+
generating: (id: string) => void;
|
|
11
|
+
}
|
|
12
|
+
export default class Responses {
|
|
13
|
+
private ee;
|
|
14
|
+
private _model;
|
|
15
|
+
private _tokeniser;
|
|
16
|
+
private _busyCount;
|
|
17
|
+
private _responses;
|
|
18
|
+
private _jobQueue;
|
|
19
|
+
private _hookedResponses;
|
|
20
|
+
private _resumeWaiters;
|
|
21
|
+
constructor(model: Model<ModelForwardAttributes, GPTConfig>, tokeniser: ITokeniser);
|
|
22
|
+
/** Number of generation jobs currently queued. */
|
|
23
|
+
get queued(): number;
|
|
24
|
+
private _processNextJob;
|
|
25
|
+
private _runForRecord;
|
|
26
|
+
/**
|
|
27
|
+
* Subscribe to response lifecycle events.
|
|
28
|
+
* @param event One of `status`, `error`, `done`, `hook`, or `generating`
|
|
29
|
+
* @param listener Callback invoked when the event fires
|
|
30
|
+
*/
|
|
31
|
+
on<E extends keyof ResponseEvents>(event: E, listener: ResponseEvents[E]): void;
|
|
32
|
+
/**
|
|
33
|
+
* Unsubscribe a previously registered event listener.
|
|
34
|
+
* @param event Event name the listener was registered for
|
|
35
|
+
* @param listener The listener to remove
|
|
36
|
+
*/
|
|
37
|
+
off<E extends keyof ResponseEvents>(event: E, listener: ResponseEvents[E]): void;
|
|
38
|
+
private generator;
|
|
39
|
+
private _cleanup;
|
|
40
|
+
/** Start a new text generation or resume a previous one.
|
|
41
|
+
* @param options Response options object
|
|
42
|
+
* @param callback Intermediate responses per token or chunk.
|
|
43
|
+
*/
|
|
44
|
+
create(options: IGenerateOptions, callback?: (response: IGeneratorResponse) => void): Promise<IGeneratorResponse>;
|
|
45
|
+
/**
|
|
46
|
+
* Retry a previous response generation, optionally truncating the conversation at `index`.
|
|
47
|
+
* @param id ID of the existing response to retry
|
|
48
|
+
* @param index Optional conversation index to truncate before regenerating
|
|
49
|
+
*/
|
|
50
|
+
retry(id: string, index?: number): Promise<IGeneratorResponse>;
|
|
51
|
+
/**
|
|
52
|
+
* Get the current generator response for a given ID.
|
|
53
|
+
* Returns `null` if no response with that ID exists.
|
|
54
|
+
* @param id Response ID
|
|
55
|
+
*/
|
|
56
|
+
getResponse(id: string): IGeneratorResponse | null;
|
|
57
|
+
/**
|
|
58
|
+
* Cancel an in-progress generation and resolve any waiting hooks.
|
|
59
|
+
* @param id Response ID to cancel
|
|
60
|
+
* @returns `true` if the response was found and cancelled, otherwise `false`
|
|
61
|
+
*/
|
|
62
|
+
cancel(id: string): boolean;
|
|
63
|
+
/**
|
|
64
|
+
* Put a response into "hooked" mode so generation pauses at the next chunk.
|
|
65
|
+
* @param id Response ID to hook
|
|
66
|
+
* @returns `true` if the response exists and was hooked, otherwise `false`
|
|
67
|
+
*/
|
|
68
|
+
hook(id: string): boolean;
|
|
69
|
+
/**
|
|
70
|
+
* Resume a previously hooked response, releasing a single paused chunk.
|
|
71
|
+
* @param id Response ID to resume
|
|
72
|
+
* @returns `true` if the resume action was performed or the response is still hooked
|
|
73
|
+
*/
|
|
74
|
+
resume(id: string): boolean;
|
|
75
|
+
/**
|
|
76
|
+
* Dispose all generators, clear queues and internal bookkeeping.
|
|
77
|
+
* This releases resources held by this Responses manager.
|
|
78
|
+
*/
|
|
79
|
+
dispose(): void;
|
|
80
|
+
}
|
|
81
|
+
export {};
|
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
import { t as e } from "../eventemitter3-D_qV3Lof.js";
|
|
2
|
+
import t from "../inference/Generator.js";
|
|
3
|
+
import { t as n } from "../v4-BK7K-jy_.js";
|
|
4
|
+
//#region lib/api/responses.ts
|
|
5
|
+
var r = class {
|
|
6
|
+
ee;
|
|
7
|
+
_model;
|
|
8
|
+
_tokeniser;
|
|
9
|
+
_busyCount = 0;
|
|
10
|
+
_responses = /* @__PURE__ */ new Map();
|
|
11
|
+
_jobQueue = [];
|
|
12
|
+
_hookedResponses = /* @__PURE__ */ new Set();
|
|
13
|
+
_resumeWaiters = /* @__PURE__ */ new Map();
|
|
14
|
+
constructor(t, n) {
|
|
15
|
+
this._model = t, this._tokeniser = n, this.ee = new e();
|
|
16
|
+
}
|
|
17
|
+
get queued() {
|
|
18
|
+
return this._jobQueue.length;
|
|
19
|
+
}
|
|
20
|
+
async _processNextJob() {
|
|
21
|
+
if (this._jobQueue.length === 0) return;
|
|
22
|
+
let e = this._jobQueue.shift();
|
|
23
|
+
try {
|
|
24
|
+
let t = await this._runForRecord(e.id, e.options, e.callback);
|
|
25
|
+
e.resolve && e.resolve(t);
|
|
26
|
+
} catch (t) {
|
|
27
|
+
e.reject && e.reject(t), this.ee.emit("error", t);
|
|
28
|
+
}
|
|
29
|
+
}
|
|
30
|
+
async _runForRecord(e, t, n) {
|
|
31
|
+
let r = t.previous_response_id ? this._responses.get(t.previous_response_id) : void 0;
|
|
32
|
+
if (r) {
|
|
33
|
+
let e = r.generator.getConversation();
|
|
34
|
+
t.input && Array.isArray(t.input) ? t.input = [...e, ...t.input] : t.input && typeof t.input == "string" ? t.input = [...e, {
|
|
35
|
+
role: t.nonConversational ? "text" : "user",
|
|
36
|
+
content: t.input
|
|
37
|
+
}] : t.input = e;
|
|
38
|
+
}
|
|
39
|
+
let i = r ? r.generator : this._responses.get(e).generator, a = n ? () => n({
|
|
40
|
+
output: i.getConversation(),
|
|
41
|
+
id: e,
|
|
42
|
+
done: !1
|
|
43
|
+
}) : null;
|
|
44
|
+
a && i.on("tokens", a);
|
|
45
|
+
let o = this._responses.get(e), s = {
|
|
46
|
+
...t,
|
|
47
|
+
_onChunk: async () => {
|
|
48
|
+
this._hookedResponses.has(e) && await new Promise((t) => {
|
|
49
|
+
let n = this._resumeWaiters.get(e) || [];
|
|
50
|
+
n.push(t), this._resumeWaiters.set(e, n), this.ee.emit("hook", e);
|
|
51
|
+
});
|
|
52
|
+
}
|
|
53
|
+
};
|
|
54
|
+
this.ee.emit("generating", t.previous_response_id || e);
|
|
55
|
+
let c = t.input ? i.generate(Array.isArray(t.input) ? t.input : [{
|
|
56
|
+
role: t.nonConversational ? "text" : "user",
|
|
57
|
+
content: t.input
|
|
58
|
+
}], s) : i.generate(s);
|
|
59
|
+
if (t.background) return c.then(() => {
|
|
60
|
+
o.done = !0, this._hookedResponses.delete(e), this._resumeWaiters.delete(e), a && i.off("tokens", a), this.ee.emit("done", e), this._processNextJob();
|
|
61
|
+
}), {
|
|
62
|
+
output: null,
|
|
63
|
+
id: e,
|
|
64
|
+
done: !1
|
|
65
|
+
};
|
|
66
|
+
let l = await c;
|
|
67
|
+
return a && i.off("tokens", a), o.done = !0, this._hookedResponses.delete(e), this._resumeWaiters.delete(e), this.ee.emit("done", e), this._processNextJob(), {
|
|
68
|
+
output: l,
|
|
69
|
+
id: e,
|
|
70
|
+
done: !0
|
|
71
|
+
};
|
|
72
|
+
}
|
|
73
|
+
on(e, t) {
|
|
74
|
+
this.ee.on(e, t);
|
|
75
|
+
}
|
|
76
|
+
off(e, t) {
|
|
77
|
+
this.ee.off(e, t);
|
|
78
|
+
}
|
|
79
|
+
generator() {
|
|
80
|
+
if (!this._model || !this._tokeniser) throw Error("model_or_tokeniser_not_initialized.");
|
|
81
|
+
let e = new t(this._model, this._tokeniser);
|
|
82
|
+
return e.on("start", () => {
|
|
83
|
+
this._busyCount === 0 && this.ee.emit("status", "busy"), this._busyCount++;
|
|
84
|
+
}), e.on("stop", () => {
|
|
85
|
+
this._busyCount--, this._busyCount === 0 && this.ee.emit("status", "ready");
|
|
86
|
+
}), e;
|
|
87
|
+
}
|
|
88
|
+
_cleanup() {
|
|
89
|
+
let e = Date.now();
|
|
90
|
+
this._responses.forEach((t, n) => {
|
|
91
|
+
t.done && e - t.timestamp > 300 * 1e3 && (t.generator.dispose(), this._responses.delete(n));
|
|
92
|
+
});
|
|
93
|
+
}
|
|
94
|
+
async create(e, t) {
|
|
95
|
+
let r = n(), i = e.previous_response_id ? this._responses.get(e.previous_response_id) : void 0;
|
|
96
|
+
i ? i.timestamp = Date.now() : this._cleanup();
|
|
97
|
+
let a = {
|
|
98
|
+
id: r,
|
|
99
|
+
generator: i ? i.generator : this.generator(),
|
|
100
|
+
done: !1,
|
|
101
|
+
timestamp: Date.now(),
|
|
102
|
+
options: e,
|
|
103
|
+
callback: t
|
|
104
|
+
};
|
|
105
|
+
if (this._responses.set(r, a), this._busyCount > 0 && !e.background) {
|
|
106
|
+
if (this._jobQueue.length > 10) throw Error("Job queue is too long, rejecting new job");
|
|
107
|
+
return new Promise((n, i) => {
|
|
108
|
+
this._jobQueue.push({
|
|
109
|
+
id: r,
|
|
110
|
+
options: e,
|
|
111
|
+
callback: t,
|
|
112
|
+
resolve: n,
|
|
113
|
+
reject: i
|
|
114
|
+
});
|
|
115
|
+
});
|
|
116
|
+
}
|
|
117
|
+
return this._runForRecord(r, e, t);
|
|
118
|
+
}
|
|
119
|
+
async retry(e, t) {
|
|
120
|
+
let n = this._responses.get(e);
|
|
121
|
+
if (n) {
|
|
122
|
+
let r = n.generator.getConversation();
|
|
123
|
+
if (t !== void 0 && t >= 0 && t < r.length && r.splice(t), n.done = !1, n.timestamp = Date.now(), this._busyCount > 0 && !n.options.background) {
|
|
124
|
+
if (this._jobQueue.length > 10) throw Error("Job queue is too long, rejecting new job");
|
|
125
|
+
return new Promise((t, r) => {
|
|
126
|
+
this._jobQueue.push({
|
|
127
|
+
id: e,
|
|
128
|
+
options: n.options,
|
|
129
|
+
callback: n.callback,
|
|
130
|
+
resolve: t,
|
|
131
|
+
reject: r
|
|
132
|
+
});
|
|
133
|
+
});
|
|
134
|
+
}
|
|
135
|
+
return this._runForRecord(e, n.options, n.callback);
|
|
136
|
+
}
|
|
137
|
+
throw Error(`No response found for id: ${e}`);
|
|
138
|
+
}
|
|
139
|
+
getResponse(e) {
|
|
140
|
+
let t = this._responses.get(e);
|
|
141
|
+
return t ? {
|
|
142
|
+
output: t.generator.getConversation(),
|
|
143
|
+
id: t.id,
|
|
144
|
+
done: t.done
|
|
145
|
+
} : null;
|
|
146
|
+
}
|
|
147
|
+
cancel(e) {
|
|
148
|
+
let t = this._responses.get(e);
|
|
149
|
+
return t ? (this._hookedResponses.delete(e), (this._resumeWaiters.get(e) || []).forEach((e) => e()), this._resumeWaiters.delete(e), t.generator.stop(), !0) : !1;
|
|
150
|
+
}
|
|
151
|
+
hook(e) {
|
|
152
|
+
return this._responses.has(e) ? (this._hookedResponses.add(e), !0) : !1;
|
|
153
|
+
}
|
|
154
|
+
resume(e) {
|
|
155
|
+
let t = this._resumeWaiters.get(e);
|
|
156
|
+
if (!t || t.length === 0) return this._hookedResponses.has(e);
|
|
157
|
+
let n = t.shift();
|
|
158
|
+
return t.length === 0 && this._resumeWaiters.delete(e), n(), !0;
|
|
159
|
+
}
|
|
160
|
+
dispose() {
|
|
161
|
+
this._resumeWaiters.forEach((e) => {
|
|
162
|
+
e.forEach((e) => e());
|
|
163
|
+
}), this._resumeWaiters.clear(), this._hookedResponses.clear(), this._responses.forEach((e) => {
|
|
164
|
+
e.generator.dispose();
|
|
165
|
+
}), this._responses.clear(), this.ee.removeAllListeners();
|
|
166
|
+
}
|
|
167
|
+
};
|
|
168
|
+
//#endregion
|
|
169
|
+
export { r as default };
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
import { ITokeniser } from '../../tokeniser/type';
|
|
2
|
+
import { default as Model, ModelForwardAttributes } from '../../models/model';
|
|
3
|
+
import { GPTConfig } from '../../models/config';
|
|
4
|
+
import { ConversationStream } from '../../data/stream';
|
|
5
|
+
import { TokenStore } from '../../training/tasks/TokenStore';
|
|
6
|
+
import { DatasetMetadata } from '../../loader/types';
|
|
7
|
+
import { TrainingLogEntry, TrainingOptions } from '../../training/types';
|
|
8
|
+
import { default as BasicTrainer } from '../../training/BasicTrainer';
|
|
9
|
+
import { Dataset } from '@tensorflow/tfjs-data';
|
|
10
|
+
import { Tensor } from '@tensorflow/tfjs-core';
|
|
11
|
+
import { AdamWOptimizer } from '../../training/AdamW';
|
|
12
|
+
export type TrainingState = 'pending' | 'running' | 'paused' | 'pausing' | 'completed' | 'cancelled' | 'cancelling' | 'error';
|
|
13
|
+
export interface ITrainingJob {
|
|
14
|
+
id: string;
|
|
15
|
+
state: TrainingState;
|
|
16
|
+
history: TrainingLogEntry[] | null;
|
|
17
|
+
progress: number;
|
|
18
|
+
remaining: number;
|
|
19
|
+
trainer: BasicTrainer;
|
|
20
|
+
options: TrainingOptions;
|
|
21
|
+
trainDataset?: Dataset<{
|
|
22
|
+
xs: Tensor;
|
|
23
|
+
ys: Tensor;
|
|
24
|
+
}>;
|
|
25
|
+
validationDataset?: Dataset<{
|
|
26
|
+
xs: Tensor;
|
|
27
|
+
ys: Tensor;
|
|
28
|
+
}>;
|
|
29
|
+
totalTokens: number;
|
|
30
|
+
datasets: DatasetMetadata[];
|
|
31
|
+
datasetId?: string;
|
|
32
|
+
breakOnLog: boolean;
|
|
33
|
+
}
|
|
34
|
+
interface TrainingEvents {
|
|
35
|
+
error: (id: string, error: Error) => void;
|
|
36
|
+
completed: (id: string) => void;
|
|
37
|
+
running: (id: string) => void;
|
|
38
|
+
paused: (id: string) => void;
|
|
39
|
+
cancelled: (id: string) => void;
|
|
40
|
+
pausing: (id: string) => void;
|
|
41
|
+
cancelling: (id: string) => void;
|
|
42
|
+
progress: (job: ITrainingJob) => void;
|
|
43
|
+
}
|
|
44
|
+
export default class Training {
|
|
45
|
+
private ee;
|
|
46
|
+
private _model;
|
|
47
|
+
private _tokeniser;
|
|
48
|
+
private _jobs;
|
|
49
|
+
constructor(model: Model<ModelForwardAttributes, GPTConfig>, tokeniser: ITokeniser);
|
|
50
|
+
private setState;
|
|
51
|
+
get activeJobs(): number;
|
|
52
|
+
get training(): boolean;
|
|
53
|
+
private assertValidJobId;
|
|
54
|
+
private assertValidOptions;
|
|
55
|
+
private assertValidDatasets;
|
|
56
|
+
private assertValidDataInput;
|
|
57
|
+
on<E extends keyof TrainingEvents>(event: E, listener: TrainingEvents[E]): void;
|
|
58
|
+
off<E extends keyof TrainingEvents>(event: E, listener: TrainingEvents[E]): void;
|
|
59
|
+
restore(options: TrainingOptions, log: TrainingLogEntry[], optimizer: AdamWOptimizer, datasets: DatasetMetadata[]): ITrainingJob;
|
|
60
|
+
private launchJob;
|
|
61
|
+
job(options: TrainingOptions, data: (ConversationStream[] | Uint16Array[] | TokenStore) | undefined, datasets: DatasetMetadata[], validation?: Uint16Array[] | TokenStore): Promise<ITrainingJob>;
|
|
62
|
+
getJob(id: string): ITrainingJob | null;
|
|
63
|
+
/** Resume a paused job and optionally change some options. */
|
|
64
|
+
resume(id: string): Promise<void>;
|
|
65
|
+
cancel(id: string): void;
|
|
66
|
+
breakpoints(id: string, enabled: boolean): void;
|
|
67
|
+
getPretrainingJob(): ITrainingJob | null;
|
|
68
|
+
dispose(): void;
|
|
69
|
+
}
|
|
70
|
+
export {};
|