From 76444f28cb5d8ddd8404488ec5d5c9142a91fa89 Mon Sep 17 00:00:00 2001 From: jialin Date: Thu, 22 Aug 2024 13:57:51 +0800 Subject: [PATCH] fix: query models sort --- src/assets/styles/common.less | 8 ++ src/components/icon-font/index.tsx | 2 +- src/locales/en-US/models.ts | 7 +- src/locales/zh-CN/models.ts | 7 +- src/pages/llmodels/apis/index.ts | 7 +- .../llmodels/components/hf-model-file.tsx | 21 ++-- src/pages/llmodels/components/model-card.tsx | 5 +- .../llmodels/components/search-model.tsx | 97 +++++++++---------- .../llmodels/components/search-result.tsx | 72 +++++++++++--- src/pages/llmodels/config/index.ts | 12 +-- src/pages/llmodels/style/hf-model-file.less | 1 + src/pages/llmodels/style/hf-model-item.less | 1 + 12 files changed, 151 insertions(+), 89 deletions(-) diff --git a/src/assets/styles/common.less b/src/assets/styles/common.less index 58e4f4ca..9a796b1f 100644 --- a/src/assets/styles/common.less +++ b/src/assets/styles/common.less @@ -26,6 +26,10 @@ margin-right: 5px; } +.m-r-2 { + margin-right: 2px; +} + .m-r-8 { margin-right: 8px; } @@ -67,6 +71,10 @@ flex-direction: column; } +.gap-5 { + gap: 5px; +} + .relative { position: relative; } diff --git a/src/components/icon-font/index.tsx b/src/components/icon-font/index.tsx index 3e7923ee..61a8c717 100644 --- a/src/components/icon-font/index.tsx +++ b/src/components/icon-font/index.tsx @@ -1,7 +1,7 @@ import { createFromIconfontCN } from '@ant-design/icons'; const IconFont = createFromIconfontCN({ - scriptUrl: '//at.alicdn.com/t/c/font_4613488_orlwzwe9x4k.js' + scriptUrl: '//at.alicdn.com/t/c/font_4613488_m179wc500g.js' }); export default IconFont; diff --git a/src/locales/en-US/models.ts b/src/locales/en-US/models.ts index 3525b6ff..c6e24385 100644 --- a/src/locales/en-US/models.ts +++ b/src/locales/en-US/models.ts @@ -22,11 +22,16 @@ export default { 'models.sort.name': 'Name', 'models.sort.size': 'Size', 'models.sort.likes': 'Likes', + 'models.sort.trending': 'Trending', 'models.sort.downloads': 'Downloads', 'models.sort.updated': 'Updated', 'models.search.result': '{count} results', 'models.data.card': 'Model Card', 'models.available.files': 'Available Files', 'models.viewin.hf': 'View in Hugging Face', - 'models.architecture': 'Architecture' + 'models.architecture': 'Architecture', + 'models.search.noresult': 'No related models found', + 'models.search.nofiles': 'No available files', + 'models.search.networkerror': 'Network connection exception!', + 'models.search.hfvisit': 'Please make sure you can visit' }; diff --git a/src/locales/zh-CN/models.ts b/src/locales/zh-CN/models.ts index a7697814..b267e698 100644 --- a/src/locales/zh-CN/models.ts +++ b/src/locales/zh-CN/models.ts @@ -22,11 +22,16 @@ export default { 'models.sort.name': '名称', 'models.sort.size': '大小', 'models.sort.likes': '喜欢', + 'models.sort.trending': '趋势', 'models.sort.downloads': '下载', 'models.sort.updated': '更新时间', 'models.search.result': '{count} 个结果', 'models.data.card': '模型简介', 'models.available.files': '可用文件', 'models.viewin.hf': '在 Hugging Face 中查看', - 'models.architecture': '架构' + 'models.architecture': '架构', + 'models.search.noresult': '未找到相关模型', + 'models.search.nofiles': '无可用文件', + 'models.search.networkerror': '网络连接异常!', + 'models.search.hfvisit': '请确保您可以访问' }; diff --git a/src/pages/llmodels/apis/index.ts b/src/pages/llmodels/apis/index.ts index a9c901df..378f92f3 100644 --- a/src/pages/llmodels/apis/index.ts +++ b/src/pages/llmodels/apis/index.ts @@ -134,6 +134,7 @@ export async function queryHuggingfaceModels( search: { query: string; tags: string[]; + sort?: string; task?: PipelineType; }; }, @@ -143,16 +144,15 @@ export async function queryHuggingfaceModels( for await (const model of listModels({ ...params, ...options, - limit: 500, + limit: 100, additionalFields: ['sha'], fetch(url: string, config: any) { try { - return fetch(url, { + return fetch(`${url}&sort=${params.search.sort}`, { ...config, signal: options.signal }); } catch (error) { - console.log('queryHuggingfaceModels error===', error); // ignore return []; } @@ -170,6 +170,7 @@ export async function queryHuggingfaceModelFiles( const result = []; for await (const fileInfo of listFiles({ ...params, + recursive: true, fetch(url: string, config: any) { try { return fetch(url, { diff --git a/src/pages/llmodels/components/hf-model-file.tsx b/src/pages/llmodels/components/hf-model-file.tsx index c39630e8..46ddd36b 100644 --- a/src/pages/llmodels/components/hf-model-file.tsx +++ b/src/pages/llmodels/components/hf-model-file.tsx @@ -1,5 +1,4 @@ import { convertFileSize } from '@/utils'; -import { SearchOutlined } from '@ant-design/icons'; import { useIntl } from '@umijs/max'; import { Col, Empty, Row, Select, Space, Spin, Tag } from 'antd'; import classNames from 'classnames'; @@ -62,7 +61,11 @@ const HFModelFile: React.FC = (props) => { signal: axiosTokenRef.current.signal } ); - const list = _.filter(res, (file: any) => { + const fileList = _.filter(res, (file: any) => { + return file.type === 'file'; + }); + + const list = _.filter(fileList, (file: any) => { return _.endsWith(file.path, '.gguf') || _.includes(file.path, '.gguf'); }); const sortList = _.sortBy(list, (item: any) => { @@ -182,16 +185,14 @@ const HFModelFile: React.FC = (props) => { })} ) : ( - !dataSource.loading && ( + !dataSource.loading && + !loadingModel && ( - } - description="No files found" + image={Empty.PRESENTED_IMAGE_SIMPLE} + description={intl.formatMessage({ + id: 'models.search.nofiles' + })} /> ) )} diff --git a/src/pages/llmodels/components/model-card.tsx b/src/pages/llmodels/components/model-card.tsx index 02d0510e..9580161c 100644 --- a/src/pages/llmodels/components/model-card.tsx +++ b/src/pages/llmodels/components/model-card.tsx @@ -19,8 +19,9 @@ const ModelCard: React.FC<{ repo: string; onCollapse: (flag: boolean) => void; collapsed: boolean; + loadingModel?: boolean; }> = (props) => { - const { repo, onCollapse, collapsed } = props; + const { repo, onCollapse, collapsed, loadingModel } = props; const intl = useIntl(); const requestSource = useRequestToken(); const [modelData, setModelData] = useState({}); @@ -138,7 +139,7 @@ const ModelCard: React.FC<{ > - README.md + README.md {collapsed ? : } diff --git a/src/pages/llmodels/components/search-model.tsx b/src/pages/llmodels/components/search-model.tsx index 2cbf57c3..fd53decc 100644 --- a/src/pages/llmodels/components/search-model.tsx +++ b/src/pages/llmodels/components/search-model.tsx @@ -1,11 +1,10 @@ -import IconFont from '@/components/icon-font'; import { BulbOutlined } from '@ant-design/icons'; import { useIntl } from '@umijs/max'; import { Button, Input, Select } from 'antd'; import _ from 'lodash'; import React, { useCallback, useEffect, useRef, useState } from 'react'; import { queryHuggingfaceModels } from '../apis'; -import { modelSourceMap, ollamaModelOptions } from '../config'; +import { ModelSortType, modelSourceMap, ollamaModelOptions } from '../config'; import SearchStyle from '../style/search-result.less'; import SearchInput from './search-input'; import SearchResult from './search-result'; @@ -17,21 +16,6 @@ interface SearchInputProps { onSelectModel: (model: any) => void; } -const sourceList = [ - { - label: ( - - ), - value: 'huggingface', - key: 'huggingface' - }, - { - label: , - value: 'ollama_library', - key: 'ollama_library' - } -]; - const SearchModel: React.FC = (props) => { console.log('SearchModel======'); const intl = useIntl(); @@ -39,24 +23,35 @@ const SearchModel: React.FC = (props) => { const [dataSource, setDataSource] = useState<{ repoOptions: any[]; loading: boolean; + networkError: boolean; + sortType: string; }>({ repoOptions: [], - loading: false + loading: false, + networkError: false, + sortType: ModelSortType.trendingScore }); const [current, setCurrent] = useState(''); - const [sortType, setSortType] = useState('downloads'); const cacheRepoOptions = useRef([]); const axiosTokenRef = useRef(null); const customOllamaModelRef = useRef(null); + const searchInputRef = useRef(''); const modelFilesSortOptions = useRef([ - { label: intl.formatMessage({ id: 'models.sort.likes' }), value: 'likes' }, + { + label: intl.formatMessage({ id: 'models.sort.trending' }), + value: ModelSortType.trendingScore + }, + { + label: intl.formatMessage({ id: 'models.sort.likes' }), + value: ModelSortType.likes + }, { label: intl.formatMessage({ id: 'models.sort.downloads' }), - value: 'downloads' + value: ModelSortType.downloads }, { label: intl.formatMessage({ id: 'models.sort.updated' }), - value: 'updatedAt' + value: ModelSortType.lastModified } ]); @@ -66,7 +61,7 @@ const SearchModel: React.FC = (props) => { }, []); const handleOnSearchRepo = useCallback( - async (text: string) => { + async (sortType?: string) => { axiosTokenRef.current?.abort?.(); axiosTokenRef.current = new AbortController(); if (dataSource.loading) return; @@ -79,7 +74,8 @@ const SearchModel: React.FC = (props) => { cacheRepoOptions.current = []; const params = { search: { - query: text, + query: searchInputRef.current || '', + sort: sortType || dataSource.sortType, tags: ['gguf'] } }; @@ -93,22 +89,22 @@ const SearchModel: React.FC = (props) => { label: item.name }; }); - const sortedList = _.sortBy( - list, - (item: any) => item[sortType] - ).reverse(); - cacheRepoOptions.current = sortedList; + + cacheRepoOptions.current = list; setDataSource({ - repoOptions: sortedList, - loading: false + repoOptions: list, + loading: false, + networkError: false, + sortType: sortType || dataSource.sortType }); setLoadingModel?.(false); - handleOnSelectModel(sortedList[0]); - } catch (error) { - console.log('queryHuggingfaceModels error===', error); + handleOnSelectModel(list[0]); + } catch (error: any) { setDataSource({ repoOptions: [], - loading: false + loading: false, + sortType: sortType || dataSource.sortType, + networkError: error?.message === 'Failed to fetch' }); setLoadingModel?.(false); handleOnSelectModel({}); @@ -120,8 +116,8 @@ const SearchModel: React.FC = (props) => { const handlerSearchModels = useCallback( async (e: any) => { - const text = e.target.value; - handleOnSearchRepo(text); + searchInputRef.current = e.target.value; + handleOnSearchRepo(); }, [handleOnSearchRepo] ); @@ -132,12 +128,14 @@ const SearchModel: React.FC = (props) => { !cacheRepoOptions.current.length && modelSource === modelSourceMap.huggingface_value ) { - handleOnSearchRepo(''); + handleOnSearchRepo(); } if (modelSourceMap.ollama_library_value === modelSource) { setDataSource({ repoOptions: ollamaModelOptions, - loading: false + loading: false, + networkError: false, + sortType: dataSource.sortType }); cacheRepoOptions.current = ollamaModelOptions; handleOnSelectModel(ollamaModelOptions[0]); @@ -151,7 +149,9 @@ const SearchModel: React.FC = (props) => { }); setDataSource({ repoOptions: list, - loading: false + loading: false, + networkError: false, + sortType: dataSource.sortType }); }; @@ -164,7 +164,9 @@ const SearchModel: React.FC = (props) => { onSourceChange?.(source); setDataSource({ repoOptions: [], - loading: false + loading: false, + networkError: false, + sortType: dataSource.sortType }); cacheRepoOptions.current = []; }; @@ -186,15 +188,7 @@ const SearchModel: React.FC = (props) => { }; const handleSortChange = (value: string) => { - const sortedList = _.sortBy( - dataSource.repoOptions, - (item: any) => item[value] - ).reverse(); - setSortType(value); - setDataSource({ - repoOptions: sortedList, - loading: false - }); + handleOnSearchRepo(value); }; const renderHFSearch = () => { return ( @@ -210,7 +204,7 @@ const SearchModel: React.FC = (props) => {