diff --git a/src/pages/playground/components/dynamic-params.tsx b/src/pages/playground/components/dynamic-params.tsx index b16f96c2..3b2e8301 100644 --- a/src/pages/playground/components/dynamic-params.tsx +++ b/src/pages/playground/components/dynamic-params.tsx @@ -1,10 +1,7 @@ -import FieldWrapper from '@/components/seal-form/field-wrapper'; -import SealInput from '@/components/seal-form/seal-input'; +import FieldComponent from '@/components/seal-form/field-component'; import SealSelect from '@/components/seal-form/seal-select'; -import { INPUT_WIDTH } from '@/constants'; -import { InfoCircleOutlined } from '@ant-design/icons'; import { useIntl } from '@umijs/max'; -import { Form, InputNumber, Slider, Tooltip } from 'antd'; +import { Form } from 'antd'; import _ from 'lodash'; import { forwardRef, @@ -16,25 +13,13 @@ import { useMemo } from 'react'; import { ParamsSchema } from '../config/types'; -import CustomLabelStyles from '../style/custom-label.less'; - -type ParamsSettingsFormProps = { - top_n?: number; - model?: string; -}; type ParamsSettingsProps = { ref?: any; parametersTitle?: React.ReactNode; - selectedModel?: string; showModelSelector?: boolean; - params?: Record; - model?: string; modelList: Global.BaseOption[]; onValuesChange?: (changeValues: any, value: Record) => void; - setParams: (params: any) => void; - onModelChange?: (model: string) => void; - globalParams?: Record; paramsConfig?: ParamsSchema[]; initialValues?: Record; extra?: React.ReactNode; @@ -43,16 +28,11 @@ type ParamsSettingsProps = { const ParamsSettings: React.FC = forwardRef( ( { - setParams, onValuesChange, - onModelChange, parametersTitle, - selectedModel, - globalParams, initialValues, paramsConfig, modelList, - params, showModelSelector = true, extra }, @@ -62,34 +42,20 @@ const ParamsSettings: React.FC = forwardRef( const [form] = Form.useForm(); const formId = useId(); + const dependFieldsValue = Form.useWatch( + paramsConfig?.flatMap((item) => item.dependencies || []), + form + ); + useImperativeHandle(ref, () => ({ form })); useEffect(() => { - let model = selectedModel || ''; - - if (showModelSelector) { - model = model || _.get(modelList, '[0].value'); - } - form.setFieldsValue({ - model: model, ...initialValues }); - setParams({ - model: model, - ...initialValues - }); - onModelChange?.(model); - }, [modelList, showModelSelector, selectedModel, initialValues]); - - const handleModelChange = useCallback( - (value: string) => { - onModelChange?.(value); - }, - [onModelChange] - ); + }, [initialValues]); const handleOnFinish = (values: any) => { console.log('handleOnFinish', values); @@ -101,172 +67,60 @@ const ParamsSettings: React.FC = forwardRef( const handleValuesChange = useCallback( (changedValues: any, allValues: any) => { - setParams?.(allValues); - onValuesChange?.(changedValues, allValues); - }, - [onValuesChange, setParams] - ); - const handleFieldValueChange = useCallback( - (val: any, field: string) => { - const values = form.getFieldsValue(); - form.setFieldsValue({ - ...values, - [field]: val - }); - setParams({ - ...values, - [field]: val - }); - onValuesChange?.( - { [field]: val }, - { - ...values, - [field]: val - } + const normalizedValues = Object.fromEntries( + Object.entries(changedValues).map(([key, value]: [string, any]) => [ + key, + value?.target?.checked ?? value?.target?.value ?? value + ]) ); + form.setFieldsValue(normalizedValues); + onValuesChange?.(normalizedValues, { + ...allValues, + ...normalizedValues + }); }, - [form, setParams, onValuesChange] + [onValuesChange] ); - useEffect(() => { - form.setFieldsValue(globalParams); - }, [globalParams]); - - const renderLabel = useCallback( - (args: { field: string; label: string; description: string }) => { - return ( - - - {args.description ? ( - - {args.label} - - - - - ) : ( - {args.label} - )} - - - handleFieldValueChange(val, args.field)} - > - - ); - }, - [form, handleFieldValueChange] - ); - - const renderDescription = useCallback( - (item: ParamsSchema) => { - if (!item.description) { - return null; - } - if (item.description.html) { - return ( -
- ); - } - return intl.formatMessage({ id: item.description.text }); - }, - [intl] - ); const renderFields = useMemo(() => { - console.log('paramsConfig++++++++++++'); - if (!paramsConfig?.length) { + if (!paramsConfig) { return null; } - return paramsConfig.map((item: ParamsSchema) => { - if (item.type === 'InputNumber') { - return ( - - - - ); - } - if (item.type === 'TextArea') { - return ( - - - - ); - } - if (item.type === 'Select') { - return ( - - - - ); - } - if (item.type === 'Slider') { - return ( - - - handleFieldValueChange(val, item.name)} - > - - - ); - } - return null; + console.log('renderFields---------'); + const formValues = form?.getFieldsValue(); + return paramsConfig?.map((item: ParamsSchema) => { + return ( + + + + ); }); - }, [paramsConfig, params, renderDescription, intl]); + }, [paramsConfig, intl, dependFieldsValue]); return (
= forwardRef( )} - + = forwardRef( ]} > = forwardRef((props, ref) => { const { modelList } = props; - const messageId = useRef(0); - const [isOpenaiCompatible, setIsOpenaiCompatible] = useState(false); - const [imageList, setImageList] = useState< - { - dataUrl: string; - height: number | string; - width: string | number; - maxHeight: string | number; - maxWidth: string | number; - uid: number; - span?: number; - loading?: boolean; - progress?: number; - }[] - >([]); const intl = useIntl(); - const [searchParams] = useSearchParams(); - const selectModel = searchParams.get('model') || ''; - const [parameters, setParams] = useState({}); const [show, setShow] = useState(false); - const [loading, setLoading] = useState(false); - const [tokenResult, setTokenResult] = useState(null); const [collapse, setCollapse] = useState(false); const scroller = useRef(null); const paramsRef = useRef(null); - const messageListLengthCache = useRef(0); - const requestToken = useRef(null); - const [currentPrompt, setCurrentPrompt] = useState(''); - const [modelMeta, setModelMeta] = useState({}); - const form = useRef(null); const inputRef = useRef(null); - const cacheFormData = useRef>({}); - const size = Form.useWatch('size', form.current?.form); - - const { initialize, updateScrollerPosition } = useOverlayScroller(); - const { initialize: innitializeParams } = useOverlayScroller(); + const { + handleOnValuesChange, + handleToggleParamsStyle, + form, + paramsConfig, + initialValues, + parameters, + isOpenaiCompatible + } = useInitImageMeta(props); + const { + loading, + tokenResult, + imageList, + promptList, + currentPrompt, + setCurrentPrompt, + handleClear, + handleStopConversation, + submitMessage + } = useTextImage({ + scroller, + paramsRef + }); useImperativeHandle(ref, () => { return { @@ -137,58 +81,10 @@ const GroundImages: React.FC = forwardRef((props, ref) => { }; }); - const removeBase64Suffix = (str: string, suffix: string) => { - return str.endsWith(suffix) ? str.slice(0, -suffix.length) : str; - }; - - const getNewImageSizeOptions = useCallback((metaData: any) => { - const { max_height, max_width } = metaData || {}; - if (!max_height || !max_width) { - return imageSizeOptions; - } - const newImageSizeOptions = imageSizeOptions.filter((item) => { - return item.width <= max_width && item.height <= max_height; - }); - if ( - !newImageSizeOptions.find( - (item) => item.width === max_width && item.height === max_height - ) - ) { - newImageSizeOptions.push({ - width: max_width, - height: max_height, - label: `${max_width}x${max_height}`, - value: `${max_width}x${max_height}` - }); - } - return newImageSizeOptions; - }, []); - - const paramsConfig = useMemo(() => { - const newImageSizeOptions = getNewImageSizeOptions(modelMeta); - let result: ParamsSchema[] = ImageParamsConfig.map((item: ParamsSchema) => { - if (item.name === 'size') { - return { - ...item, - options: newImageSizeOptions - }; - } - return item; - }); - if (!newImageSizeOptions.length) { - result = result.filter((item) => item.name !== 'size'); - } - return result; - }, [modelMeta]); - const generateNumber = (min: number, max: number) => { return Math.floor(Math.random() * (max - min + 1) + min); }; - const updateCacheFormData = (values: Record) => { - _.merge(cacheFormData.current, values); - }; - const handleRandomPrompt = useCallback(() => { const randomIndex = generateNumber(0, promptList.length - 1); const randomPrompt = promptList[randomIndex]; @@ -199,25 +95,6 @@ const GroundImages: React.FC = forwardRef((props, ref) => { }); }, []); - const setImageSize = useCallback(() => { - let size: Record = { - span: 12 - }; - if (parameters.n === 1) { - size.span = 24; - } - if (parameters.n === 2) { - size.span = 12; - } - if (parameters.n === 3) { - size.span = 12; - } - if (parameters.n === 4) { - size.span = 12; - } - return size; - }, [parameters.n]); - const finalParameters = useMemo(() => { if (parameters.size === 'custom') { return { @@ -234,6 +111,7 @@ const GroundImages: React.FC = forwardRef((props, ref) => { }, [parameters]); const viewCodeContent = useMemo(() => { + console.log('finalParameters:', finalParameters); if (isOpenaiCompatible) { return generateOpenaiImageCode({ api: CREAT_IMAGE_API, @@ -250,401 +128,30 @@ const GroundImages: React.FC = forwardRef((props, ref) => { prompt: currentPrompt } }); - }, [finalParameters, currentPrompt, parameters.size]); - - const setMessageId = () => { - messageId.current = messageId.current + 1; - return messageId.current; - }; - - const handleStopConversation = () => { - requestToken.current?.abort?.(); - setLoading(false); - }; - - const submitMessage = async (current?: { content: string }) => { - try { - await form.current?.form?.validateFields(); - if (!parameters.model) return; - const size: any = setImageSize(); - setLoading(true); - setMessageId(); - setTokenResult(null); - setCurrentPrompt(current?.content || ''); - setRouteCache(routeCachekey['/playground/text-to-image'], true); - const imgSize = _.split(finalParameters.size, 'x').map((item: number) => - _.toNumber(item) - ); - - // preview - let stream_options: Record = { - chunk_size: 16 * 1024, - chunk_results: true - }; - if (parameters.preview === 'preview') { - stream_options = { - preview: true - }; - } - - if (parameters.preview === 'preview_faster') { - stream_options = { - preview_faster: true - }; - } - - let newImageList = Array(parameters.n) - .fill({}) - .map((item, index: number) => { - return { - dataUrl: 'data:image/png;base64,', - ...size, - progress: 0, - height: imgSize[1], - width: imgSize[0], - loading: true, - progressType: 'dashboard', - preview: false, - uid: setMessageId() - }; - }); - setImageList(newImageList); - - requestToken.current?.abort?.(); - requestToken.current = new AbortController(); - - const params = { - ..._.omitBy(finalParameters, (value: string) => !value), - seed: parameters.random_seed ? generateRandomNumber() : parameters.seed, - stream: true, - stream_options: { - ...stream_options - }, - prompt: current?.content || currentPrompt || '' - }; - setParams({ - ...parameters, - seed: params.seed - }); - form.current?.form?.setFieldValue('seed', params.seed); - - const result: any = await fetchChunkedData({ - data: params, - url: `${CREAT_IMAGE_API}?t=${Date.now()}`, - signal: requestToken.current.signal - }); - if (result.error) { - setTokenResult({ - error: true, - errorMessage: extractErrorMessage(result) - }); - setImageList([]); - return; - } - - const { reader, decoder } = result; - - await readStreamData(reader, decoder, (chunk: any) => { - if (chunk?.error) { - setTokenResult({ - error: true, - errorMessage: chunk?.error?.message || chunk?.message || '' - }); - return; - } - chunk?.data?.forEach((item: any) => { - const imgItem = newImageList[item.index]; - if (item.b64_json && stream_options.chunk_results) { - imgItem.dataUrl += removeBase64Suffix(item.b64_json, ODD_STRING); - } else if (item.b64_json) { - imgItem.dataUrl = `data:image/png;base64,${removeBase64Suffix(item.b64_json, ODD_STRING)}`; - } - const progress = item.progress; - - newImageList[item.index] = { - dataUrl: imgItem.dataUrl, - height: imgSize[1], - width: imgSize[0], - maxHeight: `${imgSize[1]}px`, - maxWidth: `${imgSize[0]}px`, - uid: imgItem.uid, - span: imgItem.span, - loading: stream_options.chunk_results ? progress < 100 : false, - preview: progress >= 100, - progress: progress - }; - }); - setImageList([...newImageList]); - }); - } catch (error) { - console.log('error:', error); - requestToken.current?.abort?.(); - setImageList([]); - } finally { - setLoading(false); - setRouteCache(routeCachekey['/playground/text-to-image'], false); - } - }; - const handleClear = () => { - // setMessageId(); - // setImageList([]); - // setTokenResult(null); - }; + }, [finalParameters, isOpenaiCompatible, currentPrompt]); const handleInputChange = (e: any) => { setCurrentPrompt(e.target.value); }; - const handleSendMessage = (message: Omit) => { - const currentMessage = message.content ? message : undefined; - submitMessage(currentMessage); + const handleSendMessage = async (message: Omit) => { + try { + await form.current?.form?.validateFields(); + if (!parameters.model) return; + submitMessage(finalParameters); + setRouteCache(routeCachekey['/playground/text-to-image'], true); + } catch (error) { + // console.log('error:', error); + } finally { + console.log('finally---------'); + setRouteCache(routeCachekey['/playground/text-to-image'], false); + } }; const handleCloseViewCode = () => { setShow(false); }; - const handleToggleParamsStyle = () => { - if (isOpenaiCompatible) { - form.current?.form?.setFieldsValue({ - ...advancedFieldsDefaultValus, - ..._.pick(cacheFormData.current, _.keys(advancedFieldsDefaultValus)) - }); - setParams((pre: object) => { - return { - ..._.omit(pre, _.keys(openaiCompatibleFieldsDefaultValus)), - ...advancedFieldsDefaultValus, - ..._.pick(cacheFormData.current, _.keys(advancedFieldsDefaultValus)) - }; - }); - } else { - form.current?.form?.setFieldsValue({ - ...openaiCompatibleFieldsDefaultValus - }); - setParams((pre: object) => { - return { - ...openaiCompatibleFieldsDefaultValus, - ..._.omit(pre, _.keys(advancedFieldsDefaultValus)) - }; - }); - } - setIsOpenaiCompatible(!isOpenaiCompatible); - updateCacheFormData(parameters); - }; - - const renderExtra = useMemo(() => { - if (!isOpenaiCompatible) { - return []; - } - return ImageconstExtraConfig.map((item: ParamsSchema) => { - return ( - - - - ); - }); - }, [ImageconstExtraConfig, isOpenaiCompatible, intl]); - - const handleFieldChange = (e: any) => { - if (e.target.id.indexOf('random_seed') > -1) { - form.current?.form?.setFieldValue('random_seed', e.target.checked); - setParams((pre: object) => { - return { - ...pre, - random_seed: e.target.checked - }; - }); - } - }; - const renderAdvanced = useMemo(() => { - if (isOpenaiCompatible) { - return []; - } - const formValues = form.current?.form?.getFieldsValue(); - return ImageAdvancedParamsConfig.map((item: ParamsSchema) => { - if (item.name === 'strength') { - return null; - } - return ( - - - - ); - }); - }, [ImageAdvancedParamsConfig, isOpenaiCompatible, intl, form.current]); - - const renderCustomSize = useMemo(() => { - if (size === 'custom') { - return ImageCustomSizeConfig.map((item: ParamsSchema) => { - return ( - - - - ); - }); - } - return null; - }, [size, intl, modelMeta]); - - const handleOnModelChange = useCallback( - (val: string) => { - if (!val) return; - - const model = modelList.find((item) => item.value === val); - - setModelMeta(model?.meta || {}); - const imageSizeOptions = getNewImageSizeOptions(model?.meta); - const w = model?.meta?.default_width || 512; - const h = model?.meta?.default_height || 512; - const defaultSize = imageSizeOptions.length ? `${w}x${h}` : 'custom'; - if (!isOpenaiCompatible) { - setParams((pre: object) => { - const obj = _.merge({}, pre, _.pick(model?.meta, METAKEYS, {})); - - return { - ...obj, - size: defaultSize, - width: w, - height: h - }; - }); - form.current?.form?.setFieldsValue({ - ..._.pick(model?.meta, METAKEYS, {}), - size: defaultSize, - width: w, - height: h - }); - } - updateCacheFormData({ - ..._.pick(model?.meta, METAKEYS, {}), - size: defaultSize, - width: w, - height: h - }); - }, - [modelList, isOpenaiCompatible] - ); - - useEffect(() => { - return () => { - requestToken.current?.abort?.(); - }; - }, []); - - useEffect(() => { - if (size === 'custom') { - form.current?.form?.setFieldsValue({ - width: cacheFormData.current.width || 512, - height: cacheFormData.current.height || 512 - }); - setParams((pre: object) => { - return { - ...pre, - width: cacheFormData.current.width || 512, - height: cacheFormData.current.height || 512 - }; - }); - } - }, [size]); - - useEffect(() => { - if (scroller.current) { - initialize(scroller.current); - } - }, [scroller.current, initialize]); - - useEffect(() => { - if (paramsRef.current) { - innitializeParams(paramsRef.current); - } - }, [paramsRef.current, innitializeParams]); - - useEffect(() => { - if (loading) { - updateScrollerPosition(); - } - }, [imageList, loading]); - - useEffect(() => { - if (imageList.length > messageListLengthCache.current) { - updateScrollerPosition(); - } - messageListLengthCache.current = imageList.length; - }, [imageList.length]); - return (
@@ -771,14 +278,10 @@ const GroundImages: React.FC = forwardRef((props, ref) => {
} - onModelChange={handleOnModelChange} - setParams={setParams} + onValuesChange={handleOnValuesChange} paramsConfig={paramsConfig} initialValues={initialValues} - params={parameters} - selectedModel={selectModel} modelList={modelList} - extra={[renderCustomSize, ...renderExtra, ...renderAdvanced]} />
diff --git a/src/pages/playground/config/params-config.ts b/src/pages/playground/config/params-config.ts index 69731402..57ce953a 100644 --- a/src/pages/playground/config/params-config.ts +++ b/src/pages/playground/config/params-config.ts @@ -1,5 +1,12 @@ import { ParamsSchema } from './types'; +export interface ImageSizeItem { + label: string; + value: string; + width: number; + height: number; + locale?: boolean; +} export const imageSizeOptions: { label: string; value: string; @@ -145,6 +152,48 @@ export const ImageParamsConfig: ParamsSchema[] = [ } ]; +export const ImageCountConfig: ParamsSchema[] = [ + { + type: 'InputNumber', + name: 'n', + label: { + text: 'playground.params.counts', + isLocalized: true + }, + attrs: { + min: 1, + max: 4 + }, + rules: [ + { + required: false + } + ] + } +]; + +export const ImageSizeConfig: ParamsSchema[] = [ + { + type: 'Select', + name: 'size', + options: imageSizeOptions, + description: { + text: 'playground.params.size.description', + html: true, + isLocalized: true + }, + label: { + text: 'playground.params.size', + isLocalized: true + }, + rules: [ + { + required: false + } + ] + } +]; + export const ImageEidtParamsConfig: ParamsSchema[] = [ { type: 'InputNumber', @@ -414,6 +463,7 @@ export const ImageAdvancedParamsConfig: ParamsSchema[] = [ attrs: { min: 0 }, + dependencies: ['random_seed'], disabledConfig: { depends: ['random_seed'], when: (values: Record): boolean => values?.random_seed @@ -431,6 +481,12 @@ export const ImageAdvancedParamsConfig: ParamsSchema[] = [ text: 'playground.image.params.randomseed', isLocalized: true }, + style: { + marginBottom: 20 + }, + formItemAttrs: { + noStyle: true + }, rules: [ { required: false @@ -449,7 +505,7 @@ export const ImageCustomSizeConfig: ParamsSchema[] = [ }, attrs: { min: 256, - max: 3200, + max: 1024, step: 64, inputnumber: false }, @@ -469,7 +525,7 @@ export const ImageCustomSizeConfig: ParamsSchema[] = [ }, attrs: { min: 256, - max: 3200, + max: 1024, step: 64, inputnumber: false }, @@ -481,3 +537,119 @@ export const ImageCustomSizeConfig: ParamsSchema[] = [ ] } ]; + +export const ChatParamsConfig: ParamsSchema[] = [ + { + type: 'Slider', + name: 'temperature', + label: { + text: 'Temperature', + isLocalized: false + }, + description: { + text: 'playground.params.temperature.tips', + html: false, + isLocalized: true + }, + attrs: { + max: 2, + step: 0.1, + inputnumber: true + }, + rules: [ + { + required: false + } + ] + }, + { + type: 'Slider', + name: 'max_tokens', + label: { + text: 'Max Tokens', + isLocalized: false + }, + description: { + text: 'playground.params.maxtokens.tips', + html: false, + isLocalized: true + }, + attrs: { + max: 1024, + step: 1, + inputnumber: true + }, + rules: [ + { + required: false + } + ] + }, + { + type: 'Slider', + name: 'top_p', + label: { + text: 'Top P', + isLocalized: false + }, + description: { + text: 'playground.params.topp.tips', + html: false, + isLocalized: true + }, + attrs: { + max: 1, + step: 0.1, + inputnumber: true + }, + rules: [ + { + required: false + } + ] + }, + { + type: 'InputNumber', + name: 'seed', + label: { + text: 'Seed', + isLocalized: false + }, + description: { + text: 'playground.params.seed.tips', + html: false, + isLocalized: true + }, + attrs: { + min: 0 + }, + rules: [ + { + required: false + } + ] + }, + { + type: 'Input', + name: 'stop', + label: { + text: 'Stop Sequence', + isLocalized: false + }, + description: { + text: 'playground.params.stop.tips', + html: false, + isLocalized: true + }, + attrs: { + normalize(value: string) { + return value || null; + } + }, + rules: [ + { + required: false + } + ] + } +]; diff --git a/src/pages/playground/config/types.ts b/src/pages/playground/config/types.ts index be94461b..148ecc88 100644 --- a/src/pages/playground/config/types.ts +++ b/src/pages/playground/config/types.ts @@ -33,6 +33,7 @@ export interface ParamsSchema { text: string; isLocalized?: boolean; }; + dependencies?: string[]; style?: React.CSSProperties; options?: Global.BaseOption[]; value?: string | number | boolean | string[]; @@ -53,6 +54,7 @@ export interface ParamsSchema { }[]; placeholder?: React.ReactNode; attrs?: Record; + formItemAttrs?: Record; description?: { text: string; html?: boolean; diff --git a/src/pages/playground/hooks/use-init-meta.ts b/src/pages/playground/hooks/use-init-meta.ts index 01a4b427..af57e479 100644 --- a/src/pages/playground/hooks/use-init-meta.ts +++ b/src/pages/playground/hooks/use-init-meta.ts @@ -1,5 +1,16 @@ +import { + ImageAdvancedParamsConfig, + ImageCountConfig, + ImageCustomSizeConfig, + ImageSizeConfig, + ImageSizeItem, + ImageconstExtraConfig, + imageSizeOptions as imageSizeList +} from '@/pages/playground/config/params-config'; +import { useSearchParams } from '@umijs/max'; import _ from 'lodash'; -import { imageSizeOptions } from '../config/params-config'; +import React, { useCallback, useEffect, useRef, useState } from 'react'; +import { ParamsSchema } from '../config/types'; const LLM_METAKEYS: Record = { seed: 'seed', @@ -28,8 +39,7 @@ const llmInitialValues = { max_tokens: 1024 }; -const imgInitialValues = { - n: 1, +const advancedFieldsDefaultValus = { seed: null, sample_method: 'euler_a', cfg_scale: 4.5, @@ -40,7 +50,20 @@ const imgInitialValues = { preview: 'preview_faster' }; -export default function useInitMeta() { +const openaiCompatibleFieldsDefaultValus = { + quality: 'standard', + style: null +}; + +const imgInitialValues = { + n: 1, + size: '512x512', + width: 512, + height: 512 +}; + +// init not image meta +export const useInitLLmMeta = () => { const extractLLMMeta = (meta: any) => { const modelMeta = meta || {}; const modelMetaValue = _.pick(modelMeta, _.keys(LLM_METAKEYS)); @@ -75,13 +98,62 @@ export default function useInitMeta() { }; }; + return { extractLLMMeta }; +}; + +interface MessageProps { + modelList: Global.BaseOption[]; + loaded?: boolean; + ref?: any; +} + +export const useInitImageMeta = (props: MessageProps) => { + const { modelList } = props; + const form = useRef(null); + const [searchParams] = useSearchParams(); + const selectModel = searchParams.get('model') || ''; + const [modelMeta, setModelMeta] = useState({}); + const [isOpenaiCompatible, setIsOpenaiCompatible] = useState(false); + const [imageSizeOptions, setImageSizeOptions] = React.useState< + ImageSizeItem[] + >([]); + const [basicParamsConfig, setBasicParamsConfig] = React.useState< + ParamsSchema[] + >([...ImageCountConfig, ...ImageSizeConfig]); + + const [initialValues, setInitialValues] = useState({ + ...imgInitialValues, + ...advancedFieldsDefaultValus, + model: selectModel + }); + const [paramsConfig, setParamsConfig] = useState([ + ...ImageCountConfig, + ...ImageSizeConfig, + ...ImageAdvancedParamsConfig + ]); + + const [parameters, setParams] = useState({ + ...imgInitialValues, + ...advancedFieldsDefaultValus, + model: selectModel + }); + + const cacheFormData = React.useRef>({ + ...imgInitialValues, + ...openaiCompatibleFieldsDefaultValus, + ...advancedFieldsDefaultValus + }); + const getNewImageSizeOptions = (metaData: any) => { const { max_height, max_width } = metaData || {}; if (!max_height || !max_width) { - return imageSizeOptions; + return imageSizeList; } - const newImageSizeOptions = imageSizeOptions.filter((item) => { - return item.width <= max_width && item.height <= max_height; + const newImageSizeOptions = imageSizeList.filter((item) => { + return ( + (item.width <= max_width && item.height <= max_height) || + item.value === 'custom' + ); }); if ( !newImageSizeOptions.find( @@ -98,20 +170,217 @@ export default function useInitMeta() { return newImageSizeOptions; }; + const generateImageParamsConfig = ( + currentModel: any, + sizeOptions: ImageSizeItem[] + ) => { + if (sizeOptions.length) { + const sizeConfig = ImageSizeConfig.map((item) => { + return { + ...item, + options: sizeOptions + }; + }); + return [...ImageCountConfig, ...sizeConfig]; + } + const { max_height, max_width } = currentModel.meta || {}; + const customSizeConfig = _.cloneDeep(ImageCustomSizeConfig).map( + (item: any) => { + const max = item.name === 'height' ? max_height : max_width; + return { + ...item, + attrs: { + ...item.attrs, + max: max || item.attrs.max + } + }; + } + ); + return [...ImageCountConfig, ...customSizeConfig]; + }; + const extractIMGMeta = (meta: any) => { + const sizeOptions = getNewImageSizeOptions(meta); return { - form: _.merge({}, imgInitialValues, { - ..._.pick(meta, IMG_METAKEYS), - width: meta?.default_width || 512, - height: meta?.default_height || 512 - }), + form: _.merge( + {}, + { ...imgInitialValues, ...advancedFieldsDefaultValus }, + { + ..._.pick(meta, IMG_METAKEYS), + size: sizeOptions.length + ? `${meta?.default_width || 512}x${meta?.default_height || 512}` + : 'custom', + width: meta?.default_width || 512, + height: meta?.default_height || 512 + } + ), meta: meta, - sizeOptions: getNewImageSizeOptions(meta) + sizeOptions: sizeOptions }; }; - return { - extractLLMMeta, - extractIMGMeta + const updateCacheFormData = (values: Record) => { + _.merge(cacheFormData.current, values); }; -} + const handleToggleParamsStyle = () => { + // update values + let values: any = {}; + if (!isOpenaiCompatible) { + values = _.pick( + cacheFormData.current, + _.keys({ + ...imgInitialValues, + ...openaiCompatibleFieldsDefaultValus + }) + ); + } else { + values = _.pick( + cacheFormData.current, + _.keys({ + ...imgInitialValues, + ...advancedFieldsDefaultValus + }) + ); + } + + // update config + if (values.size === 'custom') { + const config = [ + ...basicParamsConfig, + ...ImageCustomSizeConfig, + ...(isOpenaiCompatible + ? ImageAdvancedParamsConfig + : ImageconstExtraConfig) + ]; + setParamsConfig(config); + } else { + const config = [ + ...basicParamsConfig, + ...(isOpenaiCompatible + ? ImageAdvancedParamsConfig + : ImageconstExtraConfig) + ]; + setParamsConfig(config); + } + form.current?.form?.setFieldsValue({ + ...values, + model: parameters.model + }); + setParams({ + ...values, + model: parameters.model + }); + setIsOpenaiCompatible(!isOpenaiCompatible); + updateCacheFormData({ + ...values, + model: parameters + }); + }; + + const handleOnModelChange = useCallback( + (val: string) => { + if (!val) return; + const model = modelList.find((item) => item.value === val); + const { form: initialData, sizeOptions } = extractIMGMeta(model?.meta); + const newParamsConfig = generateImageParamsConfig(model, sizeOptions); + + if (!isOpenaiCompatible) { + setParamsConfig([...newParamsConfig, ...ImageAdvancedParamsConfig]); + } else { + setParamsConfig(newParamsConfig); + } + setBasicParamsConfig(newParamsConfig); + setImageSizeOptions(sizeOptions); + setModelMeta(model?.meta || {}); + setInitialValues({ + ...initialData, + model: val + }); + setParams({ + ...initialData, + model: val + }); + updateCacheFormData(initialData); + }, + [modelList, isOpenaiCompatible] + ); + + const handleOnValuesChange = useCallback( + (changeValues: Record, allValues: Record) => { + // model change will reset all values + if (changeValues.model) { + handleOnModelChange(changeValues.model); + return; + } + + if (changeValues.size && changeValues.size === 'custom') { + setParamsConfig([ + ...basicParamsConfig, + ...ImageCustomSizeConfig, + ...(!isOpenaiCompatible + ? ImageAdvancedParamsConfig + : ImageconstExtraConfig) + ]); + setParams({ + ...allValues, + width: modelMeta.default_width || 512, + height: modelMeta.default_height || 512 + }); + form.current?.form?.setFieldsValue({ + width: modelMeta.default_width || 512, + height: modelMeta.default_height || 512 + }); + updateCacheFormData(changeValues); + } else if (changeValues.size && parameters.size === 'custom') { + setParamsConfig([ + ...basicParamsConfig, + ...(!isOpenaiCompatible + ? ImageAdvancedParamsConfig + : ImageconstExtraConfig) + ]); + setParams(allValues); + updateCacheFormData(changeValues); + } + }, + [ + handleOnModelChange, + parameters.size, + basicParamsConfig, + isOpenaiCompatible + ] + ); + + useEffect(() => { + if (!parameters.model && modelList.length) { + const model = modelList[0]?.value; + handleOnModelChange(model); + } + }, [modelList, parameters.model, handleOnModelChange]); + + return { + extractIMGMeta, + generateImageParamsConfig, + setImageSizeOptions, + setBasicParamsConfig, + updateCacheFormData, + handleOnValuesChange, + handleToggleParamsStyle, + setParams, + form, + paramsConfig, + initialValues, + parameters, + isOpenaiCompatible, + cacheFormData, + basicParamsConfig, + imageSizeOptions, + openaiCompatibleFieldsDefaultValus, + advancedFieldsDefaultValus, + ImageAdvancedParamsConfig, + imgInitialValues, + ImageCountConfig, + ImageSizeConfig, + ImageCustomSizeConfig, + ImageconstExtraConfig + }; +}; diff --git a/src/pages/playground/hooks/use-text-image.ts b/src/pages/playground/hooks/use-text-image.ts index 6fe37046..a8d02111 100644 --- a/src/pages/playground/hooks/use-text-image.ts +++ b/src/pages/playground/hooks/use-text-image.ts @@ -1,3 +1,4 @@ +import useOverlayScroller from '@/hooks/use-overlay-scroller'; import { CREAT_IMAGE_API } from '@/pages/playground/apis'; import { extractErrorMessage, promptList } from '@/pages/playground/config'; import { generateRandomNumber } from '@/utils'; @@ -6,23 +7,36 @@ import { readLargeStreamData as readStreamData } from '@/utils/fetch-chunk-data'; import _ from 'lodash'; -import { useCallback, useRef, useState } from 'react'; +import { useEffect, useRef, useState } from 'react'; const ODD_STRING = 'AAAABJRU5ErkJgg==='; -export default function useTextImage() { +export default function useTextImage({ scroller, paramsRef }: any) { const [loading, setLoading] = useState(false); const [tokenResult, setTokenResult] = useState(null); const [imageList, setImageList] = useState([]); const [currentPrompt, setCurrentPrompt] = useState(''); const messageId = useRef(0); const requestToken = useRef(null); + const { initialize } = useOverlayScroller(); + const { initialize: innitializeParams } = useOverlayScroller(); + + useEffect(() => { + if (scroller.current) { + initialize(scroller.current); + } + }, [scroller.current, initialize]); + useEffect(() => { + if (paramsRef.current) { + innitializeParams(paramsRef.current); + } + }, [paramsRef.current, innitializeParams]); const removeBase64Suffix = (str: string, suffix: string) => { return str.endsWith(suffix) ? str.slice(0, -suffix.length) : str; }; - const setImageSize = useCallback((parameters: any) => { + const setImageSize = (parameters: any) => { let size: Record = { span: 12 }; @@ -39,7 +53,7 @@ export default function useTextImage() { size.span = 12; } return size; - }, []); + }; const setMessageId = () => { messageId.current = messageId.current + 1; @@ -50,21 +64,17 @@ export default function useTextImage() { return Math.floor(Math.random() * (max - min + 1) + min); }; - const submitMessage = async (params: { - current?: { content: string }; - system?: { role: string; content: string }; - parameters: any; - }) => { - const { current, parameters } = params; + const submitMessage = async (parameters: any) => { try { if (!parameters.model) return; const size: any = setImageSize(parameters); setLoading(true); setMessageId(); setTokenResult(null); - setCurrentPrompt(current?.content || ''); - const imgSize = [parameters.width, parameters.height]; + const imgSize = _.split(parameters.size, 'x').map((item: string) => + _.toNumber(item) + ); // preview let stream_options: Record = { @@ -115,7 +125,7 @@ export default function useTextImage() { stream_options: { ...stream_options }, - prompt: current?.content + prompt: currentPrompt }; const result: any = await fetchChunkedData({ @@ -176,9 +186,9 @@ export default function useTextImage() { }; const handleClear = () => { - setMessageId(); setImageList([]); setTokenResult(null); + setCurrentPrompt(''); }; const handleStopConversation = () => { @@ -186,11 +196,20 @@ export default function useTextImage() { setLoading(false); }; + useEffect(() => { + return () => { + requestToken.current?.abort?.(); + }; + }, []); + return { loading, tokenResult, imageList, promptList, + currentPrompt, + setTokenResult, + setCurrentPrompt, handleStopConversation, generateNumber, handleClear,