@mlx-node/trl 0.0.6 → 0.0.7
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.
|
@@ -69,6 +69,10 @@ export interface GRPOTrainerConfig<T = unknown> {
|
|
|
69
69
|
topP?: number;
|
|
70
70
|
topK?: number;
|
|
71
71
|
repetitionPenalty?: number;
|
|
72
|
+
/** Presence penalty (0.0 = disabled). Subtracts a flat penalty from logits of any token in context. */
|
|
73
|
+
presencePenalty?: number;
|
|
74
|
+
/** Frequency penalty (0.0 = disabled). Subtracts penalty * count for each token in context. */
|
|
75
|
+
frequencyPenalty?: number;
|
|
72
76
|
/**
|
|
73
77
|
* Tool definitions for function calling.
|
|
74
78
|
* When provided, tools are included in the chat template so the model
|
|
@@ -1 +1 @@
|
|
|
1
|
-
{"version":3,"file":"grpo-trainer.d.ts","sourceRoot":"","sources":["../../src/trainers/grpo-trainer.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;GAwCG;AAiBH,OAAO,EACL,kBAAkB,EAClB,oBAAoB,EAIpB,WAAW,EAGX,KAAK,kBAAkB,EACvB,KAAK,mBAAmB,EACxB,KAAK,mBAAmB,IAAI,yBAAyB,EAGrD,KAAK,cAAc,EACpB,MAAM,gBAAgB,CAAC;AACxB,OAAO,EAAa,KAAK,cAAc,EAAE,MAAM,cAAc,CAAC;AAE9D,OAAO,KAAK,EAAE,WAAW,EAAE,cAAc,EAAE,cAAc,EAAE,MAAM,aAAa,CAAC;AAC/E,OAAO,EAAwB,KAAK,cAAc,EAAE,MAAM,sBAAsB,CAAC;AAGjF,OAAO,EAAE,kBAAkB,EAAE,oBAAoB,EAAE,WAAW,EAAE,MAAM,gBAAgB,CAAC;AACvF,YAAY,EACV,gBAAgB,EAChB,iBAAiB,EACjB,kBAAkB,EAClB,mBAAmB,EACnB,eAAe,EACf,0BAA0B,EAC1B,YAAY,EACZ,iBAAiB,GAClB,MAAM,gBAAgB,CAAC;AAExB;;GAEG;AACH,MAAM,WAAW,iBAAiB,CAAC,CAAC,GAAG,OAAO;IAE5C,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,SAAS,CAAC,EAAE,MAAM,CAAC;IAGnB,YAAY,CAAC,EAAE,MAAM,CAAC;IACtB,yBAAyB,CAAC,EAAE,MAAM,CAAC;IACnC,gBAAgB,CAAC,EAAE,MAAM,CAAC;IAC1B,WAAW,CAAC,EAAE,MAAM,CAAC;IAGrB,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,SAAS,CAAC,EAAE,MAAM,CAAC;IAGnB,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,QAAQ,CAAC,EAAE,MAAM,GAAG,MAAM,GAAG,SAAS,GAAG,MAAM,CAAC;IAChD,sBAAsB,CAAC,EAAE,OAAO,CAAC;IAGjC;4DACwD;IACxD,mBAAmB,CAAC,EAAE,MAAM,CAAC;IAC7B,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,iBAAiB,CAAC,EAAE,MAAM,CAAC;
|
|
1
|
+
{"version":3,"file":"grpo-trainer.d.ts","sourceRoot":"","sources":["../../src/trainers/grpo-trainer.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;GAwCG;AAiBH,OAAO,EACL,kBAAkB,EAClB,oBAAoB,EAIpB,WAAW,EAGX,KAAK,kBAAkB,EACvB,KAAK,mBAAmB,EACxB,KAAK,mBAAmB,IAAI,yBAAyB,EAGrD,KAAK,cAAc,EACpB,MAAM,gBAAgB,CAAC;AACxB,OAAO,EAAa,KAAK,cAAc,EAAE,MAAM,cAAc,CAAC;AAE9D,OAAO,KAAK,EAAE,WAAW,EAAE,cAAc,EAAE,cAAc,EAAE,MAAM,aAAa,CAAC;AAC/E,OAAO,EAAwB,KAAK,cAAc,EAAE,MAAM,sBAAsB,CAAC;AAGjF,OAAO,EAAE,kBAAkB,EAAE,oBAAoB,EAAE,WAAW,EAAE,MAAM,gBAAgB,CAAC;AACvF,YAAY,EACV,gBAAgB,EAChB,iBAAiB,EACjB,kBAAkB,EAClB,mBAAmB,EACnB,eAAe,EACf,0BAA0B,EAC1B,YAAY,EACZ,iBAAiB,GAClB,MAAM,gBAAgB,CAAC;AAExB;;GAEG;AACH,MAAM,WAAW,iBAAiB,CAAC,CAAC,GAAG,OAAO;IAE5C,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,SAAS,CAAC,EAAE,MAAM,CAAC;IAGnB,YAAY,CAAC,EAAE,MAAM,CAAC;IACtB,yBAAyB,CAAC,EAAE,MAAM,CAAC;IACnC,gBAAgB,CAAC,EAAE,MAAM,CAAC;IAC1B,WAAW,CAAC,EAAE,MAAM,CAAC;IAGrB,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,SAAS,CAAC,EAAE,MAAM,CAAC;IAGnB,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,QAAQ,CAAC,EAAE,MAAM,GAAG,MAAM,GAAG,SAAS,GAAG,MAAM,CAAC;IAChD,sBAAsB,CAAC,EAAE,OAAO,CAAC;IAGjC;4DACwD;IACxD,mBAAmB,CAAC,EAAE,MAAM,CAAC;IAC7B,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,iBAAiB,CAAC,EAAE,MAAM,CAAC;IAC3B,uGAAuG;IACvG,eAAe,CAAC,EAAE,MAAM,CAAC;IACzB,+FAA+F;IAC/F,gBAAgB,CAAC,EAAE,MAAM,CAAC;IAG1B;;;;;;;;;;;;;;;;OAgBG;IACH,KAAK,CAAC,EAAE,cAAc,EAAE,CAAC;IAEzB;;6EAEyE;IACzE,cAAc,CAAC,EAAE,OAAO,CAAC;IAGzB,UAAU,CAAC,EAAE,UAAU,GAAG,SAAS,GAAG,OAAO,CAAC;IAC9C,cAAc,CAAC,EAAE,cAAc,CAAC,CAAC,CAAC,CAAC;IACnC,eAAe,CAAC,EAAE,MAAM,CAAC;IAGzB,iBAAiB,CAAC,EAAE,MAAM,CAAC;IAG3B,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,YAAY,CAAC,EAAE,MAAM,CAAC;IACtB,YAAY,CAAC,EAAE,MAAM,CAAC;IACtB,SAAS,CAAC,EAAE,MAAM,CAAC;IACnB,UAAU,CAAC,EAAE,OAAO,CAAC;IACrB,QAAQ,CAAC,EAAE,OAAO,CAAC;IACnB,OAAO,CAAC,EAAE,MAAM,CAAC;IACjB,kFAAkF;IAClF,cAAc,CAAC,EAAE,MAAM,CAAC;IAGxB,MAAM,CAAC,EAAE,MAAM,CAAC;IAGhB,4EAA4E;IAE5E,oBAAoB,CAAC,EAAE,QAAQ,GAAG,MAAM,CAAC;IAGzC,6FAA6F;IAC7F,OAAO,CAAC,EAAE,OAAO,CAAC;IAGlB;;;;;;;;;OASG;IACH,aAAa,CAAC,EAAE,MAAM,CAAC;IAGvB;;;;;;;OAOG;IACH,eAAe,CAAC,EAAE,MAAM,CAAC;IAEzB;;;;;;;OAOG;IACH,gBAAgB,CAAC,EAAE,MAAM,CAAC;IAE1B;;;;;;OAMG;IACH,0BAA0B,CAAC,EAAE,OAAO,CAAC;IAErC;;;;;;OAMG;IACH,qBAAqB,CAAC,EAAE,OAAO,CAAC;IAEhC,0DAA0D;IAC1D,aAAa,CAAC,EAAE,KAAK,GAAG,OAAO,CAAC;IAChC,iCAAiC;IACjC,UAAU,CAAC,EAAE,MAAM,CAAC;IACpB,mCAAmC;IACnC,UAAU,CAAC,EAAE,MAAM,CAAC;IACpB,oCAAoC;IACpC,QAAQ,CAAC,EAAE,MAAM,CAAC;IAElB;;;;;;;;;;;;;OAaG;IACH,cAAc,CAAC,EAAE,MAAM,CAAC;IAGxB,sFAAsF;IACtF,WAAW,CAAC,EAAE;QACZ,+CAA+C;QAC/C,OAAO,EAAE,OAAO,CAAC;QACjB,8DAA8D;QAC9D,SAAS,CAAC,EAAE,MAAM,CAAC;QACnB,iDAAiD;QACjD,SAAS,CAAC,EAAE,MAAM,CAAC;QACnB,sDAAsD;QACtD,SAAS,CAAC,EAAE,MAAM,CAAC;QACnB,6EAA6E;QAC7E,YAAY,CAAC,EAAE,MAAM,CAAC;KACvB,CAAC;CACH;AAED;;;;;GAKG;AACH,MAAM,WAAW,eAAe;IAC9B,8CAA8C;IAC9C,IAAI,EAAE,MAAM,CAAC;IACb,yDAAyD;IACzD,WAAW,EAAE,MAAM,CAAC;IACpB,uDAAuD;IACvD,WAAW,CAAC,EAAE,MAAM,CAAC;IACrB,4DAA4D;IAC5D,qBAAqB,CAAC,EAAE,MAAM,EAAE,CAAC;CAClC;AAED;;GAEG;AACH,MAAM,WAAW,aAAa;IAC5B,IAAI,EAAE,MAAM,CAAC;IACb,KAAK,EAAE,MAAM,CAAC;IACd,SAAS,EAAE,MAAM,CAAC;IAClB,gDAAgD;IAChD,OAAO,CAAC,EAAE,eAAe,CAAC;IAC1B,kEAAkE;IAClE,iBAAiB,CAAC,EAAE,OAAO,CAAC;CAC7B;AAED;;GAEG;AACH,MAAM,WAAW,mBAAmB;IAClC,iCAAiC;IACjC,eAAe,EAAE,MAAM,EAAE,CAAC;IAC1B,uEAAuE;IACvE,YAAY,EAAE,yBAAyB,CAAC;IACxC,0DAA0D;IAC1D,WAAW,EAAE,MAAM,EAAE,CAAC;IACtB,6EAA6E;IAC7E,aAAa,EAAE,MAAM,EAAE,CAAC;CACzB;AAED;;GAEG;AACH,eAAO,MAAM,mBAAmB,EAAE,iBAuBjC,CAAC;AAEF;;GAEG;AACH,MAAM,WAAW,gBAAgB;IAC/B,0BAA0B;IAC1B,IAAI,EAAE,MAAM,CAAC;IACb,sBAAsB;IACtB,IAAI,EAAE,MAAM,CAAC;IACb,qCAAqC;IACrC,UAAU,EAAE,MAAM,CAAC;IACnB,oCAAoC;IACpC,SAAS,EAAE,MAAM,CAAC;IAClB,2BAA2B;IAC3B,aAAa,EAAE,MAAM,CAAC;IACtB,kEAAkE;IAClE,YAAY,EAAE,MAAM,CAAC;IACrB,uCAAuC;IACvC,WAAW,EAAE,MAAM,CAAC;IACpB,qCAAqC;IACrC,gBAAgB,CAAC,EAAE,OAAO,CAAC;IAC3B,+BAA+B;IAC/B,gBAAgB,CAAC,EAAE,MAAM,CAAC;IAC1B,6BAA6B;IAC7B,cAAc,CAAC,EAAE,MAAM,CAAC;IACxB,yCAAyC;IACzC,KAAK,CAAC,EAAE,MAAM,CAAC;CAChB;AAED;;GAEG;AACH,MAAM,MAAM,eAAe,GAAG,gBAAgB,CAAC;AAE/C;;;;;;;;;GASG;AACH,wBAAgB,kBAAkB,CAAC,OAAO,EAAE,cAAc,EAAE,EAAE,UAAU,SAAK,GAAG,MAAM,CAKrF;AAKD;;GAEG;AACH,qBAAa,kBAAmB,SAAQ,KAAK;aAGzB,SAAS,EAAE,MAAM;gBADjC,OAAO,EAAE,MAAM,EACC,SAAS,EAAE,MAAM;CAKpC;AAiCD;;;;;GAKG;AACH,qBAAa,WAAW,CAAC,CAAC,GAAG,OAAO;IAClC,OAAO,CAAC,MAAM,CAAqB;IACnC,OAAO,CAAC,KAAK,CAAiB;IAC9B,OAAO,CAAC,MAAM,CAAuB;IACrC,OAAO,CAAC,QAAQ,CAAC,CAAoB;IACrC,OAAO,CAAC,YAAY,CAAa;IACjC,OAAO,CAAC,WAAW,CAAa;IAChC,wEAAwE;IACxE,OAAO,CAAC,iBAAiB,CAAC,CAAS;IAGnC,OAAO,CAAC,MAAM,CAAkB;IAChC,OAAO,CAAC,aAAa,CAAkB;IACvC,OAAO,CAAC,cAAc,CAAC,CAA+B;IACtD,OAAO,CAAC,MAAM,CAAiB;IAC/B,OAAO,CAAC,iBAAiB,CAA0C;IAGnE,OAAO,CAAC,WAAW,CAAC,CAAc;IAClC,OAAO,CAAC,sBAAsB,CAAC,CAAgB;IAC/C,OAAO,CAAC,gBAAgB,CAAC,CAAS;IAClC,OAAO,CAAC,eAAe,CAAC,CAAS;IAGjC,OAAO,CAAC,kBAAkB,CAAa;IACvC,OAAO,CAAC,uBAAuB,CAAkB;IAGjD,OAAO,CAAC,sBAAsB,CAAuB;IACrD,OAAO,CAAC,sBAAsB,CAAa;IAG3C,OAAO,CAAC,eAAe,CAAC,CAAkB;IAC1C,OAAO,CAAC,qBAAqB,CAA0B;IAEvD;;;;;OAKG;gBACS,KAAK,EAAE,cAAc,EAAE,MAAM,GAAE,OAAO,CAAC,iBAAiB,CAAC,CAAC,CAAC,CAAM,EAAE,MAAM,CAAC,EAAE,cAAc;IAqFtG;;OAEG;IACH,OAAO,CAAC,iBAAiB;IAezB;;;;;;;OAOG;IACH,OAAO,CAAC,mBAAmB;IAuD3B;;OAEG;YACW,eAAe;IAuF7B;;;;;OAKG;IACH,OAAO,CAAC,qBAAqB;IA6B7B;;;;;;;;OAQG;YACW,eAAe;IAsD7B;;;;;;OAMG;IACG,4BAA4B,IAAI,OAAO,CAAC,IAAI,CAAC;IAoBnD;;OAEG;IACH,cAAc,IAAI,WAAW,GAAG,SAAS;IAIzC;;OAEG;IACH,OAAO,CAAC,kBAAkB;IAoC1B;;OAEG;YACW,aAAa;IAM3B;;;;;;;;OAQG;WACU,MAAM,CAAC,CAAC,EAAE,MAAM,EAAE,iBAAiB,CAAC,CAAC,CAAC,GAAG,OAAO,CAAC,WAAW,CAAC,CAAC,CAAC,CAAC;IAoK7E;;OAEG;IACH,MAAM,CAAC,oBAAoB,CAAC,SAAS,CAAC,EAAE,MAAM,GAAG,MAAM,GAAG,IAAI;IAmB9D;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;OA8BG;IACH,qBAAqB,CAAC,MAAM,EAAE,mBAAmB,GAAG,IAAI;IAIxD;;;;;;OAMG;IACH,iBAAiB,CAAC,EAAE,EAAE,cAAc,CAAC,CAAC,CAAC,GAAG,IAAI;IAI9C;;;;;;;;OAQG;IACG,aAAa,CAAC,OAAO,EAAE,WAAW,EAAE,EAAE,GAAG,OAAO,CAAC,mBAAmB,CAAC;IA2B3E;;;;;;OAMG;IACH,gBAAgB,CAAC,OAAO,EAAE,MAAM,EAAE,EAAE,WAAW,EAAE,MAAM,EAAE,GAAG,MAAM,EAAE;IAIpE;;;;;;;;;;;;OAYG;IACG,gBAAgB,CACpB,OAAO,EAAE,WAAW,EAAE,EAAE,EACxB,WAAW,EAAE,MAAM,EAAE,EACrB,OAAO,EAAE,CAAC,EACV,SAAS,CAAC,EAAE,MAAM,EAClB,WAAW,CAAC,EAAE,MAAM,EAAE,EACtB,aAAa,CAAC,EAAE,MAAM,EAAE,GACvB,OAAO,CAAC,YAAY,CAAC;IA+DxB;;;;;;;;;;OAUG;IACG,SAAS,CAAC,OAAO,EAAE,WAAW,EAAE,EAAE,EAAE,OAAO,CAAC,EAAE,CAAC,GAAG,OAAO,CAAC,gBAAgB,CAAC;IAKjF;;;;;;;;;;;;;OAaG;IACG,aAAa,CACjB,OAAO,EAAE,WAAW,EAAE,EAAE,EACxB,OAAO,CAAC,EAAE,CAAC,GACV,OAAO,CAAC;QAAE,OAAO,EAAE,gBAAgB,CAAC;QAAC,WAAW,EAAE,MAAM,EAAE,CAAC;QAAC,OAAO,EAAE,MAAM,EAAE,CAAC;QAAC,iBAAiB,EAAE,MAAM,EAAE,CAAA;KAAE,CAAC;IA2IhH;;;;;OAKG;IACH,aAAa,IAAI,IAAI;IAIrB;;OAEG;IACH,OAAO,IAAI,MAAM;IAIjB;;OAEG;IACH,QAAQ,IAAI,MAAM;IAIlB;;;;;;;;;;;OAWG;IACG,oBAAoB,CACxB,IAAI,EAAE,MAAM,EACZ,OAAO,EAAE;QACP,IAAI,EAAE,MAAM,CAAC;QACb,UAAU,EAAE,MAAM,CAAC;QACnB,SAAS,EAAE,MAAM,CAAC;QAClB,aAAa,EAAE,MAAM,CAAC;QACtB,YAAY,EAAE,MAAM,CAAC;QACrB,WAAW,EAAE,MAAM,CAAC;KACrB,EACD,WAAW,EAAE,MAAM,EAAE,EACrB,OAAO,EAAE,MAAM,EAAE,EACjB,OAAO,EAAE,MAAM,EAAE,GAChB,OAAO,CAAC,IAAI,CAAC;IA+ChB;;;;;;;;;;;;OAYG;IACG,KAAK,CAAC,OAAO,EAAE,cAAc,EAAE,GAAG,OAAO,CAAC,IAAI,CAAC;IA6SrD;;;;;;;;;;OAUG;IACG,cAAc,CAAC,IAAI,CAAC,EAAE,MAAM,EAAE,OAAO,CAAC,EAAE;QAAE,WAAW,CAAC,EAAE,OAAO,CAAA;KAAE,GAAG,OAAO,CAAC,MAAM,CAAC;IA4HzF;;;OAGG;IACH,OAAO,CAAC,qBAAqB;IA2C7B;;OAEG;IACH,UAAU,IAAI,IAAI;IAIlB;;;;OAIG;IACH,QAAQ,CAAC,aAAa,EAAE,MAAM,GAAG,kBAAkB;IAInD;;OAEG;IACH,KAAK,IAAI,IAAI;IAIb;;OAEG;IACH,IAAI,IAAI,IAAI,MAAM,CAEjB;IAED;;OAEG;IACH,IAAI,KAAK,IAAI,MAAM,CAElB;IAED;;OAEG;IACH,IAAI,SAAS,IAAI,MAAM,CAEtB;IAED;;OAEG;IACH,IAAI,iBAAiB,IAAI,OAAO,CAE/B;IAED;;OAEG;IACH,IAAI,WAAW,IAAI,MAAM,EAAE,CAE1B;IAED;;;;OAIG;IACH,eAAe,IAAI,kBAAkB;CAGtC;AAED;;;;;;;;;;;;;GAaG;AACH,wBAAgB,oBAAoB,IAAI,oBAAoB,CAE3D"}
|
|
@@ -213,6 +213,8 @@ export class GRPOTrainer {
|
|
|
213
213
|
topP: this.config.topP,
|
|
214
214
|
topK: this.config.topK,
|
|
215
215
|
repetitionPenalty: this.config.repetitionPenalty,
|
|
216
|
+
presencePenalty: this.config.presencePenalty,
|
|
217
|
+
frequencyPenalty: this.config.frequencyPenalty,
|
|
216
218
|
// Tool calling support
|
|
217
219
|
tools: this.config.tools,
|
|
218
220
|
enableThinking: this.config.enableThinking,
|
|
@@ -642,6 +644,7 @@ export class GRPOTrainer {
|
|
|
642
644
|
const model = await loadModel(modelPath);
|
|
643
645
|
logger.status('loading', `${modelName} loaded (${model.constructor.name})`);
|
|
644
646
|
// Create trainer with the pre-created logger
|
|
647
|
+
// @ts-expect-error
|
|
645
648
|
const trainer = new GRPOTrainer(model, config, logger);
|
|
646
649
|
// Always store the original model path (for tokenizer files when saving checkpoints)
|
|
647
650
|
trainer.originalModelPath = config.modelPath;
|
|
@@ -663,15 +666,35 @@ export class GRPOTrainer {
|
|
|
663
666
|
logger.debug(`Restored dataset metadata: size=${resumedState.dataset.size}, hash=${resumedState.dataset.contentHash}, ` +
|
|
664
667
|
`${trainer.processedBatchIndices.size} processed batches`);
|
|
665
668
|
}
|
|
666
|
-
// Restore optimizer state if available
|
|
667
|
-
|
|
669
|
+
// Restore optimizer state if available.
|
|
670
|
+
//
|
|
671
|
+
// If the checkpoint claims it has optimizer state (`hasOptimizerState`
|
|
672
|
+
// set by `saveCheckpoint` only after verifying the file exists on
|
|
673
|
+
// disk), any problem loading it is a HARD ERROR: the alternative is to
|
|
674
|
+
// silently continue with a fresh optimizer, which masks corruption and
|
|
675
|
+
// leaves the user wondering why their training dynamics drifted after
|
|
676
|
+
// a resume. Missing file => corrupt checkpoint; throw failure =>
|
|
677
|
+
// corrupt checkpoint. Either way, fail loud.
|
|
678
|
+
if (resumedState.hasOptimizerState === true) {
|
|
668
679
|
const optimizerStatePath = join(modelPath, 'optimizer_state.safetensors');
|
|
680
|
+
if (!existsSync(optimizerStatePath)) {
|
|
681
|
+
throw new Error(`Corrupt checkpoint at ${modelPath}: training_state.json declares ` +
|
|
682
|
+
`hasOptimizerState=true but ${optimizerStatePath} does not exist. ` +
|
|
683
|
+
`Refusing to silently continue with a fresh optimizer. ` +
|
|
684
|
+
`To intentionally reset optimizer state on resume, set ` +
|
|
685
|
+
`"hasOptimizerState": false in training_state.json.`);
|
|
686
|
+
}
|
|
669
687
|
try {
|
|
670
|
-
trainer.engine.loadOptimizerState(optimizerStatePath);
|
|
688
|
+
await trainer.engine.loadOptimizerState(optimizerStatePath);
|
|
671
689
|
logger.info(`Restored optimizer state from checkpoint`);
|
|
672
690
|
}
|
|
673
691
|
catch (e) {
|
|
674
|
-
|
|
692
|
+
throw new Error(`Failed to restore optimizer state from ${optimizerStatePath}: ${String(e)}. ` +
|
|
693
|
+
`Refusing to silently continue with a fresh optimizer on a checkpoint that ` +
|
|
694
|
+
`declared hasOptimizerState=true. ` +
|
|
695
|
+
`To intentionally reset optimizer state on resume, set ` +
|
|
696
|
+
`"hasOptimizerState": false in training_state.json and remove or rename ` +
|
|
697
|
+
`the optimizer_state.safetensors file.`);
|
|
675
698
|
}
|
|
676
699
|
}
|
|
677
700
|
// If resuming from a regular checkpoint (not emergency), track it as last known good
|
|
@@ -1331,15 +1354,61 @@ export class GRPOTrainer {
|
|
|
1331
1354
|
}
|
|
1332
1355
|
}
|
|
1333
1356
|
// Save optimizer state (AdamW moments + step counter)
|
|
1357
|
+
//
|
|
1358
|
+
// NOTE: `save_optimizer_state_sync` on the Rust side intentionally returns
|
|
1359
|
+
// `Ok(())` WITHOUT writing a file in two cases (see
|
|
1360
|
+
// crates/mlx-core/src/training_state.rs):
|
|
1361
|
+
// 1. No optimizer configured (SGD path — `self.optimizer.is_none()`).
|
|
1362
|
+
// 2. AdamW configured but the state map is empty because no training
|
|
1363
|
+
// step has ever run through `update_batch` (e.g. checkpoint taken
|
|
1364
|
+
// before any trainStep, or every rollout was filtered by the
|
|
1365
|
+
// degenerate-completion filter).
|
|
1366
|
+
//
|
|
1367
|
+
// Those are legitimate no-ops on the Rust side, but the TS trainer MUST
|
|
1368
|
+
// NOT lie about disk state: `hasOptimizerState` is the flag the resume
|
|
1369
|
+
// path reads to decide whether to call `loadOptimizerState`. If we set it
|
|
1370
|
+
// to `true` when no fresh file exists, the resume path will either load
|
|
1371
|
+
// a stale file from a previous save in the same directory, or crash on a
|
|
1372
|
+
// missing file — both are silent corruption paths.
|
|
1373
|
+
//
|
|
1374
|
+
// CRITICAL: checkpoint directories can be reused (e.g. emergency save into
|
|
1375
|
+
// an existing directory, or user-provided `outputDir` with a predictable
|
|
1376
|
+
// checkpoint name). An old `optimizer_state.safetensors` from a previous
|
|
1377
|
+
// save would make `existsSync` return true even when THIS save was a
|
|
1378
|
+
// no-op, so we would load stale state from a completely different step
|
|
1379
|
+
// on resume. Unlink the file up-front so that `existsSync` after the save
|
|
1380
|
+
// reflects only what the current save produced.
|
|
1381
|
+
const optimizerStatePath = join(checkpointPath, 'optimizer_state.safetensors');
|
|
1382
|
+
if (existsSync(optimizerStatePath)) {
|
|
1383
|
+
rmSync(optimizerStatePath, { force: true });
|
|
1384
|
+
}
|
|
1334
1385
|
try {
|
|
1335
|
-
this.engine.saveOptimizerState(
|
|
1336
|
-
|
|
1337
|
-
|
|
1386
|
+
await this.engine.saveOptimizerState(optimizerStatePath);
|
|
1387
|
+
const wroteOptimizerState = existsSync(optimizerStatePath);
|
|
1388
|
+
state.hasOptimizerState = wroteOptimizerState;
|
|
1389
|
+
if (!wroteOptimizerState) {
|
|
1390
|
+
// Expected when SGD is configured or no training step has populated
|
|
1391
|
+
// AdamW moments yet. Logged at info so it's visible in normal runs
|
|
1392
|
+
// without requiring debug mode — the fact that a checkpoint has no
|
|
1393
|
+
// optimizer state is useful signal for anyone debugging resume.
|
|
1394
|
+
this.logger.info(`saveOptimizerState produced no file at ${optimizerStatePath} ` +
|
|
1395
|
+
`(SGD or empty AdamW state); hasOptimizerState=false`);
|
|
1396
|
+
}
|
|
1397
|
+
// Re-write training_state.json with the accurate hasOptimizerState flag
|
|
1338
1398
|
writeFileSync(statePath, JSON.stringify(state, null, 2));
|
|
1339
1399
|
}
|
|
1340
1400
|
catch (e) {
|
|
1341
|
-
//
|
|
1401
|
+
// A thrown error from saveOptimizerState is a real save failure (not the
|
|
1402
|
+
// legitimate no-op path above). Force hasOptimizerState=false so resume
|
|
1403
|
+
// doesn't try to load a file that may be missing or partially written,
|
|
1404
|
+
// and unlink any partial file the Rust side may have left behind.
|
|
1405
|
+
// Log loud — real save failures should be visible.
|
|
1342
1406
|
this.logger.warn(`Failed to save optimizer state: ${String(e)}`);
|
|
1407
|
+
if (existsSync(optimizerStatePath)) {
|
|
1408
|
+
rmSync(optimizerStatePath, { force: true });
|
|
1409
|
+
}
|
|
1410
|
+
state.hasOptimizerState = false;
|
|
1411
|
+
writeFileSync(statePath, JSON.stringify(state, null, 2));
|
|
1343
1412
|
}
|
|
1344
1413
|
this.logger.info(`Checkpoint saved: ${checkpointPath}`);
|
|
1345
1414
|
// Track last checkpoint step for emergency save throttling
|
|
@@ -1 +1 @@
|
|
|
1
|
-
{"version":3,"file":"sft-trainer.d.ts","sourceRoot":"","sources":["../../src/trainers/sft-trainer.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;;;;;GAuBG;AAMH,OAAO,EAKL,cAAc,EAGd,KAAK,cAAc,EAEpB,MAAM,gBAAgB,CAAC;AACxB,OAAO,EAAa,KAAK,cAAc,EAAE,MAAM,cAAc,CAAC;AAE9D,OAAO,EAAE,UAAU,EAAkB,KAAK,QAAQ,EAAE,MAAM,wBAAwB,CAAC;AACnF,OAAO,KAAK,EAAE,gBAAgB,EAAE,MAAM,iBAAiB,CAAC;AAExD,OAAO,EAAwB,KAAK,cAAc,EAAE,MAAM,sBAAsB,CAAC;AAGjF,OAAO,EAAE,iBAAiB,EAAE,MAAM,gBAAgB,CAAC;AACnD,YAAY,EAAE,eAAe,EAAE,cAAc,EAAE,eAAe,EAAE,MAAM,gBAAgB,CAAC;AAEvF;;GAEG;AACH,MAAM,WAAW,gBAAgB;IAC/B,IAAI,EAAE,MAAM,CAAC;IACb,KAAK,EAAE,MAAM,CAAC;IACd,SAAS,EAAE,MAAM,CAAC;IAClB,WAAW,EAAE,KAAK,CAAC;CACpB;AAED;;GAEG;AACH,MAAM,WAAW,kBAAkB;IACjC,mBAAmB;IACnB,OAAO,EAAE,cAAc,CAAC;IACxB,oBAAoB;IACpB,KAAK,EAAE,MAAM,CAAC;CACf;AAED;;;;GAIG;AACH,qBAAa,UAAU;IACrB,OAAO,CAAC,MAAM,CAAoB;IAClC,OAAO,CAAC,KAAK,CAAiB;IAC9B,OAAO,CAAC,SAAS,CAAiB;IAClC,OAAO,CAAC,MAAM,CAAmB;IACjC,OAAO,CAAC,YAAY,CAAa;IACjC,OAAO,CAAC,WAAW,CAAa;IAChC,wEAAwE;IACxE,OAAO,CAAC,iBAAiB,CAAC,CAAS;IAGnC,OAAO,CAAC,MAAM,CAAkB;IAChC,OAAO,CAAC,aAAa,CAAkB;IACvC,OAAO,CAAC,cAAc,CAAC,CAAqB;IAC5C,OAAO,CAAC,MAAM,CAAiB;IAC/B,OAAO,CAAC,iBAAiB,CAA0C;IACnE,OAAO,CAAC,uBAAuB,CAAkB;IAEjD;;;;;;;OAOG;gBAED,KAAK,EAAE,cAAc,EACrB,SAAS,EAAE,cAAc,EACzB,MAAM,GAAE,OAAO,CAAC,gBAAgB,CAAM,EACtC,MAAM,CAAC,EAAE,cAAc;IAiDzB;;;OAGG;IACH,OAAO,CAAC,mBAAmB;IA8B3B;;OAEG;IACH,OAAO,CAAC,iBAAiB;IAezB;;OAEG;IACH,OAAO,CAAC,kBAAkB;IAmC1B;;OAEG;YACW,aAAa;IAM3B;;;;;OAKG;WACU,MAAM,CAAC,MAAM,EAAE,OAAO,CAAC,gBAAgB,CAAC,GAAG,OAAO,CAAC,UAAU,CAAC;
|
|
1
|
+
{"version":3,"file":"sft-trainer.d.ts","sourceRoot":"","sources":["../../src/trainers/sft-trainer.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;;;;;GAuBG;AAMH,OAAO,EAKL,cAAc,EAGd,KAAK,cAAc,EAEpB,MAAM,gBAAgB,CAAC;AACxB,OAAO,EAAa,KAAK,cAAc,EAAE,MAAM,cAAc,CAAC;AAE9D,OAAO,EAAE,UAAU,EAAkB,KAAK,QAAQ,EAAE,MAAM,wBAAwB,CAAC;AACnF,OAAO,KAAK,EAAE,gBAAgB,EAAE,MAAM,iBAAiB,CAAC;AAExD,OAAO,EAAwB,KAAK,cAAc,EAAE,MAAM,sBAAsB,CAAC;AAGjF,OAAO,EAAE,iBAAiB,EAAE,MAAM,gBAAgB,CAAC;AACnD,YAAY,EAAE,eAAe,EAAE,cAAc,EAAE,eAAe,EAAE,MAAM,gBAAgB,CAAC;AAEvF;;GAEG;AACH,MAAM,WAAW,gBAAgB;IAC/B,IAAI,EAAE,MAAM,CAAC;IACb,KAAK,EAAE,MAAM,CAAC;IACd,SAAS,EAAE,MAAM,CAAC;IAClB,WAAW,EAAE,KAAK,CAAC;CACpB;AAED;;GAEG;AACH,MAAM,WAAW,kBAAkB;IACjC,mBAAmB;IACnB,OAAO,EAAE,cAAc,CAAC;IACxB,oBAAoB;IACpB,KAAK,EAAE,MAAM,CAAC;CACf;AAED;;;;GAIG;AACH,qBAAa,UAAU;IACrB,OAAO,CAAC,MAAM,CAAoB;IAClC,OAAO,CAAC,KAAK,CAAiB;IAC9B,OAAO,CAAC,SAAS,CAAiB;IAClC,OAAO,CAAC,MAAM,CAAmB;IACjC,OAAO,CAAC,YAAY,CAAa;IACjC,OAAO,CAAC,WAAW,CAAa;IAChC,wEAAwE;IACxE,OAAO,CAAC,iBAAiB,CAAC,CAAS;IAGnC,OAAO,CAAC,MAAM,CAAkB;IAChC,OAAO,CAAC,aAAa,CAAkB;IACvC,OAAO,CAAC,cAAc,CAAC,CAAqB;IAC5C,OAAO,CAAC,MAAM,CAAiB;IAC/B,OAAO,CAAC,iBAAiB,CAA0C;IACnE,OAAO,CAAC,uBAAuB,CAAkB;IAEjD;;;;;;;OAOG;gBAED,KAAK,EAAE,cAAc,EACrB,SAAS,EAAE,cAAc,EACzB,MAAM,GAAE,OAAO,CAAC,gBAAgB,CAAM,EACtC,MAAM,CAAC,EAAE,cAAc;IAiDzB;;;OAGG;IACH,OAAO,CAAC,mBAAmB;IA8B3B;;OAEG;IACH,OAAO,CAAC,iBAAiB;IAezB;;OAEG;IACH,OAAO,CAAC,kBAAkB;IAmC1B;;OAEG;YACW,aAAa;IAM3B;;;;;OAKG;WACU,MAAM,CAAC,MAAM,EAAE,OAAO,CAAC,gBAAgB,CAAC,GAAG,OAAO,CAAC,UAAU,CAAC;IAiE3E;;OAEG;IACH,MAAM,CAAC,oBAAoB,CAAC,SAAS,CAAC,EAAE,MAAM,GAAG,MAAM,GAAG,IAAI;IAmB9D;;;;;OAKG;IACG,SAAS,CAAC,KAAK,EAAE,QAAQ,GAAG,OAAO,CAAC,kBAAkB,CAAC;IAqB7D;;;;OAIG;IACG,KAAK,CAAC,OAAO,EAAE,UAAU,GAAG,MAAM,GAAG,OAAO,CAAC,IAAI,CAAC;IAiLxD;;;;;OAKG;IACG,cAAc,CAAC,IAAI,CAAC,EAAE,MAAM,GAAG,OAAO,CAAC,MAAM,CAAC;IA+CpD;;OAEG;IACH,OAAO,CAAC,qBAAqB;IAiC7B;;OAEG;IACH,IAAI,IAAI,IAAI,MAAM,CAEjB;IAED;;OAEG;IACH,IAAI,KAAK,IAAI,MAAM,CAElB;IAED;;OAEG;IACH,QAAQ,IAAI,cAAc;IAU1B;;OAEG;IACH,YAAY,IAAI,cAAc;CAG/B"}
|
|
@@ -242,6 +242,7 @@ export class SFTTrainer {
|
|
|
242
242
|
const tokenizer = await Qwen3Tokenizer.fromPretrained(join(modelPath, 'tokenizer.json'));
|
|
243
243
|
logger.status('loading', `${modelName} loaded (${model.constructor.name})`);
|
|
244
244
|
// Create trainer
|
|
245
|
+
// @ts-expect-error
|
|
245
246
|
const trainer = new SFTTrainer(model, tokenizer, config, logger);
|
|
246
247
|
trainer.originalModelPath = config.modelName;
|
|
247
248
|
// Restore training state if resuming
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@mlx-node/trl",
|
|
3
|
-
"version": "0.0.
|
|
3
|
+
"version": "0.0.7",
|
|
4
4
|
"homepage": "https://github.com/mlx-node/mlx-node",
|
|
5
5
|
"bugs": {
|
|
6
6
|
"url": "https://github.com/mlx-node/mlx-node/issues"
|
|
@@ -29,12 +29,12 @@
|
|
|
29
29
|
"test:trainer": "TEST_TRAINER=1 vite test run"
|
|
30
30
|
},
|
|
31
31
|
"dependencies": {
|
|
32
|
-
"@mlx-node/core": "0.0.
|
|
33
|
-
"@mlx-node/lm": "0.0.
|
|
32
|
+
"@mlx-node/core": "0.0.7",
|
|
33
|
+
"@mlx-node/lm": "0.0.7",
|
|
34
34
|
"@std/toml": "npm:@jsr/std__toml@^1.0.11",
|
|
35
35
|
"change-case": "^5.4.4"
|
|
36
36
|
},
|
|
37
37
|
"devDependencies": {
|
|
38
|
-
"@huggingface/hub": "^2.
|
|
38
|
+
"@huggingface/hub": "^2.11.0"
|
|
39
39
|
}
|
|
40
40
|
}
|