@tuned-tensor/local 0.3.0 → 0.4.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/CHANGELOG.md +27 -0
- package/README.md +35 -0
- package/dist/active-model.d.ts +16 -0
- package/dist/active-model.js +88 -0
- package/dist/active-model.js.map +1 -0
- package/dist/artifacts.d.ts +2 -0
- package/dist/artifacts.js +2 -0
- package/dist/artifacts.js.map +1 -1
- package/dist/contracts.d.ts +176 -0
- package/dist/contracts.js +22 -0
- package/dist/contracts.js.map +1 -1
- package/dist/dataset.d.ts +9 -0
- package/dist/dataset.js +36 -0
- package/dist/dataset.js.map +1 -1
- package/dist/general-regression.d.ts +10 -0
- package/dist/general-regression.js +16 -0
- package/dist/general-regression.js.map +1 -0
- package/dist/index.d.ts +2 -0
- package/dist/index.js +163 -9
- package/dist/index.js.map +1 -1
- package/dist/model-server.d.ts +6 -1
- package/dist/model-server.js +26 -8
- package/dist/model-server.js.map +1 -1
- package/dist/orchestrator.js +151 -1
- package/dist/orchestrator.js.map +1 -1
- package/examples/general-regression.jsonl +4 -0
- package/examples/local-runner.json +5 -0
- package/package.json +1 -1
- package/training/local-runner/src/serve.py +6 -2
package/dist/model-server.js
CHANGED
|
@@ -24,7 +24,7 @@ function isLoopbackHost(host) {
|
|
|
24
24
|
function httpHost(host) {
|
|
25
25
|
return host.includes(":") && !host.startsWith("[") ? `[${host}]` : host;
|
|
26
26
|
}
|
|
27
|
-
|
|
27
|
+
function buildModelServerLaunch(args) {
|
|
28
28
|
const options = args.options ?? {};
|
|
29
29
|
const host = options.host ?? "127.0.0.1";
|
|
30
30
|
const remote = !isLoopbackHost(host);
|
|
@@ -42,7 +42,9 @@ export function buildLocalModelServerLaunch(args) {
|
|
|
42
42
|
const temperature = boundedNumber("temperature", options.temperature ?? 0, 0, 5);
|
|
43
43
|
const topP = boundedNumber("topP", options.topP ?? 1, 0, 1);
|
|
44
44
|
const maxConcurrentRequests = boundedNumber("maxConcurrentRequests", options.maxConcurrentRequests ?? 1, 1, 8, true);
|
|
45
|
-
const artifactPath =
|
|
45
|
+
const artifactPath = args.artifactUri
|
|
46
|
+
? localArtifactPath(args.artifactUri)
|
|
47
|
+
: undefined;
|
|
46
48
|
const recordedBaseModelPath = options.baseModelArtifactUri
|
|
47
49
|
? localArtifactPath(options.baseModelArtifactUri)
|
|
48
50
|
: undefined;
|
|
@@ -56,17 +58,16 @@ export function buildLocalModelServerLaunch(args) {
|
|
|
56
58
|
+ recordedBaseModelPath);
|
|
57
59
|
}
|
|
58
60
|
const localBaseModelPath = recordedBaseModelPath ?? configuredBaseModelPath;
|
|
59
|
-
resolveTrainingModel(args.
|
|
61
|
+
resolveTrainingModel(args.baseModel);
|
|
60
62
|
const entrypoint = buildBundledPythonCommand("serve.py");
|
|
61
|
-
const modelName = args.model.id;
|
|
62
63
|
const env = withBundledPythonEnvironment(withOfflineHuggingFaceCacheEnvironment({
|
|
63
64
|
...minimalMachineLearningEnvironment(process.env),
|
|
64
|
-
TT_MODEL_ARTIFACT: artifactPath,
|
|
65
|
-
TT_BASE_MODEL: localBaseModelPath ?? args.
|
|
65
|
+
...(artifactPath ? { TT_MODEL_ARTIFACT: artifactPath } : {}),
|
|
66
|
+
TT_BASE_MODEL: localBaseModelPath ?? args.baseModel,
|
|
66
67
|
...(options.baseModelRevision && !localBaseModelPath
|
|
67
68
|
? { TT_BASE_MODEL_REVISION: options.baseModelRevision }
|
|
68
69
|
: {}),
|
|
69
|
-
TT_MODEL_NAME: modelName,
|
|
70
|
+
TT_MODEL_NAME: args.modelName,
|
|
70
71
|
TT_MODEL_LOADER: "causal_lm",
|
|
71
72
|
TT_HOST: host,
|
|
72
73
|
TT_PORT: String(port),
|
|
@@ -85,10 +86,27 @@ export function buildLocalModelServerLaunch(args) {
|
|
|
85
86
|
displayCommand: entrypoint.displayCommand,
|
|
86
87
|
env,
|
|
87
88
|
url: `http://${httpHost(host)}:${port}`,
|
|
88
|
-
modelName,
|
|
89
|
+
modelName: args.modelName,
|
|
89
90
|
artifactPath,
|
|
90
91
|
};
|
|
91
92
|
}
|
|
93
|
+
export function buildLocalModelServerLaunch(args) {
|
|
94
|
+
return buildModelServerLaunch({
|
|
95
|
+
modelName: args.model.id,
|
|
96
|
+
baseModel: args.model.base_model,
|
|
97
|
+
artifactUri: args.model.artifact_uri,
|
|
98
|
+
config: args.config,
|
|
99
|
+
options: args.options,
|
|
100
|
+
});
|
|
101
|
+
}
|
|
102
|
+
export function buildLocalBaseModelServerLaunch(args) {
|
|
103
|
+
return buildModelServerLaunch({
|
|
104
|
+
modelName: `base:${args.baseModel}`,
|
|
105
|
+
baseModel: args.baseModel,
|
|
106
|
+
config: args.config,
|
|
107
|
+
options: args.options,
|
|
108
|
+
});
|
|
109
|
+
}
|
|
92
110
|
export async function serveLocalModel(launch) {
|
|
93
111
|
await new Promise((resolveServer, reject) => {
|
|
94
112
|
const child = spawn(launch.command, launch.commandArgs, {
|
package/dist/model-server.js.map
CHANGED
|
@@ -1 +1 @@
|
|
|
1
|
-
{"version":3,"file":"model-server.js","sourceRoot":"","sources":["../src/model-server.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,KAAK,EAAE,MAAM,oBAAoB,CAAC;AAC3C,OAAO,EAAE,OAAO,EAAE,MAAM,WAAW,CAAC;AACpC,OAAO,EAAE,aAAa,EAAE,MAAM,UAAU,CAAC;AAEzC,OAAO,EACL,iCAAiC,EACjC,sCAAsC,GACvC,MAAM,wBAAwB,CAAC;AAChC,OAAO,EAAE,oBAAoB,EAAE,MAAM,qBAAqB,CAAC;AAC3D,OAAO,EACL,yBAAyB,EACzB,4BAA4B,GAC7B,MAAM,qBAAqB,CAAC;AA4B7B,SAAS,iBAAiB,CAAC,GAAW;IACpC,IAAI,GAAG,CAAC,UAAU,CAAC,SAAS,CAAC;QAAE,OAAO,aAAa,CAAC,GAAG,CAAC,CAAC;IACzD,IAAI,sBAAsB,CAAC,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC;QACrC,MAAM,IAAI,KAAK,CAAC,yDAAyD,GAAG,EAAE,CAAC,CAAC;IAClF,CAAC;IACD,OAAO,OAAO,CAAC,GAAG,CAAC,CAAC;AACtB,CAAC;AAED,SAAS,aAAa,CAAC,IAAY,EAAE,KAAa,EAAE,OAAe,EAAE,OAAe,EAAE,OAAO,GAAG,KAAK;IACnG,IAAI,CAAC,MAAM,CAAC,QAAQ,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,OAAO,IAAI,KAAK,GAAG,OAAO,IAAI,CAAC,OAAO,IAAI,CAAC,MAAM,CAAC,SAAS,CAAC,KAAK,CAAC,CAAC,EAAE,CAAC;QAC3G,MAAM,IAAI,KAAK,CAAC,GAAG,IAAI,oBAAoB,OAAO,QAAQ,OAAO,EAAE,CAAC,CAAC;IACvE,CAAC;IACD,OAAO,KAAK,CAAC;AACf,CAAC;AAED,SAAS,cAAc,CAAC,IAAY;IAClC,OAAO,IAAI,KAAK,WAAW,IAAI,IAAI,KAAK,WAAW,IAAI,IAAI,KAAK,KAAK,IAAI,IAAI,KAAK,OAAO,CAAC;AAC5F,CAAC;AAED,SAAS,QAAQ,CAAC,IAAY;IAC5B,OAAO,IAAI,CAAC,QAAQ,CAAC,GAAG,CAAC,IAAI,CAAC,IAAI,CAAC,UAAU,CAAC,GAAG,CAAC,CAAC,CAAC,CAAC,IAAI,IAAI,GAAG,CAAC,CAAC,CAAC,IAAI,CAAC;AAC1E,CAAC;AAED,
|
|
1
|
+
{"version":3,"file":"model-server.js","sourceRoot":"","sources":["../src/model-server.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,KAAK,EAAE,MAAM,oBAAoB,CAAC;AAC3C,OAAO,EAAE,OAAO,EAAE,MAAM,WAAW,CAAC;AACpC,OAAO,EAAE,aAAa,EAAE,MAAM,UAAU,CAAC;AAEzC,OAAO,EACL,iCAAiC,EACjC,sCAAsC,GACvC,MAAM,wBAAwB,CAAC;AAChC,OAAO,EAAE,oBAAoB,EAAE,MAAM,qBAAqB,CAAC;AAC3D,OAAO,EACL,yBAAyB,EACzB,4BAA4B,GAC7B,MAAM,qBAAqB,CAAC;AA4B7B,SAAS,iBAAiB,CAAC,GAAW;IACpC,IAAI,GAAG,CAAC,UAAU,CAAC,SAAS,CAAC;QAAE,OAAO,aAAa,CAAC,GAAG,CAAC,CAAC;IACzD,IAAI,sBAAsB,CAAC,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC;QACrC,MAAM,IAAI,KAAK,CAAC,yDAAyD,GAAG,EAAE,CAAC,CAAC;IAClF,CAAC;IACD,OAAO,OAAO,CAAC,GAAG,CAAC,CAAC;AACtB,CAAC;AAED,SAAS,aAAa,CAAC,IAAY,EAAE,KAAa,EAAE,OAAe,EAAE,OAAe,EAAE,OAAO,GAAG,KAAK;IACnG,IAAI,CAAC,MAAM,CAAC,QAAQ,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,OAAO,IAAI,KAAK,GAAG,OAAO,IAAI,CAAC,OAAO,IAAI,CAAC,MAAM,CAAC,SAAS,CAAC,KAAK,CAAC,CAAC,EAAE,CAAC;QAC3G,MAAM,IAAI,KAAK,CAAC,GAAG,IAAI,oBAAoB,OAAO,QAAQ,OAAO,EAAE,CAAC,CAAC;IACvE,CAAC;IACD,OAAO,KAAK,CAAC;AACf,CAAC;AAED,SAAS,cAAc,CAAC,IAAY;IAClC,OAAO,IAAI,KAAK,WAAW,IAAI,IAAI,KAAK,WAAW,IAAI,IAAI,KAAK,KAAK,IAAI,IAAI,KAAK,OAAO,CAAC;AAC5F,CAAC;AAED,SAAS,QAAQ,CAAC,IAAY;IAC5B,OAAO,IAAI,CAAC,QAAQ,CAAC,GAAG,CAAC,IAAI,CAAC,IAAI,CAAC,UAAU,CAAC,GAAG,CAAC,CAAC,CAAC,CAAC,IAAI,IAAI,GAAG,CAAC,CAAC,CAAC,IAAI,CAAC;AAC1E,CAAC;AAED,SAAS,sBAAsB,CAAC,IAM/B;IACC,MAAM,OAAO,GAAG,IAAI,CAAC,OAAO,IAAI,EAAE,CAAC;IACnC,MAAM,IAAI,GAAG,OAAO,CAAC,IAAI,IAAI,WAAW,CAAC;IACzC,MAAM,MAAM,GAAG,CAAC,cAAc,CAAC,IAAI,CAAC,CAAC;IACrC,IAAI,MAAM,IAAI,CAAC,OAAO,CAAC,WAAW,EAAE,CAAC;QACnC,MAAM,IAAI,KAAK,CAAC,mEAAmE,CAAC,CAAC;IACvF,CAAC;IACD,MAAM,MAAM,GAAG,OAAO,CAAC,SAAS,CAAC,CAAC,CAAC,OAAO,CAAC,GAAG,CAAC,OAAO,CAAC,SAAS,CAAC,EAAE,IAAI,EAAE,CAAC,CAAC,CAAC,SAAS,CAAC;IACtF,IAAI,OAAO,CAAC,SAAS,IAAI,CAAC,MAAM;QAAE,MAAM,IAAI,KAAK,CAAC,GAAG,OAAO,CAAC,SAAS,cAAc,CAAC,CAAC;IACtF,IAAI,MAAM,IAAI,CAAC,MAAM,EAAE,CAAC;QACtB,MAAM,IAAI,KAAK,CAAC,oFAAoF,CAAC,CAAC;IACxG,CAAC;IACD,MAAM,IAAI,GAAG,aAAa,CAAC,MAAM,EAAE,OAAO,CAAC,IAAI,IAAI,IAAI,EAAE,CAAC,EAAE,MAAM,EAAE,IAAI,CAAC,CAAC;IAC1E,MAAM,SAAS,GAAG,aAAa,CAAC,WAAW,EAAE,OAAO,CAAC,SAAS,IAAI,GAAG,EAAE,CAAC,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;IACtF,MAAM,WAAW,GAAG,aAAa,CAAC,aAAa,EAAE,OAAO,CAAC,WAAW,IAAI,CAAC,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC;IACjF,MAAM,IAAI,GAAG,aAAa,CAAC,MAAM,EAAE,OAAO,CAAC,IAAI,IAAI,CAAC,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC;IAC5D,MAAM,qBAAqB,GAAG,aAAa,CACzC,uBAAuB,EACvB,OAAO,CAAC,qBAAqB,IAAI,CAAC,EAClC,CAAC,EACD,CAAC,EACD,IAAI,CACL,CAAC;IACF,MAAM,YAAY,GAAG,IAAI,CAAC,WAAW;QACnC,CAAC,CAAC,iBAAiB,CAAC,IAAI,CAAC,WAAW,CAAC;QACrC,CAAC,CAAC,SAAS,CAAC;IACd,MAAM,qBAAqB,GAAG,OAAO,CAAC,oBAAoB;QACxD,CAAC,CAAC,iBAAiB,CAAC,OAAO,CAAC,oBAAoB,CAAC;QACjD,CAAC,CAAC,SAAS,CAAC;IACd,MAAM,uBAAuB,GAAG,IAAI,CAAC,MAAM,CAAC,KAAK,CAAC,SAAS;QACzD,CAAC,CAAC,OAAO,CAAC,IAAI,CAAC,MAAM,CAAC,KAAK,CAAC,SAAS,CAAC;QACtC,CAAC,CAAC,SAAS,CAAC;IACd,IACE,qBAAqB;WAClB,uBAAuB;WACvB,OAAO,CAAC,qBAAqB,CAAC,KAAK,uBAAuB,EAC7D,CAAC;QACD,MAAM,IAAI,KAAK,CACb,8BAA8B,uBAAuB,6DAA6D;cAChH,qBAAqB,CACxB,CAAC;IACJ,CAAC;IACD,MAAM,kBAAkB,GAAG,qBAAqB,IAAI,uBAAuB,CAAC;IAC5E,oBAAoB,CAAC,IAAI,CAAC,SAAS,CAAC,CAAC;IACrC,MAAM,UAAU,GAAG,yBAAyB,CAAC,UAAU,CAAC,CAAC;IACzD,MAAM,GAAG,GAAG,4BAA4B,CACtC,sCAAsC,CAAC;QACrC,GAAG,iCAAiC,CAAC,OAAO,CAAC,GAAG,CAAC;QACjD,GAAG,CAAC,YAAY,CAAC,CAAC,CAAC,EAAE,iBAAiB,EAAE,YAAY,EAAE,CAAC,CAAC,CAAC,EAAE,CAAC;QAC5D,aAAa,EAAE,kBAAkB,IAAI,IAAI,CAAC,SAAS;QACnD,GAAG,CAAC,OAAO,CAAC,iBAAiB,IAAI,CAAC,kBAAkB;YAClD,CAAC,CAAC,EAAE,sBAAsB,EAAE,OAAO,CAAC,iBAAiB,EAAE;YACvD,CAAC,CAAC,EAAE,CAAC;QACP,aAAa,EAAE,IAAI,CAAC,SAAS;QAC7B,eAAe,EAAE,WAAW;QAC5B,OAAO,EAAE,IAAI;QACb,OAAO,EAAE,MAAM,CAAC,IAAI,CAAC;QACrB,SAAS,EAAE,OAAO,CAAC,MAAM,IAAI,IAAI,CAAC,MAAM,CAAC,UAAU,CAAC,SAAS,CAAC,MAAM;QACpE,aAAa,EAAE,MAAM,CAAC,SAAS,CAAC;QAChC,cAAc,EAAE,MAAM,CAAC,WAAW,CAAC;QACnC,QAAQ,EAAE,MAAM,CAAC,IAAI,CAAC;QACtB,oBAAoB,EAAE,OAAO;QAC7B,0BAA0B,EAAE,MAAM,CAAC,qBAAqB,CAAC;QACzD,GAAG,CAAC,OAAO,CAAC,YAAY,CAAC,CAAC,CAAC,EAAE,gBAAgB,EAAE,OAAO,CAAC,YAAY,EAAE,CAAC,CAAC,CAAC,EAAE,CAAC;QAC3E,GAAG,CAAC,MAAM,CAAC,CAAC,CAAC,EAAE,UAAU,EAAE,MAAM,EAAE,CAAC,CAAC,CAAC,EAAE,CAAC;KAC1C,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,CAAC,UAAU,CAAC,CACjC,CAAC;IACF,OAAO;QACL,OAAO,EAAE,UAAU,CAAC,OAAO;QAC3B,WAAW,EAAE,UAAU,CAAC,WAAW;QACnC,cAAc,EAAE,UAAU,CAAC,cAAc;QACzC,GAAG;QACH,GAAG,EAAE,UAAU,QAAQ,CAAC,IAAI,CAAC,IAAI,IAAI,EAAE;QACvC,SAAS,EAAE,IAAI,CAAC,SAAS;QACzB,YAAY;KACb,CAAC;AACJ,CAAC;AAED,MAAM,UAAU,2BAA2B,CAAC,IAI3C;IACC,OAAO,sBAAsB,CAAC;QAC5B,SAAS,EAAE,IAAI,CAAC,KAAK,CAAC,EAAE;QACxB,SAAS,EAAE,IAAI,CAAC,KAAK,CAAC,UAAU;QAChC,WAAW,EAAE,IAAI,CAAC,KAAK,CAAC,YAAY;QACpC,MAAM,EAAE,IAAI,CAAC,MAAM;QACnB,OAAO,EAAE,IAAI,CAAC,OAAO;KACtB,CAAC,CAAC;AACL,CAAC;AAED,MAAM,UAAU,+BAA+B,CAAC,IAI/C;IACC,OAAO,sBAAsB,CAAC;QAC5B,SAAS,EAAE,QAAQ,IAAI,CAAC,SAAS,EAAE;QACnC,SAAS,EAAE,IAAI,CAAC,SAAS;QACzB,MAAM,EAAE,IAAI,CAAC,MAAM;QACnB,OAAO,EAAE,IAAI,CAAC,OAAO;KACtB,CAAC,CAAC;AACL,CAAC;AAED,MAAM,CAAC,KAAK,UAAU,eAAe,CAAC,MAA8B;IAClE,MAAM,IAAI,OAAO,CAAO,CAAC,aAAa,EAAE,MAAM,EAAE,EAAE;QAChD,MAAM,KAAK,GAAG,KAAK,CAAC,MAAM,CAAC,OAAO,EAAE,MAAM,CAAC,WAAW,EAAE;YACtD,GAAG,EAAE,MAAM,CAAC,GAAG;YACf,KAAK,EAAE,SAAS;YAChB,QAAQ,EAAE,OAAO,CAAC,QAAQ,KAAK,OAAO;SACvC,CAAC,CAAC;QACH,IAAI,QAAQ,GAAG,KAAK,CAAC;QACrB,IAAI,cAA0C,CAAC;QAC/C,MAAM,gBAAgB,GAAG,CAAC,MAAsB,EAAE,EAAE;YAClD,IAAI,KAAK,CAAC,GAAG,IAAI,OAAO,CAAC,QAAQ,KAAK,OAAO,EAAE,CAAC;gBAC9C,IAAI,CAAC;oBACH,OAAO,CAAC,IAAI,CAAC,CAAC,KAAK,CAAC,GAAG,EAAE,MAAM,CAAC,CAAC;oBACjC,OAAO;gBACT,CAAC;gBAAC,MAAM,CAAC;oBACP,0DAA0D;gBAC5D,CAAC;YACH,CAAC;YACD,KAAK,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QACrB,CAAC,CAAC;QACF,MAAM,IAAI,GAAG,CAAC,MAAsB,EAAE,EAAE;YACtC,IAAI,QAAQ;gBAAE,OAAO;YACrB,QAAQ,GAAG,IAAI,CAAC;YAChB,gBAAgB,CAAC,MAAM,CAAC,CAAC;YACzB,cAAc,GAAG,UAAU,CAAC,GAAG,EAAE,CAAC,gBAAgB,CAAC,SAAS,CAAC,EAAE,KAAK,CAAC,CAAC;YACtE,cAAc,CAAC,KAAK,EAAE,CAAC;QACzB,CAAC,CAAC;QACF,MAAM,QAAQ,GAAG,GAAG,EAAE,CAAC,IAAI,CAAC,QAAQ,CAAC,CAAC;QACtC,MAAM,SAAS,GAAG,GAAG,EAAE,CAAC,IAAI,CAAC,SAAS,CAAC,CAAC;QACxC,MAAM,OAAO,GAAG,GAAG,EAAE;YACnB,OAAO,CAAC,GAAG,CAAC,QAAQ,EAAE,QAAQ,CAAC,CAAC;YAChC,OAAO,CAAC,GAAG,CAAC,SAAS,EAAE,SAAS,CAAC,CAAC;YAClC,IAAI,cAAc;gBAAE,YAAY,CAAC,cAAc,CAAC,CAAC;QACnD,CAAC,CAAC;QACF,OAAO,CAAC,IAAI,CAAC,QAAQ,EAAE,QAAQ,CAAC,CAAC;QACjC,OAAO,CAAC,IAAI,CAAC,SAAS,EAAE,SAAS,CAAC,CAAC;QACnC,KAAK,CAAC,IAAI,CAAC,OAAO,EAAE,CAAC,KAAK,EAAE,EAAE;YAC5B,OAAO,EAAE,CAAC;YACV,MAAM,CAAC,KAAK,CAAC,CAAC;QAChB,CAAC,CAAC,CAAC;QACH,KAAK,CAAC,IAAI,CAAC,OAAO,EAAE,CAAC,IAAI,EAAE,MAAM,EAAE,EAAE;YACnC,OAAO,EAAE,CAAC;YACV,IAAI,QAAQ,IAAI,MAAM,KAAK,QAAQ,IAAI,MAAM,KAAK,SAAS,EAAE,CAAC;gBAC5D,aAAa,EAAE,CAAC;YAClB,CAAC;iBAAM,IAAI,IAAI,KAAK,CAAC,EAAE,CAAC;gBACtB,aAAa,EAAE,CAAC;YAClB,CAAC;iBAAM,CAAC;gBACN,MAAM,CAAC,IAAI,KAAK,CAAC,uCAAuC,IAAI,IAAI,SAAS,GAAG,CAAC,CAAC,CAAC;YACjF,CAAC;QACH,CAAC,CAAC,CAAC;IACL,CAAC,CAAC,CAAC;AACL,CAAC"}
|
package/dist/orchestrator.js
CHANGED
|
@@ -7,7 +7,7 @@ import { basename, dirname, join, relative, resolve } from "node:path";
|
|
|
7
7
|
import { fileURLToPath } from "node:url";
|
|
8
8
|
import { defaultArtifactPrefix, fileUri, assertArtifactManifest, ARTIFACT_WORKFLOW_LOCK_FILE, claimRunArtifactDirectory, prepareRunDirectories, readJson, resolveRunArtifacts, writeArtifactManifest, writeFileAtomic, writeJsonAtomic, } from "./artifacts.js";
|
|
9
9
|
import { evalReportSchema, fineTuneRunRequestSchema, localRunnerConfigSchema, runReportSchema, trainingReportSchema, } from "./contracts.js";
|
|
10
|
-
import { buildSystemMessage, compileSpecToJsonl, examplesFromChatJsonl, examplesFromSpec, normalizeChatJsonlForRelocation, } from "./dataset.js";
|
|
10
|
+
import { buildSystemMessage, compileSpecToJsonl, evaluationSuiteFromChatJsonl, examplesFromChatJsonl, examplesFromSpec, normalizeChatJsonlForRelocation, } from "./dataset.js";
|
|
11
11
|
import { INFERENCE_PROTOCOL_VERSION, compareEvalReports, deriveSampleSeed, evaluateExamples, splitSpecExamples, } from "./evaluation.js";
|
|
12
12
|
import { launchProcessTraining } from "./process-training.js";
|
|
13
13
|
import { assertUsableModelArtifact, localModelArtifactPath } from "./model-registry.js";
|
|
@@ -15,6 +15,7 @@ import { ProcessCancelledError } from "./process-runner.js";
|
|
|
15
15
|
import { createLocalStore } from "./store.js";
|
|
16
16
|
import { withHuggingFaceCacheEnvironment } from "./huggingface-cache.js";
|
|
17
17
|
import { verifyLocalBaseModel } from "./prefetch.js";
|
|
18
|
+
import { evaluateGeneralRegressionGate } from "./general-regression.js";
|
|
18
19
|
export async function loadJsonFile(path) {
|
|
19
20
|
return JSON.parse(await readFile(path, "utf8"));
|
|
20
21
|
}
|
|
@@ -44,6 +45,15 @@ export async function loadLocalRunnerConfig(path) {
|
|
|
44
45
|
baseModel: configPathValue(config.paths.baseModel),
|
|
45
46
|
modelCache: configPathValue(config.paths.modelCache),
|
|
46
47
|
},
|
|
48
|
+
evaluation: {
|
|
49
|
+
...config.evaluation,
|
|
50
|
+
generalRegression: config.evaluation.generalRegression
|
|
51
|
+
? {
|
|
52
|
+
...config.evaluation.generalRegression,
|
|
53
|
+
dataset: configPathValue(config.evaluation.generalRegression.dataset),
|
|
54
|
+
}
|
|
55
|
+
: undefined,
|
|
56
|
+
},
|
|
47
57
|
};
|
|
48
58
|
}
|
|
49
59
|
function elapsed(started) {
|
|
@@ -198,6 +208,9 @@ export async function validateLocalFineTuneInput(input) {
|
|
|
198
208
|
const config = localRunnerConfigSchema.parse(input.config);
|
|
199
209
|
const request = fineTuneRunRequestSchema.parse(addDryRunPlaceholders(input.request, config.dryRun));
|
|
200
210
|
await validateDatasetInputs(request, config);
|
|
211
|
+
if (config.evaluation.generalRegression) {
|
|
212
|
+
await evaluationSuiteFromChatJsonl(config.evaluation.generalRegression.dataset, config.evaluation.generalRegression.systemPrompt);
|
|
213
|
+
}
|
|
201
214
|
return { request, config };
|
|
202
215
|
}
|
|
203
216
|
function artifactPrefix(request) {
|
|
@@ -378,6 +391,8 @@ async function cleanupStageArtifacts(artifacts, stage) {
|
|
|
378
391
|
await Promise.all([
|
|
379
392
|
removePrefixedArtifacts(artifacts.baselineEvalJson),
|
|
380
393
|
removePrefixedArtifacts(artifacts.candidateEvalJson),
|
|
394
|
+
removePrefixedArtifacts(artifacts.generalBaselineEvalJson),
|
|
395
|
+
removePrefixedArtifacts(artifacts.generalCandidateEvalJson),
|
|
381
396
|
rm(artifacts.trainingDir, { recursive: true, force: true }),
|
|
382
397
|
removePrefixedArtifacts(artifacts.trainingReportJson),
|
|
383
398
|
rm(resolve(artifacts.runDir, "model.tar.gz"), { force: true }),
|
|
@@ -389,6 +404,7 @@ async function cleanupStageArtifacts(artifacts, stage) {
|
|
|
389
404
|
if (stage === "baseline") {
|
|
390
405
|
await Promise.all([
|
|
391
406
|
removePrefixedArtifacts(artifacts.baselineEvalJson),
|
|
407
|
+
removePrefixedArtifacts(artifacts.generalBaselineEvalJson),
|
|
392
408
|
removeReport(),
|
|
393
409
|
]);
|
|
394
410
|
return;
|
|
@@ -399,6 +415,7 @@ async function cleanupStageArtifacts(artifacts, stage) {
|
|
|
399
415
|
removePrefixedArtifacts(artifacts.trainingReportJson),
|
|
400
416
|
rm(resolve(artifacts.runDir, "model.tar.gz"), { force: true }),
|
|
401
417
|
removePrefixedArtifacts(artifacts.candidateEvalJson),
|
|
418
|
+
removePrefixedArtifacts(artifacts.generalCandidateEvalJson),
|
|
402
419
|
removeReport(),
|
|
403
420
|
]);
|
|
404
421
|
await prepareRunDirectories(artifacts);
|
|
@@ -407,6 +424,7 @@ async function cleanupStageArtifacts(artifacts, stage) {
|
|
|
407
424
|
if (stage === "candidate") {
|
|
408
425
|
await Promise.all([
|
|
409
426
|
removePrefixedArtifacts(artifacts.candidateEvalJson),
|
|
427
|
+
removePrefixedArtifacts(artifacts.generalCandidateEvalJson),
|
|
410
428
|
removeReport(),
|
|
411
429
|
]);
|
|
412
430
|
return;
|
|
@@ -581,6 +599,12 @@ async function stageFingerprint(args) {
|
|
|
581
599
|
scoring: args.config.evaluation.scoring,
|
|
582
600
|
timeout_ms: args.config.evaluation.timeoutMs,
|
|
583
601
|
baseline_cache: args.config.evaluation.baselineCache,
|
|
602
|
+
general_regression: args.prepared.generalRegression
|
|
603
|
+
? {
|
|
604
|
+
dataset_sha256: args.prepared.generalRegression.datasetSha256,
|
|
605
|
+
system_prompt: args.prepared.generalRegression.system,
|
|
606
|
+
}
|
|
607
|
+
: null,
|
|
584
608
|
model_cache: args.config.paths.modelCache,
|
|
585
609
|
};
|
|
586
610
|
return hashJson({
|
|
@@ -695,6 +719,14 @@ async function computePreparedRun(args) {
|
|
|
695
719
|
const maxEvalExamples = config.evaluation.maxExamples ?? request.hyperparameters.max_eval_examples;
|
|
696
720
|
const evalExamplesUsed = Math.min(maxEvalExamples ?? examples.length, examples.length);
|
|
697
721
|
const fingerprints = await datasetFingerprints(request);
|
|
722
|
+
const generalRegressionConfig = config.evaluation.generalRegression;
|
|
723
|
+
const generalRegression = generalRegressionConfig
|
|
724
|
+
? {
|
|
725
|
+
datasetPath: generalRegressionConfig.dataset,
|
|
726
|
+
datasetSha256: await hashFile(generalRegressionConfig.dataset),
|
|
727
|
+
...await evaluationSuiteFromChatJsonl(generalRegressionConfig.dataset, generalRegressionConfig.systemPrompt),
|
|
728
|
+
}
|
|
729
|
+
: undefined;
|
|
698
730
|
const requestFingerprint = hashJson(request);
|
|
699
731
|
const runtimeFingerprintValue = await runtimeFingerprint();
|
|
700
732
|
const baseModelRevision = await resolveBaseModelRevision(request, config);
|
|
@@ -746,8 +778,22 @@ async function computePreparedRun(args) {
|
|
|
746
778
|
system,
|
|
747
779
|
baseModelForEvaluation,
|
|
748
780
|
maxEvalExamples,
|
|
781
|
+
generalRegression,
|
|
749
782
|
};
|
|
750
783
|
}
|
|
784
|
+
function generalRegressionEvaluationConfig(config) {
|
|
785
|
+
const { maxExamples: _maxExamples, sampleSeed: _sampleSeed, ...evaluation } = config.evaluation;
|
|
786
|
+
return { ...config, evaluation };
|
|
787
|
+
}
|
|
788
|
+
function generalRegressionRequiredPaths(prepared, kind) {
|
|
789
|
+
if (!prepared.generalRegression)
|
|
790
|
+
return [];
|
|
791
|
+
return [
|
|
792
|
+
kind === "baseline"
|
|
793
|
+
? prepared.artifacts.generalBaselineEvalJson
|
|
794
|
+
: prepared.artifacts.generalCandidateEvalJson,
|
|
795
|
+
];
|
|
796
|
+
}
|
|
751
797
|
async function prepareStage(args) {
|
|
752
798
|
const preparedExists = await pathExists(args.artifacts.stageMetadataJson)
|
|
753
799
|
&& await pathExists(args.artifacts.trainingJsonl);
|
|
@@ -795,6 +841,7 @@ async function runBaselineStage(args) {
|
|
|
795
841
|
stage: "baseline",
|
|
796
842
|
prepared: args.prepared,
|
|
797
843
|
config: args.config,
|
|
844
|
+
additionalPaths: generalRegressionRequiredPaths(args.prepared, "baseline"),
|
|
798
845
|
})) {
|
|
799
846
|
await throwIfCancelled(args.store, args.prepared.request);
|
|
800
847
|
await updateRun({
|
|
@@ -845,6 +892,36 @@ async function runBaselineStage(args) {
|
|
|
845
892
|
sampleSeed: args.prepared.metadata.eval_sample_seed,
|
|
846
893
|
shouldCancel: () => args.store.isCancellationRequested(args.prepared.request.run_id),
|
|
847
894
|
});
|
|
895
|
+
if (args.prepared.generalRegression) {
|
|
896
|
+
await updateRun({
|
|
897
|
+
store: args.store,
|
|
898
|
+
reporter: args.reporter,
|
|
899
|
+
request: args.prepared.request,
|
|
900
|
+
status: "evaluating_baseline",
|
|
901
|
+
stage: "evaluating_baseline",
|
|
902
|
+
message: "Running baseline general regression evaluation.",
|
|
903
|
+
details: {
|
|
904
|
+
examples: args.prepared.generalRegression.examples.length,
|
|
905
|
+
dataset: args.prepared.generalRegression.datasetPath,
|
|
906
|
+
},
|
|
907
|
+
});
|
|
908
|
+
await evaluateExamples({
|
|
909
|
+
kind: "baseline",
|
|
910
|
+
modelId: args.prepared.baseModelForEvaluation,
|
|
911
|
+
baseModelId: args.prepared.baseModelForEvaluation,
|
|
912
|
+
baseModelRevision: args.config.paths.baseModel
|
|
913
|
+
? undefined
|
|
914
|
+
: args.prepared.metadata.base_model_revision ?? undefined,
|
|
915
|
+
sourceFingerprint: args.prepared.metadata.base_model_fingerprint ?? undefined,
|
|
916
|
+
examples: args.prepared.generalRegression.examples,
|
|
917
|
+
system: args.prepared.generalRegression.system,
|
|
918
|
+
config: generalRegressionEvaluationConfig(args.config),
|
|
919
|
+
outputPath: args.prepared.artifacts.generalBaselineEvalJson,
|
|
920
|
+
reporter: args.runReporter,
|
|
921
|
+
evalSplit: "general_regression",
|
|
922
|
+
shouldCancel: () => args.store.isCancellationRequested(args.prepared.request.run_id),
|
|
923
|
+
});
|
|
924
|
+
}
|
|
848
925
|
await throwIfCancelled(args.store, args.prepared.request);
|
|
849
926
|
await writeStageFingerprint({ stage: "baseline", prepared: args.prepared, config: args.config });
|
|
850
927
|
return report;
|
|
@@ -943,6 +1020,7 @@ async function runCandidateStage(args) {
|
|
|
943
1020
|
additionalPaths: [
|
|
944
1021
|
args.prepared.artifacts.trainingReportJson,
|
|
945
1022
|
stageFingerprintPath(args.prepared, "train"),
|
|
1023
|
+
...generalRegressionRequiredPaths(args.prepared, "candidate"),
|
|
946
1024
|
],
|
|
947
1025
|
})) {
|
|
948
1026
|
await throwIfCancelled(args.store, args.prepared.request);
|
|
@@ -1015,6 +1093,36 @@ async function runCandidateStage(args) {
|
|
|
1015
1093
|
sampleSeed: args.prepared.metadata.eval_sample_seed,
|
|
1016
1094
|
shouldCancel: () => args.store.isCancellationRequested(args.prepared.request.run_id),
|
|
1017
1095
|
});
|
|
1096
|
+
if (args.prepared.generalRegression) {
|
|
1097
|
+
await updateRun({
|
|
1098
|
+
store: args.store,
|
|
1099
|
+
reporter: args.reporter,
|
|
1100
|
+
request: args.prepared.request,
|
|
1101
|
+
status: "evaluating_candidate",
|
|
1102
|
+
stage: "evaluating_candidate",
|
|
1103
|
+
message: "Running candidate general regression evaluation.",
|
|
1104
|
+
details: {
|
|
1105
|
+
examples: args.prepared.generalRegression.examples.length,
|
|
1106
|
+
dataset: args.prepared.generalRegression.datasetPath,
|
|
1107
|
+
},
|
|
1108
|
+
});
|
|
1109
|
+
await evaluateExamples({
|
|
1110
|
+
kind: "candidate",
|
|
1111
|
+
modelId: modelArtifact,
|
|
1112
|
+
baseModelId: args.prepared.baseModelForEvaluation,
|
|
1113
|
+
baseModelRevision: args.config.paths.baseModel
|
|
1114
|
+
? undefined
|
|
1115
|
+
: args.prepared.metadata.base_model_revision ?? undefined,
|
|
1116
|
+
adapterPath: modelArtifact,
|
|
1117
|
+
examples: args.prepared.generalRegression.examples,
|
|
1118
|
+
system: args.prepared.generalRegression.system,
|
|
1119
|
+
config: generalRegressionEvaluationConfig(args.config),
|
|
1120
|
+
outputPath: args.prepared.artifacts.generalCandidateEvalJson,
|
|
1121
|
+
reporter: args.runReporter,
|
|
1122
|
+
evalSplit: "general_regression",
|
|
1123
|
+
shouldCancel: () => args.store.isCancellationRequested(args.prepared.request.run_id),
|
|
1124
|
+
});
|
|
1125
|
+
}
|
|
1018
1126
|
await throwIfCancelled(args.store, args.prepared.request);
|
|
1019
1127
|
await writeStageFingerprint({
|
|
1020
1128
|
stage: "candidate",
|
|
@@ -1034,6 +1142,11 @@ async function runReportStage(args) {
|
|
|
1034
1142
|
if (!await pathExists(args.prepared.artifacts.trainingReportJson)) {
|
|
1035
1143
|
throw new Error("Run reporting requires training-report.json.");
|
|
1036
1144
|
}
|
|
1145
|
+
if (args.prepared.generalRegression
|
|
1146
|
+
&& (!await pathExists(args.prepared.artifacts.generalBaselineEvalJson)
|
|
1147
|
+
|| !await pathExists(args.prepared.artifacts.generalCandidateEvalJson))) {
|
|
1148
|
+
throw new Error("Run reporting requires both general regression evaluations.");
|
|
1149
|
+
}
|
|
1037
1150
|
const currentTraining = await canReuseStageArtifact({
|
|
1038
1151
|
stage: "train",
|
|
1039
1152
|
prepared: args.prepared,
|
|
@@ -1044,12 +1157,18 @@ async function runReportStage(args) {
|
|
|
1044
1157
|
stage: "baseline",
|
|
1045
1158
|
prepared: args.prepared,
|
|
1046
1159
|
config: args.config,
|
|
1160
|
+
additionalPaths: generalRegressionRequiredPaths(args.prepared, "baseline"),
|
|
1047
1161
|
});
|
|
1048
1162
|
const currentCandidate = await canReuseStageArtifact({
|
|
1049
1163
|
stage: "candidate",
|
|
1050
1164
|
prepared: args.prepared,
|
|
1051
1165
|
config: args.config,
|
|
1052
1166
|
verifyModel: true,
|
|
1167
|
+
additionalPaths: [
|
|
1168
|
+
args.prepared.artifacts.trainingReportJson,
|
|
1169
|
+
stageFingerprintPath(args.prepared, "train"),
|
|
1170
|
+
...generalRegressionRequiredPaths(args.prepared, "candidate"),
|
|
1171
|
+
],
|
|
1053
1172
|
});
|
|
1054
1173
|
if (!currentTraining || !currentBaseline || !currentCandidate) {
|
|
1055
1174
|
throw new Error("report stage inputs are stale for the current request/config. Re-run baseline, train, and candidate as needed.");
|
|
@@ -1061,6 +1180,8 @@ async function runReportStage(args) {
|
|
|
1061
1180
|
stageFingerprintPath(args.prepared, "candidate"),
|
|
1062
1181
|
args.prepared.artifacts.trainingReportJson,
|
|
1063
1182
|
stageFingerprintPath(args.prepared, "train"),
|
|
1183
|
+
...generalRegressionRequiredPaths(args.prepared, "baseline"),
|
|
1184
|
+
...generalRegressionRequiredPaths(args.prepared, "candidate"),
|
|
1064
1185
|
], { verifyModel: true });
|
|
1065
1186
|
await throwIfCancelled(args.store, args.prepared.request);
|
|
1066
1187
|
await args.store.invalidateRunOutputs(args.prepared.request.run_id, { report: true });
|
|
@@ -1078,6 +1199,27 @@ async function runReportStage(args) {
|
|
|
1078
1199
|
const candidate = evalReportSchema.parse(await readJson(args.prepared.artifacts.candidateEvalJson));
|
|
1079
1200
|
const training = trainingReportSchema.parse(await readJson(args.prepared.artifacts.trainingReportJson));
|
|
1080
1201
|
const comparison = compareEvalReports(baseline, candidate);
|
|
1202
|
+
let generalRegression;
|
|
1203
|
+
if (args.prepared.generalRegression) {
|
|
1204
|
+
const generalBaseline = evalReportSchema.parse(await readJson(args.prepared.artifacts.generalBaselineEvalJson));
|
|
1205
|
+
const generalCandidate = evalReportSchema.parse(await readJson(args.prepared.artifacts.generalCandidateEvalJson));
|
|
1206
|
+
const generalComparison = compareEvalReports(generalBaseline, generalCandidate);
|
|
1207
|
+
const policy = args.config.evaluation.generalRegression;
|
|
1208
|
+
const gate = evaluateGeneralRegressionGate(generalComparison, policy);
|
|
1209
|
+
generalRegression = {
|
|
1210
|
+
dataset_uri: fileUri(args.prepared.generalRegression.datasetPath),
|
|
1211
|
+
dataset_sha256: args.prepared.generalRegression.datasetSha256,
|
|
1212
|
+
baseline: generalBaseline,
|
|
1213
|
+
candidate: generalCandidate,
|
|
1214
|
+
comparison: generalComparison,
|
|
1215
|
+
policy: {
|
|
1216
|
+
max_score_drop: policy.maxScoreDrop,
|
|
1217
|
+
max_pass_rate_drop: policy.maxPassRateDrop,
|
|
1218
|
+
},
|
|
1219
|
+
passed: gate.passed,
|
|
1220
|
+
failures: gate.failures,
|
|
1221
|
+
};
|
|
1222
|
+
}
|
|
1081
1223
|
const completedAt = new Date().toISOString();
|
|
1082
1224
|
const duration = elapsed(args.startedPerf);
|
|
1083
1225
|
const report = runReportSchema.parse({
|
|
@@ -1091,11 +1233,18 @@ async function runReportStage(args) {
|
|
|
1091
1233
|
baseline,
|
|
1092
1234
|
candidate,
|
|
1093
1235
|
comparison,
|
|
1236
|
+
general_regression: generalRegression,
|
|
1094
1237
|
training,
|
|
1095
1238
|
artifact_uris: {
|
|
1096
1239
|
dataset: fileUri(args.prepared.artifacts.trainingJsonl),
|
|
1097
1240
|
baseline_eval: fileUri(args.prepared.artifacts.baselineEvalJson),
|
|
1098
1241
|
candidate_eval: fileUri(args.prepared.artifacts.candidateEvalJson),
|
|
1242
|
+
general_baseline_eval: generalRegression
|
|
1243
|
+
? fileUri(args.prepared.artifacts.generalBaselineEvalJson)
|
|
1244
|
+
: undefined,
|
|
1245
|
+
general_candidate_eval: generalRegression
|
|
1246
|
+
? fileUri(args.prepared.artifacts.generalCandidateEvalJson)
|
|
1247
|
+
: undefined,
|
|
1099
1248
|
report: fileUri(args.prepared.artifacts.runReportJson),
|
|
1100
1249
|
},
|
|
1101
1250
|
run_metadata: {
|
|
@@ -1132,6 +1281,7 @@ async function runReportStage(args) {
|
|
|
1132
1281
|
report_path: args.prepared.artifacts.runReportJson,
|
|
1133
1282
|
...(!isDryTraining(training) ? { model_id: `local-${args.prepared.request.run_id}` } : {}),
|
|
1134
1283
|
avg_score_delta: comparison.avg_score_delta,
|
|
1284
|
+
general_regression_passed: generalRegression?.passed,
|
|
1135
1285
|
elapsed_seconds: duration.seconds,
|
|
1136
1286
|
},
|
|
1137
1287
|
});
|