@genai-fi/nanogpt 1.0.3 → 1.1.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/dist/api/training.d.ts +5 -2
- package/dist/api/training.js +24 -15
- package/dist/inference/Generator.d.ts +1 -0
- package/dist/inference/Generator.js +6 -1
- package/dist/loader/types.d.ts +1 -0
- package/dist/training/BasicTrainer.js +3 -1
- package/dist/training/prepareData.js +9 -8
- package/package.json +1 -1
package/dist/api/training.d.ts
CHANGED
|
@@ -29,7 +29,8 @@ export interface ITrainingJob {
|
|
|
29
29
|
totalTokens: number;
|
|
30
30
|
datasets: DatasetMetadata[];
|
|
31
31
|
datasetId?: string;
|
|
32
|
-
breakOnLog:
|
|
32
|
+
breakOnLog: Set<number>;
|
|
33
|
+
resumeCounter: number;
|
|
33
34
|
}
|
|
34
35
|
interface TrainingEvents {
|
|
35
36
|
error: (id: string, error: Error) => void;
|
|
@@ -46,6 +47,7 @@ export default class Training {
|
|
|
46
47
|
private _model;
|
|
47
48
|
private _tokeniser;
|
|
48
49
|
private _jobs;
|
|
50
|
+
private _breakCounter;
|
|
49
51
|
constructor(model: Model<ModelForwardAttributes, GPTConfig>, tokeniser: ITokeniser);
|
|
50
52
|
private setState;
|
|
51
53
|
get activeJobs(): number;
|
|
@@ -63,7 +65,8 @@ export default class Training {
|
|
|
63
65
|
/** Resume a paused job and optionally change some options. */
|
|
64
66
|
resume(id: string): Promise<void>;
|
|
65
67
|
cancel(id: string): void;
|
|
66
|
-
|
|
68
|
+
addBreak(jobid: string): number | undefined;
|
|
69
|
+
deleteBreak(jobid: string, breakId: number): void;
|
|
67
70
|
getPretrainingJob(): ITrainingJob | null;
|
|
68
71
|
dispose(): void;
|
|
69
72
|
}
|
package/dist/api/training.js
CHANGED
|
@@ -10,6 +10,7 @@ var o = class {
|
|
|
10
10
|
_model;
|
|
11
11
|
_tokeniser;
|
|
12
12
|
_jobs = /* @__PURE__ */ new Map();
|
|
13
|
+
_breakCounter = 1;
|
|
13
14
|
constructor(t, n) {
|
|
14
15
|
this.ee = new e(), this._model = t, this._tokeniser = n;
|
|
15
16
|
}
|
|
@@ -19,32 +20,31 @@ var o = class {
|
|
|
19
20
|
};
|
|
20
21
|
switch (t) {
|
|
21
22
|
case "pending":
|
|
22
|
-
e.state !== "completed" && e.state !== "paused" && e.state !== "cancelled" && e.state !== "pending" && r();
|
|
23
|
+
e.state !== "completed" && e.state !== "paused" && e.state !== "cancelled" && e.state !== "pending" && r(), e.state = t;
|
|
23
24
|
break;
|
|
24
25
|
case "running":
|
|
25
|
-
e.state !== "pending" && e.state !== "paused" && r(), this.ee.emit("running", e.id);
|
|
26
|
+
e.state !== "pending" && e.state !== "paused" && r(), e.state = t, this.ee.emit("running", e.id);
|
|
26
27
|
break;
|
|
27
28
|
case "paused":
|
|
28
|
-
e.state !== "pausing" && r(), this.ee.emit("paused", e.id);
|
|
29
|
+
e.state !== "pausing" && r(), e.state = t, this.ee.emit("paused", e.id);
|
|
29
30
|
break;
|
|
30
31
|
case "completed":
|
|
31
|
-
e.state !== "running" && r(), this.ee.emit("completed", e.id);
|
|
32
|
+
e.state !== "running" && r(), e.state = t, this.ee.emit("completed", e.id);
|
|
32
33
|
break;
|
|
33
34
|
case "pausing":
|
|
34
|
-
e.state !== "running" && r(), this.ee.emit("pausing", e.id);
|
|
35
|
+
e.state !== "running" && r(), e.state = t, this.ee.emit("pausing", e.id);
|
|
35
36
|
break;
|
|
36
37
|
case "cancelling":
|
|
37
|
-
e.state !== "running" && e.state !== "paused" && e.state !== "pausing" && r(), this.ee.emit("cancelling", e.id);
|
|
38
|
+
e.state !== "running" && e.state !== "paused" && e.state !== "pausing" && r(), e.state = t, this.ee.emit("cancelling", e.id);
|
|
38
39
|
break;
|
|
39
40
|
case "cancelled":
|
|
40
|
-
e.state !== "paused" && e.state !== "cancelling" && r(), this.ee.emit("cancelled", e.id);
|
|
41
|
+
e.state !== "paused" && e.state !== "cancelling" && r(), e.state = t, this.ee.emit("cancelled", e.id);
|
|
41
42
|
break;
|
|
42
43
|
case "error":
|
|
43
|
-
n && this.ee.emit("error", e.id, n);
|
|
44
|
+
e.state = t, n && this.ee.emit("error", e.id, n);
|
|
44
45
|
break;
|
|
45
46
|
default: throw Error("invalid_state");
|
|
46
47
|
}
|
|
47
|
-
e.state = t;
|
|
48
48
|
}
|
|
49
49
|
get activeJobs() {
|
|
50
50
|
return this._jobs.values().reduce((e, t) => t.state === "running" ? e + 1 : e, 0);
|
|
@@ -102,13 +102,14 @@ var o = class {
|
|
|
102
102
|
totalTokens: t[t.length - 1]?.totalTokens || 0,
|
|
103
103
|
datasets: o,
|
|
104
104
|
datasetId: a(o),
|
|
105
|
-
breakOnLog:
|
|
105
|
+
breakOnLog: /* @__PURE__ */ new Set(),
|
|
106
|
+
resumeCounter: 0
|
|
106
107
|
};
|
|
107
108
|
return s.trainer.log = t, s.trainer.resumeFromLog(t[t.length - 1]), this._jobs.set(s.id, s), s;
|
|
108
109
|
}
|
|
109
110
|
launchJob(e) {
|
|
110
111
|
e.trainDataset && (this.setState(e, "running"), e.trainer.trainOnDataset(e.trainDataset, e.options, e.validationDataset, (t) => {
|
|
111
|
-
e.history = e.trainer.log, e.progress = t.totalTokens / e.totalTokens, e.remaining = Math.max(0, (e.totalTokens - t.totalTokens) / t.totalTokens * t.duration), e.breakOnLog && e.state === "running" && (this.setState(e, "pausing"), e.trainer.stop()), this.ee.emit("progress", e);
|
|
112
|
+
e.history = e.trainer.log, e.progress = t.totalTokens / e.totalTokens, e.remaining = Math.max(0, (e.totalTokens - t.totalTokens) / t.totalTokens * t.duration), e.breakOnLog.size > 0 && e.state === "running" && (this.setState(e, "pausing"), e.trainer.stop()), this.ee.emit("progress", e);
|
|
112
113
|
}).then(() => {
|
|
113
114
|
e.state === "running" ? this.setState(e, "completed") : e.state === "pausing" ? this.setState(e, "paused") : e.state === "cancelling" && this.setState(e, "cancelled");
|
|
114
115
|
}).catch((t) => {
|
|
@@ -135,7 +136,8 @@ var o = class {
|
|
|
135
136
|
options: e,
|
|
136
137
|
totalTokens: 0,
|
|
137
138
|
datasets: s,
|
|
138
|
-
breakOnLog:
|
|
139
|
+
breakOnLog: /* @__PURE__ */ new Set(),
|
|
140
|
+
resumeCounter: 0
|
|
139
141
|
};
|
|
140
142
|
if (!u) throw Error("invalid_previous_job_id");
|
|
141
143
|
if (this._jobs.set(u.id, u), this.setState(u, "pending"), o instanceof t) {
|
|
@@ -175,7 +177,7 @@ var o = class {
|
|
|
175
177
|
async resume(e) {
|
|
176
178
|
this.assertValidJobId(e);
|
|
177
179
|
let t = this.getJob(e);
|
|
178
|
-
t && t.state === "paused" && (this.setState(t, "pending"), this.launchJob(t));
|
|
180
|
+
t && t.state === "paused" && (t.resumeCounter++, !(t.resumeCounter < t.breakOnLog.size) && (t.resumeCounter = 0, this.setState(t, "pending"), this.launchJob(t)));
|
|
179
181
|
}
|
|
180
182
|
cancel(e) {
|
|
181
183
|
this.assertValidJobId(e);
|
|
@@ -185,9 +187,16 @@ var o = class {
|
|
|
185
187
|
return;
|
|
186
188
|
} else (t.state === "running" || t.state === "pausing") && (this.setState(t, "cancelling"), t.trainer.stop());
|
|
187
189
|
}
|
|
188
|
-
|
|
190
|
+
addBreak(e) {
|
|
191
|
+
let t = this.getJob(e);
|
|
192
|
+
if (t) {
|
|
193
|
+
let e = this._breakCounter++;
|
|
194
|
+
return t.breakOnLog.add(e), e;
|
|
195
|
+
}
|
|
196
|
+
}
|
|
197
|
+
deleteBreak(e, t) {
|
|
189
198
|
let n = this.getJob(e);
|
|
190
|
-
n &&
|
|
199
|
+
n && n.breakOnLog.delete(t);
|
|
191
200
|
}
|
|
192
201
|
getPretrainingJob() {
|
|
193
202
|
for (let e of this._jobs.values()) if (e.options.method.type === "pretraining") return e;
|
|
@@ -34,6 +34,7 @@ export default class Generator extends EE<'start' | 'stop' | 'tokens' | 'reset'>
|
|
|
34
34
|
private startTime;
|
|
35
35
|
private tokenCount;
|
|
36
36
|
constructor(model: Model<ModelForwardAttributes>, tokeniser: ITokeniser);
|
|
37
|
+
private shouldTerminate;
|
|
37
38
|
/** Generate logits and select a token. */
|
|
38
39
|
private _generateToken;
|
|
39
40
|
/** Generate multiple tokens in a loop and produce text */
|
|
@@ -40,6 +40,11 @@ var b = class extends e {
|
|
|
40
40
|
constructor(e, t) {
|
|
41
41
|
super(), this.model = e, this.tokeniser = t, this.actualTokeniser = t;
|
|
42
42
|
}
|
|
43
|
+
shouldTerminate(e, t) {
|
|
44
|
+
if (e) return !1;
|
|
45
|
+
let n = this.tokeniser.getSpecialTokenIndex("<|assistant_end|>");
|
|
46
|
+
return t === this.actualTokeniser.eosToken || t === n;
|
|
47
|
+
}
|
|
43
48
|
async _generateToken(e, t, n) {
|
|
44
49
|
let s = n?.temperature ?? 1, h = n?.topK, _ = n?.topP, v = n?.usePadding ?? !1, y = {
|
|
45
50
|
training: !1,
|
|
@@ -132,7 +137,7 @@ var b = class extends e {
|
|
|
132
137
|
S.dispose(), S = E;
|
|
133
138
|
let D = (await S.array())[0][0], O = this.actualTokeniser.decode([D]);
|
|
134
139
|
this.lastToken = D;
|
|
135
|
-
let k =
|
|
140
|
+
let k = this.shouldTerminate(n?.allowSpecial ?? !1, D), A = {
|
|
136
141
|
outputTensor: S,
|
|
137
142
|
token: D,
|
|
138
143
|
text: O,
|
package/dist/loader/types.d.ts
CHANGED
|
@@ -234,7 +234,9 @@ var p = {
|
|
|
234
234
|
if (isNaN(e) || !isFinite(e)) throw console.error("Invalid loss value:", e), console.error("Batch xs:", await n.xs.array()), console.error("Batch ys:", await n.ys.array()), console.error("State:", l), Error("Loss is NaN or Infinity");
|
|
235
235
|
console.log(`Step ${l.step}: Loss = ${e}`);
|
|
236
236
|
}
|
|
237
|
-
n.xs.dispose(), n.ys.dispose(), l.step++, l.totalSteps
|
|
237
|
+
n.xs.dispose(), n.ys.dispose(), l.step++, l.totalSteps++;
|
|
238
|
+
let u = l.step >= c;
|
|
239
|
+
r || u ? await this.performLogging(s, n.xs.shape[0], f, i) : (l.gradientNorm &&= (l.gradientNorm.dispose(), void 0), l.accuracy &&= (l.accuracy.dispose(), void 0)), s.dispose(), u && this.stop();
|
|
238
240
|
}
|
|
239
241
|
} catch (e) {
|
|
240
242
|
throw console.error("Training error:", e), r(), e;
|
|
@@ -8,13 +8,14 @@ async function a(a, o, s, c, l, u, d) {
|
|
|
8
8
|
if (d && f) throw Error("Cannot specify datasets when using LoRA fine-tuning");
|
|
9
9
|
if (!d && !f) throw Error("Must specify datasets for non-LoRA training");
|
|
10
10
|
if (d) {
|
|
11
|
-
let e =
|
|
12
|
-
for (let
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
11
|
+
let e = !1;
|
|
12
|
+
for (let t of d) t.conversational && (e = !0);
|
|
13
|
+
o.metaData.pretrainingData = d.map((e) => ({
|
|
14
|
+
id: e.id,
|
|
15
|
+
name: e.name,
|
|
16
|
+
conversational: e.conversational,
|
|
17
|
+
url: e.url
|
|
18
|
+
})), e ? o.metaData.mode = "conversational" : o.metaData.mode !== "conversational" && (o.metaData.mode = "completion");
|
|
18
19
|
} else o.metaData.mode !== "conversational" && (o.metaData.mode = "completion");
|
|
19
20
|
let p = a.maskedLoss ?? a.method.type === "supervised", m, h = u;
|
|
20
21
|
if (Array.isArray(c)) if (c[0] instanceof Uint16Array) m = c;
|
|
@@ -27,7 +28,7 @@ async function a(a, o, s, c, l, u, d) {
|
|
|
27
28
|
}
|
|
28
29
|
else m = c;
|
|
29
30
|
let g = m instanceof e ? m.getTokenCount() : m.reduce((e, t) => e + t.length, 0);
|
|
30
|
-
a.epochSteps = Math.
|
|
31
|
+
a.epochSteps = Math.max(1, Math.floor(g / ((a?.batchSize || 32) * o.config.blockSize)));
|
|
31
32
|
let _ = new n(s, o.config.blockSize);
|
|
32
33
|
if (h) {
|
|
33
34
|
let { trainDataset: e, validationDataset: t } = await r(m, h, s, _, a?.batchSize || 32);
|