refactor: 调优创建提取表单模型
将默认参数、命令构建、payload 构造逻辑抽离为 fineTuneFormModel.ts,新增 FineTuneStartPayload 类型约束启动训练接口,FineTuneCreateView 瘦身为视图层,列表微调,回归脚本适配。
This commit is contained in:
@@ -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')
|
||||
|
||||
@@ -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>) =>
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, // 训练时量化(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 <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 },
|
||||
@@ -178,7 +115,7 @@ const allParams = computed(() => {
|
||||
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_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>
|
||||
|
||||
@@ -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
|
||||
|
||||
146
frontend/src/views/fine-tune/fineTuneFormModel.ts
Normal file
146
frontend/src/views/fine-tune/fineTuneFormModel.ts
Normal 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 }
|
||||
}
|
||||
Reference in New Issue
Block a user