MCPcopy Create free account
hub / github.com/ml-explore/mlx-examples / StableDiffusionXL

Class StableDiffusionXL

stable_diffusion/stable_diffusion/__init__.py:172–306  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

170
171
172class StableDiffusionXL(StableDiffusion):
173 def __init__(self, model: str = _DEFAULT_MODEL, float16: bool = False):
174 super().__init__(model, float16)
175
176 self.sampler = SimpleEulerAncestralSampler(self.diffusion_config)
177
178 self.text_encoder_1 = self.text_encoder
179 self.tokenizer_1 = self.tokenizer
180 del self.tokenizer, self.text_encoder
181
182 self.text_encoder_2 = load_text_encoder(
183 model,
184 float16,
185 model_key="text_encoder_2",
186 )
187 self.tokenizer_2 = load_tokenizer(
188 model,
189 merges_key="tokenizer_2_merges",
190 vocab_key="tokenizer_2_vocab",
191 )
192
193 def ensure_models_are_loaded(self):
194 mx.eval(self.unet.parameters())
195 mx.eval(self.text_encoder_1.parameters())
196 mx.eval(self.text_encoder_2.parameters())
197 mx.eval(self.autoencoder.parameters())
198
199 def _get_text_conditioning(
200 self,
201 text: str,
202 n_images: int = 1,
203 cfg_weight: float = 7.5,
204 negative_text: str = "",
205 ):
206 tokens_1 = self._tokenize(
207 self.tokenizer_1,
208 text,
209 (negative_text if cfg_weight > 1 else None),
210 )
211 tokens_2 = self._tokenize(
212 self.tokenizer_2,
213 text,
214 (negative_text if cfg_weight > 1 else None),
215 )
216
217 conditioning_1 = self.text_encoder_1(tokens_1)
218 conditioning_2 = self.text_encoder_2(tokens_2)
219 conditioning = mx.concatenate(
220 [conditioning_1.hidden_states[-2], conditioning_2.hidden_states[-2]],
221 axis=-1,
222 )
223 pooled_conditioning = conditioning_2.pooled_output
224
225 if n_images > 1:
226 conditioning = mx.repeat(conditioning, n_images, axis=0)
227 pooled_conditioning = mx.repeat(pooled_conditioning, n_images, axis=0)
228
229 return conditioning, pooled_conditioning

Callers 2

txt2image.pyFile · 0.90
image2image.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected