AbstractPhil commited on
Commit
c54bfc9
·
verified ·
1 Parent(s): f10f931

Update pipeline/pipeline.py

Browse files
Files changed (1) hide show
  1. pipeline/pipeline.py +41 -44
pipeline/pipeline.py CHANGED
@@ -1,61 +1,58 @@
 
 
1
  import torch
2
  from diffusers import StableDiffusionPipeline
3
 
4
  class OmegaDiffusionPipeline(StableDiffusionPipeline):
5
- def __init__(
6
- self,
7
- vae,
8
- text_encoder,
9
- tokenizer,
10
- unet,
11
- scheduler,
12
- bridge,
13
- safety_checker=None,
14
- feature_extractor=None,
15
- **kwargs,
16
- ):
17
- super().__init__(
18
- vae=vae,
19
- text_encoder=text_encoder,
20
- tokenizer=tokenizer,
21
- unet=unet,
22
- scheduler=scheduler,
23
- safety_checker=safety_checker,
24
- feature_extractor=feature_extractor,
25
- **kwargs,
26
- )
27
  self.register_modules(bridge=bridge)
28
 
29
- # --- Monkey-patch the VAE.decode method ---
30
- orig_decode = self.vae.decode
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
 
32
- def decode_with_bridge(z_scaled, *args, return_dict=False, generator=None, **decode_kwargs):
33
- # z_scaled = z4 / scaling_factor
34
- sc = self.vae.config.scaling_factor
35
- z4 = z_scaled * sc # recover 4-channel latent
36
- z16 = self.bridge.dec(z4) # map 4→16 channels
37
- z16_scaled = z16 / sc # rescale for VAE
38
- # call original decoder
39
- return orig_decode(z16_scaled, *args, return_dict=return_dict, generator=generator, **decode_kwargs)
40
 
41
- # Replace VAE.decode at runtime
42
- self.vae.decode = decode_with_bridge
43
 
44
  @property
45
  def components(self):
46
  return {
47
- "scheduler": self.scheduler,
48
- "tokenizer": self.tokenizer,
49
- "vae": self.vae,
50
- "unet": self.unet,
51
- "text_encoder": self.text_encoder,
52
- "bridge": self.bridge,
53
- "feature_extractor": self.feature_extractor,
54
- "safety_checker": self.safety_checker,
55
- "kwargs": {},
56
  }
57
 
58
  @torch.no_grad()
59
  def __call__(self, *args, **kwargs):
60
- # let the parent handle everything, including calling vae.decode (now patched)
61
  return super().__call__(*args, **kwargs)
 
1
+ # pipeline/pipeline.py
2
+
3
  import torch
4
  from diffusers import StableDiffusionPipeline
5
 
6
  class OmegaDiffusionPipeline(StableDiffusionPipeline):
7
+ def __init__(self, vae, text_encoder, tokenizer, unet, scheduler, bridge,
8
+ safety_checker=None, feature_extractor=None, **kwargs):
9
+ super().__init__(vae=vae, text_encoder=text_encoder, tokenizer=tokenizer,
10
+ unet=unet, scheduler=scheduler,
11
+ safety_checker=safety_checker,
12
+ feature_extractor=feature_extractor, **kwargs)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  self.register_modules(bridge=bridge)
14
 
15
+ # Capture the original decode
16
+ _orig_decode = self.vae.decode
17
+ in_ch = self.unet.config.in_channels
18
+ sc = self.vae.config.scaling_factor
19
+
20
+ def _decode_with_bridge(z_scaled, *args, return_dict=False, generator=None, **decode_kwargs):
21
+ # z_scaled is latents/scaling_factor
22
+ # Reconstruct the raw latent
23
+ z = z_scaled * sc
24
+
25
+ # If it’s 4-ch, run the bridge → 16-ch
26
+ if z.shape[1] == in_ch:
27
+ z = self.bridge.dec(z)
28
+
29
+ # Rescale for the VAE’s conv_in
30
+ z_scaled2 = z / sc
31
 
32
+ # Now call the real VAE.decode on exactly 16 channels
33
+ return _orig_decode(z_scaled2, *args,
34
+ return_dict=return_dict,
35
+ generator=generator,
36
+ **decode_kwargs)
 
 
 
37
 
38
+ # Override the instance’s decode method
39
+ self.vae.decode = _decode_with_bridge
40
 
41
  @property
42
  def components(self):
43
  return {
44
+ "scheduler": self.scheduler,
45
+ "tokenizer": self.tokenizer,
46
+ "vae": self.vae,
47
+ "unet": self.unet,
48
+ "text_encoder": self.text_encoder,
49
+ "bridge": self.bridge,
50
+ "feature_extractor": self.feature_extractor,
51
+ "safety_checker": self.safety_checker,
52
+ "kwargs": {},
53
  }
54
 
55
  @torch.no_grad()
56
  def __call__(self, *args, **kwargs):
57
+ # Defer entirely to the parent, which will call our patched decode
58
  return super().__call__(*args, **kwargs)