@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.
@@ -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;
@@ -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 = !n?.allowSpecial && this.tokeniser.isSpecialToken(D), A = {
140
+ let k = this.shouldTerminate(n?.allowSpecial ?? !1, D), A = {
136
141
  outputTensor: S,
137
142
  token: D,
138
143
  text: O,
@@ -36,6 +36,7 @@ export interface DatasetMetadata {
36
36
  id: string;
37
37
  name: string;
38
38
  conversational: boolean;
39
+ url?: string;
39
40
  }
40
41
  export interface ActionLogEntry {
41
42
  action: 'pretrain' | 'generate' | 'finetune';
@@ -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;
@@ -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 = o.metaData.pretrainingData || [], t = [...e], n = !1;
12
- for (let r of d) e.some((e) => e.id === r.id) || t.push({
13
- id: r.id,
14
- name: r.name,
15
- conversational: r.conversational
16
- }), r.conversational && (n = !0);
17
- o.metaData.pretrainingData = t, n ? o.metaData.mode = "conversational" : o.metaData.mode !== "conversational" && (o.metaData.mode = "completion");
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.ceil(g / ((a?.batchSize || 32) * o.config.blockSize));
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);
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "@genai-fi/nanogpt",
3
- "version": "1.0.3",
3
+ "version": "1.1.0",
4
4
  "type": "module",
5
5
  "main": "dist/main.js",
6
6
  "types": "dist/main.d.ts",