multimodalart HF Staff commited on
Commit
c765e3a
·
verified ·
1 Parent(s): 71083a9

Fix ZeroGPU/diffusers-drift breakage: controlnet channel mismatch, import paths, CUDA fork

Browse files

Root cause: current diffusers dropped extra_condition_channels support in
FluxControlNetModel (breaks alimama-creative FLUX.1-dev-Controlnet-Inpainting-Beta,
which needs 68 input channels not 64), moved zero_module/BaseOutput import paths,
and changed FluxAttnProcessor2_0's rotary-embedding format to a (cos, sin) tuple.

Fixes:
- Wire app.py to the vendored controlnet_flux.FluxControlNetModel and
pipeline_flux_controlnet_inpaint.FluxControlNetInpaintingPipeline (already in
the repo, already used by main.py) instead of diffusers built-ins.
- Fix diffusers.models.controlnet -> diffusers.models.controlnets.controlnet import.
- Import spaces before torch (avoids CUDA init in parent process poisoning the
ZeroGPU worker fork).
- Drop unused optimum-quanto import/dependency (its CUDA extension import was
triggering a real cuInit in the parent process).
- Fix run_pipe() call signature mismatch in inpaint() (pre-existing bug).
- Swap vendored EmbedND rope for diffusers FluxPosEmbed (matching tuple format)
and strip the 3D batch dim from txt_ids/img_ids before pos_embed, mirroring
diffusers upstream, since FluxPosEmbed expects 2D ids.

Files changed (4) hide show
  1. app.py +10 -14
  2. controlnet_flux.py +9 -7
  3. requirements.txt +1 -2
  4. transformer_flux.py +7 -5
app.py CHANGED
@@ -1,14 +1,15 @@
 
1
  import gradio as gr
2
  import torch
3
- import spaces
4
- from diffusers import AutoencoderKL, FluxTransformer2DModel, FluxControlNetModel, FluxControlNetInpaintPipeline
5
  from diffusers.utils import load_image
6
  from transformer_flux import FluxTransformer2DModel
 
 
7
  from transformers import T5EncoderModel, CLIPTextModel
8
  from PIL import Image, ImageDraw
9
  import numpy as np
10
  from huggingface_hub import hf_hub_download
11
- from optimum.quanto import freeze, qfloat8, quantize
12
 
13
 
14
  controlnet = FluxControlNetModel.from_pretrained("alimama-creative/FLUX.1-dev-Controlnet-Inpainting-Beta", torch_dtype=torch.bfloat16)
@@ -16,7 +17,7 @@ transformer = FluxTransformer2DModel.from_pretrained(
16
  "black-forest-labs/FLUX.1-dev", subfolder='transformer', torch_dtype=torch.bfloat16
17
  )
18
 
19
- pipe = FluxControlNetInpaintPipeline.from_pretrained(
20
  "black-forest-labs/FLUX.1-dev",
21
  transformer=transformer,
22
  controlnet=controlnet,
@@ -165,19 +166,14 @@ def inpaint(image, width, height, overlap_percentage, num_inference_steps, resiz
165
 
166
  #generator = torch.Generator(device="cuda").manual_seed(42)
167
  result = run_pipe(
168
- prompt=final_prompt,
169
  height=height,
170
  width=width,
171
- control_image=cnet_image,
172
- control_mask=mask,
173
  num_inference_steps=num_inference_steps,
174
- #generator=generator,
175
- controlnet_conditioning_scale=0.9,
176
- guidance_scale=3.5,
177
- negative_prompt="",
178
- true_guidance_scale=3.5,
179
- ).images[0]
180
-
181
  result = result.convert("RGBA")
182
  cnet_image.paste(result, (0, 0), mask)
183
 
 
1
+ import spaces
2
  import gradio as gr
3
  import torch
4
+ from diffusers import AutoencoderKL
 
5
  from diffusers.utils import load_image
6
  from transformer_flux import FluxTransformer2DModel
7
+ from controlnet_flux import FluxControlNetModel
8
+ from pipeline_flux_controlnet_inpaint import FluxControlNetInpaintingPipeline
9
  from transformers import T5EncoderModel, CLIPTextModel
10
  from PIL import Image, ImageDraw
11
  import numpy as np
12
  from huggingface_hub import hf_hub_download
 
13
 
14
 
15
  controlnet = FluxControlNetModel.from_pretrained("alimama-creative/FLUX.1-dev-Controlnet-Inpainting-Beta", torch_dtype=torch.bfloat16)
 
17
  "black-forest-labs/FLUX.1-dev", subfolder='transformer', torch_dtype=torch.bfloat16
18
  )
19
 
20
+ pipe = FluxControlNetInpaintingPipeline.from_pretrained(
21
  "black-forest-labs/FLUX.1-dev",
22
  transformer=transformer,
23
  controlnet=controlnet,
 
166
 
167
  #generator = torch.Generator(device="cuda").manual_seed(42)
168
  result = run_pipe(
169
+ final_prompt=final_prompt,
170
  height=height,
171
  width=width,
172
+ cnet_image=cnet_image,
173
+ mask=mask,
174
  num_inference_steps=num_inference_steps,
175
+ )
176
+
 
 
 
 
 
177
  result = result.convert("RGBA")
178
  cnet_image.paste(result, (0, 0), mask)
179
 
controlnet_flux.py CHANGED
@@ -15,14 +15,15 @@ from diffusers.utils import (
15
  scale_lora_layers,
16
  unscale_lora_layers,
17
  )
18
- from diffusers.models.controlnet import BaseOutput, zero_module
 
19
  from diffusers.models.embeddings import (
20
  CombinedTimestepGuidanceTextProjEmbeddings,
21
  CombinedTimestepTextProjEmbeddings,
22
  )
23
  from diffusers.models.modeling_outputs import Transformer2DModelOutput
 
24
  from transformer_flux import (
25
- EmbedND,
26
  FluxSingleTransformerBlock,
27
  FluxTransformerBlock,
28
  )
@@ -59,9 +60,7 @@ class FluxControlNetModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
59
  self.out_channels = in_channels
60
  self.inner_dim = num_attention_heads * attention_head_dim
61
 
62
- self.pos_embed = EmbedND(
63
- dim=self.inner_dim, theta=10000, axes_dim=axes_dims_rope
64
- )
65
  text_time_guidance_cls = (
66
  CombinedTimestepGuidanceTextProjEmbeddings
67
  if guidance_embeds
@@ -293,8 +292,11 @@ class FluxControlNetModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
293
  )
294
  encoder_hidden_states = self.context_embedder(encoder_hidden_states)
295
 
296
- txt_ids = txt_ids.expand(img_ids.size(0), -1, -1)
297
- ids = torch.cat((txt_ids, img_ids), dim=1)
 
 
 
298
  image_rotary_emb = self.pos_embed(ids)
299
 
300
  block_samples = ()
 
15
  scale_lora_layers,
16
  unscale_lora_layers,
17
  )
18
+ from diffusers.utils import BaseOutput
19
+ from diffusers.models.controlnets.controlnet import zero_module
20
  from diffusers.models.embeddings import (
21
  CombinedTimestepGuidanceTextProjEmbeddings,
22
  CombinedTimestepTextProjEmbeddings,
23
  )
24
  from diffusers.models.modeling_outputs import Transformer2DModelOutput
25
+ from diffusers.models.transformers.transformer_flux import FluxPosEmbed
26
  from transformer_flux import (
 
27
  FluxSingleTransformerBlock,
28
  FluxTransformerBlock,
29
  )
 
60
  self.out_channels = in_channels
61
  self.inner_dim = num_attention_heads * attention_head_dim
62
 
63
+ self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope)
 
 
64
  text_time_guidance_cls = (
65
  CombinedTimestepGuidanceTextProjEmbeddings
66
  if guidance_embeds
 
292
  )
293
  encoder_hidden_states = self.context_embedder(encoder_hidden_states)
294
 
295
+ if txt_ids.ndim == 3:
296
+ txt_ids = txt_ids[0]
297
+ if img_ids.ndim == 3:
298
+ img_ids = img_ids[0]
299
+ ids = torch.cat((txt_ids, img_ids), dim=0)
300
  image_rotary_emb = self.pos_embed(ids)
301
 
302
  block_samples = ()
requirements.txt CHANGED
@@ -4,5 +4,4 @@ transformers
4
  safetensors
5
  accelerate
6
  sentencepiece
7
- peft
8
- optimum-quanto
 
4
  safetensors
5
  accelerate
6
  sentencepiece
7
+ peft
 
transformer_flux.py CHANGED
@@ -32,6 +32,7 @@ from diffusers.models.embeddings import (
32
  CombinedTimestepTextProjEmbeddings,
33
  )
34
  from diffusers.models.modeling_outputs import Transformer2DModelOutput
 
35
 
36
 
37
  logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -293,9 +294,7 @@ class FluxTransformer2DModel(
293
  self.config.num_attention_heads * self.config.attention_head_dim
294
  )
295
 
296
- self.pos_embed = EmbedND(
297
- dim=self.inner_dim, theta=10000, axes_dim=axes_dims_rope
298
- )
299
  text_time_guidance_cls = (
300
  CombinedTimestepGuidanceTextProjEmbeddings
301
  if guidance_embeds
@@ -417,8 +416,11 @@ class FluxTransformer2DModel(
417
  )
418
  encoder_hidden_states = self.context_embedder(encoder_hidden_states)
419
 
420
- txt_ids = txt_ids.expand(img_ids.size(0), -1, -1)
421
- ids = torch.cat((txt_ids, img_ids), dim=1)
 
 
 
422
  image_rotary_emb = self.pos_embed(ids)
423
 
424
  for index_block, block in enumerate(self.transformer_blocks):
 
32
  CombinedTimestepTextProjEmbeddings,
33
  )
34
  from diffusers.models.modeling_outputs import Transformer2DModelOutput
35
+ from diffusers.models.transformers.transformer_flux import FluxPosEmbed
36
 
37
 
38
  logger = logging.get_logger(__name__) # pylint: disable=invalid-name
 
294
  self.config.num_attention_heads * self.config.attention_head_dim
295
  )
296
 
297
+ self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_dims_rope)
 
 
298
  text_time_guidance_cls = (
299
  CombinedTimestepGuidanceTextProjEmbeddings
300
  if guidance_embeds
 
416
  )
417
  encoder_hidden_states = self.context_embedder(encoder_hidden_states)
418
 
419
+ if txt_ids.ndim == 3:
420
+ txt_ids = txt_ids[0]
421
+ if img_ids.ndim == 3:
422
+ img_ids = img_ids[0]
423
+ ids = torch.cat((txt_ids, img_ids), dim=0)
424
  image_rotary_emb = self.pos_embed(ids)
425
 
426
  for index_block, block in enumerate(self.transformer_blocks):