Spaces:
Paused
Paused
| # 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 torch | |
| from transformers import Qwen2Tokenizer, Qwen3VLModel | |
| from transformers.masking_utils import create_causal_mask | |
| from ...pipelines.ideogram4.prompt_enhancer import ( | |
| PROMPT_UPSAMPLE_TEMPERATURE, | |
| Ideogram4PromptEnhancerHead, | |
| build_caption_logits_processor, | |
| build_prompt_enhancer, | |
| generate_captions, | |
| ) | |
| from ...utils import is_outlines_available, logging | |
| from ..modular_pipeline import ModularPipelineBlocks, PipelineState | |
| from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam | |
| from .modular_pipeline import Ideogram4ModularPipeline | |
| logger = logging.get_logger(__name__) # pylint: disable=invalid-name | |
| # 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) | |
| # auto_docstring | |
| class Ideogram4PromptUpsampleStep(ModularPipelineBlocks): | |
| """ | |
| Optional step that rewrites the prompt(s) into Ideogram4's native structured JSON caption when | |
| `prompt_upsampling=True` (the format the model is trained on). Requires a generative `text_encoder` (a | |
| `Qwen3VLForConditionalGeneration`); install `outlines` for schema-constrained captions. | |
| Components: | |
| text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`Qwen2Tokenizer`): The tokenizer paired | |
| with the text encoder. prompt_enhancer_head (`Ideogram4PromptEnhancerHead`): LM head grafted onto the text | |
| encoder for prompt upsampling. | |
| Inputs: | |
| prompt (`str`): | |
| The prompt or prompts to guide image generation. | |
| prompt_upsampling (`bool`, *optional*, defaults to False): | |
| If True, rewrite the prompt into Ideogram4's native JSON caption before encoding. | |
| prompt_upsampling_temperature (`float`, *optional*, defaults to 1.0): | |
| Sampling temperature for prompt upsampling. | |
| height (`int`, *optional*): | |
| The height in pixels of the generated image. | |
| width (`int`, *optional*): | |
| The width in pixels of the generated image. | |
| max_sequence_length (`int`, *optional*, defaults to 2048): | |
| Maximum sequence length for prompt encoding. | |
| generator (`Generator`, *optional*): | |
| Torch generator for deterministic generation. | |
| Outputs: | |
| prompt (`list`): | |
| The (possibly upsampled) prompt forwarded to the text encoder. | |
| """ | |
| model_name = "ideogram4" | |
| def __init__(self): | |
| # 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 | |
| super().__init__() | |
| def description(self) -> str: | |
| return ( | |
| "Optional step that rewrites the prompt(s) into Ideogram4's native structured JSON caption when " | |
| "`prompt_upsampling=True` (the format the model is trained on). Requires a generative `text_encoder` " | |
| "(a `Qwen3VLForConditionalGeneration`); install `outlines` for schema-constrained captions." | |
| ) | |
| def expected_components(self) -> list[ComponentSpec]: | |
| return [ | |
| ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), | |
| ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer paired with the text encoder."), | |
| ComponentSpec( | |
| "prompt_enhancer_head", | |
| Ideogram4PromptEnhancerHead, | |
| description="LM head grafted onto the text encoder for prompt upsampling.", | |
| ), | |
| ] | |
| def inputs(self) -> list[InputParam]: | |
| return [ | |
| InputParam.template("prompt", required=True), | |
| InputParam( | |
| name="prompt_upsampling", | |
| type_hint=bool, | |
| default=False, | |
| description="If True, rewrite the prompt into Ideogram4's native JSON caption before encoding.", | |
| ), | |
| InputParam( | |
| name="prompt_upsampling_temperature", | |
| type_hint=float, | |
| default=PROMPT_UPSAMPLE_TEMPERATURE, | |
| description="Sampling temperature for prompt upsampling.", | |
| ), | |
| InputParam.template("height"), | |
| InputParam.template("width"), | |
| InputParam.template("max_sequence_length", default=2048), | |
| InputParam.template("generator"), | |
| ] | |
| def intermediate_outputs(self) -> list[OutputParam]: | |
| return [ | |
| OutputParam( | |
| name="prompt", | |
| type_hint=list, | |
| description="The (possibly upsampled) prompt forwarded to the text encoder.", | |
| ), | |
| ] | |
| def __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: | |
| block_state = self.get_block_state(state) | |
| if block_state.prompt_upsampling: | |
| if components.prompt_enhancer_head is None: | |
| raise ValueError( | |
| "Prompt upsampling requires the `prompt_enhancer_head` component, which is not loaded. Load an " | |
| "`Ideogram4PromptEnhancerHead` and add it to the pipeline." | |
| ) | |
| if self._prompt_enhancer is None: | |
| self._prompt_enhancer = build_prompt_enhancer(components.text_encoder, components.prompt_enhancer_head) | |
| if self._caption_logits_processor is None and is_outlines_available(): | |
| self._caption_logits_processor = build_caption_logits_processor( | |
| self._prompt_enhancer, components.tokenizer | |
| ) | |
| if self._caption_logits_processor is None: | |
| logger.warning_once( | |
| "`outlines` is not installed; prompt upsampling runs unconstrained and may not return " | |
| "schema-valid JSON. Install with `pip install outlines` for structured captions." | |
| ) | |
| height = block_state.height or components.default_height | |
| width = block_state.width or components.default_width | |
| block_state.prompt = generate_captions( | |
| self._prompt_enhancer, | |
| components.tokenizer, | |
| self._caption_logits_processor, | |
| block_state.prompt, | |
| height, | |
| width, | |
| temperature=block_state.prompt_upsampling_temperature, | |
| max_new_tokens=block_state.max_sequence_length, | |
| generator=block_state.generator, | |
| device=components._execution_device, | |
| ) | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| # auto_docstring | |
| class Ideogram4TextEncoderStep(ModularPipelineBlocks): | |
| """ | |
| Text encoder step that tokenizes the prompt(s) and runs the Qwen3-VL text encoder, returning the per-token text | |
| features (concatenated from a fixed set of activation layers). Only the text tokens are encoded; the packed image | |
| tokens are appended later (the encoder is causal with image after text, so they never affect the text features). | |
| Components: | |
| text_encoder (`Qwen3VLModel`): The Qwen3-VL text encoder. tokenizer (`Qwen2Tokenizer`): The tokenizer paired | |
| with the text encoder. | |
| Inputs: | |
| prompt (`str`): | |
| The prompt or prompts to guide image generation. | |
| max_sequence_length (`int`, *optional*, defaults to 2048): | |
| Maximum sequence length for prompt encoding. | |
| Outputs: | |
| text_features (`Tensor`): | |
| Per-prompt text features (B, max_sequence_length, llm_features_dim), padding zeroed. | |
| text_lengths (`list`): | |
| Per-prompt real text-token counts, used to lay out the packed sequence. | |
| """ | |
| model_name = "ideogram4" | |
| def description(self) -> str: | |
| return ( | |
| "Text encoder step that tokenizes the prompt(s) and runs the Qwen3-VL text encoder, returning the " | |
| "per-token text features (concatenated from a fixed set of activation layers). Only the text tokens are " | |
| "encoded; the packed image tokens are appended later (the encoder is causal with image after text, so " | |
| "they never affect the text features)." | |
| ) | |
| def expected_components(self) -> list[ComponentSpec]: | |
| return [ | |
| ComponentSpec("text_encoder", Qwen3VLModel, description="The Qwen3-VL text encoder."), | |
| ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer paired with the text encoder."), | |
| ] | |
| def inputs(self) -> list[InputParam]: | |
| return [ | |
| InputParam.template("prompt", required=True), | |
| InputParam.template("max_sequence_length", default=2048), | |
| ] | |
| def intermediate_outputs(self) -> list[OutputParam]: | |
| return [ | |
| OutputParam( | |
| name="text_features", | |
| type_hint=torch.Tensor, | |
| description="Per-prompt text features (B, max_sequence_length, llm_features_dim), padding zeroed.", | |
| ), | |
| OutputParam( | |
| name="text_lengths", | |
| type_hint=list, | |
| description="Per-prompt real text-token counts, used to lay out the packed sequence.", | |
| ), | |
| ] | |
| # Copied from diffusers.pipelines.ideogram4.pipeline_ideogram4.Ideogram4Pipeline._get_text_encoder_hidden_states | |
| 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 __call__(self, components: Ideogram4ModularPipeline, state: PipelineState) -> PipelineState: | |
| block_state = self.get_block_state(state) | |
| device = components._execution_device | |
| tokenizer = components.tokenizer | |
| max_text_tokens = block_state.max_sequence_length | |
| prompts = [block_state.prompt] if isinstance(block_state.prompt, str) else list(block_state.prompt) | |
| batch_size = len(prompts) | |
| # Tokenize each chat-formatted prompt and left-pad to `max_sequence_length`. | |
| token_ids = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) | |
| attention_mask = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) | |
| text_position_ids = torch.zeros(batch_size, max_text_tokens, dtype=torch.long) | |
| text_lengths = [] | |
| for b, text_prompt in enumerate(prompts): | |
| messages = [{"role": "user", "content": [{"type": "text", "text": text_prompt}]}] | |
| text = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=False) | |
| toks = tokenizer(text, return_tensors="pt", add_special_tokens=False)["input_ids"][0] | |
| n = int(toks.shape[0]) | |
| if n > max_text_tokens: | |
| raise ValueError(f"prompt has {n} tokens, exceeds max_sequence_length={max_text_tokens}") | |
| text_lengths.append(n) | |
| offset = max_text_tokens - n | |
| token_ids[b, offset:] = toks | |
| attention_mask[b, offset:] = 1 | |
| text_position_ids[b, offset:] = torch.arange(n) | |
| token_ids = token_ids.to(device) | |
| attention_mask = attention_mask.to(device) | |
| text_position_ids = text_position_ids.to(device) | |
| # Run the text encoder, tapping the activation-layer hidden states, then concatenate them into per-token | |
| # text features (padding zeroed). | |
| selected = self._get_text_encoder_hidden_states( | |
| components.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_text_tokens, -1) | |
| text_features = (text_features * attention_mask.to(text_features.dtype).unsqueeze(-1)).to(torch.float32) | |
| block_state.text_features = text_features | |
| block_state.text_lengths = text_lengths | |
| self.set_block_state(state, block_state) | |
| return components, state | |