mirror of
https://github.com/vladmandic/automatic
synced 2026-08-27 15:41:00 +02:00
3c2236c4ed
Signed-off-by: Vladimir Mandic <mandic00@live.com>
777 lines
36 KiB
Python
777 lines
36 KiB
Python
# Copyright 2026 Ideogram AI and The HuggingFace Team. All rights reserved.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
import os
|
|
import math
|
|
from typing import Any, Callable
|
|
|
|
import torch
|
|
from transformers import AutoTokenizer, PreTrainedModel
|
|
from transformers.masking_utils import create_causal_mask
|
|
|
|
from diffusers.image_processor import VaeImageProcessor
|
|
from diffusers.models.autoencoders import AutoencoderKLFlux2
|
|
from diffusers.models.transformers.transformer_ideogram4 import (
|
|
IMAGE_POSITION_OFFSET,
|
|
LLM_TOKEN_INDICATOR,
|
|
OUTPUT_IMAGE_INDICATOR,
|
|
SEQUENCE_PADDING_INDICATOR,
|
|
Ideogram4Transformer2DModel,
|
|
)
|
|
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
|
from diffusers.utils.torch_utils import randn_tensor
|
|
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
|
from diffusers.pipelines.ideogram4.pipeline_output import Ideogram4PipelineOutput
|
|
from diffusers.pipelines.ideogram4.prompt_enhancer import (
|
|
PROMPT_UPSAMPLE_TEMPERATURE,
|
|
Ideogram4PromptEnhancerHead,
|
|
build_caption_logits_processor,
|
|
build_prompt_enhancer,
|
|
generate_captions,
|
|
)
|
|
from modules.logger import log
|
|
|
|
|
|
debug_prompt = log.trace if os.environ.get('SD_PROMPT_DEBUG', None) is not None else lambda *args, **kwargs: None
|
|
|
|
|
|
# Hidden states of these Qwen3-VL decoder layers are concatenated to form the per-token
|
|
# text conditioning consumed by the Ideogram4 transformer.
|
|
QWEN3_VL_ACTIVATION_LAYERS = (0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 35)
|
|
|
|
|
|
EXAMPLE_DOC_STRING = """
|
|
Examples:
|
|
```py
|
|
>>> import torch
|
|
>>> from diffusers import Ideogram4Pipeline
|
|
|
|
>>> pipe = Ideogram4Pipeline.from_pretrained("ideogram-ai/ideogram-v4", torch_dtype=torch.bfloat16)
|
|
>>> pipe.to("cuda")
|
|
|
|
>>> prompt = "A photo of a cat holding a sign that says hello world"
|
|
>>> # The defaults are the recommended settings for best quality.
|
|
>>> image = pipe(prompt, height=2048, width=2048, generator=torch.Generator("cuda").manual_seed(0)).images[0]
|
|
>>> image.save("ideogram4.png")
|
|
```
|
|
"""
|
|
|
|
def prompt_to_json(prompt):
|
|
"""Normalize a JSON caption to the compact form Ideogram 4 trained on, or wrap plain text.
|
|
Ideogram 4 expects a structured JSON caption serialized compactly. A valid JSON prompt is
|
|
re-serialized to that form; a plain-text prompt is wrapped in a minimal caption so it stays
|
|
in distribution instead of tripping the weight-baked "blocked by safety filter" placeholder.
|
|
"""
|
|
import json
|
|
if isinstance(prompt, list):
|
|
return [prompt_to_json(p) for p in prompt]
|
|
if not isinstance(prompt, str) or len(prompt) == 0:
|
|
return prompt
|
|
try:
|
|
return json.dumps(json.loads(prompt), ensure_ascii=False, separators=(',', ':'))
|
|
except ValueError:
|
|
caption = {
|
|
'high_level_description': prompt,
|
|
'compositional_deconstruction': {
|
|
'background': prompt,
|
|
'elements': [],
|
|
},
|
|
}
|
|
return json.dumps(caption, ensure_ascii=False, separators=(',', ':'))
|
|
|
|
|
|
def _logit_normal_sigmas(
|
|
num_inference_steps: int,
|
|
mu: float,
|
|
std: float = 1.0,
|
|
logsnr_min: float = -15.0,
|
|
logsnr_max: float = 18.0,
|
|
device: torch.device | None = None,
|
|
) -> torch.Tensor:
|
|
r"""
|
|
Build a length-`num_inference_steps` sigma schedule using the Ideogram4 logit-normal flow-matching schedule.
|
|
|
|
Sigmas are returned in `[0, 1]` in decreasing order (sigma close to 1 corresponds to pure noise, sigma close to 0
|
|
to clean data), matching diffusers conventions.
|
|
|
|
The Ideogram4 schedule applies `sigma(s) = 1 - logit_normal_cdf_inverse(1 - s)` to `s = linspace(0, 1, N + 1)` and
|
|
keeps the first `N` entries; a terminal zero is appended downstream by the scheduler.
|
|
"""
|
|
intervals = torch.linspace(0.0, 1.0, num_inference_steps + 1, dtype=torch.float64)
|
|
# Apply the inverse CDF of a normal then push through the logistic to obtain a logit-normal CDF inverse.
|
|
z = torch.special.ndtri(intervals)
|
|
y = mu + std * z
|
|
t = 1.0 - torch.special.expit(y)
|
|
t_min = 1.0 / (1.0 + math.exp(0.5 * logsnr_max))
|
|
t_max = 1.0 / (1.0 + math.exp(0.5 * logsnr_min))
|
|
t = t.clamp(t_min, t_max)
|
|
# Convert from model time (0 = noise, 1 = data) to diffusers sigma (1 = noise, 0 = data) and reverse.
|
|
sigmas = (1.0 - t).flip(0)
|
|
# Drop the trailing 0; FlowMatchEulerDiscreteScheduler.set_timesteps appends one back internally.
|
|
sigmas = sigmas[:-1].to(dtype=torch.float32, device=device)
|
|
return sigmas
|
|
|
|
|
|
def _resolution_aware_mu(
|
|
height: int,
|
|
width: int,
|
|
base_mu: float,
|
|
base_resolution: tuple[int, int] = (512, 512),
|
|
) -> float:
|
|
"""Shift the schedule mean as a function of image resolution."""
|
|
num_pixels = height * width
|
|
base_pixels = base_resolution[0] * base_resolution[1]
|
|
return base_mu + 0.5 * math.log(num_pixels / base_pixels)
|
|
|
|
|
|
def _expand_tensor_to_effective_batch(
|
|
tensor: torch.Tensor,
|
|
batch_size: int,
|
|
num_per_prompt: int,
|
|
tensor_name: str | None = None,
|
|
) -> torch.Tensor:
|
|
"""Replicate `tensor` along dim 0 from `batch_size` (or 1) to `batch_size * num_per_prompt`."""
|
|
target_batch_size = batch_size * num_per_prompt
|
|
|
|
if tensor.shape[0] == target_batch_size:
|
|
return tensor
|
|
|
|
if tensor.shape[0] == 1:
|
|
repeat_by = target_batch_size
|
|
elif tensor.shape[0] == batch_size:
|
|
repeat_by = num_per_prompt
|
|
else:
|
|
tensor_name = f"`{tensor_name}`" if tensor_name is not None else "Tensor"
|
|
raise ValueError(
|
|
f"{tensor_name} batch size must be 1, `batch_size` ({batch_size}), or "
|
|
f"`batch_size * num_*_per_prompt` ({target_batch_size}), but got {tensor.shape[0]}."
|
|
)
|
|
|
|
return torch.repeat_interleave(tensor, repeats=repeat_by, dim=0, output_size=tensor.shape[0] * repeat_by)
|
|
|
|
|
|
def transformer_compute_dtype(module: torch.nn.Module) -> torch.dtype:
|
|
"""Activation dtype of a possibly-quantized transformer: sub-16-bit
|
|
floating params (e.g. fp8) surface through `module.dtype` but are a
|
|
storage format; the dequantizer's result dtype is the compute dtype."""
|
|
dtype = module.dtype
|
|
if not dtype.is_floating_point or torch.finfo(dtype).bits >= 16:
|
|
return dtype
|
|
for m in module.modules():
|
|
dequantizer = getattr(m, 'sdnq_dequantizer', None)
|
|
if dequantizer is not None:
|
|
return dequantizer.result_dtype
|
|
return dtype
|
|
|
|
|
|
class Ideogram4Pipeline(DiffusionPipeline):
|
|
r"""
|
|
Text-to-image pipeline for Ideogram4.
|
|
|
|
Ideogram4 is a flow-matching model trained with asymmetric classifier-free guidance: a `transformer` consumes
|
|
text-conditioned features alongside the image latents, while a separate `unconditional_transformer` denoises with
|
|
zeroed text features. The two velocity predictions are linearly blended each step.
|
|
|
|
Args:
|
|
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
|
Flow-matching scheduler. The pipeline overrides the default sigma schedule with a resolution-aware
|
|
logit-normal schedule.
|
|
vae ([`AutoencoderKLFlux2`]):
|
|
Variational auto-encoder used to decode latents back into images.
|
|
text_encoder ([`PreTrainedModel`]):
|
|
Multimodal text encoder. The pipeline consumes hidden states from a fixed set of intermediate decoder
|
|
layers (see `QWEN3_VL_ACTIVATION_LAYERS`).
|
|
tokenizer ([`AutoTokenizer`]):
|
|
Tokenizer paired with `text_encoder`.
|
|
transformer ([`Ideogram4Transformer2DModel`]):
|
|
Conditional flow-matching transformer.
|
|
unconditional_transformer ([`Ideogram4Transformer2DModel`]):
|
|
Unconditional (asymmetric-CFG) flow-matching transformer.
|
|
"""
|
|
|
|
model_cpu_offload_seq = "prompt_enhancer_head->text_encoder->transformer->unconditional_transformer->vae"
|
|
_optional_components = ["prompt_enhancer_head"]
|
|
_callback_tensor_inputs = ["latents"]
|
|
|
|
def __init__(
|
|
self,
|
|
scheduler: FlowMatchEulerDiscreteScheduler,
|
|
vae: AutoencoderKLFlux2,
|
|
text_encoder: PreTrainedModel,
|
|
tokenizer: AutoTokenizer,
|
|
transformer: Ideogram4Transformer2DModel,
|
|
unconditional_transformer: Ideogram4Transformer2DModel,
|
|
prompt_enhancer_head: Ideogram4PromptEnhancerHead | None = None,
|
|
) -> None:
|
|
super().__init__()
|
|
|
|
self.register_modules(
|
|
scheduler=scheduler,
|
|
vae=vae,
|
|
text_encoder=text_encoder,
|
|
tokenizer=tokenizer,
|
|
transformer=transformer,
|
|
unconditional_transformer=unconditional_transformer,
|
|
prompt_enhancer_head=prompt_enhancer_head,
|
|
)
|
|
|
|
self.vae_scale_factor = (
|
|
2 ** (len(self.vae.config.block_out_channels) - 1) if getattr(self, "vae", None) is not None else 8
|
|
)
|
|
# Ideogram4 patchifies the VAE output by a factor of 2 before feeding into the transformer.
|
|
self.patch_size = 2
|
|
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * self.patch_size)
|
|
|
|
# Built lazily on first upsample: the head-less encoder body + `prompt_enhancer_head`, combined.
|
|
self.prompt_enhancer = None
|
|
# Outlines logits processor for schema-constrained captions; built lazily on first upsample.
|
|
self.caption_logits_processor = None
|
|
|
|
def upsample_prompt(self, prompt, height=1024, width=1024, temperature=1.0, max_new_tokens=1024, generator=None, device=None) -> list[str]:
|
|
"""Rewrite each prompt into Ideogram4's native structured JSON caption.
|
|
Requires the optional `prompt_enhancer_head` component, which is grafted onto the shared `text_encoder` body to
|
|
make it generative. Generation is schema-constrained when `outlines` is installed, otherwise it runs unconstrained.
|
|
"""
|
|
if self.prompt_enhancer_head is None:
|
|
return prompt
|
|
from installer import install
|
|
from pipelines.ideogram.patch_qwen import hijack_qwen3vl, restore_qwen3vl
|
|
install('outlines')
|
|
if self.prompt_enhancer is None:
|
|
self.prompt_enhancer = build_prompt_enhancer(self.text_encoder, self.prompt_enhancer_head)
|
|
if self.caption_logits_processor is None:
|
|
self.caption_logits_processor = build_caption_logits_processor(self.prompt_enhancer, self.tokenizer)
|
|
|
|
log.debug(f'Encode: enhancer={self.prompt_enhancer.__class__.__name__} processor={self.caption_logits_processor.__class__.__name__} max={max_new_tokens} device={device}')
|
|
self.prompt_enhancer.to(device)
|
|
hijack_qwen3vl()
|
|
try:
|
|
caption = generate_captions(
|
|
self.prompt_enhancer,
|
|
self.tokenizer,
|
|
self.caption_logits_processor,
|
|
prompt,
|
|
height,
|
|
width,
|
|
temperature=temperature,
|
|
max_new_tokens=max_new_tokens,
|
|
generator=generator,
|
|
device=device,
|
|
)
|
|
finally:
|
|
restore_qwen3vl()
|
|
self.prompt_enhancer.to('cpu')
|
|
debug_prompt(f'Prompt: input="{prompt}"')
|
|
debug_prompt(f'Prompt: enhanced="{caption}"')
|
|
return caption
|
|
|
|
@staticmethod
|
|
def _prepare_ids(
|
|
text_lengths: list[int],
|
|
grid_h: int,
|
|
grid_w: int,
|
|
max_text_tokens: int,
|
|
device: torch.device,
|
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
"""Build the packed `[left-pad][text][image]` layout from the per-prompt text lengths and the image grid.
|
|
|
|
Returns `position_ids` (3-axis MRoPE), `segment_ids` (block-diagonal attention) and `indicator` (per-token
|
|
text/image/pad role).
|
|
"""
|
|
batch_size = len(text_lengths)
|
|
num_image_tokens = grid_h * grid_w
|
|
total_seq_len = max_text_tokens + num_image_tokens
|
|
|
|
# Image position ids (t=0, h, w); offset keeps them disjoint from text positions.
|
|
h_idx = torch.arange(grid_h).view(-1, 1).expand(grid_h, grid_w).reshape(-1)
|
|
w_idx = torch.arange(grid_w).view(1, -1).expand(grid_h, grid_w).reshape(-1)
|
|
t_idx = torch.zeros_like(h_idx)
|
|
image_pos = torch.stack([t_idx, h_idx, w_idx], dim=1) + IMAGE_POSITION_OFFSET
|
|
|
|
position_ids = torch.zeros(batch_size, total_seq_len, 3, dtype=torch.long)
|
|
segment_ids = torch.full((batch_size, total_seq_len), SEQUENCE_PADDING_INDICATOR, dtype=torch.long)
|
|
indicator = torch.zeros(batch_size, total_seq_len, dtype=torch.long)
|
|
|
|
for b, num_text in enumerate(text_lengths):
|
|
offset = max_text_tokens - num_text
|
|
|
|
text_pos = torch.arange(num_text)
|
|
text_pos_3d = torch.stack([text_pos, text_pos, text_pos], dim=1)
|
|
position_ids[b, offset : offset + num_text] = text_pos_3d
|
|
position_ids[b, offset + num_text :] = image_pos
|
|
|
|
indicator[b, offset : offset + num_text] = LLM_TOKEN_INDICATOR
|
|
indicator[b, offset + num_text :] = OUTPUT_IMAGE_INDICATOR
|
|
|
|
segment_ids[b, offset : offset + num_text + num_image_tokens] = 1
|
|
|
|
return position_ids.to(device), segment_ids.to(device), indicator.to(device)
|
|
|
|
@staticmethod
|
|
def _get_text_encoder_hidden_states(
|
|
text_encoder,
|
|
token_ids: torch.Tensor,
|
|
attention_mask: torch.Tensor,
|
|
pos_2d: torch.Tensor,
|
|
) -> list[torch.Tensor]:
|
|
"""Run the text encoder's decoder layers, returning the hidden states tapped at each activation layer."""
|
|
|
|
language_model = text_encoder.language_model
|
|
|
|
inputs_embeds = language_model.embed_tokens(token_ids)
|
|
|
|
position_ids_4d = pos_2d[None, ...].expand(4, pos_2d.shape[0], -1)
|
|
text_position_ids = position_ids_4d[0]
|
|
mrope_position_ids = position_ids_4d[1:]
|
|
|
|
causal_mask = create_causal_mask(
|
|
config=language_model.config,
|
|
inputs_embeds=inputs_embeds,
|
|
attention_mask=attention_mask,
|
|
past_key_values=None,
|
|
position_ids=text_position_ids,
|
|
)
|
|
position_embeddings = language_model.rotary_emb(inputs_embeds, mrope_position_ids)
|
|
|
|
tap_set = set(QWEN3_VL_ACTIVATION_LAYERS)
|
|
captured: dict[int, torch.Tensor] = {}
|
|
hidden_states = inputs_embeds
|
|
for layer_idx, decoder_layer in enumerate(language_model.layers):
|
|
hidden_states = decoder_layer(
|
|
hidden_states,
|
|
attention_mask=causal_mask,
|
|
position_ids=text_position_ids,
|
|
past_key_values=None,
|
|
position_embeddings=position_embeddings,
|
|
)
|
|
if layer_idx in tap_set:
|
|
captured[layer_idx] = hidden_states
|
|
|
|
return [captured[i] for i in QWEN3_VL_ACTIVATION_LAYERS]
|
|
|
|
def _encode_prompt(
|
|
self,
|
|
prompt: str | list[str],
|
|
grid_h: int,
|
|
grid_w: int,
|
|
max_sequence_length: int,
|
|
device: torch.device,
|
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
"""Prepare the conditioning for the packed text+image sequence (one entry per prompt).
|
|
|
|
Returns a flat tuple `(prompt_embeds, position_ids, segment_ids, indicator)`. The unconditional branch carries
|
|
no text, so the pipeline builds its (zeroed) inputs directly rather than encoding a negative prompt.
|
|
"""
|
|
prompts = [prompt] if isinstance(prompt, str) else list(prompt)
|
|
batch_size = len(prompts)
|
|
num_image_tokens = grid_h * grid_w
|
|
|
|
# Tokenize each chat-formatted prompt and left-pad to `max_sequence_length`. Only the text region is fed to
|
|
# the encoder: the packed image tokens come after the text and the encoder is causal, so they never affect it.
|
|
token_ids = torch.zeros(batch_size, max_sequence_length, dtype=torch.long)
|
|
attention_mask = torch.zeros(batch_size, max_sequence_length, dtype=torch.long)
|
|
text_position_ids = torch.zeros(batch_size, max_sequence_length, dtype=torch.long)
|
|
text_lengths = []
|
|
for b, text_prompt in enumerate(prompts):
|
|
messages = [{"role": "user", "content": [{"type": "text", "text": text_prompt}]}]
|
|
text = self.tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)
|
|
toks = self.tokenizer(text, return_tensors="pt", add_special_tokens=False)["input_ids"][0]
|
|
n = int(toks.shape[0])
|
|
if n > max_sequence_length:
|
|
raise ValueError(f"prompt has {n} tokens, exceeds max_sequence_length={max_sequence_length}")
|
|
text_lengths.append(n)
|
|
offset = max_sequence_length - n
|
|
token_ids[b, offset:] = toks
|
|
attention_mask[b, offset:] = 1
|
|
text_position_ids[b, offset:] = torch.arange(n)
|
|
|
|
text_encoder_device = next(self.text_encoder.parameters()).device if len(list(self.text_encoder.parameters())) > 0 else device
|
|
token_ids = token_ids.to(text_encoder_device)
|
|
attention_mask = attention_mask.to(text_encoder_device)
|
|
text_position_ids = text_position_ids.to(text_encoder_device)
|
|
|
|
# Concatenate the tapped activation-layer hidden states into per-token text features, zeroing padding.
|
|
selected = self._get_text_encoder_hidden_states(
|
|
self.text_encoder, token_ids, attention_mask, text_position_ids
|
|
)
|
|
text_features = torch.stack(selected, dim=0).permute(1, 2, 3, 0).reshape(batch_size, max_sequence_length, -1)
|
|
text_features = (text_features * attention_mask.to(text_features.dtype).unsqueeze(-1)).to(torch.float32).to(device)
|
|
|
|
position_ids, segment_ids, indicator = self._prepare_ids(
|
|
text_lengths, grid_h, grid_w, max_sequence_length, device
|
|
)
|
|
|
|
# Pack the text features into the full sequence; image positions carry no text features.
|
|
image_feature_padding = torch.zeros(
|
|
batch_size, num_image_tokens, text_features.shape[-1], dtype=text_features.dtype, device=device
|
|
)
|
|
prompt_embeds = torch.cat([text_features, image_feature_padding], dim=1)
|
|
return prompt_embeds, position_ids, segment_ids, indicator
|
|
|
|
def encode_prompt(
|
|
self,
|
|
prompt: str | list[str],
|
|
grid_h: int,
|
|
grid_w: int,
|
|
max_sequence_length: int,
|
|
device: torch.device,
|
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
from modules import shared, devices
|
|
if not shared.opts.model_ideogram4_enable_pe:
|
|
prompt = prompt_to_json(prompt)
|
|
self.text_encoder.to(self._execution_device)
|
|
try:
|
|
prompt = self._encode_prompt(prompt, grid_h, grid_w, max_sequence_length, device)
|
|
finally:
|
|
if shared.opts.diffusers_offload_mode != 'none':
|
|
self.text_encoder.to(devices.cpu)
|
|
return prompt
|
|
|
|
def prepare_latents(
|
|
self,
|
|
batch_size: int,
|
|
num_image_tokens: int,
|
|
latent_dim: int,
|
|
dtype: torch.dtype,
|
|
device: torch.device,
|
|
generator: torch.Generator | list[torch.Generator] | None,
|
|
latents: torch.Tensor | None = None,
|
|
) -> torch.Tensor:
|
|
shape = (batch_size, num_image_tokens, latent_dim)
|
|
if latents is None:
|
|
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
|
else:
|
|
if latents.shape != shape:
|
|
raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")
|
|
latents = latents.to(device=device, dtype=dtype)
|
|
return latents
|
|
|
|
@property
|
|
def guidance_scale(self) -> float | None:
|
|
return self._guidance_scale
|
|
|
|
@property
|
|
def num_timesteps(self) -> int:
|
|
return self._num_timesteps
|
|
|
|
@property
|
|
def interrupt(self) -> bool:
|
|
return self._interrupt
|
|
|
|
def check_inputs(
|
|
self,
|
|
prompt,
|
|
height,
|
|
width,
|
|
num_inference_steps,
|
|
guidance_scale,
|
|
guidance_schedule,
|
|
callback_on_step_end_tensor_inputs=None,
|
|
):
|
|
if prompt is None:
|
|
raise ValueError("`prompt` must be provided.")
|
|
if not isinstance(prompt, (str, list)):
|
|
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
|
|
|
if (
|
|
height % (self.vae_scale_factor * self.patch_size) != 0
|
|
or width % (self.vae_scale_factor * self.patch_size) != 0
|
|
):
|
|
raise ValueError(
|
|
f"`height` ({height}) and `width` ({width}) must both be divisible by {self.vae_scale_factor * self.patch_size} "
|
|
f"(vae_scale_factor * patch_size)."
|
|
)
|
|
|
|
# Guidance is controlled by either a constant `guidance_scale` or a per-step `guidance_schedule`; exactly
|
|
# one must be set (the `guidance_schedule` default makes the no-arg call use the recommended schedule).
|
|
if guidance_scale is not None:
|
|
guidance_schedule = [guidance_scale] * num_inference_steps
|
|
if guidance_scale is None and guidance_schedule is None:
|
|
guidance_schedule = [7.0] * int(num_inference_steps * 0.9) + [3.0] * (num_inference_steps - int(num_inference_steps * 0.9)) # 90% of steps with 7.0, last 10% with 3.0
|
|
|
|
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs):
|
|
raise ValueError(
|
|
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found "
|
|
f"{[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
|
)
|
|
|
|
@torch.no_grad()
|
|
def __call__(
|
|
self,
|
|
prompt: str | list[str] | None = None,
|
|
height: int = 2048,
|
|
width: int = 2048,
|
|
num_inference_steps: int = 48,
|
|
guidance_scale: float | None = None,
|
|
guidance_schedule: list[float] | torch.Tensor | None = None,
|
|
mu: float = 0.0,
|
|
std: float = 1.5,
|
|
prompt_upsampling: bool = False,
|
|
prompt_upsampling_temperature: float = PROMPT_UPSAMPLE_TEMPERATURE,
|
|
max_sequence_length: int = 2048,
|
|
num_images_per_prompt: int = 1,
|
|
generator: torch.Generator | list[torch.Generator] | None = None,
|
|
latents: torch.Tensor | None = None,
|
|
output_type: str = "pil",
|
|
return_dict: bool = True,
|
|
callback_on_step_end: Callable[["Ideogram4Pipeline", int, int, dict[str, Any]], dict[str, Any]] | None = None,
|
|
callback_on_step_end_tensor_inputs: list[str] = ["latents"],
|
|
) -> Ideogram4PipelineOutput | tuple[Any]:
|
|
r"""
|
|
Run text-to-image generation.
|
|
|
|
Args:
|
|
prompt (`str` or `list[str]`):
|
|
Prompt(s) to guide image generation.
|
|
height (`int`, *optional*, defaults to 2048):
|
|
Output image height in pixels; must be a multiple of `vae_scale_factor * patch_size`.
|
|
width (`int`, *optional*, defaults to 2048):
|
|
Output image width in pixels; must be a multiple of `vae_scale_factor * patch_size`.
|
|
num_inference_steps (`int`, *optional*, defaults to 48):
|
|
Number of flow-matching steps. The default is the recommended setting for best quality.
|
|
guidance_scale (`float`, *optional*):
|
|
Constant classifier-free guidance scale applied at every step. The conditional and unconditional
|
|
velocity predictions are blended as `v = guidance_scale * v_pos + (1 - guidance_scale) * v_neg`.
|
|
Mutually exclusive with `guidance_schedule` (setting both raises). Defaults to `None`.
|
|
guidance_schedule (`list[float]` or `torch.Tensor`, *optional*):
|
|
Per-step guidance scale schedule; must have length `num_inference_steps`. The first entry corresponds
|
|
to the first step (largest noise level). Mutually exclusive with `guidance_scale`; exactly one must be
|
|
set. Defaults to the recommended schedule (7.0 for the main steps, dropping to 3.0 for the final 3
|
|
"polish" steps). To use a constant scale instead, pass `guidance_scale` and `guidance_schedule=None`.
|
|
mu (`float`, *optional*, defaults to 0.0):
|
|
Base mean of the logit-normal flow-matching schedule. The schedule mean is shifted by half the log of
|
|
the resolution ratio relative to 512x512.
|
|
std (`float`, *optional*, defaults to 1.5):
|
|
Standard deviation of the logit-normal flow-matching schedule.
|
|
prompt_upsampling (`bool`, *optional*, defaults to `False`):
|
|
If `True`, rewrite `prompt` into Ideogram4's native structured JSON caption via
|
|
[`~Ideogram4Pipeline.upsample_prompt`] before encoding. Requires the optional `prompt_enhancer_head`
|
|
component; install `outlines` for schema-constrained captions. `generator` is reused to make the
|
|
upsampling reproducible.
|
|
prompt_upsampling_temperature (`float`, *optional*, defaults to 1.0):
|
|
Sampling temperature for prompt upsampling when `prompt_upsampling=True`.
|
|
max_sequence_length (`int`, *optional*, defaults to 2048):
|
|
Maximum number of text tokens per prompt.
|
|
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
|
Number of images to generate per prompt.
|
|
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
|
Generator(s) used to make sampling deterministic.
|
|
latents (`torch.Tensor`, *optional*):
|
|
Pre-generated noise of shape `(batch_size, num_image_tokens, latent_dim)`.
|
|
output_type (`str`, *optional*, defaults to `"pil"`):
|
|
One of `"pil"`, `"np"`, `"pt"`, or `"latent"`.
|
|
return_dict (`bool`, *optional*, defaults to `True`):
|
|
Whether to return an [`~pipelines.ideogram4.Ideogram4PipelineOutput`].
|
|
callback_on_step_end (`Callable`, *optional*):
|
|
Callback invoked at the end of every denoising step.
|
|
callback_on_step_end_tensor_inputs (`list[str]`, *optional*):
|
|
Names of tensors to expose to the callback via `callback_kwargs`.
|
|
|
|
Examples:
|
|
|
|
Returns:
|
|
[`~pipelines.ideogram4.Ideogram4PipelineOutput`] or `tuple`.
|
|
"""
|
|
self.check_inputs(
|
|
prompt=prompt,
|
|
height=height,
|
|
width=width,
|
|
num_inference_steps=num_inference_steps,
|
|
guidance_scale=guidance_scale,
|
|
guidance_schedule=guidance_schedule,
|
|
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
|
)
|
|
|
|
batch_size = 1
|
|
if isinstance(prompt, list):
|
|
batch_size = len(prompt)
|
|
else:
|
|
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
|
|
|
device = self._execution_device
|
|
self._guidance_scale = guidance_scale # pylint: disable=attribute-defined-outside-init
|
|
self._interrupt = False # pylint: disable=attribute-defined-outside-init
|
|
|
|
# 0. Optionally rewrite the prompt(s) into Ideogram4's native structured JSON caption.
|
|
if prompt_upsampling:
|
|
prompt = self.upsample_prompt(
|
|
prompt,
|
|
height=height,
|
|
width=width,
|
|
temperature=prompt_upsampling_temperature,
|
|
max_new_tokens=max_sequence_length - 64, # leave room for json structure tokens
|
|
generator=generator,
|
|
device=device,
|
|
)
|
|
|
|
# 1. Image grid (drives both the packed layout and the latent shape).
|
|
grid_h, grid_w = (
|
|
height // (self.vae_scale_factor * self.patch_size),
|
|
width // (self.vae_scale_factor * self.patch_size),
|
|
)
|
|
num_image_tokens = grid_h * grid_w
|
|
|
|
# 2. Encode prompts into the packed conditioning (one entry per prompt).
|
|
llm_features, position_ids, segment_ids, indicator = self.encode_prompt(
|
|
prompt=prompt,
|
|
grid_h=grid_h,
|
|
grid_w=grid_w,
|
|
max_sequence_length=max_sequence_length,
|
|
device=device,
|
|
)
|
|
|
|
# 3. Replicate the conditioning for num_images_per_prompt.
|
|
llm_features = _expand_tensor_to_effective_batch(llm_features, batch_size, num_images_per_prompt)
|
|
position_ids = _expand_tensor_to_effective_batch(position_ids, batch_size, num_images_per_prompt)
|
|
segment_ids = _expand_tensor_to_effective_batch(segment_ids, batch_size, num_images_per_prompt)
|
|
indicator = _expand_tensor_to_effective_batch(indicator, batch_size, num_images_per_prompt)
|
|
|
|
# 4. Unconditional (image-only) branch, derived from the conditioning: zeroed text features and the
|
|
# image-region slices of the layout.
|
|
neg_llm_features = torch.zeros(
|
|
batch_size * num_images_per_prompt,
|
|
num_image_tokens,
|
|
llm_features.shape[-1],
|
|
dtype=llm_features.dtype,
|
|
device=device,
|
|
)
|
|
neg_position_ids = position_ids[:, max_sequence_length:]
|
|
neg_segment_ids = segment_ids[:, max_sequence_length:]
|
|
neg_indicator = indicator[:, max_sequence_length:]
|
|
|
|
# 4. Set up the resolution-aware logit-normal schedule on the scheduler.
|
|
schedule_mu = _resolution_aware_mu(height=height, width=width, base_mu=mu)
|
|
sigmas = _logit_normal_sigmas(num_inference_steps, schedule_mu, std=std, device=device)
|
|
self.scheduler.set_timesteps(sigmas=sigmas.tolist(), device=device) # pylint: disable=unexpected-keyword-arg
|
|
timesteps = self.scheduler.timesteps
|
|
self._num_timesteps = len(timesteps) # pylint: disable=attribute-defined-outside-init
|
|
|
|
# 5. Resolve the per-step guidance schedule (a constant `guidance_scale` broadcasts to every step, otherwise
|
|
# use the provided `guidance_schedule`, validated by `check_inputs`) and the tensor of per-step weights `gw`.
|
|
if guidance_scale is not None:
|
|
guidance_schedule = [float(guidance_scale)] * num_inference_steps
|
|
if guidance_schedule is None:
|
|
guidance_schedule = [7.0] * int(num_inference_steps * 0.9) + [3.0] * (num_inference_steps - int(num_inference_steps * 0.9)) # 90% of steps with 7.0, last 10% with 3.0
|
|
gw = torch.as_tensor(guidance_schedule, dtype=torch.float32, device=device)
|
|
|
|
# 6. Prepare latents in the packed (B, num_image_tokens, latent_dim) layout.
|
|
latent_dim = self.transformer.config.in_channels
|
|
latents = self.prepare_latents(
|
|
batch_size=batch_size * num_images_per_prompt,
|
|
num_image_tokens=num_image_tokens,
|
|
latent_dim=latent_dim,
|
|
dtype=torch.float32,
|
|
device=device,
|
|
generator=generator,
|
|
latents=latents,
|
|
)
|
|
|
|
# 7. Padding for the text region of the conditional packed sequence (image latents are appended after it).
|
|
max_text_tokens = max_sequence_length
|
|
text_z_padding = torch.zeros(
|
|
batch_size * num_images_per_prompt,
|
|
max_text_tokens,
|
|
latent_dim,
|
|
dtype=torch.float32,
|
|
device=device,
|
|
)
|
|
|
|
# The transformers run in their loaded compute dtype; cast the (otherwise float32) text features to match.
|
|
# `latents` stay float32 for scheduler precision and are cast per-step at the transformer call below.
|
|
cond_dtype = transformer_compute_dtype(self.transformer)
|
|
uncond_dtype = transformer_compute_dtype(self.unconditional_transformer) if self.unconditional_transformer is not None else cond_dtype
|
|
llm_features = llm_features.to(cond_dtype)
|
|
neg_llm_features = neg_llm_features.to(uncond_dtype)
|
|
|
|
# 8. Denoising loop. The scheduler stores `num_train_timesteps`-scaled timesteps; convert back to model time.
|
|
num_train_timesteps = self.scheduler.config.num_train_timesteps # pylint: disable=no-member
|
|
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
|
for i, t in enumerate(timesteps):
|
|
if self.interrupt:
|
|
continue
|
|
|
|
# Map sigma-domain timestep to model time `t` in [0, 1] (0 = noise, 1 = clean data).
|
|
t_model = 1.0 - (t.float() / num_train_timesteps)
|
|
t_model = t_model.expand(batch_size * num_images_per_prompt).to(cond_dtype)
|
|
|
|
# Conditional pass operates on the full packed sequence.
|
|
pos_z = torch.cat([text_z_padding, latents], dim=1).to(cond_dtype)
|
|
pos_out = self.transformer(
|
|
hidden_states=pos_z,
|
|
timestep=t_model,
|
|
encoder_hidden_states=llm_features,
|
|
position_ids=position_ids,
|
|
segment_ids=segment_ids,
|
|
indicator=indicator,
|
|
return_dict=False,
|
|
)[0]
|
|
# Velocity (and guidance) is computed in float32 for scheduler precision; the transformers
|
|
# return their compute dtype, so cast the predicted velocities up here.
|
|
pos_v = pos_out[:, max_text_tokens:].to(torch.float32)
|
|
|
|
# Unconditional pass uses image-only positions with zeroed text features.
|
|
self._guidance_scale = guidance_schedule[i] # pylint: disable=attribute-defined-outside-init
|
|
gw_i = gw[i]
|
|
if gw[i] > 1.0:
|
|
uncond_transformer = self.unconditional_transformer if self.unconditional_transformer is not None else self.transformer
|
|
neg_v = uncond_transformer(
|
|
hidden_states=latents.to(uncond_dtype),
|
|
timestep=t_model,
|
|
encoder_hidden_states=neg_llm_features,
|
|
position_ids=neg_position_ids,
|
|
segment_ids=neg_segment_ids,
|
|
indicator=neg_indicator,
|
|
return_dict=False,
|
|
)[0].to(torch.float32)
|
|
v = gw_i * pos_v + (1.0 - gw_i) * neg_v
|
|
else:
|
|
v = gw_i * pos_v
|
|
|
|
latents = self.scheduler.step(-v, t, latents, return_dict=False)[0]
|
|
|
|
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)
|
|
|
|
progress_bar.update()
|
|
|
|
# 9. Decode: unpatch the latents, denormalize with the VAE batch-norm stats, and decode through the VAE.
|
|
if output_type == "latent":
|
|
image = latents
|
|
else:
|
|
z = latents
|
|
# VAE bn stores per-channel statistics on the packed-channel latent space (ae_channels * patch ** 2).
|
|
bn_mean = self.vae.bn.running_mean.view(1, 1, -1).to(device=z.device, dtype=z.dtype)
|
|
bn_std = torch.sqrt(self.vae.bn.running_var + self.vae.config.batch_norm_eps).view(1, 1, -1)
|
|
bn_std = bn_std.to(device=z.device, dtype=z.dtype)
|
|
z = z * bn_std + bn_mean
|
|
|
|
patch = self.patch_size
|
|
ae_channels = z.shape[-1] // (patch * patch)
|
|
z = z.view(batch_size * num_images_per_prompt, grid_h, grid_w, patch, patch, ae_channels)
|
|
z = z.permute(0, 5, 1, 3, 2, 4).contiguous()
|
|
z = z.view(batch_size * num_images_per_prompt, ae_channels, grid_h * patch, grid_w * patch)
|
|
|
|
decoded = self.vae.decode(z.to(self.vae.dtype), return_dict=False)[0]
|
|
image = self.image_processor.postprocess(decoded.float(), output_type=output_type)
|
|
|
|
self.maybe_free_model_hooks()
|
|
|
|
if not return_dict:
|
|
return (image,)
|
|
return Ideogram4PipelineOutput(images=image)
|