feat: search draft models from remote
This commit is contained in:
@@ -255,20 +255,22 @@ export async function queryModelScopeModelFiles(
|
|||||||
// list models from huggingface
|
// list models from huggingface
|
||||||
export async function queryHuggingfaceModels(
|
export async function queryHuggingfaceModels(
|
||||||
params: {
|
params: {
|
||||||
|
limit?: number;
|
||||||
search: {
|
search: {
|
||||||
query: string;
|
query: string;
|
||||||
tags: string[];
|
tags?: string[];
|
||||||
sort?: string;
|
sort?: string;
|
||||||
task?: PipelineType;
|
task?: PipelineType;
|
||||||
};
|
};
|
||||||
},
|
},
|
||||||
options?: any
|
options?: any
|
||||||
) {
|
) {
|
||||||
|
console.log('params', params);
|
||||||
const result = [];
|
const result = [];
|
||||||
for await (const model of listModels({
|
for await (const model of listModels({
|
||||||
...params,
|
...params,
|
||||||
...options,
|
...options,
|
||||||
limit: 500,
|
limit: params.limit || 500,
|
||||||
additionalFields: ['sha', 'tags'],
|
additionalFields: ['sha', 'tags'],
|
||||||
fetch(_url: string, config: any) {
|
fetch(_url: string, config: any) {
|
||||||
const url = params.search.sort
|
const url = params.search.sort
|
||||||
|
|||||||
@@ -349,7 +349,8 @@ export const DO_NOT_TRIGGER_CHECK_COMPATIBILITY = [
|
|||||||
'gpu_selector.gpu_ids',
|
'gpu_selector.gpu_ids',
|
||||||
'run_command',
|
'run_command',
|
||||||
'image_name',
|
'image_name',
|
||||||
'extended_kv_cache.enabled'
|
'extended_kv_cache.enabled',
|
||||||
|
'speculative_config.draft_model'
|
||||||
];
|
];
|
||||||
|
|
||||||
// ignore to compare old and new data when these fields change in updating model
|
// ignore to compare old and new data when these fields change in updating model
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import SealSelect from '@/components/seal-form/seal-select';
|
|||||||
import { useIntl } from '@umijs/max';
|
import { useIntl } from '@umijs/max';
|
||||||
import { Form, Select } from 'antd';
|
import { Form, Select } from 'antd';
|
||||||
import React from 'react';
|
import React from 'react';
|
||||||
import { deployFormKeyMap } from '../config';
|
import { DeployFormKeyMap } from '../config';
|
||||||
import { useCatalogFormContext, useFormContext } from '../config/form-context';
|
import { useCatalogFormContext, useFormContext } from '../config/form-context';
|
||||||
import KVCacheForm from './kv-cache';
|
import KVCacheForm from './kv-cache';
|
||||||
import SpeculativeDecode from './speculative-decode';
|
import SpeculativeDecode from './speculative-decode';
|
||||||
@@ -41,7 +41,7 @@ const Performance: React.FC = () => {
|
|||||||
return (
|
return (
|
||||||
<>
|
<>
|
||||||
<div data-field="extended_kv_cache.enabled"></div>
|
<div data-field="extended_kv_cache.enabled"></div>
|
||||||
{formKey === deployFormKeyMap.catalog && (
|
{formKey === DeployFormKeyMap.CATALOG && (
|
||||||
<Form.Item name="mode">
|
<Form.Item name="mode">
|
||||||
<SealSelect
|
<SealSelect
|
||||||
onChange={onModeChange}
|
onChange={onModeChange}
|
||||||
|
|||||||
@@ -4,10 +4,18 @@ import SealInputNumber from '@/components/seal-form/input-number';
|
|||||||
import SealInput from '@/components/seal-form/seal-input';
|
import SealInput from '@/components/seal-form/seal-input';
|
||||||
import SealSelect from '@/components/seal-form/seal-select';
|
import SealSelect from '@/components/seal-form/seal-select';
|
||||||
import useAppUtils from '@/hooks/use-app-utils';
|
import useAppUtils from '@/hooks/use-app-utils';
|
||||||
|
import useDeferredRequest from '@/hooks/use-deferred-request';
|
||||||
import { useIntl } from '@umijs/max';
|
import { useIntl } from '@umijs/max';
|
||||||
import { Form } from 'antd';
|
import { Form } from 'antd';
|
||||||
|
import _ from 'lodash';
|
||||||
import { useEffect, useRef, useState } from 'react';
|
import { useEffect, useRef, useState } from 'react';
|
||||||
import { queryDraftModelList } from '../apis';
|
import {
|
||||||
|
queryDraftModelList,
|
||||||
|
queryHuggingfaceModels,
|
||||||
|
queryModelScopeModels
|
||||||
|
} from '../apis';
|
||||||
|
import { modelSourceMap } from '../config';
|
||||||
|
import { useFormContext } from '../config/form-context';
|
||||||
import { FormData } from '../config/types';
|
import { FormData } from '../config/types';
|
||||||
|
|
||||||
const AlgorithmMap = {
|
const AlgorithmMap = {
|
||||||
@@ -18,6 +26,7 @@ const AlgorithmMap = {
|
|||||||
|
|
||||||
const SpeculativeDecode = () => {
|
const SpeculativeDecode = () => {
|
||||||
const intl = useIntl();
|
const intl = useIntl();
|
||||||
|
const { source } = useFormContext();
|
||||||
const { getRuleMessage } = useAppUtils();
|
const { getRuleMessage } = useAppUtils();
|
||||||
const form = Form.useFormInstance();
|
const form = Form.useFormInstance();
|
||||||
const speculativeEnabled = Form.useWatch(
|
const speculativeEnabled = Form.useWatch(
|
||||||
@@ -26,9 +35,12 @@ const SpeculativeDecode = () => {
|
|||||||
);
|
);
|
||||||
const algorithm = Form.useWatch(['speculative_config', 'algorithm'], form);
|
const algorithm = Form.useWatch(['speculative_config', 'algorithm'], form);
|
||||||
const [draftModelList, setDraftModelList] = useState<
|
const [draftModelList, setDraftModelList] = useState<
|
||||||
Global.BaseOption<string>[]
|
Global.BaseOption<string, { options: Global.BaseOption<string>[] }>[]
|
||||||
>([]);
|
>([]);
|
||||||
|
const presetDraftModelListRef = useRef<Global.BaseOption<string>[]>([]);
|
||||||
const speculativeConfigRef = useRef<any>({});
|
const speculativeConfigRef = useRef<any>({});
|
||||||
|
const axiosTokenRef = useRef<AbortController | null>(null);
|
||||||
|
const [loading, setLoading] = useState(false);
|
||||||
|
|
||||||
const fetchDraftModels = async () => {
|
const fetchDraftModels = async () => {
|
||||||
const response = await queryDraftModelList({
|
const response = await queryDraftModelList({
|
||||||
@@ -39,9 +51,128 @@ const SpeculativeDecode = () => {
|
|||||||
label: item.name,
|
label: item.name,
|
||||||
value: item.name
|
value: item.name
|
||||||
}));
|
}));
|
||||||
|
presetDraftModelListRef.current = options;
|
||||||
setDraftModelList(options);
|
setDraftModelList(options);
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const getHuggingfaceModels = async (query: string) => {
|
||||||
|
if (axiosTokenRef.current) {
|
||||||
|
axiosTokenRef.current.abort();
|
||||||
|
}
|
||||||
|
axiosTokenRef.current = new AbortController();
|
||||||
|
try {
|
||||||
|
const params = {
|
||||||
|
limit: 10,
|
||||||
|
search: {
|
||||||
|
query: query
|
||||||
|
}
|
||||||
|
};
|
||||||
|
setLoading(true);
|
||||||
|
const data = await queryHuggingfaceModels(params, {
|
||||||
|
signal: axiosTokenRef.current.signal
|
||||||
|
});
|
||||||
|
const list = _.map(data || [], (item: any) => {
|
||||||
|
return {
|
||||||
|
value: item.name,
|
||||||
|
label: item.name
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
const catalogModelList =
|
||||||
|
presetDraftModelListRef.current.length > 0
|
||||||
|
? [
|
||||||
|
{
|
||||||
|
label: `${intl.formatMessage({ id: 'models.form.source' })}: ${intl.formatMessage({ id: 'menu.models.modelCatalog' })}`,
|
||||||
|
title: `${intl.formatMessage({ id: 'models.form.source' })}: ${intl.formatMessage({ id: 'menu.models.modelCatalog' })}`,
|
||||||
|
options: presetDraftModelListRef.current || []
|
||||||
|
}
|
||||||
|
]
|
||||||
|
: [];
|
||||||
|
|
||||||
|
setDraftModelList([
|
||||||
|
...catalogModelList,
|
||||||
|
{
|
||||||
|
label: `${intl.formatMessage({ id: 'models.form.source' })}: Hugging Face`,
|
||||||
|
title: `${intl.formatMessage({ id: 'models.form.source' })}: Hugging Face`,
|
||||||
|
options: list
|
||||||
|
}
|
||||||
|
]);
|
||||||
|
} catch (error) {
|
||||||
|
setDraftModelList(presetDraftModelListRef.current);
|
||||||
|
} finally {
|
||||||
|
setLoading(false);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const getModelScopeModels = async (query: string) => {
|
||||||
|
if (axiosTokenRef.current) {
|
||||||
|
axiosTokenRef.current.abort();
|
||||||
|
}
|
||||||
|
axiosTokenRef.current = new AbortController();
|
||||||
|
try {
|
||||||
|
const params = {
|
||||||
|
Name: query,
|
||||||
|
PageSize: 10,
|
||||||
|
PageNumber: 1,
|
||||||
|
tasks: []
|
||||||
|
};
|
||||||
|
setLoading(true);
|
||||||
|
const data = await queryModelScopeModels(params, {
|
||||||
|
signal: axiosTokenRef.current.signal
|
||||||
|
});
|
||||||
|
const list = _.map(
|
||||||
|
_.get(data, 'Data.Model.Models') || [],
|
||||||
|
(item: any) => {
|
||||||
|
return {
|
||||||
|
label: `${item.Path}/${item.Name}`,
|
||||||
|
value: `${item.Path}/${item.Name}`
|
||||||
|
};
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
const catalogModelList =
|
||||||
|
presetDraftModelListRef.current.length > 0
|
||||||
|
? [
|
||||||
|
{
|
||||||
|
label: `${intl.formatMessage({ id: 'models.form.source' })}: ${intl.formatMessage({ id: 'menu.models.modelCatalog' })}`,
|
||||||
|
title: `${intl.formatMessage({ id: 'models.form.source' })}: ${intl.formatMessage({ id: 'menu.models.modelCatalog' })}`,
|
||||||
|
options: presetDraftModelListRef.current || []
|
||||||
|
}
|
||||||
|
]
|
||||||
|
: [];
|
||||||
|
|
||||||
|
setDraftModelList([
|
||||||
|
...catalogModelList,
|
||||||
|
{
|
||||||
|
label: `${intl.formatMessage({ id: 'models.form.source' })}: ModelScope`,
|
||||||
|
title: `${intl.formatMessage({ id: 'models.form.source' })}: ModelScope`,
|
||||||
|
options: list
|
||||||
|
}
|
||||||
|
]);
|
||||||
|
} catch (error) {
|
||||||
|
setDraftModelList(presetDraftModelListRef.current);
|
||||||
|
} finally {
|
||||||
|
setLoading(false);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const handleOnSearch = async (value: string) => {
|
||||||
|
if (!value) {
|
||||||
|
setDraftModelList(presetDraftModelListRef.current);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if (source === modelSourceMap.huggingface_value) {
|
||||||
|
await getHuggingfaceModels(value);
|
||||||
|
} else if (source === modelSourceMap.modelscope_value) {
|
||||||
|
await getModelScopeModels(value);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const { run: onSearch } = useDeferredRequest(
|
||||||
|
(value: string) => handleOnSearch(value),
|
||||||
|
150
|
||||||
|
);
|
||||||
|
|
||||||
const handleSpeculativeEnabledChange = (e: any) => {
|
const handleSpeculativeEnabledChange = (e: any) => {
|
||||||
if (e.target.checked) {
|
if (e.target.checked) {
|
||||||
form.setFieldValue('speculative_config', {
|
form.setFieldValue('speculative_config', {
|
||||||
@@ -129,6 +260,7 @@ const SpeculativeDecode = () => {
|
|||||||
id: 'models.form.draftModel.tips'
|
id: 'models.form.draftModel.tips'
|
||||||
})}
|
})}
|
||||||
options={draftModelList}
|
options={draftModelList}
|
||||||
|
onSearch={onSearch}
|
||||||
></AutoComlete>
|
></AutoComlete>
|
||||||
</Form.Item>
|
</Form.Item>
|
||||||
)}
|
)}
|
||||||
|
|||||||
Reference in New Issue
Block a user