import ModalFooter from '@/components/modal-footer'; import { PageActionType } from '@/config/types'; import { CloseOutlined } from '@ant-design/icons'; import { useIntl } from '@umijs/max'; import { Button, Drawer } from 'antd'; import _ from 'lodash'; import { FC, useCallback, useEffect, useMemo, useRef, useState } from 'react'; import styled from 'styled-components'; import { backendOptionsMap, getSourceRepoConfigValue, modelSourceMap } from '../config'; import { FormContext } from '../config/form-context'; import { FormData, SourceType } from '../config/types'; import { checkOnlyAscendNPU, useCheckCompatibility, useSelectModel } from '../hooks'; import ColumnWrapper from './column-wrapper'; import CompatibilityAlert from './compatible-alert'; import DataForm from './data-form'; import HFModelFile from './hf-model-file'; import ModelCard from './model-card'; import OllamaTips from './ollama-tips'; import SearchModel from './search-model'; import Separator from './separator'; import TitleWrapper from './title-wrapper'; const resetFieldsByModel = [ 'cpu_offloading', 'distributed_inference_across_workers', 'backend_version', 'backend_parameters', '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; gpuOptions: any[]; modelFileOptions: any[]; initialValues?: any; deploymentType?: 'modelList' | 'modelFiles'; 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 } = props || {}; const SEARCH_SOURCE = [ modelSourceMap.huggingface_value, modelSourceMap.modelscope_value ]; const { handleShowCompatibleAlert, setWarningStatus, handleBackendChangeBefore, cancelEvaluate, handleOnValuesChange: handleOnValuesChangeBefore, handleEvaluateOnChange, warningStatus, submitAnyway } = useCheckCompatibility(); const { onSelectModel } = useSelectModel({ gpuOptions: props.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 isHolderRef = useRef<{ model: boolean; file: boolean; }>({ model: false, file: false }); const evaluateStateRef = useRef<{ state: EvaluateProccessType }>({ state: 'form' }); const requestModelIdRef = useRef(0); /** * Update the request model id to distinguish * the evaluate request. */ const updateRequestModelId = () => { requestModelIdRef.current += 1; }; /** * * @param state target to distinguish the evaluate state */ const setEvaluteState = (state: EvaluateProccessType) => { evaluateStateRef.current.state = state; }; /** * * @param flag set the evaluate status of the model or file */ const setIsHolderRef = (flag: Record) => { isHolderRef.current = { ...isHolderRef.current, ...flag }; }; const handleOnValuesChange = (data: { changedValues: any; allValues: any; source: SourceType; }) => { setEvaluteState(EvaluateProccess.form); handleOnValuesChangeBefore(data); }; const getDefaultSpec = (item: any) => { const defaultSpec = item.evaluateResult?.default_spec || {}; return _.omit(defaultSpec, [ 'cpu_offloading', 'distributed_inference_across_workers' ]); }; const getCategory = (item: any) => { const categories = item.evaluateResult?.default_spec?.categories || []; if (Array.isArray(categories)) { return categories?.[0] || null; } return categories || null; }; const handleSelectModelFile = async (item: any, evaluate?: boolean) => { form.current?.form?.resetFields(resetFieldsByFile); const modelInfo = onSelectModel(selectedModel, props.source); /** display the selected model file information, but not * unitl the evaluate result is ready */ form.current?.setFieldsValue?.({ file_name: item.fakeName, ...modelInfo, categories: getCategory(item) }); await new Promise((resolve) => { setTimeout(() => { resolve(true); }, 0); }); if (item.fakeName) { const currentModelId = requestModelIdRef.current; setEvaluteState(EvaluateProccess.file); const evaluateRes = await handleEvaluateOnChange?.({ changedValues: {}, allValues: form.current?.form?.getFieldsValue?.(), source: props.source }); if (currentModelId !== requestModelIdRef.current) { // if the request model id has changed, do not update the form return; } const defaultSpec = getDefaultSpec({ evaluateResult: evaluateRes }); /** * do not reset backend_parameters when select a model file */ const formBackendParameters = form.current?.getFieldValue?.('backend_parameters') || []; form.current?.setFieldsValue?.({ file_name: item.fakeName, ...defaultSpec, ...modelInfo, backend_parameters: formBackendParameters.length > 0 ? formBackendParameters : defaultSpec.backend_parameters || [], categories: getCategory(item) }); } }; const handleOnSelectModel = (item: any, evaluate?: boolean) => { /** * evaluate: false means select a new model * evaluate: true means select a model file from the evaluate result */ updateRequestModelId(); if (!evaluate) { setEvaluteState(EvaluateProccess.model); setSelectedModel(item); form.current?.form?.resetFields(resetFieldsByModel); const modelInfo = onSelectModel(item, props.source); form.current?.setFieldsValue?.({ ...modelInfo, categories: getCategory(item) }); } if (!item.isGGUF) { setIsGGUF(false); const modelInfo = onSelectModel(item, props.source); if ( !isHolderRef.current.model && evaluateStateRef.current.state === EvaluateProccess.model ) { handleShowCompatibleAlert(item.evaluateResult); form.current?.setFieldsValue?.({ ...getDefaultSpec(item), ...modelInfo, categories: getCategory(item) }); } } }; 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) => { setIsGGUF(flag); await new Promise((resolve) => { setTimeout(() => { resolve(true); }, 0); }); if (flag) { modelFileRef.current?.fetchModelFiles?.(); } }; 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 handleOnOpen = () => { if (props.deploymentType === 'modelFiles') { form.current?.form?.setFieldsValue({ ...props.initialValues }); handleOnValuesChange?.({ changedValues: {}, allValues: props.initialValues, source: source }); } else { let backend = checkOnlyAscendNPU(props.gpuOptions) ? backendOptionsMap.ascendMindie : backendOptionsMap.vllm; if (source === modelSourceMap.ollama_library_value) { backend = backendOptionsMap.llamaBox; } form.current?.setFieldValue?.('backend', backend); setIsGGUF(source === modelSourceMap.ollama_library_value); } }; const showExtraButton = useMemo(() => { return warningStatus.show && warningStatus.type !== 'success'; }, [warningStatus.show, warningStatus.type]); const displayEvaluateStatus = (data: { show?: boolean; flag: Record; }) => { setIsHolderRef(data.flag); setWarningStatus({ show: isHolderRef.current.model || isHolderRef.current.file, title: '', type: 'transition', message: intl.formatMessage({ id: 'models.form.evaluating' }) }); }; useEffect(() => { if (open) { handleOnOpen(); } else { cancelEvaluate(); } return () => { setSelectedModel({}); setWarningStatus({ show: false, title: '', message: [] }); }; }, [open, props.gpuOptions.length]); return ( {title} } open={open} onClose={handleCancel} destroyOnClose={true} closeIcon={false} maskClosable={false} keyboard={false} zIndex={2000} styles={{ body: { height: 'calc(100vh - 57px)', padding: '16px 0', overflowX: 'hidden' }, content: { borderRadius: '6px 0 0 6px' } }} width={width} footer={false} >
{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} > } > <> {source === modelSourceMap.ollama_library_value && ( )} {SEARCH_SOURCE.includes(source) && deploymentType === 'modelList' && ( {intl.formatMessage({ id: 'models.form.configurations' })} )}
); }; export default AddModal;