import ModalFooter from '@/components/modal-footer'; import GSDrawer from '@/components/scroller-modal/gs-drawer'; import { PageActionType } from '@/config/types'; import useDeferredRequest from '@/hooks/use-deferred-request'; import { ProviderValueMap } from '@/pages/cluster-management/config'; import { useIntl } from '@umijs/max'; import { Button } from 'antd'; import _ from 'lodash'; import { FC, useCallback, useEffect, useMemo, useRef, useState } from 'react'; import styled from 'styled-components'; import { backendOptionsMap, defaultFormValues, getSourceRepoConfigValue, modelSourceMap } from '../config'; import { FormContext } from '../config/form-context'; import { FormData, SourceType } from '../config/types'; import { MessageStatus, WarningStausOptions, checkOnlyAscendNPU, useCheckCompatibility, useSelectModel } from '../hooks'; import ColumnWrapper from './column-wrapper'; import CompatibilityAlert from './compatible-alert'; import DataForm from './data-form'; import GGUFResult from './gguf-result'; import ModelCard from './model-card'; import SearchModel from './search-model'; import Separator from './separator'; import TitleWrapper from './title-wrapper'; const resetFieldsByModel = ['backend_version', 'backend_parameters', 'env']; const pickFieldsFromSpec = ['backend_version', 'backend_parameters', 'env']; const dropFieldsFromForm = ['name', 'file_name', 'repo_id', 'backend']; const resetFields = ['worker_selector', 'env']; const resetFieldsByFile = [ 'cpu_offloading', 'distributed_inference_across_workers' ]; const ModalFooterStyle = { padding: '16px 24px', display: 'flex', justifyContent: 'flex-end' }; const ColWrapper = styled.div` display: flex; flex: 1; maxwidth: 33.33%; `; const FormWrapper = styled.div` display: flex; flex: 1; maxwidth: 100%; `; type AddModalProps = { title: string; hasLinuxWorker?: boolean; action: PageActionType; open: boolean; source: SourceType; isGGUF?: boolean; width?: string | number; initialValues?: any; deploymentType?: 'modelList' | 'modelFiles'; clusterList: Global.BaseOption< number, { provider: string; state: string | number } >[]; onOk: (values: FormData) => void; onCancel: () => void; }; type EvaluateProccessType = 'model' | 'file' | 'form'; const EvaluateProccess: Record = { model: 'model', file: 'file', form: 'form' }; const AddModal: FC = (props) => { const { title, open, onOk, onCancel, hasLinuxWorker, source, action, width = 600, deploymentType = 'modelList', initialValues, clusterList } = props || {}; const SEARCH_SOURCE = [ modelSourceMap.huggingface_value, modelSourceMap.modelscope_value ]; const { handleShowCompatibleAlert, setWarningStatus, handleBackendChangeBefore, cancelEvaluate, unlockWarningStatus, handleOnValuesChange: handleOnValuesChangeBefore, clearCahceFormValues, warningStatus, submitAnyway } = useCheckCompatibility(); const { onSelectModel } = useSelectModel({ gpuOptions: [] }); const form = useRef({}); const intl = useIntl(); const [selectedModel, setSelectedModel] = useState({}); const [collapsed, setCollapsed] = useState(false); const [isGGUF, setIsGGUF] = useState(false); const modelFileRef = useRef(null); const evaluateStateRef = useRef<{ state: EvaluateProccessType; requestModelId: number; }>({ state: 'form', requestModelId: 0 }); const requestModelIdRef = useRef(0); const currentSelectedModel = useRef({}); const { run: fetchModelFiles } = useDeferredRequest( () => modelFileRef.current?.fetchModelFiles?.(), 100 ); const updateSelectedModel = (model: any) => { currentSelectedModel.current = model; setSelectedModel(model); }; /** * Update the request model id to distinguish * the evaluate request. */ const updateRequestModelId = () => { requestModelIdRef.current += 1; return requestModelIdRef.current; }; /** * * @param state target to distinguish the evaluate state, current evaluate state * can be 'model', 'file' or 'form'. */ const setEvaluteState = (state: { state: EvaluateProccessType; requestModelId: number; }) => { evaluateStateRef.current = state; }; const handleOnValuesChange = (data: { changedValues: any; allValues: any; source: SourceType; }) => { setEvaluteState({ state: EvaluateProccess.form, requestModelId: updateRequestModelId() }); handleOnValuesChangeBefore(data); }; const getDefaultSpec = (item: any) => { const defaultSpec = _.pick( item.evaluateResult?.default_spec, pickFieldsFromSpec ); return defaultSpec; }; const getCategory = (item: any) => { const categories = item.evaluateResult?.default_spec?.categories || []; if (Array.isArray(categories)) { return categories?.[0] || null; } return categories || null; }; const handleCancelFiles = () => { cancelEvaluate(); modelFileRef.current?.cancelRequest(); }; const generateNameValue = ( item: any, modelName: string, manual?: boolean ) => { if (item.name === currentSelectedModel.current.name) { return manual ? modelName : form.current?.getFieldValue?.('name'); } return modelName; }; const currentModelDuringEvaluate = (item: any) => { return ( evaluateStateRef.current.state === EvaluateProccess.form && item.name === currentSelectedModel.current.name ); }; const handleOnSelectModel = async (item: any, manual?: boolean) => { // If the item is empty or the same as the selected model, do nothing handleCancelFiles(); if ( _.isEmpty(item) || (item.isGGUF === selectedModel.isGGUF && item.name === selectedModel.name) ) { return; } console.log('isgguf==================> select 1', item.isGGUF); console.log('handleOnSelectModel:', item, selectedModel); setIsGGUF(item.isGGUF); clearCahceFormValues(); unlockWarningStatus(); setEvaluteState({ state: EvaluateProccess.model, requestModelId: updateRequestModelId() }); updateSelectedModel(item); // TODO form.current?.resetFields(resetFields); const modelInfo = onSelectModel(item, props.source); form.current?.setFieldsValue?.({ ...defaultFormValues, ...modelInfo, categories: getCategory(item) }); setWarningStatus( { show: true, title: '', type: 'transition', message: intl.formatMessage({ id: 'models.form.evaluating' }) }, { override: true } ); if (item.isGGUF) { fetchModelFiles(); } }; const handleOnSelectModelAfterEvaluate = (item: any, manual?: boolean) => { if (currentModelDuringEvaluate(item)) { return; } if (manual) { form.current?.resetFields(resetFields); } console.log('isgguf==================> select 2', item.isGGUF); // If the item is empty setIsGGUF(item.isGGUF); updateSelectedModel(item); setEvaluteState({ state: EvaluateProccess.model, requestModelId: updateRequestModelId() }); handleCancelFiles(); const modelInfo = onSelectModel(item, props.source); if ( evaluateStateRef.current.state === EvaluateProccess.model && item.evaluated ) { handleShowCompatibleAlert(item.evaluateResult); const newFormValues = { ...(manual ? { ...defaultFormValues } : _.omit(form.current?.form?.getFieldsValue?.(), [ ...dropFieldsFromForm ])), ...getDefaultSpec(item), ...modelInfo, name: generateNameValue(item, modelInfo.name, manual), categories: getCategory(item) }; console.log('newFormValues:', newFormValues); form.current?.form?.setFieldsValue?.(newFormValues); handleOnValuesChangeBefore({ changedValues: {}, allValues: newFormValues, source: props.source }); } }; const handleOnOk = async (allValues: FormData) => { const result = getSourceRepoConfigValue(props.source, allValues).values; onOk(result); }; const handleSubmitAnyway = async () => { submitAnyway.current = true; form.current?.submit?.(); }; const handleSumit = () => { form.current?.submit?.(); }; const handleSetIsGGUF = async (flag: boolean) => { console.log('isgguf==================>', flag); setIsGGUF(flag); }; const handleBackendChange = async (backend: string) => { if (backend === backendOptionsMap.llamaBox) { setIsGGUF(true); } else { setIsGGUF(false); } const data = form.current.form.getFieldsValue?.(); const res = handleBackendChangeBefore(data); if (res.show) { return; } if (data.local_path || props.source !== modelSourceMap.local_path_value) { handleOnValuesChange?.({ changedValues: {}, allValues: backend === backendOptionsMap.llamaBox ? data : _.omit(data, [ 'cpu_offloading', 'distributed_inference_across_workers' ]), source: props.source }); } }; const onValuesChange = async (changedValues: any, allValues: any) => { handleOnValuesChange?.({ changedValues, allValues, source: props.source }); }; const handleCancel = useCallback(() => { onCancel?.(); }, [onCancel]); const initClusterId = () => { const cluster_id = clusterList?.find((item) => item.provider === ProviderValueMap.Custom) ?.value || clusterList?.[0]?.value; return cluster_id; }; const handleOnOpen = () => { if (props.deploymentType === 'modelFiles') { form.current?.form?.setFieldsValue({ ...props.initialValues, cluster_id: initClusterId() }); handleOnValuesChange?.({ changedValues: {}, allValues: { ...props.initialValues, cluster_id: initClusterId() }, source: source }); } else { let backend = checkOnlyAscendNPU([]) ? backendOptionsMap.ascendMindie : backendOptionsMap.vllm; form.current?.setFieldsValue?.({ backend, cluster_id: initClusterId() }); } }; const showExtraButton = useMemo(() => { return warningStatus.show && warningStatus.type !== 'success'; }, [warningStatus.show, warningStatus.type]); // This is only a placeholder for querying the model or file during the transition period. const displayEvaluateStatus = ( params: MessageStatus, options?: WarningStausOptions ) => { setWarningStatus( { show: params.show, title: '', type: 'transition', message: intl.formatMessage({ id: 'models.form.evaluating' }) }, options ); }; useEffect(() => { if (open) { handleOnOpen(); form.current?.getGPUOptionList?.({ clusterId: initClusterId() }); } else { cancelEvaluate(); clearCahceFormValues(); } return () => { setSelectedModel({}); setWarningStatus({ show: false, title: '', message: [] }); }; }, [open, clusterList]); return (
{SEARCH_SOURCE.includes(props.source) && deploymentType === 'modelList' && ( <> {isGGUF && } )} { setWarningStatus({ show: false, message: '' }); }} warningStatus={warningStatus} contentStyle={{ paddingInline: 0 }} > {intl.formatMessage({ id: 'models.form.submit.anyway' })} ) } style={ModalFooterStyle} > } > <> {SEARCH_SOURCE.includes(source) && deploymentType === 'modelList' && ( {intl.formatMessage({ id: 'models.form.configurations' })} )}
); }; export default AddModal;