@genai-fi/nanogpt 1.0.4 → 1.1.1

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.
@@ -66,6 +66,7 @@ export default class Responses {
66
66
  * @returns `true` if the response exists and was hooked, otherwise `false`
67
67
  */
68
68
  hook(id: string): boolean;
69
+ unhook(id: string): void;
69
70
  /**
70
71
  * Resume a previously hooked response, releasing a single paused chunk.
71
72
  * @param id Response ID to resume
@@ -151,6 +151,9 @@ var r = class {
151
151
  hook(e) {
152
152
  return this._responses.has(e) ? (this._hookedResponses.add(e), !0) : !1;
153
153
  }
154
+ unhook(e) {
155
+ this._hookedResponses.delete(e), (this._resumeWaiters.get(e) || []).forEach((e) => e()), this._resumeWaiters.delete(e);
156
+ }
154
157
  resume(e) {
155
158
  let t = this._resumeWaiters.get(e);
156
159
  if (!t || t.length === 0) return this._hookedResponses.has(e);
@@ -29,7 +29,8 @@ export interface ITrainingJob {
29
29
  totalTokens: number;
30
30
  datasets: DatasetMetadata[];
31
31
  datasetId?: string;
32
- breakOnLog: boolean;
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
- breakpoints(id: string, enabled: boolean): void;
68
+ addBreak(jobid: string): number | undefined;
69
+ deleteBreak(jobid: string, breakId: number): void;
67
70
  getPretrainingJob(): ITrainingJob | null;
68
71
  dispose(): void;
69
72
  }
@@ -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: !1
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: !1
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
- breakpoints(e, t) {
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 && (n.breakOnLog = t);
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++, r ? 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(), l.step >= c && this.stop();
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.ceil(g / ((a?.batchSize || 32) * o.config.blockSize));
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);
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "@genai-fi/nanogpt",
3
- "version": "1.0.4",
3
+ "version": "1.1.1",
4
4
  "type": "module",
5
5
  "main": "dist/main.js",
6
6
  "types": "dist/main.d.ts",