Skip to content

Commit 92a76dd

Browse files
authored
Merge pull request #233 from modelstudioai/feat/finetune-hyperparameters
feat(finetune): expose advanced training parameters
2 parents cd68e85 + fa9af0f commit 92a76dd

5 files changed

Lines changed: 301 additions & 31 deletions

File tree

‎packages/commands/src/commands/finetune/create.ts‎

Lines changed: 168 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -264,12 +264,104 @@ const COMMON_FLAGS = {
264264

265265
/**
266266
* Text flags: text models consume the full hyper-parameter surface — training
267-
* type selection plus n_epochs / batch_size / learning_rate / max_length (see
267+
* type selection, base parameters and explicit LoRA/evaluation/save settings (see
268268
* resolveTextHyperParameters). Only text exposes --training-type because only
269269
* text models support types other than the sft-lora default.
270270
*/
271271
const TEXT_FLAGS = {
272272
...COMMON_FLAGS,
273+
jobName: {
274+
type: "string",
275+
valueHint: "<value>",
276+
description: {
277+
"en-US": "Training job display name (job_name)",
278+
"zh-CN": "训练任务显示名称(job_name)",
279+
},
280+
},
281+
priority: {
282+
type: "string",
283+
valueHint: "<value>",
284+
description: {
285+
"en-US": "Requested scheduling priority; verify the service response",
286+
"zh-CN": "请求的调度优先级;请核对服务端回执",
287+
},
288+
choices: ["L0", "L1", "L2", "L3"] as const,
289+
},
290+
evalSteps: {
291+
type: "number",
292+
valueHint: "<value>",
293+
description: { "en-US": "Validation interval in training steps", "zh-CN": "训练验证间隔步数" },
294+
},
295+
loraAlpha: {
296+
type: "number",
297+
valueHint: "<value>",
298+
description: { "en-US": "LoRA scaling coefficient", "zh-CN": "LoRA 缩放系数" },
299+
},
300+
loraDropout: {
301+
type: "number",
302+
valueHint: "<value>",
303+
description: { "en-US": "LoRA dropout probability", "zh-CN": "LoRA 丢弃率" },
304+
},
305+
loraRank: {
306+
type: "number",
307+
valueHint: "<value>",
308+
description: { "en-US": "LoRA matrix rank", "zh-CN": "LoRA 矩阵秩" },
309+
},
310+
lrSchedulerType: {
311+
type: "string",
312+
valueHint: "<value>",
313+
description: {
314+
"en-US": "Learning rate scheduler supported by the selected model",
315+
"zh-CN": "所选模型支持的学习率调度策略",
316+
},
317+
},
318+
saveStrategy: {
319+
type: "string",
320+
valueHint: "<value>",
321+
description: { "en-US": "Checkpoint saving strategy", "zh-CN": "Checkpoint 保存策略" },
322+
choices: ["epoch", "steps"] as const,
323+
},
324+
saveTotalLimit: {
325+
type: "number",
326+
valueHint: "<value>",
327+
description: {
328+
"en-US": "Maximum number of saved checkpoints",
329+
"zh-CN": "最多保存的 Checkpoint 数量",
330+
},
331+
},
332+
saveSteps: {
333+
type: "number",
334+
valueHint: "<value>",
335+
description: {
336+
"en-US": "Checkpoint saving interval for strategy=steps",
337+
"zh-CN": "按 steps 保存时的间隔",
338+
},
339+
},
340+
split: {
341+
type: "number",
342+
valueHint: "<value>",
343+
description: {
344+
"en-US": "Training fraction when no validation dataset is supplied",
345+
"zh-CN": "未指定验证集时训练集所占比例",
346+
},
347+
},
348+
maxSplitValDatasetSample: {
349+
type: "number",
350+
valueHint: "<value>",
351+
description: {
352+
"en-US": "Maximum automatically split validation samples",
353+
"zh-CN": "自动切分验证集的样本数量上限",
354+
},
355+
},
356+
dataAugmentation: {
357+
type: "string",
358+
valueHint: "<value>",
359+
description: {
360+
"en-US": "Mix platform training data (true or false)",
361+
"zh-CN": "是否混入平台训练数据(true 或 false)",
362+
},
363+
choices: ["true", "false"] as const,
364+
},
273365
trainingType: {
274366
type: "string",
275367
valueHint: "<t>",
@@ -348,7 +440,7 @@ const IMAGE_FLAGS = {
348440
} satisfies FlagsDef;
349441

350442
const TEXT_USAGE =
351-
"--base-model <model> --datasets <id|path,...> [--validations <id|path,...>] [--model-name <name>] [--suffix <text>] [--n-epochs <n>] [--batch-size <n>] [--learning-rate <str>] [--max-length <n>] [--training-type <sft|sft-lora|dpo|dpo-lora|cpt>]";
443+
"--base-model <model> --datasets <id|path,...> [--validations <id|path,...>] [--job-name <name>] [--priority <L0|L1|L2|L3>] [--model-name <name>] [--suffix <text>] [--n-epochs <n>] [--batch-size <n>] [--learning-rate <str>] [--max-length <n>] [--training-type <sft|sft-lora|dpo|dpo-lora|cpt>]";
352444

353445
const AUDIO_USAGE =
354446
"--base-model <model> --datasets <id|path> [--validations <id|path>] [--model-name <name>] [--suffix <text>]";
@@ -573,6 +665,77 @@ async function runCreate<F extends FlagsDef>(
573665
flags as Record<string, unknown>,
574666
) as FineTuneHyperParameters;
575667

668+
if (commandModality === "text") {
669+
const extraParameters: Record<string, string> = {
670+
evalSteps: "eval_steps",
671+
loraAlpha: "lora_alpha",
672+
loraDropout: "lora_dropout",
673+
loraRank: "lora_rank",
674+
lrSchedulerType: "lr_scheduler_type",
675+
saveStrategy: "save_strategy",
676+
saveTotalLimit: "save_total_limit",
677+
saveSteps: "save_steps",
678+
split: "split",
679+
maxSplitValDatasetSample: "max_split_val_dataset_sample",
680+
};
681+
for (const [flagName, parameterName] of Object.entries(extraParameters)) {
682+
const value = flags[flagName];
683+
if (value !== undefined) hp[parameterName] = value;
684+
}
685+
for (const parameterName of [
686+
"eval_steps",
687+
"lora_alpha",
688+
"lora_rank",
689+
"save_total_limit",
690+
"save_steps",
691+
"max_split_val_dataset_sample",
692+
]) {
693+
const value = hp[parameterName];
694+
if (
695+
value !== undefined &&
696+
(typeof value !== "number" || !Number.isInteger(value) || value <= 0)
697+
) {
698+
throw new BailianError(
699+
`${parameterName} must be a positive integer. / 必须为正整数。`,
700+
ExitCode.USAGE,
701+
);
702+
}
703+
}
704+
if (
705+
hp.lora_dropout !== undefined &&
706+
(typeof hp.lora_dropout !== "number" ||
707+
!Number.isFinite(hp.lora_dropout) ||
708+
hp.lora_dropout < 0 ||
709+
hp.lora_dropout >= 1)
710+
) {
711+
throw new BailianError(
712+
"lora_dropout must be in [0, 1). / LoRA 丢弃率必须在 [0, 1) 内。",
713+
ExitCode.USAGE,
714+
);
715+
}
716+
if (
717+
hp.split !== undefined &&
718+
(typeof hp.split !== "number" ||
719+
!Number.isFinite(hp.split) ||
720+
hp.split <= 0 ||
721+
hp.split >= 1 ||
722+
flags.validations)
723+
) {
724+
throw new BailianError(
725+
"split must be in (0, 1) and cannot be combined with validations. / 切分比例须在 (0, 1) 内,且不能与独立验证集同时设置。",
726+
ExitCode.USAGE,
727+
);
728+
}
729+
if (flags.dataAugmentation !== undefined)
730+
hp.data_augmentation = flags.dataAugmentation === "true";
731+
if (hp.save_strategy === "steps" && hp.save_steps === undefined) {
732+
throw new BailianError(
733+
"save_strategy=steps requires --save-steps. / 按步保存时必须指定 --save-steps。",
734+
ExitCode.USAGE,
735+
);
736+
}
737+
}
738+
576739
// Restore the batch-size clamping warning that was lost when the logic moved
577740
// into profiles. The profile silently clamps to [8, 1024]; surface it here
578741
// so the user has an audit trail. Skip modalities that bypass the batch_size
@@ -699,6 +862,8 @@ async function runCreate<F extends FlagsDef>(
699862
if (validationFileIds && validationFileIds.length > 0) {
700863
body.validation_file_ids = validationFileIds;
701864
}
865+
if (typeof flags.jobName === "string") body.job_name = flags.jobName;
866+
if (typeof flags.priority === "string") body.priority = flags.priority;
702867
if (modelName) body.model_name = modelName;
703868
if (suffix) body.finetuned_output_suffix = suffix;
704869

@@ -741,6 +906,7 @@ export const finetuneTextCreate = defineCommand({
741906
"--base-model qwen3-8b --datasets ./train.jsonl --validations ./eval.jsonl",
742907
"--base-model qwen3-8b --datasets file-aaa,./extra.jsonl",
743908
"--base-model qwen3-8b --datasets ./train.jsonl --training-type sft",
909+
"--base-model qwen3-8b --datasets ./jev-train.jsonl --job-name jev-train-v1 --priority L0 --lora-rank 8 --lora-alpha 16 --lora-dropout 0.1 --lr-scheduler-type linear --eval-steps 50 --save-strategy epoch --save-total-limit 3 --split 0.9 --max-split-val-dataset-sample 1000 --data-augmentation false --dry-run",
744910
'--base-model qwen3-8b --datasets file-xxx --learning-rate "1.6e-5" --n-epochs 4',
745911
"--base-model qwen3-8b --datasets file-xxx --output json",
746912
"--base-model qwen3-8b --datasets file-xxx --dry-run",

‎packages/commands/src/commands/finetune/get.ts‎

Lines changed: 8 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -38,20 +38,19 @@ export default defineCommand({
3838
}
3939

4040
const hyperParameters = job.hyper_parameters;
41-
const hyperParts: string[] = [];
42-
if (hyperParameters?.n_epochs !== undefined)
43-
hyperParts.push(`n_epochs=${hyperParameters.n_epochs}`);
44-
if (hyperParameters?.batch_size !== undefined)
45-
hyperParts.push(`batch_size=${hyperParameters.batch_size}`);
46-
if (hyperParameters?.learning_rate !== undefined)
47-
hyperParts.push(`learning_rate=${hyperParameters.learning_rate}`);
48-
if (hyperParameters?.max_length !== undefined)
49-
hyperParts.push(`max_length=${hyperParameters.max_length}`);
41+
const hyperParts = Object.entries(hyperParameters ?? {}).map(
42+
([parameterName, value]) =>
43+
`${parameterName}=${typeof value === "string" ? value : JSON.stringify(value)}`,
44+
);
5045

5146
const usageTokens = typeof job.usage === "number" ? job.usage : undefined;
5247

5348
const item: Record<string, unknown> = {
5449
job_id: job.job_id ?? jobId,
50+
job_name: job.job_name ?? "",
51+
priority: job.priority ?? "",
52+
hyper_parameters: hyperParameters ?? {},
53+
max_output_cnt: job.max_output_cnt ?? null,
5554
base_model: job.model ?? "",
5655
status: job.status ?? "",
5756
training_type: job.training_type ?? "",

‎packages/commands/tests/e2e/finetune.e2e.test.ts‎

Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -539,3 +539,89 @@ describe.skipIf(!isDashScopeE2EReady())("e2e: finetune (DashScope)", () => {
539539
}
540540
}, 60_000);
541541
});
542+
543+
describe("finetune complete recipe (offline)", () => {
544+
const recipeArgs = [
545+
"finetune",
546+
"text",
547+
"create",
548+
"--base-model",
549+
"qwen3-4b-instruct-2507",
550+
"--datasets",
551+
"file-jev",
552+
"--job-name",
553+
"jev-train",
554+
"--model-name",
555+
"jev-model",
556+
"--priority",
557+
"L0",
558+
"--n-epochs",
559+
"1",
560+
"--batch-size",
561+
"8",
562+
"--learning-rate",
563+
"5e-5",
564+
"--max-length",
565+
"32768",
566+
"--eval-steps",
567+
"50",
568+
"--lora-alpha",
569+
"16",
570+
"--lora-dropout",
571+
"0.1",
572+
"--lora-rank",
573+
"8",
574+
"--lr-scheduler-type",
575+
"linear",
576+
"--save-strategy",
577+
"epoch",
578+
"--save-total-limit",
579+
"3",
580+
"--split",
581+
"0.9",
582+
"--max-split-val-dataset-sample",
583+
"1000",
584+
"--data-augmentation",
585+
"false",
586+
"--dry-run",
587+
"--output",
588+
"json",
589+
];
590+
test("preserves names, priority and every explicit hyperparameter", async () => {
591+
const result = await runCommandE2e(FINETUNE_ROUTES, recipeArgs);
592+
expect(result.exitCode, result.stderr).toBe(0);
593+
const response = parseStdoutJson<{ body: Record<string, unknown> }>(result.stdout);
594+
expect(response.body).toMatchObject({
595+
job_name: "jev-train",
596+
model_name: "jev-model",
597+
priority: "L0",
598+
});
599+
expect(response.body.hyper_parameters).toEqual({
600+
n_epochs: 1,
601+
batch_size: 8,
602+
learning_rate: "5e-5",
603+
max_length: 32768,
604+
eval_steps: 50,
605+
lora_alpha: 16,
606+
lora_dropout: 0.1,
607+
lora_rank: 8,
608+
lr_scheduler_type: "linear",
609+
save_strategy: "epoch",
610+
save_total_limit: 3,
611+
split: 0.9,
612+
max_split_val_dataset_sample: 1000,
613+
data_augmentation: false,
614+
});
615+
});
616+
test.each([
617+
["--lora-rank", "0"],
618+
["--lora-dropout", "1"],
619+
["--split", "1"],
620+
["--priority", "L9"],
621+
])("rejects invalid %s before submission", async (flagName, invalidValue) => {
622+
const invalidArgs = [...recipeArgs];
623+
invalidArgs[invalidArgs.indexOf(flagName) + 1] = invalidValue;
624+
const result = await runCommandE2e(FINETUNE_ROUTES, invalidArgs);
625+
expect(result.exitCode).toBe(2);
626+
});
627+
});

‎packages/core/src/finetune/types.ts‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,8 @@ export interface CreateFineTuneRequest {
3939
hyper_parameters?: FineTuneHyperParameters;
4040
/** Display name for the job (optional, server generates if omitted). */
4141
job_name?: string;
42+
/** Requested scheduling priority; the service determines the effective priority. */
43+
priority?: string;
4244
/** Output model name. Either bring your own or let the server generate one. */
4345
model_name?: string;
4446
/** Suffix appended by the platform; field is `finetuned_output_suffix` (NOT `suffix`). */

0 commit comments

Comments
 (0)