refactor: adjust playground pages dir

This commit is contained in:
jialin
2026-03-17 16:35:55 +08:00
committed by jialin
parent 64955c26a7
commit 9447088f89
25 changed files with 262 additions and 317 deletions
@@ -1,733 +0,0 @@
import AlertInfo from '@/components/alert-info';
import ScatterChart from '@/components/echarts/scatter';
import HighlightCode from '@/components/highlight-code';
import IconFont from '@/components/icon-font';
import SealInputNumber from '@/components/seal-form/input-number';
import useOverlayScroller from '@/hooks/use-overlay-scroller';
import useRequestToken from '@/hooks/use-request-token';
import ResizeContainer from '@/pages/_components/terminal-tabs/resize-container';
import {
ClearOutlined,
PlusOutlined,
QuestionCircleOutlined,
SendOutlined
} from '@ant-design/icons';
import { useIntl } from '@umijs/max';
import { Button, Checkbox, Form, Segmented, Spin, Tabs, Tooltip } from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { EMBEDDING_API, handleEmbedding } from '../apis';
import { extractErrorMessage } from '../config';
import { embeddingSamples } from '../config/samples';
import { LLM_METAKEYS } from '../hooks/config';
import useEmbeddingWorker from '../hooks/use-embedding-worker';
import { useInitLLmMeta } from '../hooks/use-init-meta';
import '../style/ground-llm.less';
import '../style/rerank.less';
import { generateEmbeddingCode } from '../view-code/embedding';
import DynamicParams from './dynamic-params';
import FileList from './file-list';
import InputList from './input-list';
import RightContainer from './right-container';
import TokenUsage from './token-usage';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundEmbedding: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const intl = useIntl();
const { workerRef, createWorker, postMessage, terminateWorker } =
useEmbeddingWorker();
const requestSource = useRequestToken();
const [show, setShow] = useState(false);
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [collapse, setCollapse] = useState(false);
const contentRef = useRef<any>('');
const scroller = useRef<any>(null);
const inputListRef = useRef<any>(null);
const messageListLengthCache = useRef<number>(0);
const requestToken = useRef<any>(null);
const [fileList, setFileList] = useState<
{ text: string; name: string; uid: number | string }[]
>([]);
const [outputType, setOutputType] = useState<string>('chart');
const [outputHeight, setOutputHeight] = useState<number>(180);
const [embeddingData, setEmbeddingData] = useState<{
code: string;
copyValue: string;
}>({
code: '',
copyValue: ''
});
const [lessTwoInput, setLessTwoInput] = useState<boolean>(false);
const multiplePasteEnable = useRef<boolean>(true);
const selectionTextRef = useRef<any>(null);
const [textList, setTextList] = useState<
{ text: string; dataUrl?: string; uid: number | string; name: string }[]
>([
{
text: '',
uid: -1,
name: '',
dataUrl: ''
},
{
text: '',
uid: -2,
name: '',
dataUrl: ''
}
]);
const [scatterData, setScatterData] = useState<any[]>([]);
const resizeRef = useRef<any>(null);
const resizeMaxHeight = 400;
const { initialize, updateScrollerPosition: updateDocumentScrollerPosition } =
useOverlayScroller();
const {
handleOnValuesChange,
formRef,
paramsConfig,
initialValues,
parameters,
paramsRef,
modelMeta,
formFields
} = useInitLLmMeta(
{
modelList,
isChat: true
},
{
defaultValues: {},
defaultParamsConfig: [],
metaKeys: LLM_METAKEYS
}
);
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
calculateNewMaxFromBoundary: (maxWidth?: number, maxHeight?: number) => {
resizeRef.current?.container?.calculateNewMaxFromBoundary();
},
collapse: collapse
};
});
const viewCodeContent = useMemo(() => {
console.log('viewCodeContent:', embeddingData.copyValue);
return generateEmbeddingCode({
api: EMBEDDING_API,
parameters: {
..._.pick(parameters, ['model', ..._.split(formFields, ',')]),
input: [
...textList.map((item) => item.text).filter((item) => item),
...fileList.map((item) => item.text).filter((item) => item)
]
}
});
}, [parameters, formFields, textList, fileList]);
const inputEmpty = useMemo(() => {
const list = [...textList, ...fileList];
return list.length < 2;
}, [textList, fileList]);
const setMessageId = () => {
return inputListRef.current?.setMessageId?.();
};
const handleStopConversation = () => {
requestToken.current?.cancel?.();
setLoading(false);
};
const submitMessage = async (current?: { role: string; content: string }) => {
try {
await formRef.current?.form.validateFields();
if (!parameters.model) return;
const validTextList = textList.filter(
(item) => item.text || item.dataUrl
);
const validFileList = fileList.filter((item) => item.text);
const inputList = [
...validTextList.map((item) => item.text || item.dataUrl || ''),
...validFileList.map((item) => item.text || '')
];
if (inputList.length < 2) {
setLessTwoInput(true);
return;
}
setTextList(validTextList);
setFileList(validFileList);
setLessTwoInput(false);
setLoading(true);
setMessageId();
setTokenResult(null);
requestToken.current?.cancel?.();
requestToken.current = requestSource();
contentRef.current = current?.content || '';
const result: any = await handleEmbedding(
{
model: parameters.model,
encoding_format: 'float',
input: inputList
},
{
token: requestToken.current.token
}
);
setTokenResult(result.usage);
const embeddingsList = result.data || [];
createWorker();
workerRef.current!.onmessage = (event: MessageEvent) => {
const { scatterData, embeddingData } = event.data;
setScatterData(scatterData);
setEmbeddingData(embeddingData);
setLoading(false);
};
postMessage({
embeddings: embeddingsList,
textList: textList,
fileList: fileList
});
} catch (error: any) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(error.response)
});
setLoading(false);
}
};
const handleSendMessage = () => {
submitMessage();
};
const handleCloseViewCode = () => {
setShow(false);
};
const handleScaleOutputSize = (
e: any,
direction: string,
ref: any,
d: any
) => {
console.log('handleScaleOutputSize', e, direction, ref, d);
if (
d.height + outputHeight <= resizeMaxHeight &&
d.height + outputHeight >= 180
) {
setOutputHeight(d.height + outputHeight);
}
};
const handleScaleResize = () => {
const height = resizeRef.current?.container?.state?.height;
if (height) {
setOutputHeight(height);
}
};
const handleDeleteFile = (uid: number | string) => {
setFileList((preList) => {
return preList.filter((item) => item.uid !== uid);
});
};
const handleAddText = () => {
inputListRef.current?.handleAdd();
};
const handleTextListChange = (
list: { text: string; uid: number | string; name: string }[]
) => {
setTextList(list);
};
const handleOnUploadImage = (
list: { uid: number | string; dataUrl: string }[],
index: number
) => {
// replace the text with dataUrl in the textList
setTextList((preList) => {
const newList = [...preList];
const current = newList[index];
if (current) {
newList[index] = {
...current,
text: '',
uid: list[0].uid,
dataUrl: list[0].dataUrl
};
}
return newList;
});
};
const handleOnDeleteImage = (item: {
text: string;
uid: number | string;
name: string;
dataUrl?: string;
}) => {
// replace the dataUrl with empty text in the textList
setTextList((preList) => {
const newList = [...preList];
const current = newList.find((i) => i.uid === item.uid);
if (current) {
current.text = '';
current.dataUrl = '';
}
return newList;
});
};
const handleonSelect = useCallback(
(data: {
start: number;
end: number;
beforeText: string;
afterText: string;
index: number;
}) => {
selectionTextRef.current = data;
},
[]
);
const handleOnPaste = useCallback(
(e: any, index: number) => {
if (!multiplePasteEnable.current) return;
const text = e.clipboardData.getData('text');
if (text) {
const dataLlist = text.split('\n').map((item: string) => {
return {
text: item?.trim(),
name: '',
uid: setMessageId()
};
});
dataLlist[0].text = `${selectionTextRef.current?.beforeText || ''}${dataLlist[0].text}${selectionTextRef.current?.afterText || ''}`;
const result = [
...textList.slice(0, index),
...dataLlist,
...textList.slice(index + 1)
].filter((item) => item.text);
setTextList(result);
}
},
[textList]
);
const handleClearDocuments = () => {
setTextList([
{
text: '',
uid: -1,
name: ''
},
{
text: '',
uid: -2,
name: ''
}
]);
setFileList([]);
setScatterData([]);
setTokenResult(null);
setLessTwoInput(false);
setEmbeddingData({
code: '',
copyValue: ''
});
};
const handleOutputTypeChange = (value: string) => {
setOutputType(value);
};
const renderExtra = useMemo(() => {
if (modelMeta?.n_ctx && modelMeta?.n_slot) {
return (
<Form.Item>
<SealInputNumber
disabled
label="Max Tokens"
value={_.floor(_.divide(modelMeta?.n_ctx, modelMeta?.n_slot))}
></SealInputNumber>
</Form.Item>
);
}
return [];
}, [modelMeta]);
const outputItems = useMemo(() => {
return [
{
key: 'chart',
label: 'Chart',
children: (
<ScatterChart
key={collapse ? 'collapse' : 'expand'}
seriesData={scatterData}
height={outputHeight}
width="100%"
xAxisData={[]}
></ScatterChart>
)
},
{
key: 'json',
label: 'JSON',
children: (
<div
style={{
backgroundColor: 'var(--ant-color-bg-container)',
borderRadius: 'var(--border-radius-base)',
overflow: 'hidden'
}}
>
<HighlightCode
height={outputHeight - 32}
code={embeddingData.code}
copyValue={embeddingData.copyValue}
lang="json"
copyable={true}
style={{ marginBottom: 0 }}
></HighlightCode>
</div>
)
}
];
}, [outputHeight, collapse, scatterData, embeddingData]);
const onValuesChange = useCallback(
(changeValues: Record<string, any>, allValues: Record<string, any>) => {
if (changeValues.model) {
setScatterData([]);
setTokenResult(null);
}
handleOnValuesChange(changeValues, allValues);
},
[]
);
useEffect(() => {
if (scroller.current) {
initialize(scroller.current);
}
}, [initialize]);
useEffect(() => {
if (textList.length + fileList.length > messageListLengthCache.current) {
updateDocumentScrollerPosition();
}
messageListLengthCache.current = textList.length + fileList.length;
}, [textList.length, fileList.length]);
useEffect(() => {
if (intl.locale || 'en-US') {
const sample = embeddingSamples[intl.locale];
if (sample) {
setTextList(
sample.map((item: string, index: number) => ({
text: item,
uid: setMessageId(),
name: `Document ${index + 1}`
}))
);
}
}
}, []);
return (
<div className="ground-left-wrapper rerank">
<div className="ground-left">
<div
className="center"
ref={scroller}
style={{ height: 'auto', maxHeight: '100%' }}
>
<div className="documents">
<div className="flex-between m-b-8 doc-header">
<h3 className="m-l-10 flex-between flex-center font-size-14 line-24 m-b-0">
<div className="flex gap-20">
<span>
{intl.formatMessage({
id: 'playground.embedding.documents'
})}
</span>
</div>
</h3>
<div className="flex-center gap-10">
<Tooltip
title={intl.formatMessage({
id: 'playground.input.multiplePaste.tips'
})}
>
<Checkbox
defaultChecked={multiplePasteEnable.current}
onChange={(e: any) => {
multiplePasteEnable.current = e.target.checked;
}}
>
{intl.formatMessage({
id: 'playground.input.multiplePaste'
})}
<QuestionCircleOutlined className="m-l-4" />
</Checkbox>
</Tooltip>
<Button
size="middle"
onClick={handleAddText}
disabled={loading}
>
<PlusOutlined />
{intl.formatMessage({ id: 'playground.embedding.addtext' })}
</Button>
<Button
icon={<ClearOutlined />}
size="middle"
disabled={loading}
onClick={handleClearDocuments}
>
{intl.formatMessage({ id: 'common.button.clear' })}
</Button>
{!loading ? (
<Tooltip
title={
<span>
{intl.formatMessage({ id: 'common.button.submit' })}
</span>
}
>
<Button
size="middle"
type="primary"
disabled={inputEmpty}
onClick={handleSendMessage}
icon={
<SendOutlined rotate={0} className="font-size-14" />
}
style={{ width: 60 }}
></Button>
</Tooltip>
) : (
<Tooltip
title={intl.formatMessage({ id: 'common.button.stop' })}
>
<Button
style={{ width: 60 }}
size="middle"
type="primary"
onClick={handleStopConversation}
icon={
<IconFont
type="icon-stop1"
className="font-size-12"
></IconFont>
}
></Button>
</Tooltip>
)}
</div>
</div>
<div className="docs-wrapper">
<InputList
ref={inputListRef}
textList={textList}
onChange={handleTextListChange}
onSelect={handleonSelect}
onPaste={handleOnPaste}
onUploadImage={handleOnUploadImage}
onDeleteImage={handleOnDeleteImage}
></InputList>
{fileList.length > 0 && (
<div style={{ marginTop: 8 }}>
<FileList
fileList={fileList}
textListCount={textList.length || 0}
onDelete={handleDeleteFile}
></FileList>
</div>
)}
{lessTwoInput && (
<div className="m-t-16">
<AlertInfo
type="danger"
message={intl.formatMessage({
id: 'playground.documents.verify.embedding'
})}
></AlertInfo>
</div>
)}
<TokenUsage
tokenResult={tokenResult}
className="m-t-16"
></TokenUsage>
</div>
</div>
</div>
<div
className="ground-left-footer"
style={{
width: '100%',
padding: '0 32px 16px'
}}
>
<h3 className="m-l-10 flex-between flex-center font-size-14 line-24 m-b-16">
<div className="flex gap-16">
<span className="flex-center">
{intl.formatMessage({ id: 'playground.embedding.output' })}
<Tooltip
title={
<span className="flex-column">
<span>
1.
{intl.formatMessage({
id: 'playground.embedding.pcatips1'
})}
</span>
<span>
2.{' '}
{intl.formatMessage({
id: 'playground.embedding.pcatips2'
})}
</span>
</span>
}
>
<QuestionCircleOutlined className="m-l-4" />
</Tooltip>
</span>
<AlertInfo
type="danger"
style={{ marginRight: 16 }}
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
<Segmented
onChange={handleOutputTypeChange}
value={outputType}
options={[
{
label: intl.formatMessage({
id: 'playground.embedding.chart'
}),
value: 'chart'
},
{ label: 'JSON', value: 'json' }
]}
></Segmented>
</h3>
<div className="embed-chart">
{loading && (
<div
style={{
position: 'absolute',
top: 0,
left: 0,
right: 0,
bottom: 0,
zIndex: 10,
display: 'flex',
justifyContent: 'center',
alignItems: 'center'
}}
>
<Spin spinning={true}></Spin>
</div>
)}
<ResizeContainer
ref={resizeRef}
maxHeight={resizeMaxHeight}
minHeight={180}
defaultHeight={180}
onResize={handleScaleResize}
onResizeStop={handleScaleOutputSize}
>
<div
style={{
border: '1px solid var(--ant-color-border)',
borderRadius: 'var(--border-radius-base)',
width: '100%'
}}
className="scatter"
>
<Tabs
defaultActiveKey={outputType}
activeKey={outputType}
centered
renderTabBar={() => <></>}
items={outputItems}
></Tabs>
</div>
</ResizeContainer>
</div>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={formRef}
onValuesChange={onValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
extra={renderExtra}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundEmbedding;
@@ -1,288 +0,0 @@
import { setRouteCache } from '@/atoms/route-cache';
import AlertInfo from '@/components/alert-info';
import IconFont from '@/components/icon-font';
import routeCachekey from '@/config/route-cachekey';
import ThumbImg from '@/pages/playground/components/thumb-img';
import { generateRandomNumber } from '@/utils';
import { FileImageOutlined } from '@ant-design/icons';
import { useIntl } from '@umijs/max';
import { Button, Tooltip } from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { CREAT_IMAGE_API } from '../apis';
import { useInitImageMeta } from '../hooks/use-init-meta';
import useTextImage from '../hooks/use-text-image';
import '../style/ground-llm.less';
import '../style/system-message-wrap.less';
import { generateImageCode, generateOpenaiImageCode } from '../view-code/image';
import DynamicParams from './dynamic-params';
import MessageInput from './message-input';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundImages: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const intl = useIntl();
const [show, setShow] = useState(false);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const paramsRef = useRef<any>(null);
const inputRef = useRef<any>(null);
const {
handleOnValuesChange,
handleToggleParamsStyle,
setParams,
form,
formFields,
paramsConfig,
initialValues,
parameters,
isOpenaiCompatible
} = useInitImageMeta(props, {
type: 'create'
});
const {
loading,
tokenResult,
imageList,
promptList,
currentPrompt,
setCurrentPrompt,
handleClear,
handleStopConversation,
submitMessage
} = useTextImage({
scroller,
paramsRef,
chunkFields: ['stream_options', 'chunk_results'],
API: CREAT_IMAGE_API
});
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const generateNumber = (min: number, max: number) => {
return Math.floor(Math.random() * (max - min + 1) + min);
};
const handleRandomPrompt = useCallback(() => {
const randomIndex = generateNumber(0, promptList.length - 1);
const randomPrompt = promptList[randomIndex];
inputRef.current?.handleInputChange({
target: {
value: randomPrompt
}
});
}, []);
const finalParameters = useMemo(() => {
if (parameters.size === 'custom') {
return {
..._.omit(parameters, ['width', 'height', 'preview', 'random_seed']),
size:
parameters.width && parameters.height
? `${parameters.width}x${parameters.height}`
: ''
};
}
return {
..._.omit(parameters, ['width', 'height', 'random_seed', 'preview'])
};
}, [parameters]);
const viewCodeContent = useMemo(() => {
if (isOpenaiCompatible) {
return generateOpenaiImageCode({
api: CREAT_IMAGE_API,
parameters: {
...finalParameters,
prompt: currentPrompt
}
});
}
return generateImageCode({
api: CREAT_IMAGE_API,
parameters: {
...finalParameters,
prompt: currentPrompt
}
});
}, [finalParameters, isOpenaiCompatible, currentPrompt]);
const handleInputChange = (e: any) => {
setCurrentPrompt(e.target.value);
};
const generateParams = () => {
const params = {
..._.omitBy(finalParameters, (value: string) => !value),
seed: parameters.random_seed
? generateRandomNumber()
: parameters.seed || null,
stream: false,
prompt: currentPrompt
};
return params;
};
const handleSendMessage = async () => {
try {
await form.current?.form?.validateFields();
if (!parameters.model) return;
const params = generateParams();
console.log('generateParams:', params);
setParams({
...parameters,
seed: params.seed
});
form.current?.form?.setFieldValue('seed', params.seed);
console.log('params:', params, parameters);
setRouteCache(routeCachekey['/playground/text-to-image'], true);
await submitMessage(params);
} catch (error) {
// console.log('error:', error);
} finally {
console.log('finally---------');
setRouteCache(routeCachekey['/playground/text-to-image'], false);
}
};
const handleCloseViewCode = useCallback(() => {
setShow(false);
}, []);
return (
<div className="ground-left-wrapper">
<div className="ground-left">
<div
className="message-list-wrap"
ref={scroller}
style={{ paddingBottom: 16 }}
>
<>
<div className="content" style={{ height: '100%' }}>
<ThumbImg
style={{
padding: 0,
height: '100%',
justifyContent: 'center',
flexDirection: 'column',
flexWrap: 'unset',
alignItems: 'center'
}}
autoBgColor={false}
editable={false}
dataList={imageList}
responseable={true}
gutter={[8, 16]}
autoSize={true}
></ThumbImg>
{!imageList.length && (
<div className="flex-column font-size-14 flex-center gap-20 justify-center hold-wrapper">
<span>
<FileImageOutlined className="font-size-32 text-secondary" />
</span>
<span>
{intl.formatMessage({
id: 'playground.params.empty.tips'
})}
</span>
</div>
)}
</div>
</>
</div>
{tokenResult && (
<div style={{ height: 40 }}>
<AlertInfo
type="danger"
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
)}
<div className="ground-left-footer">
<MessageInput
ref={inputRef}
placeholer={intl.formatMessage({
id: 'playground.input.prompt.holder'
})}
actions={[]}
defaultSize={{
minRows: 5,
maxRows: 5
}}
title={intl.formatMessage({ id: 'playground.image.prompt' })}
loading={loading}
disabled={!parameters.model}
isEmpty={!imageList.length}
handleSubmit={handleSendMessage}
handleAbortFetch={handleStopConversation}
onInputChange={handleInputChange}
shouldResetMessage={false}
clearAll={handleClear}
tools={
<>
<Tooltip
title={intl.formatMessage({
id: 'playground.image.prompt.random'
})}
>
<Button
onClick={handleRandomPrompt}
size="middle"
type="text"
icon={<IconFont type="icon-random"></IconFont>}
></Button>
</Tooltip>
</>
}
/>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={form}
formFields={formFields}
onValuesChange={handleOnValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundImages;
@@ -1,202 +0,0 @@
import { Spin } from 'antd';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { CHAT_API } from '../apis';
import { Roles, generateMessagesByListContent } from '../config';
import { ChatParamsConfig } from '../config/params-config';
import { MessageItem, MessageItemAction } from '../config/types';
import { LLM_METAKEYS, llmInitialValues } from '../hooks/config';
import useChatCompletion from '../hooks/use-chat-completion';
import { useInitLLmMeta } from '../hooks/use-init-meta';
import '../style/ground-llm.less';
import '../style/system-message-wrap.less';
import { generateLLMCode } from '../view-code/llm';
import DynamicParams from './dynamic-params';
import MessageInput from './message-input';
import MessageContent from './multiple-chat/message-content';
import SystemMessage from './multiple-chat/system-message';
import ReferenceParams from './reference-params';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundLeft: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const [systemMessage, setSystemMessage] = useState('');
const [show, setShow] = useState(false);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const [actions, setActions] = useState<MessageItemAction[]>([
'upload',
'delete',
'copy',
'edit'
]);
const {
submitMessage,
handleStopConversation,
handleAddNewMessage,
handleClear,
setMessageList,
tokenResult,
messageList,
loading
} = useChatCompletion(scroller);
const {
handleOnValuesChange,
formRef,
paramsRef,
paramsConfig,
initialValues,
parameters
} = useInitLLmMeta(
{ modelList, isChat: true },
{
defaultValues: {
...llmInitialValues,
model: modelList[0]?.value
},
defaultParamsConfig: ChatParamsConfig,
metaKeys: LLM_METAKEYS
}
);
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const viewCodeContent = useMemo(() => {
const resultList = systemMessage
? [{ role: Roles.System, content: systemMessage }]
: [];
const list = generateMessagesByListContent([...messageList]);
return generateLLMCode({
api: CHAT_API,
parameters: {
...parameters,
messages: [...resultList, ...list]
}
});
}, [messageList, systemMessage, parameters]);
const generateValidMessage = (message: Omit<MessageItem, 'uid'>) => {
if (!message.content && !message.imgs?.length && !message.audio?.length) {
return undefined;
}
return message;
};
const handleSendMessage = async (message: Omit<MessageItem, 'uid'>) => {
const currentMessage = generateValidMessage(message);
submitMessage({
system: systemMessage
? { role: Roles.System, content: systemMessage }
: undefined,
current: currentMessage,
parameters
});
};
const handleCloseViewCode = () => {
setShow(false);
};
return (
<div className="ground-left-wrapper">
<div className="ground-left">
<div className="message-list-wrap" ref={scroller}>
<>
<div
style={{
marginBottom: 20
}}
>
<SystemMessage
style={{
borderRadius: 'var(--border-radius-mini)',
overflow: 'hidden'
}}
systemMessage={systemMessage}
setSystemMessage={setSystemMessage}
></SystemMessage>
</div>
<div className="content">
<MessageContent
messageList={messageList}
setMessageList={setMessageList}
editable={true}
loading={loading}
actions={actions}
/>
{loading && (
<Spin size="small">
<div style={{ height: '46px' }}></div>
</Spin>
)}
</div>
</>
</div>
{tokenResult && !loading && (
<div style={{ height: 40 }}>
<ReferenceParams usage={tokenResult}></ReferenceParams>
</div>
)}
<div className="ground-left-footer">
<MessageInput
defaultSize={{
minRows: 5,
maxRows: 5
}}
actions={['clear', 'layout', 'role', 'upload', 'add', 'paste']}
defaultChecked={false}
loading={loading}
disabled={!parameters.model}
isEmpty={!messageList.length}
handleSubmit={handleSendMessage}
addMessage={handleAddNewMessage}
handleAbortFetch={handleStopConversation}
clearAll={handleClear}
/>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={formRef}
onValuesChange={handleOnValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
showModelSelector={true}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
></ViewCommonCode>
</div>
);
});
export default GroundLeft;
@@ -1,660 +0,0 @@
import AlertInfo from '@/components/alert-info';
import SealInputNumber from '@/components/seal-form/input-number';
import useOverlayScroller from '@/hooks/use-overlay-scroller';
import useRequestToken from '@/hooks/use-request-token';
import {
ClearOutlined,
PlusOutlined,
QuestionCircleOutlined,
SendOutlined
} from '@ant-design/icons';
import { useIntl } from '@umijs/max';
import {
Button,
Checkbox,
Form,
Input,
Spin,
Tag,
Tooltip,
Typography
} from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import styled from 'styled-components';
import { RERANKER_API, rerankerQuery } from '../apis';
import { extractErrorMessage } from '../config';
import { rerankerSamples } from '../config/samples';
import { ParamsSchema } from '../config/types';
import { LLM_METAKEYS } from '../hooks/config';
import { useInitLLmMeta } from '../hooks/use-init-meta';
import useRerankerResponse from '../reranker/hooks/use-reranker-response';
import '../style/ground-llm.less';
import '../style/rerank.less';
import '../style/system-message-wrap.less';
import { generateRerankCode } from '../view-code/rerank';
import DynamicParams from './dynamic-params';
import InputList from './input-list';
import RightContainer from './right-container';
import TokenUsage from './token-usage';
import ViewCommonCode from './view-common-code';
const { Text } = Typography;
const SearchInputWrapper = styled.div`
margin: 16px 32px 10px;
position: relative;
`;
const ValidText = styled(Text)`
position: absolute;
bottom: -20px;
left: 0;
`;
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const fieldConfig: ParamsSchema[] = [
{
type: 'InputNumber',
name: 'top_n',
label: {
text: 'Top N',
isLocalized: false
},
attrs: {
min: 1
},
rules: [
{
required: true,
message: 'Top N is required'
}
]
}
];
const GroundReranker: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const { handleSGlangResponse } = useRerankerResponse();
const intl = useIntl();
const requestSource = useRequestToken();
const [show, setShow] = useState(false);
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const inputListRef = useRef<any>(null);
const messageListLengthCache = useRef<number>(0);
const requestToken = useRef<any>(null);
const multiplePasteEnable = useRef<boolean>(true);
const [isEmptyText, setIsEmptyText] = useState<boolean>(false);
const [fileList, setFileList] = useState<
{
text: string;
name: string;
uid: number | string;
score?: number;
showExtra?: boolean;
percent?: number;
rank?: number;
}[]
>([]);
const [isEmptyQuery, setIsEmptyQuery] = useState<boolean>(false);
const [textList, setTextList] = useState<
{
text: string;
uid: number | string;
name: string;
score?: number;
showExtra?: boolean;
dataUrl?: string;
percent?: number;
rank?: number;
}[]
>([
{
text: '',
uid: -1,
name: '',
dataUrl: ''
},
{
text: '',
uid: -2,
name: '',
dataUrl: ''
}
]);
const [queryValue, setQueryValue] = useState<string>('');
const selectionTextRef = useRef<any>(null);
const { initialize, updateScrollerPosition: updateDocumentScrollerPosition } =
useOverlayScroller();
const {
handleOnValuesChange,
formRef,
paramsConfig,
initialValues,
parameters,
paramsRef,
modelMeta,
formFields
} = useInitLLmMeta(
{
modelList,
isChat: true
},
{
defaultValues: { top_n: 3 },
defaultParamsConfig: fieldConfig,
metaKeys: LLM_METAKEYS
}
);
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const setMessageId = () => {
const uid = inputListRef.current?.setMessageId();
return uid;
};
useEffect(() => {
if (intl.locale || 'en-US') {
const sample = rerankerSamples[intl.locale];
if (sample) {
setTextList(
sample.documents.map((item: string, index: number) => ({
text: item,
uid: setMessageId(),
name: `Document ${index + 1}`,
percent: undefined,
score: undefined,
rank: undefined
}))
);
setQueryValue(sample.query);
}
}
}, []);
const viewCodeContent = useMemo(() => {
return generateRerankCode({
api: RERANKER_API,
parameters: {
..._.pick(parameters, ['model', ..._.split(formFields, ',')]),
query: queryValue,
documents: [...textList, ...fileList]
.map((item) => item.text)
.filter((text) => text)
}
});
}, [parameters, formFields, queryValue, textList, fileList]);
// [0.1, 1.0]
const normalizValue = (data: { min: number; max: number; value: number }) => {
const range = [0.5, 1.0];
const [a, b] = range;
const { min, max, value } = data;
if (isNaN(value) || isNaN(min) || isNaN(max) || min > max) {
return 0;
}
if (min === max) {
return 100;
}
const res = a + ((value - min) * (b - a)) / (max - min);
return res * 100;
};
const renderPercent = (data: any) => {
if (!data.showExtra || !data.percent) {
return null;
}
const percent = data.percent;
return (
<div className="rank-wrapper">
<div className="percent-wrapper">
<div
className="pregress-bar"
style={{
backgroundImage: `linear-gradient(90deg, var(--ant-blue-5) 0%, var(--ant-blue-2) 100%)`,
width: `${percent}%`,
height: '4px',
borderRadius: '2px'
}}
></div>
</div>
<span className="flex-center hover-hidden rank-tag">
<Tag color={'geekblue'} variant="filled">
{intl.formatMessage({ id: 'playground.rerank.rank' })}: {data.rank}
</Tag>
<Tag color={'gold'} variant="filled">
{intl.formatMessage({ id: 'playground.rerank.score' })}:{' '}
{_.round(data.score, 2)}
</Tag>
</span>
</div>
);
};
const submitMessage = async (query: string) => {
try {
setIsEmptyQuery(!queryValue);
setTokenResult(null);
await formRef.current?.form.validateFields();
if (!parameters.model || !queryValue) return;
const documentList: any[] = [...textList, ...fileList];
const validDocus = documentList.filter((item) => item.text);
if (!validDocus.length) {
setIsEmptyText(true);
return;
}
setIsEmptyText(false);
setLoading(true);
setMessageId();
requestToken.current?.cancel?.();
requestToken.current = requestSource();
const filledList = textList.filter((item) => item.text);
setTextList(filledList);
const res: any = await rerankerQuery(
{
model: parameters.model,
top_n: parameters.top_n,
query: query,
documents: filledList.map((item) => item.text)
},
{
token: requestToken.current.token
}
);
// detect response type
const result = Array.isArray(res) ? handleSGlangResponse(res) : res;
setMessageId();
setTokenResult(result.usage);
const sortList = _.sortBy(
result.results || [],
(item: any) => item.relevance_score
);
const maxValue = sortList[sortList.length - 1].relevance_score;
const minValue = sortList[0].relevance_score;
// reset state
let newTextList = filledList.map((item) => {
item.percent = undefined;
item.score = undefined;
item.rank = undefined;
return item;
});
result.results?.forEach((item: any, sIndex: number) => {
newTextList[item.index] = {
...newTextList[item.index],
uid: setMessageId(),
rank: sIndex + 1,
score: item.relevance_score,
showExtra: true,
percent: normalizValue({
min: minValue,
max: maxValue,
value: item.relevance_score
})
};
});
newTextList = _.sortBy(newTextList, 'rank');
setTextList(newTextList);
} catch (error: any) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(error.response)
});
} finally {
setLoading(false);
}
};
const handleSearch = (val: string, event: any, action: any) => {
if (action.source === 'clear') {
return;
}
submitMessage(val);
};
const handleQueryChange = (e: any) => {
setQueryValue(e.target.value);
setIsEmptyQuery(!e.target.value);
};
const handleCloseViewCode = () => {
setShow(false);
};
const handleAddText = () => {
inputListRef.current?.handleAdd();
};
const handleTextListChange = (
list: { text: string; uid: number | string; name: string }[]
) => {
const newList = list?.map((item: any) => {
item.percent = undefined;
item.score = undefined;
item.rank = undefined;
return item;
});
setTextList(newList);
};
const handleonSelect = (data: {
start: number;
end: number;
beforeText: string;
afterText: string;
index: number;
}) => {
selectionTextRef.current = data;
};
const handleOnUploadImage = (
list: { uid: number | string; dataUrl: string }[],
index: number
) => {
// replace the text with dataUrl in the textList
setTextList((preList) => {
const newList = [...preList];
const current = newList[index];
if (current) {
newList[index] = {
...current,
text: '',
uid: list[0].uid,
dataUrl: list[0].dataUrl
};
}
return newList;
});
};
const handleOnDeleteImage = (item: {
text: string;
uid: number | string;
name: string;
dataUrl?: string;
}) => {
// replace the dataUrl with empty text in the textList
setTextList((preList) => {
const newList = [...preList];
const current = newList.find((i) => i.uid === item.uid);
if (current) {
current.text = '';
current.dataUrl = '';
}
return newList;
});
};
const handleOnPaste = (e: any, index: number) => {
if (!multiplePasteEnable.current) return;
const text = e.clipboardData.getData('text');
if (text) {
const dataLlist = text.split('\n').map((item: string) => {
return {
text: item?.trim(),
name: '',
uid: setMessageId(),
percent: undefined,
score: undefined,
rank: undefined
};
});
dataLlist[0].text = `${selectionTextRef.current?.beforeText || ''}${dataLlist[0].text}${selectionTextRef.current?.afterText || ''}`;
const result = [
...textList.slice(0, index),
...dataLlist,
...textList.slice(index + 1)
].filter((item) => item.text);
setTextList(result);
}
};
const renderExtra = useMemo(() => {
if (modelMeta?.n_ctx && modelMeta?.n_slot) {
return (
<Form.Item>
<SealInputNumber
disabled
label="Max Tokens"
value={_.floor(_.divide(modelMeta?.n_ctx, modelMeta?.n_slot))}
></SealInputNumber>
</Form.Item>
);
}
return null;
}, [modelMeta]);
const handleClearDocuments = () => {
setTextList([
{
text: '',
uid: setMessageId(),
name: ''
},
{
text: '',
uid: setMessageId(),
name: ''
}
]);
setFileList([]);
setTokenResult(null);
setIsEmptyText(false);
};
const onValuesChange = useCallback((changedValues: any, allValues: any) => {
if (changedValues.model) {
setTokenResult(null);
}
handleOnValuesChange(changedValues, allValues);
}, []);
useEffect(() => {
if (scroller.current) {
initialize(scroller.current);
}
}, [initialize]);
useEffect(() => {
if (textList.length + fileList.length > messageListLengthCache.current) {
updateDocumentScrollerPosition();
}
messageListLengthCache.current = textList.length + fileList.length;
}, [textList.length, fileList.length]);
return (
<div className="ground-left-wrapper rerank">
<div className="ground-left">
<div className="ground-left-footer">
<h3
className="m-l-10 flex-between flex-center font-size-14 line-24 m-b-0"
style={{ padding: '0 32px', marginTop: 16 }}
>
<span>{intl.formatMessage({ id: 'playground.rerank.query' })}</span>
</h3>
<SearchInputWrapper>
<Input.Search
allowClear
value={queryValue}
onSearch={handleSearch}
onChange={handleQueryChange}
enterButton={
<Tooltip
title={
<span>
{intl.formatMessage({ id: 'common.button.submit' })}
</span>
}
>
<span className="full-wrap">
<SendOutlined rotate={0} className="font-size-14" />
</span>
</Tooltip>
}
placeholder={intl.formatMessage({
id: 'playground.rerank.query.holder'
})}
></Input.Search>
{isEmptyQuery && (
<ValidText type="danger">
{intl.formatMessage({ id: 'playground.rerank.query.validate' })}
</ValidText>
)}
</SearchInputWrapper>
</div>
<div className="center" ref={scroller} style={{ marginTop: 16 }}>
<div className="documents">
<div
className="flex-between m-b-8 doc-header"
style={{ marginTop: 0 }}
>
<h3 className="m-l-10 flex-between flex-center font-size-14 line-24 m-b-0">
<span>
{intl.formatMessage({ id: 'playground.embedding.documents' })}
</span>
</h3>
<div className="flex-center gap-10">
<Tooltip
title={intl.formatMessage({
id: 'playground.input.multiplePaste.tips'
})}
>
<Checkbox
defaultChecked={multiplePasteEnable.current}
onChange={(e: any) => {
multiplePasteEnable.current = e.target.checked;
}}
>
{intl.formatMessage({
id: 'playground.input.multiplePaste'
})}
<QuestionCircleOutlined className="m-l-4" />
</Checkbox>
</Tooltip>
<Button size="middle" onClick={handleAddText}>
<PlusOutlined />
{intl.formatMessage({ id: 'playground.embedding.addtext' })}
</Button>
<Button
icon={<ClearOutlined />}
size="middle"
disabled={loading}
onClick={handleClearDocuments}
>
{intl.formatMessage({ id: 'common.button.clear' })}
</Button>
</div>
</div>
<div className="docs-wrapper">
<InputList
ref={inputListRef}
textList={textList}
showLabel={false}
height={46}
onChange={handleTextListChange}
extra={renderPercent}
onSelect={handleonSelect}
onUploadImage={handleOnUploadImage}
onPaste={handleOnPaste}
onDeleteImage={handleOnDeleteImage}
></InputList>
{isEmptyText && (
<div className="m-t-16">
<AlertInfo
type="danger"
message={intl.formatMessage({
id: 'playground.documents.verify.rerank'
})}
></AlertInfo>
</div>
)}
<TokenUsage
tokenResult={tokenResult}
className="m-t-16"
></TokenUsage>
</div>
</div>
<div></div>
<div
className="message-list-wrap"
style={{ paddingInline: 0, paddingTop: 0 }}
>
<>
<div className="content">
{loading && (
<Spin size="small">
<div style={{ height: '46px' }}></div>
</Spin>
)}
</div>
</>
</div>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={formRef}
onValuesChange={onValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
extra={renderExtra}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundReranker;
@@ -1,521 +0,0 @@
import { setRouteCache } from '@/atoms/route-cache';
import AlertInfo from '@/components/alert-info';
import AudioAnimation from '@/components/audio-animation';
import AudioPlayer from '@/components/audio-player';
import CopyButton from '@/components/copy-button';
import IconFont from '@/components/icon-font';
import UploadAudio from '@/components/upload-audio';
import routeCachekey from '@/config/route-cachekey';
import { HEADER_HEIGHT } from '@/config/settings';
import useOverlayScroller from '@/hooks/use-overlay-scroller';
import { useCancelToken } from '@/hooks/use-request-token';
import { readAudioFile } from '@/utils/load-audio-file';
import { SendOutlined } from '@ant-design/icons';
import { useIntl, useSearchParams } from '@umijs/max';
import { Button, Spin, Tooltip } from 'antd';
import _ from 'lodash';
import React, {
forwardRef,
useCallback,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { AUDIO_SPEECH_TO_TEXT_API, speechToText } from '../apis';
import {
SpeechToTextFormat,
defaultLanguages,
extractErrorMessage
} from '../config';
import { allLanguages } from '../config/languages';
import { RealtimeParamsConfig as paramsConfig } from '../config/params-config';
import { ParamsSchema } from '../config/types';
import '../style/ground-llm.less';
import '../style/speech-to-text.less';
import '../style/system-message-wrap.less';
import { speechToTextCode } from '../view-code/audio';
import AudioInput from './audio-input';
import DynamicParams from './dynamic-params';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
ref?: any;
}
const GroundSTT: React.FC<MessageProps> = forwardRef((props, ref) => {
const intl = useIntl();
const { modelList } = props;
const messageId = useRef<number>(0);
const [messageList, setMessageList] = useState<
{ uid: number; content: string }[]
>([]);
const [searchParams] = useSearchParams();
const modelType = searchParams.get('type') || '';
const selectModel = searchParams.get('model')
? modelType === 'stt' && searchParams.get('model')
: '';
const defaultModel = selectModel || modelList[0]?.value || '';
const [parameters, setParams] = useState<any>({
model: defaultModel,
language: 'auto'
});
const [show, setShow] = useState(false);
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const paramsRef = useRef<any>(null);
const [audioPermissionOn, setAudioPermissionOn] = useState(true);
const [audioData, setAudioData] = useState<any>(null);
const [audioChunks, setAudioChunks] = useState<any>({
data: [],
analyser: null
});
const [isRecording, setIsRecording] = useState(false);
const formRef = useRef<any>(null);
const { updateCancelToken, getCanceltToken, cancelRequest } =
useCancelToken();
const { initialize, updateScrollerPosition } = useOverlayScroller();
const { initialize: innitializeParams } = useOverlayScroller();
const [modelMeta, setModelMeta] = useState<any>(null);
const [fieldsConfig, setFieldsConfig] =
useState<ParamsSchema[]>(paramsConfig);
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const setMessageId = () => {
messageId.current = messageId.current + 1;
};
const viewCodeContent = useMemo(() => {
return speechToTextCode({
api: AUDIO_SPEECH_TO_TEXT_API,
parameters: {
...parameters
}
});
}, [parameters]);
const handleStopConversation = () => {
cancelRequest();
setLoading(false);
};
const submitMessage = async () => {
try {
await formRef.current?.form.validateFields();
if (!parameters.model) return;
setLoading(true);
setMessageId();
setTokenResult(null);
setMessageList([]);
updateCancelToken();
setRouteCache(routeCachekey['/playground/speech'], true);
const params = {
...parameters,
file: new File([audioData.data], audioData.name, {
type: audioData.type
})
};
const result: any = await speechToText(
{
data: params
},
{
cancelToken: getCanceltToken()
}
);
if (
(result?.status_code && result?.status_code !== 200) ||
result?.error
) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(result)
});
return;
}
setMessageList([
{
content: result.text,
uid: messageId.current
}
]);
} catch (error: any) {
console.log('error:', error);
const res = error?.response?.data;
if (res?.error || (res?.status_code && res?.status_code !== 200)) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(res)
});
}
} finally {
setLoading(false);
setIsRecording(false);
setRouteCache(routeCachekey['/playground/speech'], false);
}
};
const handleClear = () => {
setMessageId();
setMessageList([]);
setTokenResult(null);
};
const handleCloseViewCode = () => {
setShow(false);
};
const handleOnAudioData = useCallback(
(data: {
chunks: Blob[];
url: string;
name: string;
duration: number;
type: string;
}) => {
setAudioData(() => {
return {
url: data.url,
name: data.name,
data: data.chunks,
type: data.type,
duration: data.duration
};
});
},
[]
);
const handleOnAudioPermission = useCallback((permission: boolean) => {
setAudioPermissionOn(permission);
}, []);
const handleUploadChange = useCallback(
async (data: { file: any; fileList: any }) => {
try {
const res = await readAudioFile(data.file);
setAudioData(res);
setTokenResult(null);
} catch (error) {}
},
[]
);
const handleOnAnalyse = useCallback((data: any, analyser: any) => {
setAudioChunks((pre: any) => {
return {
data: data,
analyser: analyser
};
});
}, []);
const handleOnRecord = useCallback((val: boolean) => {
setIsRecording(val);
setAudioData(null);
setTokenResult(null);
setMessageList([]);
}, []);
const handleOnGenerate = async () => {
if (loading) {
handleStopConversation();
return;
}
submitMessage();
};
const renderAniamtion = () => {
if (!audioPermissionOn) {
return (
<div className="tips-text">
<IconFont type={'icon-audio'} style={{ fontSize: 20 }}></IconFont>
<span>
{intl.formatMessage({ id: 'playground.audio.enablemic' })}
</span>
</div>
);
}
if (isRecording) {
return (
<AudioAnimation
fixedHeight={true}
height={82}
width={500}
analyserData={audioChunks}
></AudioAnimation>
);
}
return (
<div className="tips-text">
<IconFont type={'icon-audio'} style={{ fontSize: 18 }}></IconFont>
<span>
{intl.formatMessage({ id: 'playground.audio.speechtotext.tips' })}
</span>
</div>
);
};
const handleSelectModel = (model: string) => {
if (!model) return;
const selected = modelList.find((item) => item.value === model);
setModelMeta(selected?.meta || {});
const languages = selected?.meta?.languages || [];
let currentLanguage = [...defaultLanguages];
if (languages.length > 0) {
// sort languages based on the order in the model meta
currentLanguage = [];
languages.forEach((langCode: string) => {
const langItem = allLanguages.find((item) => item.value === langCode);
if (langItem) {
currentLanguage.push(langItem);
}
});
const newConfig = paramsConfig.map((item) => {
const oItem = _.cloneDeep(item);
if (item.name === 'language') {
return {
...oItem,
options: currentLanguage
};
}
return oItem;
});
setFieldsConfig(newConfig);
}
setParams((pre: any) => {
return {
...pre,
language:
selected?.meta?.language || currentLanguage[0]?.value || 'auto',
model: model
};
});
};
const handleOnValuesChange = (changedValues: any, allValues: any) => {
if (changedValues.model) {
handleSelectModel(changedValues.model);
} else {
setParams(allValues);
}
};
useEffect(() => {
if (scroller.current) {
initialize(scroller.current);
}
}, [initialize]);
useEffect(() => {
if (paramsRef.current) {
innitializeParams(paramsRef.current);
}
}, [innitializeParams]);
useEffect(() => {
if (loading) {
updateScrollerPosition();
}
}, [messageList, loading]);
useEffect(() => {
const defaultModel = selectModel || modelList[0]?.value || '';
handleSelectModel(defaultModel);
}, [modelList, selectModel]);
return (
<div
className="ground-left-wrapper"
style={{
height: `calc(100vh - ${HEADER_HEIGHT}px)`
}}
>
<div className="ground-left">
<div className="ground-left-footer" style={{ flex: 1 }}>
<div className="speech-to-text">
<div className="speech-box">
{!isRecording && (
<UploadAudio
type="default"
accept={SpeechToTextFormat.join(', ')}
onChange={handleUploadChange}
></UploadAudio>
)}
<AudioInput
type="default"
voiceActivity={true}
onAudioData={handleOnAudioData}
onAudioPermission={handleOnAudioPermission}
onAnalyse={handleOnAnalyse}
onRecord={handleOnRecord}
></AudioInput>
</div>
{audioData ? (
<div className="flex-between flex-center justify-center relative">
<div style={{ width: 600 }}>
<AudioPlayer
url={audioData.url}
name={audioData.name}
duration={audioData.duration}
extra={
<Tooltip
title={
loading
? intl.formatMessage({
id: 'common.button.stop'
})
: intl.formatMessage({
id: 'playground.audio.button.generate'
})
}
>
{
<Button
disabled={!audioData}
type="primary"
size="middle"
shape="circle"
onClick={handleOnGenerate}
icon={
loading ? (
<IconFont
type="icon-stop1"
className="font-size-14"
></IconFont>
) : (
<SendOutlined></SendOutlined>
)
}
></Button>
}
</Tooltip>
}
></AudioPlayer>
</div>
</div>
) : (
renderAniamtion()
)}
</div>
</div>
<div
style={{
flex: 1,
display: 'flex',
flexDirection: 'column',
justifyContent: 'space-between',
overflow: 'auto'
}}
>
<div
className="message-list-wrap"
style={{
flex: 1,
position: 'relative'
}}
>
{messageList?.length > 0 && (
<span
style={{
position: 'absolute',
top: 20,
right: 32,
zIndex: 10
}}
>
<CopyButton
text={messageList[0]?.content}
type="link"
></CopyButton>
</span>
)}
<div
className="content"
style={{ height: '100%', overflow: 'auto' }}
ref={scroller}
>
<div>
{!tokenResult && (
<div
style={{
padding: '8px 14px',
lineHeight: '20px',
display: 'flex',
justifyContent: 'center',
wordBreak: 'break-word'
}}
>
{messageList.length ? (
messageList[0]?.content
) : (
<span className="text-tertiary">
{intl.formatMessage({
id: 'playground.audio.generating.tips'
})}
</span>
)}
</div>
)}
{tokenResult && (
<div style={{ height: 40 }}>
<AlertInfo
type="danger"
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
)}
</div>
</div>
{loading && (
<div style={{ width: '100%', flex: 1 }}>
<Spin size="small">
<div style={{ height: '46px' }}></div>
</Spin>
</div>
)}
</div>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={formRef}
onValuesChange={handleOnValuesChange}
paramsConfig={fieldsConfig}
initialValues={parameters}
modelList={modelList}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundSTT;
@@ -1,515 +0,0 @@
import { setRouteCache } from '@/atoms/route-cache';
import AlertInfo from '@/components/alert-info';
import IconFont from '@/components/icon-font';
import AutoComplete from '@/components/seal-form/auto-complete';
import FieldComponent from '@/components/seal-form/field-component';
import SealSelect from '@/components/seal-form/seal-select';
import SpeechContent from '@/components/speech-content';
import routeCachekey from '@/config/route-cachekey';
import useOverlayScroller from '@/hooks/use-overlay-scroller';
import CollapsePanel from '@/pages/_components/collapse-panel';
import { getLocale, useIntl, useSearchParams } from '@umijs/max';
import { Form, Spin } from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { AUDIO_TEXT_TO_SPEECH_API, CHAT_API, textToSpeech } from '../apis';
import { RefAudioFormItem } from '../audio/form';
import {
TTSParamsConfig as paramsConfig,
TTSAdvancedParamsConfig
} from '../audio/params-config';
import { extractErrorMessage } from '../config';
import { MessageItem, ParamsSchema } from '../config/types';
import '../style/ground-llm.less';
import '../style/system-message-wrap.less';
import { TextToSpeechCode } from '../view-code/audio';
import DynamicParams from './dynamic-params';
import MessageInput from './message-input';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
const MetaFields = [
'task_type',
'language',
'instructions',
'max_new_tokens',
'ref_audio'
];
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundTTS: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const messageId = useRef<number>(0);
const [messageList, setMessageList] = useState<
{
input: string;
voice: string;
format: string;
speed: number;
uid: number;
autoplay: boolean;
audioUrl: string;
}[]
>([]);
const locale = getLocale();
const intl = useIntl();
const [searchParams] = useSearchParams();
const modelType = searchParams.get('type') || '';
const selectModel = searchParams.get('model')
? modelType === 'tts' && searchParams.get('model')
: '';
const [parameters, setParams] = useState<any>({
model: selectModel,
voice: '',
response_format: 'mp3'
});
const [show, setShow] = useState(false);
const [loading, setLoading] = useState(false);
const [tokenResult, setTokenResult] = useState<any>(null);
const [collapse, setCollapse] = useState(false);
const controllerRef = useRef<any>(null);
const scroller = useRef<any>(null);
const paramsRef = useRef<any>(null);
const checkvalueRef = useRef<any>(true);
const [currentPrompt, setCurrentPrompt] = useState<string>('');
const [voiceDataList, setVoiceList] = useState<Global.BaseOption<string>[]>(
[]
);
const [modelMeta, setModelMeta] = useState<any>({});
const formRef = useRef<any>(null);
const { initialize } = useOverlayScroller();
const { initialize: innitializeParams } = useOverlayScroller();
const [activeKey, setActiveKey] = useState<string | string[]>(
'advanced_config'
);
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const defaultModel = useMemo(() => {
return selectModel || modelList[0]?.value || '';
}, [modelList]);
const dropEmptyFields = (parameters: Record<string, any>) => {
const fields = [
'task_type',
'instructions',
'max_new_tokens',
'ref_audio',
'ref_text',
'language',
'x_vector_only_mode'
];
const newParams = { ...parameters };
return _.omitBy(newParams, (value: any, key: string) => {
return fields.includes(key) && !value;
});
};
const viewCodeContent = useMemo(() => {
return TextToSpeechCode({
api: AUDIO_TEXT_TO_SPEECH_API,
parameters: {
...dropEmptyFields(parameters),
input: currentPrompt
}
});
}, [parameters, currentPrompt]);
const sortVoiceList = useCallback(
(locale: string, voiceDataList: Global.BaseOption<string>[]) => {
const lang = locale === 'en-US' ? 'english' : 'chinese';
const list = voiceDataList.sort((a, b) => {
const aContains = a.value.toLowerCase().includes(lang) ? 1 : 0;
const bContains = b.value.toLowerCase().includes(lang) ? 1 : 0;
return bContains - aContains;
});
return list;
},
[]
);
const voiceList = useMemo(() => {
if (!voiceDataList.length) return [];
const newList = sortVoiceList(locale, voiceDataList);
return newList;
}, [locale, voiceDataList, sortVoiceList]);
useEffect(() => {
const newList = sortVoiceList(locale, voiceDataList);
setParams((pre: any) => {
return {
...pre,
voice: newList[0]?.value
};
});
formRef.current?.form.setFieldValue('voice', newList[0]?.value);
}, [locale, voiceDataList, sortVoiceList]);
const setMessageId = () => {
messageId.current = messageId.current + 1;
};
const handleStopConversation = () => {
controllerRef.current?.abort?.();
setLoading(false);
};
const handleInputChange = (e: any) => {
setCurrentPrompt(e.target.value);
};
const submitMessage = async (current?: { role: string; content: string }) => {
try {
await formRef.current?.form.validateFields();
if (!parameters.model) return;
setLoading(true);
setMessageId();
setTokenResult(null);
setCurrentPrompt(current?.content || '');
setMessageList([]);
setRouteCache(routeCachekey['/playground/speech'], true);
controllerRef.current?.abort?.();
controllerRef.current = new AbortController();
const signal = controllerRef.current.signal;
const params = {
...dropEmptyFields(parameters),
input: current?.content || currentPrompt
};
const res: any = await textToSpeech({
data: params,
url: CHAT_API,
signal
});
setParams(params);
console.log('result:', res);
if ((res?.status_code && res?.status_code !== 200) || res?.error) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(res)
});
setMessageList([]);
return;
}
setMessageList([
{
input: current?.content || currentPrompt,
voice: parameters.voice,
format: parameters.response_format,
speed: parameters.speed,
uid: messageId.current,
autoplay: checkvalueRef.current,
audioUrl: res.url
}
]);
} catch (error: any) {
const res = error?.response?.data;
console.log('error:', error);
if (res?.error) {
setTokenResult({
error: true,
errorMessage: extractErrorMessage(res)
});
}
} finally {
setLoading(false);
setRouteCache(routeCachekey['/playground/speech'], false);
}
};
const handleClear = () => {
setMessageId();
setMessageList([]);
setTokenResult(null);
};
const handleSendMessage = (message: Omit<MessageItem, 'uid'>) => {
submitMessage(message);
};
const handleCloseViewCode = () => {
setShow(false);
};
const handleSelectModel = async (value: string) => {
if (!value) {
return;
}
const model = modelList.find((item) => item.value === value);
const list = _.map(model?.meta?.voices || [], (item: any) => {
return {
label: item,
value: item
};
});
const newList = sortVoiceList(locale, list);
setVoiceList(newList);
setModelMeta(model?.meta || {});
setParams((pre: any) => {
return {
...pre,
..._.pick(model?.meta || {}, MetaFields),
task_type: model?.meta?.task_type,
model: value,
language: model?.meta?.languages?.[0] || '',
voice: newList[0]?.value
};
});
};
const handleOnValuesChange = useCallback(
(changeValues: Record<string, any>, allValues: Record<string, any>) => {
if (changeValues.model) {
handleSelectModel(changeValues.model);
} else {
setParams(allValues);
}
},
[handleSelectModel]
);
useEffect(() => {
if (paramsRef.current) {
innitializeParams(paramsRef.current);
}
}, [innitializeParams]);
const handleOnCheckChange = (e: any) => {
checkvalueRef.current = e.target.checked;
};
const handleOnCollapse = (keys: string | string[]) => {
setActiveKey(keys);
};
const renderAdvancedFields = () => {
const formItems = TTSAdvancedParamsConfig.map((item: ParamsSchema) => {
const comProps = {
...item.attrs,
label: item.label.isLocalized
? intl.formatMessage({ id: item.label.text })
: item.label.text
};
return (
<>
<Form.Item
name={item.name}
rules={item.rules}
key={item.name}
{...item.formItemAttrs}
>
<FieldComponent
{...comProps}
description={
item.description?.isLocalized
? intl.formatMessage({ id: item.description.text })
: item.description?.text
}
onChange={null}
{..._.omit(item, [
'name',
'rules',
'disabledConfig',
'description'
])}
{...item.initAttrs?.(modelMeta)}
></FieldComponent>
</Form.Item>
</>
);
});
return (
<CollapsePanel
activeKey={activeKey}
onChange={handleOnCollapse}
accordion={false}
items={[
{
key: 'advanced_config',
label: intl.formatMessage({ id: 'resources.form.advanced' }),
forceRender: true,
children: (
<>
{formItems}
<RefAudioFormItem />
</>
)
}
]}
></CollapsePanel>
);
};
const renderExtra = () => {
return paramsConfig.map((item: ParamsSchema) => {
const comProps = {
...item.attrs,
options: item.name === 'voice' ? voiceList : item.options,
label: item.label.isLocalized
? intl.formatMessage({ id: item.label.text })
: item.label.text
};
return (
<>
<Form.Item name={item.name} rules={item.rules} key={item.name}>
{item.type === 'AutoComplete' ? (
<AutoComplete {...comProps} />
) : (
<SealSelect {...comProps}></SealSelect>
)}
</Form.Item>
</>
);
});
};
useEffect(() => {
if (defaultModel && modelList.length) {
handleSelectModel(defaultModel);
}
}, [defaultModel, modelList.length]);
useEffect(() => {
if (scroller.current) {
initialize(scroller.current);
}
}, [initialize]);
useEffect(() => {
if (paramsRef.current) {
innitializeParams(paramsRef.current);
}
}, [innitializeParams]);
return (
<div className="ground-left-wrapper">
<div className="ground-left">
<div className="message-list-wrap">
<div
style={{
height: '100%',
display: 'flex',
justifyContent: 'center',
alignItems: 'center'
}}
>
<div className="content" style={{ maxWidth: 1000 }}>
{messageList.length ? (
<SpeechContent dataList={messageList} loading={loading} />
) : (
<div className="flex-column font-size-14 flex-center gap-20">
<span>
<IconFont
type="icon-audio "
className="font-size-32 text-secondary"
></IconFont>
</span>
<span>
{intl.formatMessage({
id: 'playground.audio.texttospeech.tips'
})}
</span>
</div>
)}
{loading && (
<Spin size="small">
<div style={{ height: '46px' }}></div>
</Spin>
)}
</div>
</div>
</div>
{tokenResult && (
<div style={{ height: 40 }}>
<AlertInfo
type="danger"
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
)}
<div className="ground-left-footer">
<MessageInput
actions={['check']}
checkLabel={intl.formatMessage({
id: 'playground.toolbar.autoplay'
})}
placeholer={intl.formatMessage({
id: 'playground.input.text.holder'
})}
defaultSize={{
minRows: 5,
maxRows: 5
}}
title={intl.formatMessage({ id: 'playground.audio.textinput' })}
onCheck={handleOnCheckChange}
loading={loading}
disabled={!parameters.model}
isEmpty={true}
handleSubmit={handleSendMessage}
handleAbortFetch={handleStopConversation}
onInputChange={handleInputChange}
clearAll={handleClear}
shouldResetMessage={false}
/>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={formRef}
meta={modelMeta}
onValuesChange={handleOnValuesChange}
initialValues={parameters}
modelList={modelList}
extra={[
<>
{renderExtra()}
{renderAdvancedFields()}
</>
]}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundTTS;
@@ -1,263 +0,0 @@
import { setRouteCache } from '@/atoms/route-cache';
import AlertInfo from '@/components/alert-info';
import routeCachekey from '@/config/route-cachekey';
import { VideoCameraOutlined } from '@ant-design/icons';
import { useIntl } from '@umijs/max';
import { Spin } from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { CREATE_VIDEO_API } from '../apis';
import { useInitVideoMeta } from '../hooks/use-init-video-meta';
import useTextVideo from '../hooks/use-text-video';
import '../style/ground-llm.less';
import '../style/system-message-wrap.less';
import { generateCode } from '../view-code/video';
import DynamicParams from './dynamic-params';
import MessageInput from './message-input';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundVideo: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const intl = useIntl();
const [show, setShow] = useState(false);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const paramsRef = useRef<any>(null);
const inputRef = useRef<any>(null);
const {
handleOnValuesChange,
handleToggleParamsStyle,
setParams,
form,
formFields,
paramsConfig,
initialValues,
parameters,
isOpenaiCompatible
} = useInitVideoMeta(props, {
type: 'create'
});
const {
loading,
tokenResult,
videoList,
promptList,
currentPrompt,
setCurrentPrompt,
handleClear,
handleStopConversation,
submitMessage
} = useTextVideo({
scroller,
paramsRef,
API: CREATE_VIDEO_API
});
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const generateNumber = (min: number, max: number) => {
return Math.floor(Math.random() * (max - min + 1) + min);
};
const handleRandomPrompt = useCallback(() => {
const randomIndex = generateNumber(0, promptList.length - 1);
const randomPrompt = promptList[randomIndex];
inputRef.current?.handleInputChange({
target: {
value: randomPrompt
}
});
}, []);
const finalParameters = useMemo(() => {
if (parameters.size === 'custom') {
return {
..._.omit(parameters, ['width', 'height', 'random_seed', 'seed']),
size:
parameters.width && parameters.height
? `${parameters.width}x${parameters.height}`
: ''
};
}
return {
..._.omit(parameters, ['width', 'height', 'random_seed', 'seed'])
};
}, [parameters]);
const viewCodeContent = useMemo(() => {
return generateCode({
api: CREATE_VIDEO_API,
isFormdata: true,
parameters: {
...finalParameters,
prompt: currentPrompt
}
});
}, [finalParameters, isOpenaiCompatible, currentPrompt]);
const handleInputChange = (e: any) => {
setCurrentPrompt(e.target.value);
};
const generateParams = () => {
const params = {
..._.omitBy(finalParameters, (value: string) => !value),
prompt: currentPrompt
};
return params;
};
const handleSendMessage = async () => {
try {
await form.current?.form?.validateFields();
if (!parameters.model) return;
const params = generateParams();
console.log('generateParams:', params);
setParams({
...parameters
});
console.log('params:', params, parameters);
setRouteCache(routeCachekey['/playground/video'], true);
await submitMessage(params);
} catch (error) {
// console.log('error:', error);
} finally {
console.log('finally---------');
setRouteCache(routeCachekey['/playground/video'], false);
}
};
const handleCloseViewCode = useCallback(() => {
setShow(false);
}, []);
return (
<div className="ground-left-wrapper">
<div className="ground-left">
<div
className="message-list-wrap"
ref={scroller}
style={{ paddingBottom: 16 }}
>
<>
<div className="content" style={{ height: '100%' }}>
{videoList.length > 0 && (
<div
style={{
width: '100%',
maxWidth: 720,
margin: '24px auto',
aspectRatio: '16 / 9',
background: '#000',
borderRadius: 8,
overflow: 'hidden',
boxShadow: '0 4px 24px rgba(0,0,0,0.12)'
}}
>
<Spin spinning={loading}>
<video
src={videoList[0]?.dataUrl}
poster="https://placehold.co/640x360.png?text=GPUStack"
controls
disablePictureInPicture
style={{
width: '100%',
height: '100%',
objectFit: 'contain'
}}
/>
</Spin>
</div>
)}
{!videoList.length && (
<div className="flex-column font-size-14 flex-center gap-20 justify-center hold-wrapper">
<VideoCameraOutlined className="font-size-32 text-secondary" />
<span>
{intl.formatMessage({
id: 'playground.video.empty.tips'
})}
</span>
</div>
)}
</div>
</>
</div>
{tokenResult && (
<div style={{ height: 40 }}>
<AlertInfo
type="danger"
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
)}
<div className="ground-left-footer">
<MessageInput
ref={inputRef}
placeholer={intl.formatMessage({
id: 'playground.input.prompt.holder'
})}
actions={[]}
defaultSize={{
minRows: 5,
maxRows: 5
}}
title={intl.formatMessage({ id: 'playground.image.prompt' })}
loading={loading}
disabled={!parameters.model}
isEmpty={!videoList.length}
handleSubmit={handleSendMessage}
handleAbortFetch={handleStopConversation}
onInputChange={handleInputChange}
shouldResetMessage={false}
clearAll={handleClear}
/>
</div>
</div>
<RightContainer collapsed={collapse}>
<DynamicParams
ref={form}
onValuesChange={handleOnValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundVideo;
@@ -1,553 +0,0 @@
import { setRouteCache } from '@/atoms/route-cache';
import AlertInfo from '@/components/alert-info';
import SingleImage from '@/components/auto-image/single-image';
import IconFont from '@/components/icon-font';
import CanvasImageEditor from '@/components/image-editor';
import { processImage } from '@/components/image-editor/extract-image-colors';
import routeCachekey from '@/config/route-cachekey';
import UploadImg from '@/pages/playground/components/upload-img';
import { base64ToFile, generateRandomNumber } from '@/utils';
import { useIntl } from '@umijs/max';
import { Divider } from 'antd';
import _ from 'lodash';
import 'overlayscrollbars/overlayscrollbars.css';
import React, {
forwardRef,
useCallback,
useEffect,
useImperativeHandle,
useMemo,
useRef,
useState
} from 'react';
import { EDIT_IMAGE_API } from '../apis';
import { EDIT_IMAGE_ACCEPT, scaleImageSize } from '../config';
import { useInitImageMeta } from '../hooks/use-init-meta';
import useTextImage from '../hooks/use-text-image';
import '../style/ground-llm.less';
import '../style/system-message-wrap.less';
import { generateImageCode, generateOpenaiImageCode } from '../view-code/image';
import DynamicParams from './dynamic-params';
import MessageInput from './message-input';
import RightContainer from './right-container';
import ViewCommonCode from './view-common-code';
interface MessageProps {
modelList: Global.BaseOption<string>[];
loaded?: boolean;
ref?: any;
}
const GroundImages: React.FC<MessageProps> = forwardRef((props, ref) => {
const { modelList } = props;
const intl = useIntl();
const [show, setShow] = useState(false);
const [collapse, setCollapse] = useState(false);
const scroller = useRef<any>(null);
const paramsRef = useRef<any>(null);
const inputRef = useRef<any>(null);
const [image, setImage] = useState<string>('');
const [mask, setMask] = useState<string | null>(null);
const [uploadList, setUploadList] = useState<any[]>([]);
const [maskUpload, setMaskUpload] = useState<any[]>([]);
const [imageStatus, setImageStatus] = useState<{
isOriginal: boolean;
isResetNeeded: boolean;
width: number;
height: number;
}>({
isOriginal: false,
isResetNeeded: false,
width: 512,
height: 512
});
const doneImage = useRef<boolean>(false);
const [activeImgUid, setActiveImgUid] = useState<number>(0);
const imageEditorRef = useRef<any>(null);
const {
handleOnValuesChange,
handleToggleParamsStyle,
setParams,
updateCacheFormData,
setInitialValues,
updateParamsConfig,
setParamsConfig,
form,
modelMeta,
formFields,
paramsConfig,
initialValues,
parameters,
isOpenaiCompatible
} = useInitImageMeta(props, {
type: 'edit'
});
const {
loading,
tokenResult,
imageList,
currentPrompt,
setImageList,
setCurrentPrompt,
handleStopConversation,
submitMessage
} = useTextImage({
scroller,
paramsRef,
chunkFields: ['stream_options_chunk_result'],
API: EDIT_IMAGE_API
});
useImperativeHandle(ref, () => {
return {
viewCode() {
setShow(true);
},
setCollapse() {
setCollapse(!collapse);
},
collapse: collapse
};
});
const finalParameters = useMemo(() => {
if (parameters.size === 'custom') {
return {
..._.omit(parameters, ['width', 'height', 'preview', 'random_seed']),
image: null,
mask: null,
size:
parameters.width && parameters.height
? `${parameters.width}x${parameters.height}`
: ''
};
}
return {
image: null,
mask: null,
..._.omit(parameters, ['width', 'height', 'random_seed', 'preview'])
};
}, [parameters]);
const viewCodeContent = useMemo(() => {
if (isOpenaiCompatible) {
return generateOpenaiImageCode({
api: EDIT_IMAGE_API,
edit: true,
isFormdata: true,
parameters: {
...finalParameters,
prompt: currentPrompt
}
});
}
return generateImageCode({
api: EDIT_IMAGE_API,
isFormdata: true,
edit: true,
parameters: {
...finalParameters,
prompt: currentPrompt
}
});
}, [finalParameters, currentPrompt, isOpenaiCompatible]);
const handleClear = () => {
setCurrentPrompt('');
};
const handleInputChange = (e: any) => {
setCurrentPrompt(e.target.value);
};
const generateParams = () => {
const params = {
..._.omitBy(finalParameters, (value: string) => !value),
seed: parameters.random_seed
? generateRandomNumber()
: parameters.seed || null,
stream: false,
prompt: currentPrompt
};
return params;
};
const handleSendMessage = async () => {
try {
await form.current?.form?.validateFields();
if (!parameters.model) return;
const params = generateParams();
setParams({
...parameters,
seed: params.seed
});
form.current?.form?.setFieldValue('seed', params.seed);
setRouteCache(routeCachekey['/playground/text-to-image'], true);
await submitMessage({
...params,
image: base64ToFile(_.get(uploadList, '0.dataUrl'), 'image'),
mask: mask ? base64ToFile(mask, 'mask') : null
});
} catch (error) {
// console.log('error:', error);
} finally {
console.log('finally---------');
setRouteCache(routeCachekey['/playground/text-to-image'], false);
}
};
const handleCloseViewCode = () => {
setShow(false);
};
const handleOnScaleImageSize = useCallback(
(data: { rawWidth: number; rawHeight: number }) => {
let { width, height } = scaleImageSize({
width: data.rawWidth,
height: data.rawHeight
});
const { max_width: maxWidth, max_height: maxHeight } = modelMeta;
// update width, height
if (maxWidth) {
width = Math.max(Math.min(width, maxWidth), 512);
}
if (maxHeight) {
height = Math.max(Math.min(height, maxHeight), 512);
}
const newParamsConfig = updateParamsConfig({
size: 'custom',
isOpenaiCompatible
});
setParamsConfig(newParamsConfig);
const newParameters = {
...parameters,
size: 'custom',
width: width || 512,
height: height || 512
};
setParams(newParameters);
updateCacheFormData({
size: 'custom',
width: width || 512,
height: height || 512
});
setInitialValues(newParameters);
},
[parameters, modelMeta, isOpenaiCompatible]
);
const handleUpdateImageList = useCallback(
(base64List: any[]) => {
const currentImg = _.get(base64List, '[0]', {});
const img = _.get(currentImg, 'dataUrl', '');
handleOnScaleImageSize(currentImg);
setUploadList(base64List);
setImage(img);
setActiveImgUid(_.get(base64List, '[0].uid', ''));
setImageStatus({
isOriginal: true,
isResetNeeded: true,
width: _.get(currentImg, 'width', 512),
height: _.get(currentImg, 'height', 512)
});
setImageList([]);
},
[handleOnScaleImageSize]
);
const handleUpdateMaskList = useCallback(async (base64List: any[]) => {
const mask = _.get(base64List, '[0].dataUrl', '');
const maskColors = await processImage(mask);
console.log('maskColors:', maskColors);
imageEditorRef.current?.loadMaskPixs(maskColors || []);
}, []);
const handleClearUploadMask = useCallback(() => {
setMaskUpload([]);
setMask(null);
imageEditorRef.current?.clearMask();
}, []);
const handleOnSave = useCallback(
(data: { img: string; mask: string | null }) => {
setImageStatus((pre) => {
return {
...pre,
isResetNeeded: false
};
});
setMask(data.mask || maskUpload[0]?.dataUrl || null);
setImage(data.img);
},
[]
);
const renderImageEditor = useMemo(() => {
if (image) {
return (
<CanvasImageEditor
ref={imageEditorRef}
imguid={activeImgUid}
imageStatus={imageStatus}
imageSrc={image}
loading={loading}
disabled={loading || !imageStatus.isOriginal}
onSave={handleOnSave}
clearUploadMask={handleClearUploadMask}
handleUpdateImageList={handleUpdateImageList}
handleUpdateMaskList={handleUpdateMaskList}
maskUpload={maskUpload}
accept={EDIT_IMAGE_ACCEPT}
></CanvasImageEditor>
);
}
return (
<>
<UploadImg
accept={EDIT_IMAGE_ACCEPT}
drag={true}
multiple={false}
handleUpdateImgList={handleUpdateImageList}
>
<div
className="flex-column flex-center gap-10 justify-center"
style={{ width: 155, height: 155 }}
>
<IconFont
type="icon-upload_image"
className="font-size-24"
></IconFont>
<span>
{intl.formatMessage({ id: 'playground.image.edit.tips' })}
</span>
</div>
</UploadImg>
</>
);
}, [
intl,
image,
loading,
maskUpload,
imageStatus,
handleOnSave,
handleUpdateImageList
]);
const handleOnImgClick = useCallback(
(item: any, isOrigin: boolean) => {
if (item.progress < 100 && !isOrigin) {
return;
}
if (item.uid === activeImgUid) {
return;
}
setActiveImgUid(item.uid);
setImage(item.dataUrl);
setImageStatus({
isOriginal: isOrigin,
isResetNeeded: false,
width: item.width,
height: item.height
});
},
[activeImgUid]
);
useEffect(() => {
if (imageList.length > 0) {
const doneImg = imageList.find((item) => item.progress === 100);
if (doneImg && !doneImage.current) {
doneImage.current = true;
handleOnImgClick(doneImg, false);
}
}
}, [imageList, handleOnImgClick]);
const renderOriginImage = useMemo(() => {
if (!uploadList.length) {
return null;
}
return (
<>
<SingleImage
{...uploadList[0]}
height={125}
maxHeight={125}
preview={false}
loading={false}
autoSize={false}
editable={false}
autoBgColor={false}
onClick={() => handleOnImgClick(uploadList[0], true)}
label={
<span>{intl.formatMessage({ id: 'playground.image.origin' })}</span>
}
></SingleImage>
</>
);
}, [uploadList, intl, handleOnImgClick]);
const renderMaskImage = useMemo(() => {
if (!maskUpload.length) {
return null;
}
return (
<>
<SingleImage
{...maskUpload[0]}
height={125}
maxHeight={125}
preview={false}
loading={false}
autoSize={false}
editable={true}
autoBgColor={false}
onDelete={() => handleClearUploadMask()}
label={
<span>{intl.formatMessage({ id: 'playground.image.mask' })}</span>
}
></SingleImage>
</>
);
}, [maskUpload, intl]);
return (
<div className="ground-left-wrapper">
<div className="ground-left">
<div className="message-list-wrap" style={{ paddingBottom: 16 }}>
<>
<div className="content" style={{ height: '100%' }}>
{
<div className="flex-column font-size-14 flex-center gap-20 justify-center hold-wrapper">
{renderImageEditor}
</div>
}
</div>
</>
</div>
<div className="ground-left-footer" style={{ padding: 10 }}>
{tokenResult && (
<div style={{ height: 40 }}>
<AlertInfo
type="danger"
message={tokenResult?.errorMessage}
></AlertInfo>
</div>
)}
<div
style={{
display: 'flex',
justifyContent: 'center',
alignItems: 'center'
}}
>
<div className="m-r-10">{renderMaskImage}</div>
{renderOriginImage}
{imageList.length > 0 && (
<>
<Divider
orientation="vertical"
style={{
margin: '0 30px',
height: 80
}}
></Divider>
<div
style={{
display: 'flex',
justifyContent: 'center',
alignItems: 'center',
gap: 10
}}
>
{_.map(imageList, (item: any, index: number) => {
return (
<div
style={{
height: 125,
width: 125,
maxHeight: 125
}}
key={item.uid}
>
<SingleImage
{...item}
height={125}
width={125}
maxHeight={125}
preview={item.preview}
loading={item.loading}
autoSize={false}
editable={false}
autoBgColor={false}
onClick={() => handleOnImgClick(item, false)}
></SingleImage>
</div>
);
})}
</div>
</>
)}
</div>
</div>
</div>
<RightContainer
collapsed={collapse}
footer={
<div style={{ width: 389 }}>
<MessageInput
actions={[]}
defaultSize={{
minRows: 5,
maxRows: 5
}}
ref={inputRef}
placeholer={intl.formatMessage({
id: 'playground.input.prompt.holder'
})}
title={
<span className="font-600">
{intl.formatMessage({ id: 'playground.image.prompt' })}
</span>
}
loading={loading}
disabled={!parameters.model}
isEmpty={!imageList.length}
handleSubmit={handleSendMessage}
handleAbortFetch={handleStopConversation}
onInputChange={handleInputChange}
shouldResetMessage={false}
clearAll={handleClear}
/>
</div>
}
>
<DynamicParams
ref={form}
formFields={formFields}
onValuesChange={handleOnValuesChange}
paramsConfig={paramsConfig}
initialValues={initialValues}
modelList={modelList}
/>
</RightContainer>
<ViewCommonCode
open={show}
viewCodeContent={viewCodeContent}
onCancel={handleCloseViewCode}
title={intl.formatMessage({ id: 'playground.viewcode' })}
></ViewCommonCode>
</div>
);
});
export default GroundImages;
@@ -24,9 +24,9 @@ import React, {
} from 'react';
import 'simplebar-react/dist/simplebar.min.css';
import { CHAT_API } from '../../apis';
import { ChatParamsConfig } from '../../chat/params-config';
import { Roles, generateMessagesByListContent } from '../../config';
import CompareContext from '../../config/compare-context';
import { ChatParamsConfig } from '../../config/params-config';
import { MessageItem, ModelSelectionItem } from '../../config/types';
import { LLM_METAKEYS, llmInitialValues } from '../../hooks/config';
import useChatCompletion from '../../hooks/use-chat-completion';