refactor: create image hooks

This commit is contained in:
jialin
2025-02-23 13:32:03 +08:00
parent f8c7f88ac2
commit bb8b00e547
6 changed files with 593 additions and 775 deletions
+286 -17
View File
@@ -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
};
};
+33 -14
View File
@@ -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,