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 { SwapOutlined } from '@ant-design/icons'; import { useIntl } from '@umijs/max'; import { Button, Divider, Tooltip } from 'antd'; import classNames from 'classnames'; 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 { scaleImageSize } from '../config'; import { useInitImageMeta } from '../hooks/use-init-meta'; import useTextImage from '../hooks/use-text-image'; import '../style/ground-left.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 ViewCommonCode from './view-common-code'; interface MessageProps { modelList: Global.BaseOption[]; loaded?: boolean; ref?: any; } const GroundImages: React.FC = forwardRef((props, ref) => { const { modelList } = props; const intl = useIntl(); const [show, setShow] = useState(false); const [collapse, setCollapse] = useState(false); const scroller = useRef(null); const paramsRef = useRef(null); const inputRef = useRef(null); const [image, setImage] = useState(''); const [mask, setMask] = useState(null); const [uploadList, setUploadList] = useState([]); const [maskUpload, setMaskUpload] = useState([]); const [imageStatus, setImageStatus] = useState<{ isOriginal: boolean; isResetNeeded: boolean; width: number; height: number; }>({ isOriginal: false, isResetNeeded: false, width: 512, height: 512 }); const doneImage = useRef(false); const [activeImgUid, setActiveImgUid] = useState(0); const imageEditorRef = useRef(null); const { handleOnValuesChange, handleToggleParamsStyle, setParams, updateCacheFormData, setInitialValues, updateParamsConfig, setParamsConfig, form, modelMeta, watchFields, 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, parameters.size]); const handleClear = () => { setCurrentPrompt(''); }; const handleInputChange = (e: any) => { setCurrentPrompt(e.target.value); }; const generateParams = () => { // preview let stream_options: Record = { stream_options_chunk_size: 16 * 1024, stream_options_chunk_result: true }; if (parameters.preview === 'preview') { stream_options = { stream_options_preview: true }; } if (parameters.preview === 'preview_faster') { stream_options = { stream_options_preview_faster: true }; } const params = { ..._.omitBy(finalParameters, (value: string) => !value), seed: parameters.random_seed ? generateRandomNumber() : parameters.seed, stream: true, ...stream_options, 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) => { // setMaskUpload(base64List); const mask = _.get(base64List, '[0].dataUrl', ''); const maskColors = await processImage(mask); console.log('maskColors:', maskColors); imageEditorRef.current?.loadMaskPixs(maskColors || []); // setMask(mask); }, []); 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 ( } disabled={loading} handleUpdateImgList={handleUpdateMaskList} size="middle" accept="image/*" > } > ); } return ( <>
{intl.formatMessage({ id: 'playground.image.edit.tips' })}
); }, [ intl, image, loading, maskUpload, imageStatus, handleOnSave, handleUpdateImageList ]); const handleOnImgClick = useCallback((item: any, isOrigin: boolean) => { if (item.progress < 100 && !isOrigin) { return; } setActiveImgUid(item.uid); setImage(item.dataUrl); setImageStatus({ isOriginal: isOrigin, isResetNeeded: false, width: item.width, height: item.height }); }, []); 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 ( <> handleOnImgClick(uploadList[0], true)} label={ {intl.formatMessage({ id: 'playground.image.origin' })} } > ); }, [uploadList, intl, handleOnImgClick]); const renderMaskImage = useMemo(() => { if (!maskUpload.length) { return null; } return ( <> handleClearUploadMask()} label={ {intl.formatMessage({ id: 'playground.image.mask' })} } > ); }, [maskUpload, intl]); return (
<>
{
{renderImageEditor}
}
{tokenResult && (
)}
{renderMaskImage}
{renderOriginImage} {imageList.length > 0 && ( <>
{_.map(imageList, (item: any, index: number) => { return (
handleOnImgClick(item, false)} >
); })}
)}
{intl.formatMessage({ id: 'playground.parameters' })}
} onValuesChange={handleOnValuesChange} paramsConfig={paramsConfig} initialValues={initialValues} modelList={modelList} />
{intl.formatMessage({ id: 'playground.image.prompt' })} } loading={loading} disabled={!parameters.model} isEmpty={!imageList.length} handleSubmit={handleSendMessage} handleAbortFetch={handleStopConversation} onInputChange={handleInputChange} shouldResetMessage={false} clearAll={handleClear} />
); }); export default GroundImages;