GT4SD/diffusers
0
1import logging2import pathlib3import gradio as gr4import pandas as pd5from gt4sd.algorithms.generation.diffusion import (6 DiffusersGenerationAlgorithm,7 DDPMGenerator,8 DDIMGenerator,9 ScoreSdeGenerator,10 LDMTextToImageGenerator,11 LDMGenerator,12 StableDiffusionGenerator,13)14from gt4sd.algorithms.registry import ApplicationsRegistry15 16logger = logging.getLogger(__name__)17logger.addHandler(logging.NullHandler())18 19 20def run_inference(model_type: str, prompt: str):21 22 if prompt == "":23 config = eval(f"{model_type}()")24 else:25 config = eval(f'{model_type}(prompt="{prompt}")')26 if config.modality != "token2image" and prompt != "":27 raise ValueError(28 f"{model_type} is an unconditional generative model, please remove prompt (not={prompt})"29 )30 model = DiffusersGenerationAlgorithm(config)31 image = list(model.sample(1))[0]32 33 return image34 35 36if __name__ == "__main__":37 38 # Preparation (retrieve all available algorithms)39 all_algos = ApplicationsRegistry.list_available()40 algos = [41 x["algorithm_application"]42 for x in list(filter(lambda x: "Diff" in x["algorithm_name"], all_algos))43 ]44 algos = [a for a in algos if not "GeoDiff" in a]45 46 # Load metadata47 metadata_root = pathlib.Path(__file__).parent.joinpath("model_cards")48 49 examples = pd.read_csv(metadata_root.joinpath("examples.csv"), header=None).fillna(50 ""51 )52 53 with open(metadata_root.joinpath("article.md"), "r") as f:54 article = f.read()55 with open(metadata_root.joinpath("description.md"), "r") as f:56 description = f.read()57 58 demo = gr.Interface(59 fn=run_inference,60 title="Diffusion-based image generators",61 inputs=[62 gr.Dropdown(63 algos, label="Diffusion model", value="StableDiffusionGenerator"64 ),65 gr.Textbox(label="Text prompt", placeholder="A blue tree", lines=1),66 ],67 outputs=gr.Image(type="pil"),68 article=article,69 description=description,70 examples=examples.values.tolist(),71 )72 demo.launch(debug=True, show_error=True)73 