| 17 | |
| 18 | |
| 19 | class StableDiffusion: |
| 20 | def __init__(self, model: str = _DEFAULT_MODEL, float16: bool = False): |
| 21 | self.dtype = mx.float16 if float16 else mx.float32 |
| 22 | self.diffusion_config = load_diffusion_config(model) |
| 23 | self.unet = load_unet(model, float16) |
| 24 | self.text_encoder = load_text_encoder(model, float16) |
| 25 | self.autoencoder = load_autoencoder(model, False) |
| 26 | self.sampler = SimpleEulerSampler(self.diffusion_config) |
| 27 | self.tokenizer = load_tokenizer(model) |
| 28 | |
| 29 | def ensure_models_are_loaded(self): |
| 30 | mx.eval(self.unet.parameters()) |
| 31 | mx.eval(self.text_encoder.parameters()) |
| 32 | mx.eval(self.autoencoder.parameters()) |
| 33 | |
| 34 | def _tokenize(self, tokenizer, text: str, negative_text: Optional[str] = None): |
| 35 | # Tokenize the text |
| 36 | tokens = [tokenizer.tokenize(text)] |
| 37 | if negative_text is not None: |
| 38 | tokens += [tokenizer.tokenize(negative_text)] |
| 39 | lengths = [len(t) for t in tokens] |
| 40 | N = max(lengths) |
| 41 | tokens = [t + [0] * (N - len(t)) for t in tokens] |
| 42 | tokens = mx.array(tokens) |
| 43 | |
| 44 | return tokens |
| 45 | |
| 46 | def _get_text_conditioning( |
| 47 | self, |
| 48 | text: str, |
| 49 | n_images: int = 1, |
| 50 | cfg_weight: float = 7.5, |
| 51 | negative_text: str = "", |
| 52 | ): |
| 53 | # Tokenize the text |
| 54 | tokens = self._tokenize( |
| 55 | self.tokenizer, text, (negative_text if cfg_weight > 1 else None) |
| 56 | ) |
| 57 | |
| 58 | # Compute the features |
| 59 | conditioning = self.text_encoder(tokens).last_hidden_state |
| 60 | |
| 61 | # Repeat the conditioning for each of the generated images |
| 62 | if n_images > 1: |
| 63 | conditioning = mx.repeat(conditioning, n_images, axis=0) |
| 64 | |
| 65 | return conditioning |
| 66 | |
| 67 | def _denoising_step( |
| 68 | self, x_t, t, t_prev, conditioning, cfg_weight: float = 7.5, text_time=None |
| 69 | ): |
| 70 | x_t_unet = mx.concatenate([x_t] * 2, axis=0) if cfg_weight > 1 else x_t |
| 71 | t_unet = mx.broadcast_to(t, [len(x_t_unet)]) |
| 72 | eps_pred = self.unet( |
| 73 | x_t_unet, t_unet, encoder_x=conditioning, text_time=text_time |
| 74 | ) |
| 75 | |
| 76 | if cfg_weight > 1: |
no outgoing calls
no test coverage detected