yuaiyu/Versatile-Diffusion
0
1from __future__ import annotations2 3import ast4import csv5import inspect6import os7import subprocess8import tempfile9import threading10import warnings11from pathlib import Path12from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Tuple13 14import matplotlib15import matplotlib.pyplot as plt16import numpy as np17import PIL18import PIL.Image19 20import gradio21from gradio import components, processing_utils, routes, utils22from gradio.context import Context23from gradio.documentation import document, set_documentation_group24from gradio.flagging import CSVLogger25 26if TYPE_CHECKING: # Only import for type checking (to avoid circular imports).27 from gradio.components import IOComponent28 29CACHED_FOLDER = "gradio_cached_examples"30LOG_FILE = "log.csv"31 32def create_myexamples(33 examples: List[Any] | List[List[Any]] | str,34 inputs: IOComponent | List[IOComponent],35 outputs: IOComponent | List[IOComponent] | None = None,36 fn: Callable | None = None,37 cache_examples: bool = False,38 examples_per_page: int = 10,39 _api_mode: bool = False,40 label: str | None = None,41 elem_id: str | None = None,42 run_on_click: bool = False,43 preprocess: bool = True,44 postprocess: bool = True,45 batch: bool = False,):46 """Top-level synchronous function that creates Examples. Provided for backwards compatibility, i.e. so that gr.Examples(...) can be used to create the Examples component."""47 examples_obj = MyExamples(48 examples=examples,49 inputs=inputs,50 outputs=outputs,51 fn=fn,52 cache_examples=cache_examples,53 examples_per_page=examples_per_page,54 _api_mode=_api_mode,55 label=label,56 elem_id=elem_id,57 run_on_click=run_on_click,58 preprocess=preprocess,59 postprocess=postprocess,60 batch=batch,61 _initiated_directly=False,62 )63 utils.synchronize_async(examples_obj.create)64 return examples_obj65 66class MyExamples(gradio.helpers.Examples):67 def __init__(68 self,69 examples: List[Any] | List[List[Any]] | str,70 inputs: IOComponent | List[IOComponent],71 outputs: IOComponent | List[IOComponent] | None = None,72 fn: Callable | None = None,73 cache_examples: bool = False,74 examples_per_page: int = 10,75 _api_mode: bool = False,76 label: str | None = "Examples",77 elem_id: str | None = None,78 run_on_click: bool = False,79 preprocess: bool = True,80 postprocess: bool = True,81 batch: bool = False,82 _initiated_directly: bool = True,):83 84 if _initiated_directly:85 warnings.warn(86 "Please use gr.Examples(...) instead of gr.examples.Examples(...) to create the Examples.",87 )88 89 if cache_examples and (fn is None or outputs is None):90 raise ValueError("If caching examples, `fn` and `outputs` must be provided")91 92 if not isinstance(inputs, list):93 inputs = [inputs]94 if outputs and not isinstance(outputs, list):95 outputs = [outputs]96 97 working_directory = Path().absolute()98 99 if examples is None:100 raise ValueError("The parameter `examples` cannot be None")101 elif isinstance(examples, list) and (102 len(examples) == 0 or isinstance(examples[0], list)103 ):104 pass105 elif (106 isinstance(examples, list) and len(inputs) == 1107 ): # If there is only one input component, examples can be provided as a regular list instead of a list of lists108 examples = [[e] for e in examples]109 elif isinstance(examples, str):110 if not Path(examples).exists():111 raise FileNotFoundError(112 "Could not find examples directory: " + examples113 )114 working_directory = examples115 if not (Path(examples) / LOG_FILE).exists():116 if len(inputs) == 1:117 examples = [[e] for e in os.listdir(examples)]118 else:119 raise FileNotFoundError(120 "Could not find log file (required for multiple inputs): "121 + LOG_FILE122 )123 else:124 with open(Path(examples) / LOG_FILE) as logs:125 examples = list(csv.reader(logs))126 examples = [127 examples[i][: len(inputs)] for i in range(1, len(examples))128 ] # remove header and unnecessary columns129 130 else:131 raise ValueError(132 "The parameter `examples` must either be a string directory or a list"133 "(if there is only 1 input component) or (more generally), a nested "134 "list, where each sublist represents a set of inputs."135 )136 137 input_has_examples = [False] * len(inputs)138 for example in examples:139 for idx, example_for_input in enumerate(example):140 # if not (example_for_input is None):141 if True:142 try:143 input_has_examples[idx] = True144 except IndexError:145 pass # If there are more example components than inputs, ignore. This can sometimes be intentional (e.g. loading from a log file where outputs and timestamps are also logged)146 147 inputs_with_examples = [148 inp for (inp, keep) in zip(inputs, input_has_examples) if keep149 ]150 non_none_examples = [151 [ex for (ex, keep) in zip(example, input_has_examples) if keep]152 for example in examples153 ]154 155 self.examples = examples156 self.non_none_examples = non_none_examples157 self.inputs = inputs158 self.inputs_with_examples = inputs_with_examples159 self.outputs = outputs160 self.fn = fn161 self.cache_examples = cache_examples162 self._api_mode = _api_mode163 self.preprocess = preprocess164 self.postprocess = postprocess165 self.batch = batch166 167 with utils.set_directory(working_directory):168 self.processed_examples = [169 [170 component.postprocess(sample)171 for component, sample in zip(inputs, example)172 ]173 for example in examples174 ]175 self.non_none_processed_examples = [176 [ex for (ex, keep) in zip(example, input_has_examples) if keep]177 for example in self.processed_examples178 ]179 if cache_examples:180 for example in self.examples:181 if len([ex for ex in example if ex is not None]) != len(self.inputs):182 warnings.warn(183 "Examples are being cached but not all input components have "184 "example values. This may result in an exception being thrown by "185 "your function. If you do get an error while caching examples, make "186 "sure all of your inputs have example values for all of your examples "187 "or you provide default values for those particular parameters in your function."188 )189 break190 191 with utils.set_directory(working_directory):192 self.dataset = components.Dataset(193 components=inputs_with_examples,194 samples=non_none_examples,195 type="index",196 label=label,197 samples_per_page=examples_per_page,198 elem_id=elem_id,199 )200 201 self.cached_folder = Path(CACHED_FOLDER) / str(self.dataset._id)202 self.cached_file = Path(self.cached_folder) / "log.csv"203 self.cache_examples = cache_examples204 self.run_on_click = run_on_click205 206from gradio import utils, processing_utils207from PIL import Image as _Image208from pathlib import Path209import numpy as np210 211def customized_postprocess(self, y):212 if y is None:213 return None214 215 if isinstance(y, dict):216 if self.tool == "sketch" and self.source in ["upload", "webcam"]:217 y, mask = y["image"], y["mask"]218 if y is None:219 return None220 elif isinstance(y, np.ndarray):221 im = processing_utils.encode_array_to_base64(y)222 elif isinstance(y, _Image.Image):223 im = processing_utils.encode_pil_to_base64(y)224 elif isinstance(y, (str, Path)):225 im = processing_utils.encode_url_or_file_to_base64(y)226 else:227 raise ValueError("Cannot process this value as an Image")228 im = self._format_image(im)229 230 if mask is None:231 return im232 elif isinstance(y, np.ndarray):233 mask_im = processing_utils.encode_array_to_base64(mask)234 elif isinstance(y, _Image.Image):235 mask_im = processing_utils.encode_pil_to_base64(mask)236 elif isinstance(y, (str, Path)):237 mask_im = processing_utils.encode_url_or_file_to_base64(mask)238 else:239 raise ValueError("Cannot process this value as an Image")240 241 return {"image": im, "mask" : mask_im,}242 243 elif isinstance(y, np.ndarray):244 return processing_utils.encode_array_to_base64(y)245 elif isinstance(y, _Image.Image):246 return processing_utils.encode_pil_to_base64(y)247 elif isinstance(y, (str, Path)):248 return processing_utils.encode_url_or_file_to_base64(y)249 else:250 raise ValueError("Cannot process this value as an Image")251 252# def customized_as_example(self, input_data=None):253# if input_data is None:254# return str('assets/demo/misc/noimage.jpg')255# elif isinstance(input_data, dict):256# im = np.array(PIL.Image.open(input_data["image"])).astype(float)257# mask = np.array(PIL.Image.open(input_data["mask"])).astype(float)/255258# imm = (im * (1-mask)).astype(np.uint8)259# import time260# ctime = int(time.time()*100)261# impath = 'assets/demo/temp/temp_{}.png'.format(ctime)262# PIL.Image.fromarray(imm).save(impath)263# return str(utils.abspath(impath))264# else:265# return str(utils.abspath(input_data))266 267def customized_as_example(self, input_data=None):268 if input_data is None:269 return str('assets/demo/misc/noimage.jpg')270 else:271 return str(utils.abspath(input_data))272 