Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
170 changes: 168 additions & 2 deletions packages/commands/src/commands/finetune/create.ts
Original file line number Diff line number Diff line change
Expand Up @@ -264,12 +264,104 @@ const COMMON_FLAGS = {

/**
* Text flags: text models consume the full hyper-parameter surface — training
* type selection plus n_epochs / batch_size / learning_rate / max_length (see
* type selection, base parameters and explicit LoRA/evaluation/save settings (see
* resolveTextHyperParameters). Only text exposes --training-type because only
* text models support types other than the sft-lora default.
*/
const TEXT_FLAGS = {
...COMMON_FLAGS,
jobName: {
type: "string",
valueHint: "<value>",
description: {
"en-US": "Training job display name (job_name)",
"zh-CN": "训练任务显示名称(job_name)",
},
},
priority: {
type: "string",
valueHint: "<value>",
description: {
"en-US": "Requested scheduling priority; verify the service response",
"zh-CN": "请求的调度优先级;请核对服务端回执",
},
choices: ["L0", "L1", "L2", "L3"] as const,
},
evalSteps: {
type: "number",
valueHint: "<value>",
description: { "en-US": "Validation interval in training steps", "zh-CN": "训练验证间隔步数" },
},
loraAlpha: {
type: "number",
valueHint: "<value>",
description: { "en-US": "LoRA scaling coefficient", "zh-CN": "LoRA 缩放系数" },
},
loraDropout: {
type: "number",
valueHint: "<value>",
description: { "en-US": "LoRA dropout probability", "zh-CN": "LoRA 丢弃率" },
},
loraRank: {
type: "number",
valueHint: "<value>",
description: { "en-US": "LoRA matrix rank", "zh-CN": "LoRA 矩阵秩" },
},
lrSchedulerType: {
type: "string",
valueHint: "<value>",
description: {
"en-US": "Learning rate scheduler supported by the selected model",
"zh-CN": "所选模型支持的学习率调度策略",
},
},
saveStrategy: {
type: "string",
valueHint: "<value>",
description: { "en-US": "Checkpoint saving strategy", "zh-CN": "Checkpoint 保存策略" },
choices: ["epoch", "steps"] as const,
},
saveTotalLimit: {
type: "number",
valueHint: "<value>",
description: {
"en-US": "Maximum number of saved checkpoints",
"zh-CN": "最多保存的 Checkpoint 数量",
},
},
saveSteps: {
type: "number",
valueHint: "<value>",
description: {
"en-US": "Checkpoint saving interval for strategy=steps",
"zh-CN": "按 steps 保存时的间隔",
},
},
split: {
type: "number",
valueHint: "<value>",
description: {
"en-US": "Training fraction when no validation dataset is supplied",
"zh-CN": "未指定验证集时训练集所占比例",
},
},
maxSplitValDatasetSample: {
type: "number",
valueHint: "<value>",
description: {
"en-US": "Maximum automatically split validation samples",
"zh-CN": "自动切分验证集的样本数量上限",
},
},
dataAugmentation: {
type: "string",
valueHint: "<value>",
description: {
"en-US": "Mix platform training data (true or false)",
"zh-CN": "是否混入平台训练数据(true 或 false)",
},
choices: ["true", "false"] as const,
},
trainingType: {
type: "string",
valueHint: "<t>",
Expand Down Expand Up @@ -348,7 +440,7 @@ const IMAGE_FLAGS = {
} satisfies FlagsDef;

const TEXT_USAGE =
"--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>]";
"--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>]";

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

if (commandModality === "text") {
const extraParameters: Record<string, string> = {
evalSteps: "eval_steps",
loraAlpha: "lora_alpha",
loraDropout: "lora_dropout",
loraRank: "lora_rank",
lrSchedulerType: "lr_scheduler_type",
saveStrategy: "save_strategy",
saveTotalLimit: "save_total_limit",
saveSteps: "save_steps",
split: "split",
maxSplitValDatasetSample: "max_split_val_dataset_sample",
};
for (const [flagName, parameterName] of Object.entries(extraParameters)) {
const value = flags[flagName];
if (value !== undefined) hp[parameterName] = value;
}
for (const parameterName of [
"eval_steps",
"lora_alpha",
"lora_rank",
"save_total_limit",
"save_steps",
"max_split_val_dataset_sample",
]) {
const value = hp[parameterName];
if (
value !== undefined &&
(typeof value !== "number" || !Number.isInteger(value) || value <= 0)
) {
throw new BailianError(
`${parameterName} must be a positive integer. / 必须为正整数。`,
ExitCode.USAGE,
);
}
}
if (
hp.lora_dropout !== undefined &&
(typeof hp.lora_dropout !== "number" ||
!Number.isFinite(hp.lora_dropout) ||
hp.lora_dropout < 0 ||
hp.lora_dropout >= 1)
) {
throw new BailianError(
"lora_dropout must be in [0, 1). / LoRA 丢弃率必须在 [0, 1) 内。",
ExitCode.USAGE,
);
}
if (
hp.split !== undefined &&
(typeof hp.split !== "number" ||
!Number.isFinite(hp.split) ||
hp.split <= 0 ||
hp.split >= 1 ||
flags.validations)
) {
throw new BailianError(
"split must be in (0, 1) and cannot be combined with validations. / 切分比例须在 (0, 1) 内,且不能与独立验证集同时设置。",
ExitCode.USAGE,
);
}
if (flags.dataAugmentation !== undefined)
hp.data_augmentation = flags.dataAugmentation === "true";
if (hp.save_strategy === "steps" && hp.save_steps === undefined) {
throw new BailianError(
"save_strategy=steps requires --save-steps. / 按步保存时必须指定 --save-steps。",
ExitCode.USAGE,
);
}
}

// Restore the batch-size clamping warning that was lost when the logic moved
// into profiles. The profile silently clamps to [8, 1024]; surface it here
// so the user has an audit trail. Skip modalities that bypass the batch_size
Expand Down Expand Up @@ -699,6 +862,8 @@ async function runCreate<F extends FlagsDef>(
if (validationFileIds && validationFileIds.length > 0) {
body.validation_file_ids = validationFileIds;
}
if (typeof flags.jobName === "string") body.job_name = flags.jobName;
if (typeof flags.priority === "string") body.priority = flags.priority;
if (modelName) body.model_name = modelName;
if (suffix) body.finetuned_output_suffix = suffix;

Expand Down Expand Up @@ -741,6 +906,7 @@ export const finetuneTextCreate = defineCommand({
"--base-model qwen3-8b --datasets ./train.jsonl --validations ./eval.jsonl",
"--base-model qwen3-8b --datasets file-aaa,./extra.jsonl",
"--base-model qwen3-8b --datasets ./train.jsonl --training-type sft",
"--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",
'--base-model qwen3-8b --datasets file-xxx --learning-rate "1.6e-5" --n-epochs 4',
"--base-model qwen3-8b --datasets file-xxx --output json",
"--base-model qwen3-8b --datasets file-xxx --dry-run",
Expand Down
17 changes: 8 additions & 9 deletions packages/commands/src/commands/finetune/get.ts
Original file line number Diff line number Diff line change
Expand Up @@ -38,20 +38,19 @@ export default defineCommand({
}

const hyperParameters = job.hyper_parameters;
const hyperParts: string[] = [];
if (hyperParameters?.n_epochs !== undefined)
hyperParts.push(`n_epochs=${hyperParameters.n_epochs}`);
if (hyperParameters?.batch_size !== undefined)
hyperParts.push(`batch_size=${hyperParameters.batch_size}`);
if (hyperParameters?.learning_rate !== undefined)
hyperParts.push(`learning_rate=${hyperParameters.learning_rate}`);
if (hyperParameters?.max_length !== undefined)
hyperParts.push(`max_length=${hyperParameters.max_length}`);
const hyperParts = Object.entries(hyperParameters ?? {}).map(
([parameterName, value]) =>
`${parameterName}=${typeof value === "string" ? value : JSON.stringify(value)}`,
);

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

const item: Record<string, unknown> = {
job_id: job.job_id ?? jobId,
job_name: job.job_name ?? "",
priority: job.priority ?? "",
hyper_parameters: hyperParameters ?? {},
max_output_cnt: job.max_output_cnt ?? null,
base_model: job.model ?? "",
status: job.status ?? "",
training_type: job.training_type ?? "",
Expand Down
86 changes: 86 additions & 0 deletions packages/commands/tests/e2e/finetune.e2e.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -539,3 +539,89 @@ describe.skipIf(!isDashScopeE2EReady())("e2e: finetune (DashScope)", () => {
}
}, 60_000);
});

describe("finetune complete recipe (offline)", () => {
const recipeArgs = [
"finetune",
"text",
"create",
"--base-model",
"qwen3-4b-instruct-2507",
"--datasets",
"file-jev",
"--job-name",
"jev-train",
"--model-name",
"jev-model",
"--priority",
"L0",
"--n-epochs",
"1",
"--batch-size",
"8",
"--learning-rate",
"5e-5",
"--max-length",
"32768",
"--eval-steps",
"50",
"--lora-alpha",
"16",
"--lora-dropout",
"0.1",
"--lora-rank",
"8",
"--lr-scheduler-type",
"linear",
"--save-strategy",
"epoch",
"--save-total-limit",
"3",
"--split",
"0.9",
"--max-split-val-dataset-sample",
"1000",
"--data-augmentation",
"false",
"--dry-run",
"--output",
"json",
];
test("preserves names, priority and every explicit hyperparameter", async () => {
const result = await runCommandE2e(FINETUNE_ROUTES, recipeArgs);
expect(result.exitCode, result.stderr).toBe(0);
const response = parseStdoutJson<{ body: Record<string, unknown> }>(result.stdout);
expect(response.body).toMatchObject({
job_name: "jev-train",
model_name: "jev-model",
priority: "L0",
});
expect(response.body.hyper_parameters).toEqual({
n_epochs: 1,
batch_size: 8,
learning_rate: "5e-5",
max_length: 32768,
eval_steps: 50,
lora_alpha: 16,
lora_dropout: 0.1,
lora_rank: 8,
lr_scheduler_type: "linear",
save_strategy: "epoch",
save_total_limit: 3,
split: 0.9,
max_split_val_dataset_sample: 1000,
data_augmentation: false,
});
});
test.each([
["--lora-rank", "0"],
["--lora-dropout", "1"],
["--split", "1"],
["--priority", "L9"],
])("rejects invalid %s before submission", async (flagName, invalidValue) => {
const invalidArgs = [...recipeArgs];
invalidArgs[invalidArgs.indexOf(flagName) + 1] = invalidValue;
const result = await runCommandE2e(FINETUNE_ROUTES, invalidArgs);
expect(result.exitCode).toBe(2);
});
});
2 changes: 2 additions & 0 deletions packages/core/src/finetune/types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,8 @@ export interface CreateFineTuneRequest {
hyper_parameters?: FineTuneHyperParameters;
/** Display name for the job (optional, server generates if omitted). */
job_name?: string;
/** Requested scheduling priority; the service determines the effective priority. */
priority?: string;
/** Output model name. Either bring your own or let the server generate one. */
model_name?: string;
/** Suffix appended by the platform; field is `finetuned_output_suffix` (NOT `suffix`). */
Expand Down
Loading
Loading