From 735a8a71f56290c85fc0ba375f8b35c62c78eddd Mon Sep 17 00:00:00 2001 From: caoxiaozhu Date: Mon, 13 Jul 2026 15:28:17 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E8=B0=83=E4=BC=98=E5=88=9B?= =?UTF-8?q?=E5=BB=BA=E6=8F=90=E5=8F=96=E8=A1=A8=E5=8D=95=E6=A8=A1=E5=9E=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 将默认参数、命令构建、payload 构造逻辑抽离为 fineTuneFormModel.ts,新增 FineTuneStartPayload 类型约束启动训练接口,FineTuneCreateView 瘦身为视图层,列表微调,回归脚本适配。 --- .../regression-fine-tune-create-ui.mjs | 27 ++- frontend/src/api/modules/fineTune.ts | 4 +- frontend/src/types/index.ts | 31 +++ .../views/fine-tune/FineTuneCreateView.vue | 222 +++++------------- .../src/views/fine-tune/FineTuneListView.vue | 2 +- .../src/views/fine-tune/fineTuneFormModel.ts | 146 ++++++++++++ 6 files changed, 262 insertions(+), 170 deletions(-) create mode 100644 frontend/src/views/fine-tune/fineTuneFormModel.ts diff --git a/frontend/scripts/regression-fine-tune-create-ui.mjs b/frontend/scripts/regression-fine-tune-create-ui.mjs index 7513ef8..742eb9c 100644 --- a/frontend/scripts/regression-fine-tune-create-ui.mjs +++ b/frontend/scripts/regression-fine-tune-create-ui.mjs @@ -1,10 +1,12 @@ import { readFileSync } from 'node:fs' import { resolve } from 'node:path' import { fileURLToPath } from 'node:url' +import ts from 'typescript' const root = resolve(fileURLToPath(new URL('..', import.meta.url))) const source = readFileSync(resolve(root, 'src/views/fine-tune/FineTuneCreateView.vue'), 'utf8') const modelDialogSource = readFileSync(resolve(root, 'src/components/ModelSelectDialog.vue'), 'utf8') +const formModelSource = readFileSync(resolve(root, 'src/views/fine-tune/fineTuneFormModel.ts'), 'utf8') function assert(condition, message) { if (!condition) { @@ -24,11 +26,32 @@ assert(modelDialogSource.includes('width="860px"'), 'Model dialog should use a c assert(modelDialogSource.includes('model-series-list'), 'Model dialog should include a model series list column') assert(modelDialogSource.includes('model-version-list'), 'Model dialog should include a snapshot/version list column') assert(modelDialogSource.includes('handleConfirm'), 'Model dialog should confirm the selected model before updating the form') -assert(source.includes("auto_merge: false"), 'Auto merge should default to disabled') +assert(formModelSource.includes('auto_merge: false'), 'Auto merge should default to disabled') assert(source.includes('v-if="form.train_type === \'SFT\'"'), 'Merge model settings should only be visible for SFT') assert(source.includes('content-position="left">合并模型'), 'SFT form should include a merge model section below data configuration') assert(source.includes('v-model="form.auto_merge"'), 'Merge model section should provide an auto merge selector') assert(source.includes('label="自动合并权重并保存"'), 'Auto merge selector should use a clear visible label') -assert((source.match(/auto_merge: form\.train_type === 'SFT' && form\.auto_merge/g) || []).length === 2, 'Auto merge should be sent for SFT when creating and starting the task') +assert(formModelSource.includes("auto_merge: form.train_type === 'SFT' && form.auto_merge"), 'Auto merge should be normalized by the shared payload builder') +assert(source.includes('const payload = buildFineTunePayload(form, selectedGpus.value)'), 'Create and start should share one normalized payload') +assert(source.includes('startFineTune({ ...payload, task_id: taskId })'), 'Start request should reuse the normalized payload') +assert(source.includes('Object.assign(form, DEFAULT_TRAINING_PARAMS)'), 'Reset should reuse the canonical defaults') +assert(!source.includes('const taskData = {'), 'The duplicated create payload should be removed') +assert(!source.includes('const createRes: any'), 'Create response should use the API return type') +assert(!source.includes('(check as any).exists'), 'Name-check response should use its API return type') +assert(source.includes('任务名校验失败'), 'Name-check failures should be visible and block submission') + +const runnableSource = ts.transpileModule( + formModelSource.replace("import type { FineTuneStartPayload, FineTuneTask } from '@/types'", ''), + { compilerOptions: { module: ts.ModuleKind.ESNext, target: ts.ScriptTarget.ES2022 } }, +).outputText +const model = await import(`data:text/javascript;base64,${Buffer.from(runnableSource).toString('base64')}`) +const defaults = model.createDefaultFineTuneForm() +const customized = { ...defaults, train_type: 'DPO', train_method: 'full', auto_merge: true, quantization_bit: 4 } +const payload = model.buildFineTunePayload(customized, [0, 2]) +assert(payload.auto_merge === false, 'Non-SFT tasks must never enable auto merge') +assert(payload.quantization_bit === 0, 'Full fine-tuning must never send QLoRA quantization') +assert(payload.gpus.join(',') === '0,2', 'Selected GPUs should be preserved in the shared payload') +const resetDefaults = { ...model.DEFAULT_TRAINING_PARAMS } +assert(resetDefaults.lora_alpha === defaults.lora_alpha, 'Reset and initial defaults must share LoRA values') console.log('fine-tune create UI regression checks passed') diff --git a/frontend/src/api/modules/fineTune.ts b/frontend/src/api/modules/fineTune.ts index eb69c22..8b50b23 100644 --- a/frontend/src/api/modules/fineTune.ts +++ b/frontend/src/api/modules/fineTune.ts @@ -1,5 +1,5 @@ import { get, post, put, del } from '../request' -import type { FineTuneTask, TrainingProgress } from '@/types' +import type { FineTuneStartPayload, FineTuneTask, TrainingProgress } from '@/types' /** 训练任务列表 */ export const getFineTuneList = () => get('/fine-tune') @@ -16,7 +16,7 @@ export const createFineTune = (data: Partial) => post<{ id: string | number }>('/fine-tune', data) /** 启动训练(第二步) */ -export const startFineTune = (data: any) => post('/fine-tune/start', data) +export const startFineTune = (data: FineTuneStartPayload) => post('/fine-tune/start', data) /** 更新训练任务 */ export const updateFineTune = (id: string | number, data: Partial) => diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 744097a..a4174df 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -139,6 +139,37 @@ export interface FineTuneTask { create_time?: string } +export type FineTuneStartPayload = Omit< + FineTuneTask, + 'id' | 'status' | 'progress' | 'process_id' | 'train_duration' | 'create_time' +> & { + task_id: string | number + description: string + template: string + train_method: TrainMethodType + train_dataset_id: string | number + auto_merge: boolean + output_model_name: string + gpus: number[] + batch_size: number + learning_rate: number + n_epochs: number + save_steps: number + lr_scheduler_type: string + max_length: number + warmup_ratio: number + weight_decay: number + lora_alpha: number + lora_dropout: number + lora_rank: number + quantization_bit: number + export_quantized: boolean + quant_method: string + quant_bits: number + quant_group_size: number + export_format: string +} + export interface TrainingProgress { status?: string progress?: number diff --git a/frontend/src/views/fine-tune/FineTuneCreateView.vue b/frontend/src/views/fine-tune/FineTuneCreateView.vue index 1462606..01530de 100644 --- a/frontend/src/views/fine-tune/FineTuneCreateView.vue +++ b/frontend/src/views/fine-tune/FineTuneCreateView.vue @@ -14,6 +14,14 @@ import { getModelList } from '@/api/modules/model' import { getDatasetList } from '@/api/modules/dataset' import { getSystemInfo } from '@/api/modules/system' import { TEMPLATE_GROUPS, LR_SCHEDULER_OPTIONS, QUANTIZATION_BIT_OPTIONS, QUANT_METHOD_OPTIONS, GGUF_FORMAT_OPTIONS } from '@/constants' +import { + DEFAULT_TRAINING_PARAMS, + buildFineTuneCommand, + buildFineTunePayload, + createDefaultFineTuneForm, + toCreateFineTunePayload, +} from './fineTuneFormModel' +import type { FineTuneFormModel } from './fineTuneFormModel' import type { ModelItem, DatasetItem, GpuInfo } from '@/types' const router = useRouter() @@ -26,36 +34,7 @@ const gpus = ref([]) const selectedGpus = ref([]) const modelDialogVisible = ref(false) -const form = reactive({ - name: '', - description: '', - train_type: 'SFT' as 'SFT' | 'DPO' | 'CPT', - base_model: '' as string | number, - template: 'qwen', - train_method: 'lora' as 'lora' | 'full', - train_dataset_id: '' as string | number, - auto_merge: false, - // 训练参数 - batch_size: 1, - learning_rate: 0.0001, - n_epochs: 1, - save_steps: 100, - lr_scheduler_type: 'cosine', - max_length: 512, - warmup_ratio: 0.05, - weight_decay: 0.01, - // LoRA 参数 - lora_alpha: 16, // 修复原项目 lora_alpha 默认值不一致 bug - lora_dropout: 0.1, - lora_rank: 8, - // 量化参数 - quantization_bit: 0, // 训练时量化(QLoRA):0=不量化 - export_quantized: false, // 训练后是否导出量化模型 - quant_method: 'bnb', // 导出量化方法 - quant_bits: 4, // 导出量化位数 - quant_group_size: 128, // 分组大小(GPTQ/AWQ) - export_format: 'Q4_K_M', // GGUF 导出格式 -}) +const form = reactive(createDefaultFineTuneForm()) const rules: FormRules = { name: [ @@ -78,45 +57,8 @@ const selectedModel = computed(() => models.value.find((model) => model.id === f const modelDialogTitle = computed(() => selectedModel.value?.name || '') -/** 训练命令实时预览 */ -const commandPreview = computed(() => { - const gpuIds = selectedGpus.value.length ? selectedGpus.value.join(',') : '0' - let cmd = `CUDA_VISIBLE_DEVICES=${gpuIds} llamafactory-cli train \\\n` - cmd += ` --stage ${form.train_type === 'DPO' ? 'dpo' : form.train_type === 'CPT' ? 'cpt' : 'sft'} \\\n` - cmd += ` --do_train \\\n` - cmd += ` --model_name_or_path \\\n` - cmd += ` --dataset \\\n` - cmd += ` --template ${form.template} \\\n` - cmd += ` --finetuning_type ${form.train_method} \\\n` - cmd += ` --output_dir ./saves/${form.name || 'output'} \\\n` - cmd += ` --per_device_train_batch_size ${form.batch_size} \\\n` - cmd += ` --learning_rate ${form.learning_rate} \\\n` - cmd += ` --num_train_epochs ${form.n_epochs} \\\n` - cmd += ` --save_steps ${form.save_steps} \\\n` - cmd += ` --lr_scheduler_type ${form.lr_scheduler_type} \\\n` - cmd += ` --cutoff_len ${form.max_length} \\\n` - cmd += ` --warmup_ratio ${form.warmup_ratio} \\\n` - cmd += ` --weight_decay ${form.weight_decay}` - if (showLoraParams.value) { - cmd += ` \\\n --lora_alpha ${form.lora_alpha}` - cmd += ` \\\n --lora_dropout ${form.lora_dropout}` - cmd += ` \\\n --lora_rank ${form.lora_rank}` - } - if (showLoraParams.value && form.quantization_bit) { - cmd += ` \\\n --quantization_bit ${form.quantization_bit}` - } - if (form.export_quantized) { - const method = form.quant_method - const bits = form.quant_bits - cmd += ` \\\n # 训练后导出量化模型:${method} ${bits}bit` - if (method === 'gguf') { - cmd += ` \\\n # export_format=${form.export_format}` - } else if (method === 'gptq' || method === 'awq') { - cmd += ` \\\n # group_size=${form.quant_group_size}` - } - } - return cmd -}) +/** 训练命令与提交载荷共用同一份表单模型。 */ +const commandPreview = computed(() => buildFineTuneCommand(form, selectedGpus.value)) /** GPU 多选切换 */ function toggleGpu(index: number) { @@ -140,31 +82,26 @@ function handleModelConfirm(modelId: string | number) { } function resetParams() { - Object.assign(form, { - batch_size: 1, - learning_rate: 0.0001, - n_epochs: 1, - save_steps: 100, - lr_scheduler_type: 'cosine', - max_length: 512, - warmup_ratio: 0.05, - weight_decay: 0.01, - lora_alpha: 16, - lora_dropout: 0.1, - lora_rank: 8, - quantization_bit: 0, - export_quantized: false, - quant_method: 'bnb', - quant_bits: 4, - quant_group_size: 128, - export_format: 'Q4_K_M', - }) + Object.assign(form, DEFAULT_TRAINING_PARAMS) } const isParamsExpanded = ref(false) +interface ParameterDefinition { + key: keyof FineTuneFormModel + name: string + desc: string + hint: string + type: 'number' | 'select' + min?: number + max?: number + step?: number + precision?: number + options?: Array<{ label: string; value: string | number }> +} + const allParams = computed(() => { - const params = [ + const params: ParameterDefinition[] = [ { key: 'batch_size', name: 'batch_size', desc: '批次大小,代表模型训练过程中,模型更新一次参数所需要的数据样本数。', hint: '[1, 64], step:1', type: 'number', min: 1, max: 64, step: 1 }, { key: 'learning_rate', name: 'learning_rate', desc: '学习率,代表每次更新数据的增量参数权重比例。', hint: '[0.000001, 1]', type: 'number', min: 0.000001, max: 1, step: 0.00001, precision: 6 }, { key: 'n_epochs', name: 'n_epochs', desc: '循环次数,代表模型训练过程中模型学习数据集的次数,可理解为看几遍数据,一般建议的范围是 1-3 遍即可,可依据需求进行调整', hint: '[1, 100], step:1', type: 'number', min: 1, max: 100, step: 1 }, @@ -176,9 +113,9 @@ const allParams = computed(() => { ] if (showLoraParams.value) { params.push( - { key: 'lora_alpha', name: 'lora_alpha', desc: 'LoRA 缩放系数。', hint: '16/32/64/128', type: 'select', options: [{label:'16',value:16},{label:'32',value:32},{label:'64',value:64},{label:'128',value:128}] }, - { key: 'lora_rank', name: 'lora_rank', desc: 'LoRA 秩大小,控制低秩矩阵的维度。', hint: '8/16/32/64', type: 'select', options: [{label:'8',value:8},{label:'16',value:16},{label:'32',value:32},{label:'64',value:64}] }, - { key: 'lora_dropout', name: 'lora_dropout', desc: 'LoRA 层的 dropout 比例。', hint: '[0, 1]', type: 'number', min: 0, max: 1, step: 0.05, precision: 2 } + { key: 'lora_alpha', name: 'lora_alpha', desc: 'LoRA 缩放系数。', hint: '16/32/64/128', type: 'select', options: [{ label: '16', value: 16 }, { label: '32', value: 32 }, { label: '64', value: 64 }, { label: '128', value: 128 }] }, + { key: 'lora_rank', name: 'lora_rank', desc: 'LoRA 秩大小,控制低秩矩阵的维度。', hint: '8/16/32/64', type: 'select', options: [{ label: '8', value: 8 }, { label: '16', value: 16 }, { label: '32', value: 32 }, { label: '64', value: 64 }] }, + { key: 'lora_dropout', name: 'lora_dropout', desc: 'LoRA 层的 dropout 比例。', hint: '[0, 1]', type: 'number', min: 0, max: 1, step: 0.05, precision: 2 }, ) } return params @@ -188,6 +125,15 @@ const visibleParams = computed(() => { return isParamsExpanded.value ? allParams.value : allParams.value.slice(0, 3) }) +function parameterValue(key: keyof FineTuneFormModel) { + const value = form[key] + return typeof value === 'boolean' ? Number(value) : value +} + +function updateParameterValue(key: keyof FineTuneFormModel, value: string | number | undefined) { + if (value !== undefined) Reflect.set(form, key, value) +} + async function loadModels() { try { models.value = (await getModelList()) || [] @@ -225,88 +171,32 @@ async function handleSubmit() { } submitting.value = true try { - // 任务名查重 - const check = await checkFineTuneName(form.name).catch(() => ({ exists: false })) - if ((check as any).exists) { + let check: { exists: boolean } + try { + check = await checkFineTuneName(form.name) + } catch { + ElMessage.error('任务名校验失败,请稍后重试') + return + } + if (check.exists) { ElMessage.error('任务名称已存在,请更换') - submitting.value = false return } - // 第一步:创建任务记录 - const taskData = { - name: form.name, - description: form.description, - base_model: form.base_model, - template: form.template, - train_type: form.train_type, - train_method: form.train_method, - gpus: selectedGpus.value, - train_dataset_id: form.train_dataset_id, - auto_merge: form.train_type === 'SFT' && form.auto_merge, - output_model_name: form.name, - batch_size: form.batch_size, - learning_rate: form.learning_rate, - n_epochs: form.n_epochs, - save_steps: form.save_steps, - lr_scheduler_type: form.lr_scheduler_type, - max_length: form.max_length, - warmup_ratio: form.warmup_ratio, - weight_decay: form.weight_decay, - lora_alpha: form.lora_alpha, - lora_dropout: form.lora_dropout, - lora_rank: form.lora_rank, - quantization_bit: form.train_method === 'lora' ? form.quantization_bit : 0, - export_quantized: form.export_quantized, - quant_method: form.export_quantized ? form.quant_method : '', - quant_bits: form.export_quantized ? form.quant_bits : 0, - quant_group_size: form.export_quantized ? form.quant_group_size : 0, - export_format: form.export_quantized && form.quant_method === 'gguf' ? form.export_format : '', - status: 'pending', - progress: 0, - } - const createRes: any = await createFineTune(taskData) - const taskId = createRes?.id || createRes + const payload = buildFineTunePayload(form, selectedGpus.value) + const createRes = await createFineTune(toCreateFineTunePayload(payload)) + const taskId = createRes.id - // 第二步:启动训练 try { - await startFineTune({ - task_id: taskId, - name: form.name, - base_model: form.base_model, - template: form.template, - train_type: form.train_type, - train_method: form.train_method, - train_dataset_id: form.train_dataset_id, - auto_merge: form.train_type === 'SFT' && form.auto_merge, - output_model_name: form.name, - gpus: selectedGpus.value, - batch_size: form.batch_size, - learning_rate: form.learning_rate, - n_epochs: form.n_epochs, - save_steps: form.save_steps, - lr_scheduler_type: form.lr_scheduler_type, - max_length: form.max_length, - warmup_ratio: form.warmup_ratio, - weight_decay: form.weight_decay, - lora_alpha: form.lora_alpha, - lora_dropout: form.lora_dropout, - lora_rank: form.lora_rank, - quantization_bit: form.train_method === 'lora' ? form.quantization_bit : 0, - export_quantized: form.export_quantized, - quant_method: form.export_quantized ? form.quant_method : '', - quant_bits: form.export_quantized ? form.quant_bits : 0, - quant_group_size: form.export_quantized ? form.quant_group_size : 0, - export_format: form.export_quantized && form.quant_method === 'gguf' ? form.export_format : '', - }) + await startFineTune({ ...payload, task_id: taskId }) ElMessage.success('训练任务已创建并启动') - } catch (e) { - // 启动失败,回写状态 + } catch { await updateFineTune(taskId, { status: 'failed' }) + ElMessage.error('任务已创建,但训练启动失败') } router.push('/fine-tune') } catch { - // ignore + ElMessage.error('训练任务创建失败,请稍后重试') } finally { submitting.value = false } @@ -427,15 +317,17 @@ onMounted(() => {
diff --git a/frontend/src/views/fine-tune/FineTuneListView.vue b/frontend/src/views/fine-tune/FineTuneListView.vue index e7aa5e7..4617a6f 100644 --- a/frontend/src/views/fine-tune/FineTuneListView.vue +++ b/frontend/src/views/fine-tune/FineTuneListView.vue @@ -37,7 +37,7 @@ const filteredList = computed(() => { if (filters.value.trainType.length && !filters.value.trainType.includes(row.train_type)) { return false } - if (filters.value.trainMethod.length && !filters.value.trainMethod.includes(row.train_method)) { + if (filters.value.trainMethod.length && (!row.train_method || !filters.value.trainMethod.includes(row.train_method))) { return false } return true diff --git a/frontend/src/views/fine-tune/fineTuneFormModel.ts b/frontend/src/views/fine-tune/fineTuneFormModel.ts new file mode 100644 index 0000000..8d776b2 --- /dev/null +++ b/frontend/src/views/fine-tune/fineTuneFormModel.ts @@ -0,0 +1,146 @@ +import type { FineTuneStartPayload, FineTuneTask } from '@/types' + +export type FineTuneFormModel = { + name: string + description: string + train_type: 'SFT' | 'DPO' | 'CPT' + base_model: string | number + template: string + train_method: 'lora' | 'full' + train_dataset_id: string | number + auto_merge: boolean + batch_size: number + learning_rate: number + n_epochs: number + save_steps: number + lr_scheduler_type: string + max_length: number + warmup_ratio: number + weight_decay: number + lora_alpha: number + lora_dropout: number + lora_rank: number + quantization_bit: number + export_quantized: boolean + quant_method: string + quant_bits: number + quant_group_size: number + export_format: string +} + +export const DEFAULT_TRAINING_PARAMS = { + batch_size: 1, + learning_rate: 0.0001, + n_epochs: 1, + save_steps: 100, + lr_scheduler_type: 'cosine', + max_length: 512, + warmup_ratio: 0.05, + weight_decay: 0.01, + lora_alpha: 16, + lora_dropout: 0.1, + lora_rank: 8, + quantization_bit: 0, + export_quantized: false, + quant_method: 'bnb', + quant_bits: 4, + quant_group_size: 128, + export_format: 'Q4_K_M', +} as const + +export function createDefaultFineTuneForm(): FineTuneFormModel { + return { + name: '', + description: '', + train_type: 'SFT', + base_model: '', + template: 'qwen', + train_method: 'lora', + train_dataset_id: '', + auto_merge: false, + ...DEFAULT_TRAINING_PARAMS, + } +} + +export function buildFineTunePayload( + form: FineTuneFormModel, + gpus: number[], +): Omit { + return { + name: form.name, + description: form.description, + base_model: form.base_model, + template: form.template, + train_type: form.train_type, + train_method: form.train_method, + gpus: [...gpus], + train_dataset_id: form.train_dataset_id, + auto_merge: form.train_type === 'SFT' && form.auto_merge, + output_model_name: form.name, + batch_size: form.batch_size, + learning_rate: form.learning_rate, + n_epochs: form.n_epochs, + save_steps: form.save_steps, + lr_scheduler_type: form.lr_scheduler_type, + max_length: form.max_length, + warmup_ratio: form.warmup_ratio, + weight_decay: form.weight_decay, + lora_alpha: form.lora_alpha, + lora_dropout: form.lora_dropout, + lora_rank: form.lora_rank, + quantization_bit: form.train_method === 'lora' ? form.quantization_bit : 0, + export_quantized: form.export_quantized, + quant_method: form.export_quantized ? form.quant_method : '', + quant_bits: form.export_quantized ? form.quant_bits : 0, + quant_group_size: form.export_quantized ? form.quant_group_size : 0, + export_format: form.export_quantized && form.quant_method === 'gguf' ? form.export_format : '', + } +} + +export function buildFineTuneCommand(form: FineTuneFormModel, gpus: number[]) { + const gpuIds = gpus.length ? gpus.join(',') : '0' + const stage = form.train_type === 'DPO' ? 'dpo' : form.train_type === 'CPT' ? 'cpt' : 'sft' + const lines = [ + `CUDA_VISIBLE_DEVICES=${gpuIds} llamafactory-cli train`, + ` --stage ${stage}`, + ' --do_train', + ' --model_name_or_path ', + ' --dataset ', + ` --template ${form.template}`, + ` --finetuning_type ${form.train_method}`, + ` --output_dir ./saves/${form.name || 'output'}`, + ` --per_device_train_batch_size ${form.batch_size}`, + ` --learning_rate ${form.learning_rate}`, + ` --num_train_epochs ${form.n_epochs}`, + ` --save_steps ${form.save_steps}`, + ` --lr_scheduler_type ${form.lr_scheduler_type}`, + ` --cutoff_len ${form.max_length}`, + ` --warmup_ratio ${form.warmup_ratio}`, + ` --weight_decay ${form.weight_decay}`, + ] + + if (form.train_method === 'lora') { + lines.push( + ` --lora_alpha ${form.lora_alpha}`, + ` --lora_dropout ${form.lora_dropout}`, + ` --lora_rank ${form.lora_rank}`, + ) + if (form.quantization_bit) lines.push(` --quantization_bit ${form.quantization_bit}`) + } + + if (form.export_quantized) { + lines.push(` # 训练后导出量化模型:${form.quant_method} ${form.quant_bits}bit`) + if (form.quant_method === 'gguf') lines.push(` # export_format=${form.export_format}`) + if (form.quant_method === 'gptq' || form.quant_method === 'awq') { + lines.push(` # group_size=${form.quant_group_size}`) + } + } + + return lines.join(' \\\n') +} + +export function toCreateFineTunePayload( + payload: Omit, +): Partial { + return { ...payload, status: 'pending', progress: 0 } +}