@genai-fi/nanogpt 1.0.0 → 1.0.2
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.js +21 -25
- package/dist/training/prepareData.d.ts +1 -1
- package/dist/training/prepareData.js +21 -21
- package/dist/training/tasks/tokenStream.d.ts +1 -1
- package/dist/training/tasks/tokenStream.js +12 -12
- package/dist/training/validation.d.ts +2 -2
- package/dist/training/validation.js +9 -1
- package/package.json +1 -1
package/dist/api/training.js
CHANGED
|
@@ -115,17 +115,17 @@ var o = class {
|
|
|
115
115
|
this.setState(e, "error", t);
|
|
116
116
|
}));
|
|
117
117
|
}
|
|
118
|
-
async job(e,
|
|
119
|
-
if (this.assertValidOptions(e), this.assertValidDataInput(
|
|
120
|
-
let
|
|
121
|
-
for (let t of
|
|
118
|
+
async job(e, o = [], s, c) {
|
|
119
|
+
if (this.assertValidOptions(e), this.assertValidDataInput(o), this.assertValidDatasets(s), c && c instanceof t && c.tokeniserId !== this._tokeniser.id) throw Error("tokeniser_mismatch");
|
|
120
|
+
let l = Array.from(this._jobs.keys());
|
|
121
|
+
for (let t of l) {
|
|
122
122
|
let n = this._jobs.get(t);
|
|
123
123
|
if (n && n.id !== e.previous_job_id) {
|
|
124
124
|
if (n.options.method.type === e.method.type && n.state !== "running") n.trainer.dispose(), this._jobs.delete(n.id);
|
|
125
125
|
else if (n.options.method.type === e.method.type && n.state === "running") throw Error("training_in_progress");
|
|
126
126
|
}
|
|
127
127
|
}
|
|
128
|
-
let
|
|
128
|
+
let u = e.previous_job_id ? this.getJob(e.previous_job_id) : {
|
|
129
129
|
id: n(),
|
|
130
130
|
state: "pending",
|
|
131
131
|
trainer: r(this._model, this._tokeniser, e),
|
|
@@ -134,40 +134,36 @@ var o = class {
|
|
|
134
134
|
remaining: 0,
|
|
135
135
|
options: e,
|
|
136
136
|
totalTokens: 0,
|
|
137
|
-
datasets:
|
|
137
|
+
datasets: s,
|
|
138
138
|
breakOnLog: !1
|
|
139
139
|
};
|
|
140
|
-
if (!
|
|
141
|
-
if (this._jobs.set(
|
|
142
|
-
if (
|
|
143
|
-
let e = /* @__PURE__ */ Error("dataset_mismatch");
|
|
144
|
-
throw this.setState(l, "error", e), e;
|
|
145
|
-
}
|
|
146
|
-
if (a.tokeniserId !== this._tokeniser.id) {
|
|
140
|
+
if (!u) throw Error("invalid_previous_job_id");
|
|
141
|
+
if (this._jobs.set(u.id, u), this.setState(u, "pending"), o instanceof t) {
|
|
142
|
+
if (u.datasetId && o.datasetId !== u.datasetId && console.warn("TokenStore dataset ID does not match job dataset ID, proceeding anyway"), o.tokeniserId !== this._tokeniser.id) {
|
|
147
143
|
let e = /* @__PURE__ */ Error("tokeniser_mismatch");
|
|
148
|
-
throw this.setState(
|
|
144
|
+
throw this.setState(u, "error", e), e;
|
|
149
145
|
}
|
|
150
|
-
|
|
146
|
+
u.datasetId = o.datasetId;
|
|
151
147
|
}
|
|
152
|
-
if (!
|
|
153
|
-
let t = await i(e, this._model, this._tokeniser,
|
|
154
|
-
|
|
148
|
+
if (!u.trainDataset) try {
|
|
149
|
+
let t = await i(e, this._model, this._tokeniser, o, u.datasetId ?? a(s), c, s);
|
|
150
|
+
u.trainDataset = t.trainDataset, u.validationDataset = t.validationDataset, u.totalTokens = t.totalTokens;
|
|
155
151
|
} catch (e) {
|
|
156
|
-
throw this.setState(
|
|
152
|
+
throw this.setState(u, "error", /* @__PURE__ */ Error("prepare_data_failed")), u.trainer.dispose(), this._jobs.delete(u.id), e;
|
|
157
153
|
}
|
|
158
154
|
if (e.previous_job_id) {
|
|
159
|
-
this.assertValidOptions(e),
|
|
155
|
+
this.assertValidOptions(e), u.options = e;
|
|
160
156
|
try {
|
|
161
|
-
|
|
157
|
+
u.trainer.configure(e);
|
|
162
158
|
} catch (e) {
|
|
163
|
-
throw this.setState(
|
|
159
|
+
throw this.setState(u, "error", /* @__PURE__ */ Error("trainer_configuration_failed")), u.trainer.dispose(), this._jobs.delete(u.id), e;
|
|
164
160
|
}
|
|
165
161
|
} else try {
|
|
166
|
-
|
|
162
|
+
u.trainer.configure(e);
|
|
167
163
|
} catch (e) {
|
|
168
|
-
throw this.setState(
|
|
164
|
+
throw this.setState(u, "error", /* @__PURE__ */ Error("trainer_configuration_failed")), u.trainer.dispose(), this._jobs.delete(u.id), e;
|
|
169
165
|
}
|
|
170
|
-
return this.launchJob(
|
|
166
|
+
return this.launchJob(u), u;
|
|
171
167
|
}
|
|
172
168
|
getJob(e) {
|
|
173
169
|
return this.assertValidJobId(e), this._jobs.get(e) ?? null;
|
|
@@ -18,5 +18,5 @@ interface PrepareDataResult {
|
|
|
18
18
|
totalTokens: number;
|
|
19
19
|
}
|
|
20
20
|
/** Take our training options, model, tokeniser, and tasks, and prepare the training and validation datasets in Tensorflow format. */
|
|
21
|
-
export default function prepareData(options: TrainingOptions, model: Model<ModelForwardAttributes>, tokeniser: ITokeniser, tasks: ConversationStream[] | Uint16Array[] | TokenStore, validation?: Uint16Array[] | TokenStore, datasets?: DatasetMetadata[]): Promise<PrepareDataResult>;
|
|
21
|
+
export default function prepareData(options: TrainingOptions, model: Model<ModelForwardAttributes>, tokeniser: ITokeniser, tasks: ConversationStream[] | Uint16Array[] | TokenStore, datasetId: string, validation?: Uint16Array[] | TokenStore, datasets?: DatasetMetadata[]): Promise<PrepareDataResult>;
|
|
22
22
|
export {};
|
|
@@ -3,45 +3,45 @@ import { tokensFromStreams as t } from "./tasks/tokenStream.js";
|
|
|
3
3
|
import { t as n } from "../DatasetBuilder-DU1G1OKX.js";
|
|
4
4
|
import { createTrainValidationDatasets as r, storeFromArray as i } from "./validation.js";
|
|
5
5
|
//#region lib/training/prepareData.ts
|
|
6
|
-
async function a(a, o, s, c, l, u) {
|
|
7
|
-
let
|
|
8
|
-
if (
|
|
9
|
-
if (!
|
|
10
|
-
if (
|
|
6
|
+
async function a(a, o, s, c, l, u, d) {
|
|
7
|
+
let f = a.loraName || a.loraConfig;
|
|
8
|
+
if (d && f) throw Error("Cannot specify datasets when using LoRA fine-tuning");
|
|
9
|
+
if (!d && !f) throw Error("Must specify datasets for non-LoRA training");
|
|
10
|
+
if (d) {
|
|
11
11
|
let e = o.metaData.pretrainingData || [], t = [...e], n = !1;
|
|
12
|
-
for (let r of
|
|
12
|
+
for (let r of d) e.some((e) => e.id === r.id) || t.push({
|
|
13
13
|
id: r.id,
|
|
14
14
|
name: r.name,
|
|
15
15
|
conversational: r.conversational
|
|
16
16
|
}), r.conversational && (n = !0);
|
|
17
17
|
o.metaData.pretrainingData = t, n ? o.metaData.mode = "conversational" : o.metaData.mode !== "conversational" && (o.metaData.mode = "completion");
|
|
18
18
|
} else o.metaData.mode !== "conversational" && (o.metaData.mode = "completion");
|
|
19
|
-
let
|
|
20
|
-
if (Array.isArray(c)) if (c[0] instanceof Uint16Array)
|
|
19
|
+
let p = a.maskedLoss ?? a.method.type === "supervised", m, h = u;
|
|
20
|
+
if (Array.isArray(c)) if (c[0] instanceof Uint16Array) m = c;
|
|
21
21
|
else {
|
|
22
|
-
let e = await t(c, s, {
|
|
23
|
-
masking:
|
|
22
|
+
let e = await t(c, s, l, {
|
|
23
|
+
masking: p,
|
|
24
24
|
validationSplit: a.validationSplit
|
|
25
25
|
});
|
|
26
|
-
|
|
26
|
+
m = e.trainingTokens, u || (h = e.validationTokens);
|
|
27
27
|
}
|
|
28
|
-
else
|
|
29
|
-
let
|
|
30
|
-
a.epochSteps = Math.ceil(
|
|
31
|
-
let
|
|
32
|
-
if (
|
|
33
|
-
let { trainDataset: e, validationDataset: t } = await r(
|
|
28
|
+
else m = c;
|
|
29
|
+
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
|
+
let _ = new n(s, o.config.blockSize);
|
|
32
|
+
if (h) {
|
|
33
|
+
let { trainDataset: e, validationDataset: t } = await r(m, h, s, _, a?.batchSize || 32);
|
|
34
34
|
return {
|
|
35
35
|
trainDataset: e,
|
|
36
36
|
validationDataset: t,
|
|
37
|
-
totalTokens:
|
|
37
|
+
totalTokens: g
|
|
38
38
|
};
|
|
39
39
|
} else {
|
|
40
|
-
let t =
|
|
40
|
+
let t = m instanceof e ? m : await i(m, s);
|
|
41
41
|
return {
|
|
42
|
-
trainDataset: (await
|
|
42
|
+
trainDataset: (await _.createTextDataset(t, a)).dataset,
|
|
43
43
|
validationDataset: void 0,
|
|
44
|
-
totalTokens:
|
|
44
|
+
totalTokens: g
|
|
45
45
|
};
|
|
46
46
|
}
|
|
47
47
|
}
|
|
@@ -10,7 +10,7 @@ interface TokensFromTasksOptions {
|
|
|
10
10
|
validationSeed?: string | number;
|
|
11
11
|
cb?: (tokens: number) => void;
|
|
12
12
|
}
|
|
13
|
-
export declare function tokensFromStreams(tasks: ConversationStream[], tokenizer: ITokeniser, options?: TokensFromTasksOptions): Promise<{
|
|
13
|
+
export declare function tokensFromStreams(tasks: ConversationStream[], tokenizer: ITokeniser, datasetId: string, options?: TokensFromTasksOptions): Promise<{
|
|
14
14
|
trainingTokens: TokenStore;
|
|
15
15
|
validationTokens?: TokenStore;
|
|
16
16
|
}>;
|
|
@@ -22,24 +22,24 @@ function r(e, t, n, r, i, a) {
|
|
|
22
22
|
} else n.set(e, r.offset), s && !Array.isArray(o) && s.set(o.mask.map((e) => +!!e), r.offset), r.offset += e.length;
|
|
23
23
|
}
|
|
24
24
|
}
|
|
25
|
-
async function i(i, a, o) {
|
|
25
|
+
async function i(i, a, o, s) {
|
|
26
26
|
await t("training-tokens");
|
|
27
|
-
let
|
|
27
|
+
let c = await e("training-tokens", a.id, o, s);
|
|
28
28
|
await t("validation-tokens");
|
|
29
|
-
let
|
|
29
|
+
let l = s?.validationSplit && s.validationSplit > 0 ? await e("validation-tokens", a.id, o, s) : void 0, u = [new Uint16Array(c.shardSize)], d = s?.masking ? [new Uint8Array(c.shardSize)] : null, f = {
|
|
30
30
|
offset: 0,
|
|
31
31
|
total: 0
|
|
32
|
-
},
|
|
32
|
+
}, p = s?.validationSplit && s.validationSplit > 0 ? [new Uint16Array(l.shardSize)] : void 0, m = s?.masking && p ? [new Uint8Array(l.shardSize)] : null, h = {
|
|
33
33
|
offset: 0,
|
|
34
34
|
total: 0
|
|
35
|
-
},
|
|
36
|
-
for (;
|
|
37
|
-
let t =
|
|
38
|
-
r(e, n, a, i, g.shardSize,
|
|
39
|
-
},
|
|
40
|
-
return
|
|
41
|
-
trainingTokens:
|
|
42
|
-
validationTokens:
|
|
35
|
+
}, g = 0, _ = s?.cb, v = s?.validationSeed === void 0 ? Math.random : n(s.validationSeed);
|
|
36
|
+
for (; g < i.length;) await i[g++].begin((e) => {
|
|
37
|
+
let t = s?.validationSplit && s.validationSplit > 0 && v() < s.validationSplit, n = t ? p : u, i = t ? h : f, o = t ? m : d, g = t ? l : c;
|
|
38
|
+
r(e, n, a, i, g.shardSize, o || void 0), n.length > 1 && (g.appendShard(n[0], o ? o[0] : void 0), n.shift(), o && o.shift());
|
|
39
|
+
}, _ ? () => _(f.total) : void 0);
|
|
40
|
+
return u.length === 1 && (u[0] = u[0].subarray(0, f.offset), c.appendShard(u[0], d ? d[0].subarray(0, f.offset) : void 0)), p && p.length === 1 && (p[0] = p[0].subarray(0, h.offset), l.appendShard(p[0], m ? m[0].subarray(0, h.offset) : void 0)), await c.finish(), l && await l.finish(), {
|
|
41
|
+
trainingTokens: c,
|
|
42
|
+
validationTokens: p ? l : void 0
|
|
43
43
|
};
|
|
44
44
|
}
|
|
45
45
|
//#endregion
|
|
@@ -9,11 +9,11 @@ export declare function createTrainValidationDatasets(trainingTokens: Uint16Arra
|
|
|
9
9
|
xs: Tensor;
|
|
10
10
|
ys: Tensor;
|
|
11
11
|
}>;
|
|
12
|
-
validationDataset
|
|
12
|
+
validationDataset?: Dataset<{
|
|
13
13
|
xs: Tensor;
|
|
14
14
|
ys: Tensor;
|
|
15
15
|
}>;
|
|
16
16
|
size: number;
|
|
17
|
-
validationState
|
|
17
|
+
validationState?: DatasetState;
|
|
18
18
|
trainState: DatasetState;
|
|
19
19
|
}>;
|
|
@@ -7,7 +7,15 @@ async function n(e, n) {
|
|
|
7
7
|
}), await r.finish(), r;
|
|
8
8
|
}
|
|
9
9
|
async function r(t, r, i, a, o) {
|
|
10
|
-
let s = t instanceof e ? t : await n(t, i), c = r instanceof e ? r : await n(r, i), l = s.getTokenCount(), { dataset: u, state: d } = await a.createTextDataset(s, { batchSize: o })
|
|
10
|
+
let s = t instanceof e ? t : await n(t, i), c = r instanceof e ? r : await n(r, i), l = s.getTokenCount(), { dataset: u, state: d } = await a.createTextDataset(s, { batchSize: o });
|
|
11
|
+
if (c.getTokenCount() === 0) return {
|
|
12
|
+
trainDataset: u,
|
|
13
|
+
validationDataset: void 0,
|
|
14
|
+
size: l,
|
|
15
|
+
validationState: void 0,
|
|
16
|
+
trainState: d
|
|
17
|
+
};
|
|
18
|
+
let { dataset: f, state: p } = await a.createTextDataset(c, {
|
|
11
19
|
batchSize: o,
|
|
12
20
|
shuffleFirst: !0
|
|
13
21
|
});
|