mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
add base flex.2 support
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -21,6 +21,7 @@ ignore-paths=/usr/lib/.*$,
|
||||
modules/intel/ipex,
|
||||
modules/intel/openvino,
|
||||
modules/k-diffusion,
|
||||
modules/flex2,
|
||||
modules/ldsr,
|
||||
modules/meissonic,
|
||||
modules/mod,
|
||||
|
||||
@@ -8,6 +8,7 @@ exclude = [
|
||||
"modules/control/proc",
|
||||
"modules/control/units",
|
||||
"modules/freescale",
|
||||
"modules/flex2",
|
||||
"modules/ggml",
|
||||
"modules/hidiffusion",
|
||||
"modules/hijack",
|
||||
|
||||
+6
-1
@@ -1,10 +1,11 @@
|
||||
# Change Log for SD.Next
|
||||
|
||||
## Update for 2025-04-23
|
||||
## Update for 2025-04-24
|
||||
|
||||
- **Features**
|
||||
- [Nunchaku](https://github.com/mit-han-lab/nunchaku) inference engine with custom **SVDQuant** 4-bit execution
|
||||
highly experimental and with limited support, but when it works, its magic: **Flux.1 at 6.0 it/s** *(not sec/it)*!
|
||||
basically, it can speed up supported models by 2-5x by using custom quantization and execution engine
|
||||
see [Nunchaku Wiki](https://github.com/vladmandic/sdnext/wiki/Nunchaku) for installation guide and list of supported models & features
|
||||
- [FramePack](https://github.com/vladmandic/sd-extension-framepack) based on **HunyuanVideo-I2V**
|
||||
full support and much more for **Lllyasviel** [FramePack](https://lllyasviel.github.io/frame_pack_gitpage/)
|
||||
@@ -16,6 +17,10 @@
|
||||
- custom models: e.g. replace llama with one of your choice
|
||||
- multiple video codecs and with hw acceleration, raw export, frame export, frame interpolation
|
||||
- quantization support, new offloading, more configuration options, cross-platform, etc.
|
||||
- [Ostris Flex.2 Preview](https://huggingface.co/ostris/Flex.2-preview)
|
||||
more than a FLUX.1 finetune, FLEX.2 is created from *Flux.1 Schnell -> OpenFlux.1 -> Flex.1-alpha -> Flex.2-preview*
|
||||
and it has universal control and inpainting support built in!
|
||||
available via *networks -> models -> reference*
|
||||
- [LTXVideo 0.9.6](https://github.com/Lightricks/LTX-Video?tab=readme-ov-file) **T2V** and **I2V**
|
||||
in both **Standard** and **Distilled** variants
|
||||
available in *video tab*
|
||||
|
||||
@@ -180,6 +180,13 @@
|
||||
"extras": "sampler: Default, cfg_scale: 3.5"
|
||||
},
|
||||
|
||||
"Ostris Flex.2 Preview": {
|
||||
"path": "ostris/Flex.2-preview",
|
||||
"preview": "ostris--Flex.2-preview.jpg",
|
||||
"desc": "Open Source 8B parameter Text to Image Diffusion Model with universal control and inpainting support built in. Early access preview release. The next version of Flex.1-alpha",
|
||||
"skip": true,
|
||||
"extras": "sampler: Default, cfg_scale: 3.5"
|
||||
},
|
||||
"Ostris Flex.1 Alpha": {
|
||||
"path": "ostris/Flex.1-alpha",
|
||||
"preview": "ostris--Flex.1-alpha.jpg",
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 58 KiB |
@@ -0,0 +1,435 @@
|
||||
from diffusers import FluxControlPipeline, FluxTransformer2DModel
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
import torch
|
||||
|
||||
from diffusers.image_processor import PipelineImageInput
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
|
||||
from diffusers.pipelines.flux.pipeline_flux import calculate_shift, retrieve_timesteps, XLA_AVAILABLE
|
||||
|
||||
|
||||
class Flex2Pipeline(FluxControlPipeline):
|
||||
def __init__(
|
||||
self,
|
||||
scheduler,
|
||||
vae,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
text_encoder_2,
|
||||
tokenizer_2,
|
||||
transformer,
|
||||
):
|
||||
super().__init__(scheduler, vae, text_encoder, tokenizer, text_encoder_2, tokenizer_2, transformer)
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
pooled_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
max_sequence_length=None,
|
||||
inpaint_image=None,
|
||||
inpaint_mask=None,
|
||||
control_image=None,
|
||||
):
|
||||
super().check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
if inpaint_image is not None and inpaint_mask is None:
|
||||
raise ValueError(
|
||||
"If `inpaint_image` is passed, `inpaint_mask` must be passed as well. "
|
||||
"Please make sure to pass both `inpaint_image` and `inpaint_mask`."
|
||||
)
|
||||
if inpaint_mask is not None and inpaint_image is None:
|
||||
raise ValueError(
|
||||
"If `inpaint_mask` is passed, `inpaint_image` must be passed as well. "
|
||||
"Please make sure to pass both `inpaint_image` and `inpaint_mask`."
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
inpaint_image: Optional[PipelineImageInput] = None,
|
||||
inpaint_mask: Optional[PipelineImageInput] = None,
|
||||
control_image: Optional[PipelineImageInput] = None,
|
||||
control_strength: Optional[float] = 1.0,
|
||||
control_stop: Optional[float] = 1.0,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
guidance_scale: float = 3.5,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead
|
||||
inpaint_image (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,:
|
||||
`List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`):
|
||||
The image to be inpainted.
|
||||
inpaint_mask (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,:
|
||||
`List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`):
|
||||
A black and white mask to be used for inpainting. The white pixels are the areas to be inpainted, while the
|
||||
black pixels are the areas to be kept.
|
||||
control_image (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,:
|
||||
`List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`):
|
||||
The control image (line, depth, pose, etc.) to be used for the generation. The control image
|
||||
control_strength (`float`, *optional*, defaults to 1.0):
|
||||
The strength of the control image. The higher the value, the more the control image will be used to
|
||||
guide the generation. The lower the value, the less the control image will be used to guide the
|
||||
generation.
|
||||
control_stop (`float`, *optional*, defaults to 1.0):
|
||||
The percentage of the generation to drop out the control. 0.0 to 1.0. 0.5 mean the control will be dropped
|
||||
out at 50% of the generation.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, *optional*, defaults to 3.5):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
|
||||
joint_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
|
||||
is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
|
||||
images.
|
||||
"""
|
||||
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._joint_attention_kwargs = joint_attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 3. Prepare text embeddings
|
||||
lora_scale = (
|
||||
self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None
|
||||
)
|
||||
(
|
||||
prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
text_ids,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels // 4
|
||||
|
||||
# only prepare latents for non controls
|
||||
# (16 + 1 + 16 )
|
||||
num_control_channels = 33
|
||||
num_channels_latents = num_channels_latents - num_control_channels
|
||||
|
||||
control_latents = None
|
||||
inpaint_latents = None
|
||||
inpaint_latents_mask = None
|
||||
|
||||
latent_height = height // self.vae_scale_factor
|
||||
latent_width = width // self.vae_scale_factor
|
||||
|
||||
# process the control and inpaint channels
|
||||
|
||||
if control_image is None:
|
||||
control_latents = torch.zeros(
|
||||
batch_size * num_images_per_prompt,
|
||||
16,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
else:
|
||||
control_image = self.prepare_image(
|
||||
image=control_image,
|
||||
width=width,
|
||||
height=height,
|
||||
batch_size=batch_size * num_images_per_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
control_image = self.vae.encode(control_image).latent_dist.sample(generator=generator)
|
||||
control_latents = (control_image - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
|
||||
# apply control strength
|
||||
control_latents = control_latents * control_strength
|
||||
|
||||
if inpaint_image is None and inpaint_mask is None:
|
||||
inpaint_latents = torch.zeros(
|
||||
batch_size * num_images_per_prompt,
|
||||
16,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
inpaint_latents_mask = torch.ones(
|
||||
batch_size * num_images_per_prompt,
|
||||
1,
|
||||
latent_height,
|
||||
latent_width,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
else:
|
||||
inpaint_image = self.prepare_image(
|
||||
image=inpaint_image,
|
||||
width=width,
|
||||
height=height,
|
||||
batch_size=batch_size * num_images_per_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
inpaint_image = self.vae.encode(inpaint_image).latent_dist.sample(generator=generator)
|
||||
inpaint_latents = (inpaint_image - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
height_inpaint_image, width_inpaint_image = control_image.shape[2:]
|
||||
|
||||
inpaint_mask = self.prepare_image(
|
||||
image=inpaint_mask,
|
||||
width=width,
|
||||
height=height,
|
||||
batch_size=batch_size * num_images_per_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
# mask is 3 ch -1 to 1. make it 1ch, 0 to 1
|
||||
inpaint_mask = inpaint_mask[:, 0:1, :, :] * 0.5 + 0.5
|
||||
# resize to match height_inpaint_image and width_inpaint_image
|
||||
inpaint_latents_mask = F.interpolate(inpaint_mask, size=(height_inpaint_image, width_inpaint_image), mode="bilinear", align_corners=False)
|
||||
|
||||
# apply inverted mask to inpaint latents
|
||||
inpaint_latents = inpaint_latents * (1 - inpaint_latents_mask)
|
||||
|
||||
# concat the latent controls on the channel dimension every step
|
||||
latent_controls = torch.cat([inpaint_latents, inpaint_latents_mask, control_latents], dim=1)
|
||||
latent_no_controls = torch.cat([inpaint_latents, inpaint_latents_mask, torch.zeros_like(control_latents)], dim=1)
|
||||
|
||||
# pack the controls
|
||||
height_latent_controls, width_latent_controls = latent_controls.shape[2:]
|
||||
packed_latent_controls = self._pack_latents(
|
||||
latent_controls,
|
||||
batch_size * num_images_per_prompt,
|
||||
num_control_channels,
|
||||
height_latent_controls,
|
||||
width_latent_controls,
|
||||
)
|
||||
packed_latent_no_controls = self._pack_latents(
|
||||
latent_no_controls,
|
||||
batch_size * num_images_per_prompt,
|
||||
num_control_channels,
|
||||
height_latent_controls,
|
||||
width_latent_controls,
|
||||
)
|
||||
|
||||
latents, latent_image_ids = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 5. Prepare timesteps
|
||||
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.get("base_image_seq_len", 256),
|
||||
self.scheduler.config.get("max_image_seq_len", 4096),
|
||||
self.scheduler.config.get("base_shift", 0.5),
|
||||
self.scheduler.config.get("max_shift", 1.15),
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# handle guidance
|
||||
if self.transformer.config.guidance_embeds:
|
||||
guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
control_cutoff = int(len(timesteps) * control_stop)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
control_latents = packed_latent_controls if i < control_cutoff else packed_latent_no_controls
|
||||
|
||||
latent_model_input = torch.cat([latents, control_latents], dim=2)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=pooled_prompt_embeds,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
if torch.backends.mps.is_available():
|
||||
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)
|
||||
latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = self.image_processor.postprocess(image, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return FluxPipelineOutput(images=image)
|
||||
@@ -0,0 +1,88 @@
|
||||
import os
|
||||
import transformers
|
||||
import diffusers
|
||||
from huggingface_hub import auth_check
|
||||
from modules import shared, devices, sd_models, model_quant, modelloader, sd_hijack_te
|
||||
|
||||
|
||||
def load_transformer(repo_id, diffusers_load_config={}):
|
||||
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Transformer', device_map=True)
|
||||
fn = None
|
||||
|
||||
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
|
||||
from modules import sd_unet
|
||||
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
|
||||
shared.log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
|
||||
return None
|
||||
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
|
||||
|
||||
if fn is not None and 'gguf' in fn.lower():
|
||||
shared.log.error('Load model: type=HiDream format="gguf" unsupported')
|
||||
transformer = None
|
||||
from modules import ggml
|
||||
transformer = ggml.load_gguf(fn, cls=diffusers.HiDreamImageTransformer2DModel, compute_dtype=devices.dtype)
|
||||
elif fn is not None and 'safetensors' in fn.lower():
|
||||
shared.log.debug(f'Load model: type=FLEX transformer="{repo_id}" quant="{model_quant.get_quant(repo_id)}" args={load_args}')
|
||||
transformer = diffusers.FluxTransformer2DModel.from_single_file(fn, cache_dir=shared.opts.hfcache_dir, **load_args)
|
||||
# elif model_quant.check_nunchaku('Transformer'):
|
||||
# shared.log.error(f'Load model: type=HiDream transformer="{repo_id}" quant="Nunchaku" unsupported')
|
||||
# transformer = None
|
||||
else:
|
||||
shared.log.debug(f'Load model: type=FLEX transformer="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
|
||||
transformer = diffusers.FluxTransformer2DModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder="transformer",
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
**load_args,
|
||||
**quant_args,
|
||||
)
|
||||
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
|
||||
sd_models.move_model(transformer, devices.cpu)
|
||||
return transformer
|
||||
|
||||
|
||||
def load_text_encoders(repo_id, diffusers_load_config={}):
|
||||
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='TE', device_map=True)
|
||||
shared.log.debug(f'Load model: type=FLEX t5="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
|
||||
text_encoder_2 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder="text_encoder_2",
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
**load_args,
|
||||
**quant_args,
|
||||
)
|
||||
if shared.opts.diffusers_offload_mode != 'none' and text_encoder_2 is not None:
|
||||
sd_models.move_model(text_encoder_2, devices.cpu)
|
||||
return text_encoder_2
|
||||
|
||||
|
||||
def load_flex(checkpoint_info, diffusers_load_config={}):
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info.name)
|
||||
login = modelloader.hf_login()
|
||||
try:
|
||||
auth_check(repo_id)
|
||||
except Exception as e:
|
||||
shared.log.error(f'Load model: repo="{repo_id}" login={login} {e}')
|
||||
return False
|
||||
|
||||
transformer = load_transformer(repo_id, diffusers_load_config)
|
||||
text_encoder_2 = load_text_encoders(repo_id, diffusers_load_config)
|
||||
|
||||
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
|
||||
shared.log.debug(f'Load model: type=FLEX model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
|
||||
|
||||
from modules.flex2 import Flex2Pipeline
|
||||
pipe = Flex2Pipeline.from_pretrained(
|
||||
repo_id,
|
||||
# custom_pipeline=repo_id,
|
||||
transformer=transformer,
|
||||
text_encoder_2=text_encoder_2,
|
||||
cache_dir=shared.opts.diffusers_dir,
|
||||
**load_args,
|
||||
)
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
del text_encoder_2
|
||||
del transformer
|
||||
|
||||
devices.torch_gc()
|
||||
return pipe
|
||||
@@ -27,7 +27,7 @@ def get_model_type(pipe):
|
||||
model_type = 'sc'
|
||||
elif "AuraFlow" in name:
|
||||
model_type = 'auraflow'
|
||||
elif "Flux" in name:
|
||||
elif "Flux" in name or "Flex.1" or "Flex.2":
|
||||
model_type = 'f1'
|
||||
elif "Lumina2" in name:
|
||||
model_type = 'lumina2'
|
||||
|
||||
@@ -379,12 +379,16 @@ def get_reference_opts(name: str, quiet=False):
|
||||
model_opts = {}
|
||||
name = name.replace('Diffusers/', 'huggingface/')
|
||||
for k, v in shared.reference_models.items():
|
||||
model_name = os.path.splitext(v.get('path', '').split('@')[0])[0]
|
||||
model_name = v.get('path', '')
|
||||
if k == name or model_name == name:
|
||||
model_opts = v
|
||||
break
|
||||
model_name = model_name.replace('huggingface/', '')
|
||||
if k == name or model_name == name:
|
||||
model_name_split = os.path.splitext(model_name.split('@')[0])[0]
|
||||
if k == name or model_name_split == name:
|
||||
model_opts = v
|
||||
break
|
||||
model_name_replace = model_name.replace('huggingface/', '')
|
||||
if k == name or model_name_replace == name:
|
||||
model_opts = v
|
||||
break
|
||||
if not model_opts:
|
||||
|
||||
@@ -96,6 +96,8 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
guess = 'FLUX'
|
||||
if size > 11000 and size < 16000:
|
||||
warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB')
|
||||
if 'flex.2' in f.lower():
|
||||
guess = 'FLEX'
|
||||
# guess for diffusers
|
||||
index = os.path.join(f, 'model_index.json')
|
||||
if os.path.exists(index) and os.path.isfile(index):
|
||||
@@ -103,7 +105,7 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
cls = index.get('_class_name', None)
|
||||
if cls is not None:
|
||||
pipeline = getattr(diffusers, cls)
|
||||
if 'Flux' in pipeline.__name__:
|
||||
if 'Flux' in pipeline.__name__ and guess != 'FLEX':
|
||||
guess = 'FLUX'
|
||||
if 'StableDiffusion3' in pipeline.__name__:
|
||||
guess = 'Stable Diffusion 3'
|
||||
|
||||
@@ -309,6 +309,9 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='
|
||||
elif model_type in ['FLUX']:
|
||||
from modules.model_flux import load_flux
|
||||
sd_model = load_flux(checkpoint_info, diffusers_load_config)
|
||||
elif model_type in ['FLEX']:
|
||||
from modules.model_flex import load_flex
|
||||
sd_model = load_flex(checkpoint_info, diffusers_load_config)
|
||||
elif model_type in ['Lumina 2']:
|
||||
from modules.model_lumina import load_lumina2
|
||||
sd_model = load_lumina2(checkpoint_info, diffusers_load_config)
|
||||
|
||||
@@ -20,6 +20,7 @@ pipelines = {
|
||||
'HunyuanDiT': getattr(diffusers, 'HunyuanDiTPipeline', None),
|
||||
'DeepFloyd IF': getattr(diffusers, 'IFPipeline', None),
|
||||
'FLUX': getattr(diffusers, 'FluxPipeline', None),
|
||||
'FLEX': getattr(diffusers, 'AutoPipelineForText2Image', None),
|
||||
'Sana': getattr(diffusers, 'SanaPipeline', None),
|
||||
'Lumina-Next': getattr(diffusers, 'LuminaText2ImgPipeline', None),
|
||||
'Lumina 2': getattr(diffusers, 'Lumina2Text2ImgPipeline', None),
|
||||
|
||||
+1
-1
Submodule wiki updated: 3c26f78069...f07c96ffce
Reference in New Issue
Block a user