refactor: create image hooks
This commit is contained in:
@@ -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<string, any> = {
|
||||
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<string>[];
|
||||
loaded?: boolean;
|
||||
ref?: any;
|
||||
}
|
||||
|
||||
export const useInitImageMeta = (props: MessageProps) => {
|
||||
const { modelList } = props;
|
||||
const form = useRef<any>(null);
|
||||
const [searchParams] = useSearchParams();
|
||||
const selectModel = searchParams.get('model') || '';
|
||||
const [modelMeta, setModelMeta] = useState<any>({});
|
||||
const [isOpenaiCompatible, setIsOpenaiCompatible] = useState<boolean>(false);
|
||||
const [imageSizeOptions, setImageSizeOptions] = React.useState<
|
||||
ImageSizeItem[]
|
||||
>([]);
|
||||
const [basicParamsConfig, setBasicParamsConfig] = React.useState<
|
||||
ParamsSchema[]
|
||||
>([...ImageCountConfig, ...ImageSizeConfig]);
|
||||
|
||||
const [initialValues, setInitialValues] = useState<any>({
|
||||
...imgInitialValues,
|
||||
...advancedFieldsDefaultValus,
|
||||
model: selectModel
|
||||
});
|
||||
const [paramsConfig, setParamsConfig] = useState<ParamsSchema[]>([
|
||||
...ImageCountConfig,
|
||||
...ImageSizeConfig,
|
||||
...ImageAdvancedParamsConfig
|
||||
]);
|
||||
|
||||
const [parameters, setParams] = useState<any>({
|
||||
...imgInitialValues,
|
||||
...advancedFieldsDefaultValus,
|
||||
model: selectModel
|
||||
});
|
||||
|
||||
const cacheFormData = React.useRef<Record<string, any>>({
|
||||
...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<string, any>) => {
|
||||
_.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<string, any>, allValues: Record<string, any>) => {
|
||||
// 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
|
||||
};
|
||||
};
|
||||
|
||||
@@ -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<any>(null);
|
||||
const [imageList, setImageList] = useState<any[]>([]);
|
||||
const [currentPrompt, setCurrentPrompt] = useState('');
|
||||
const messageId = useRef<number>(0);
|
||||
const requestToken = useRef<any>(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<string, string | number> = {
|
||||
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<string, any> = {
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user