chansung/sd-kerascv
0
1from typing import Dict, List, Any2import sys3import base644import logging5import keras_cv6 7class EndpointHandler():8 def __init__(self, path="", version="2"):9 self.sd = self._instantiate_stable_diffusion(version)10 11 if isinstance(self.sd, str):12 sys.exit(self.sd)13 else:14 self.sd.text_to_image("test prompt", batch_size=1)15 logging.warning(f"Stable Diffusion v{version} is fully loaded")16 17 def _instantiate_stable_diffusion(self, version: str):18 if version is "1.4":19 return keras_cv.models.StableDiffusion(img_width=512, img_height=512)20 elif version is "2":21 return keras_cv.models.StableDiffusionV2(img_width=512, img_height=512)22 else:23 return f"v{version} is not supported"24 25 def __call__(self, data: Dict[str, Any]) -> str:26 prompt = data.pop("inputs", data)27 batch_size = data.pop("batch_size", 1)28 29 images = self.sd.text_to_image(prompt, batch_size=batch_size)30 return base64.b64encode(images.tobytes()).decode()31 