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

Class StableDiffusion

stable_diffusion/stable_diffusion/__init__.py:19–169  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17
18
19class 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:

Callers 2

txt2image.pyFile · 0.90
image2image.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected