@genai-fi/nanogpt 1.0.4 → 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
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;
|
|
@@ -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;
|
|
@@ -28,7 +28,7 @@ async function a(a, o, s, c, l, u, d) {
|
|
|
28
28
|
}
|
|
29
29
|
else m = c;
|
|
30
30
|
let g = m instanceof e ? m.getTokenCount() : m.reduce((e, t) => e + t.length, 0);
|
|
31
|
-
a.epochSteps = Math.
|
|
31
|
+
a.epochSteps = Math.max(1, Math.floor(g / ((a?.batchSize || 32) * o.config.blockSize)));
|
|
32
32
|
let _ = new n(s, o.config.blockSize);
|
|
33
33
|
if (h) {
|
|
34
34
|
let { trainDataset: e, validationDataset: t } = await r(m, h, s, _, a?.batchSize || 32);
|