refactor: 调优创建提取表单模型

将默认参数、命令构建、payload 构造逻辑抽离为 fineTuneFormModel.ts,新增 FineTuneStartPayload 类型约束启动训练接口,FineTuneCreateView 瘦身为视图层,列表微调,回归脚本适配。
This commit is contained in:
caoxiaozhu
2026-07-13 15:28:17 +08:00
parent bff07b7a31
commit 735a8a71f5
6 changed files with 262 additions and 170 deletions

View File

@@ -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<FineTuneTask[]>('/fine-tune')
@@ -16,7 +16,7 @@ export const createFineTune = (data: Partial<FineTuneTask>) =>
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<FineTuneTask>) =>

View File

@@ -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

View File

@@ -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<GpuInfo[]>([])
const selectedGpus = ref<number[]>([])
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, // 训练时量化QLoRA0=不量化
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 <base_model_path> \\\n`
cmd += ` --dataset <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(() => {
<div class="param-col config">
<el-input-number
v-if="param.type === 'number'"
v-model="(form as any)[param.key]"
:model-value="Number(parameterValue(param.key))"
:min="param.min" :max="param.max" :step="param.step" :precision="param.precision"
controls-position="right"
style="width: 200px"
@update:model-value="updateParameterValue(param.key, $event)"
/>
<el-select
v-else-if="param.type === 'select'"
v-model="(form as any)[param.key]"
:model-value="parameterValue(param.key)"
style="width: 200px"
@update:model-value="updateParameterValue(param.key, $event)"
>
<el-option v-for="o in param.options" :key="o.value" :label="o.label" :value="o.value" />
</el-select>

View File

@@ -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

View File

@@ -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<FineTuneStartPayload, 'task_id'> {
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 <base_model_path>',
' --dataset <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<FineTuneStartPayload, 'task_id'>,
): Partial<FineTuneTask> {
return { ...payload, status: 'pending', progress: 0 }
}