Skip to content

ernie-image model/pipeline review #13577

Description

@hlky

ERNIE-Image Model And Pipeline Issue Report

Tested on commit 0f1abc4ae8b0eb2a3b40e82a310507281144c423. I followed .ai/review-rules.md and the referenced model, pipeline, modular pipeline, and parity-testing rules.

Duplicate check: I searched existing Issues and PRs for the ERNIE-Image-specific findings. I found no existing reports for findings 1, 2, 3, 5, or 6. Finding 4 is already covered by PR #13532: #13532

Test coverage check: ERNIE-Image has transformer tests at tests/models/transformers/test_models_transformer_ernie_image.py and modular pipeline tests at tests/modular_pipelines/ernie_image/test_modular_pipeline_ernie_image.py. I did not find standard ErnieImagePipeline tests under tests/pipelines/.

1. [P2] use_pe=False Still Runs Modular Prompt Enhancer

AutoPipelineBlocks selects on trigger input presence, not truthiness, so use_pe=False still selects the prompt enhancer.

defselect_block(self, **kwargs) ->str|None:
"""Select block based on which trigger input is present (not None)."""
fortrigger_input, block_nameinzip(self.block_trigger_inputs, self.block_names):
iftrigger_inputisnotNoneandkwargs.get(trigger_input) isnotNone:
returnblock_name
returnNone

fromdiffusersimportErnieImageAutoBlocksblocks=ErnieImageAutoBlocks()
print("without use_pe:", "prompt_enhancer"inblocks.get_execution_blocks(prompt="x").sub_blocks)
print("with use_pe=False:", "prompt_enhancer"inblocks.get_execution_blocks(prompt="x", use_pe=False).sub_blocks)

Small fix:

-from ..modular_pipeline import AutoPipelineBlocks, SequentialPipelineBlocks+from ..modular_pipeline import ConditionalPipelineBlocks, SequentialPipelineBlocks-class ErnieImageAutoPromptEnhancerStep(AutoPipelineBlocks):+class ErnieImageAutoPromptEnhancerStep(ConditionalPipelineBlocks):
block_trigger_inputs = ["use_pe"]
++ def select_block(self, use_pe=None):+ return "prompt_enhancer" if use_pe else None

2. [P2] Standard Decode Uses Hardcoded BN Epsilon

The standard pipeline uses 1e-5 instead of self.vae.config.batch_norm_eps. AutoencoderKLFlux2 defaults batch_norm_eps to 1e-4, and both modular ERNIE and Flux2 use the config value.

bn_mean=self.vae.bn.running_mean.view(1, -1, 1, 1).to(device)
bn_std=torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) +1e-5).to(device)
latents=latents*bn_std+bn_mean

bn_mean=vae.bn.running_mean.view(1, -1, 1, 1).to(device=device, dtype=latents.dtype)
bn_std=torch.sqrt(vae.bn.running_var.view(1, -1, 1, 1) +vae.config.batch_norm_eps).to(
device=device, dtype=latents.dtype
)

latents_bn_mean=self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype)
latents_bn_std=torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) +self.vae.config.batch_norm_eps).to(
latents.device, latents.dtype
)

importinspectimporttorchfromdiffusersimportAutoencoderKLFlux2eps=inspect.signature(AutoencoderKLFlux2.__init__).parameters["batch_norm_eps"].defaultprint("configured eps:", eps)
print("hardcoded std:", torch.sqrt(torch.ones(1) +1e-5).item())
print("configured std:", torch.sqrt(torch.ones(1) +eps).item())

Small fix:

- bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(device)- bn_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + 1e-5).to(device)+ bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(device=device, dtype=latents.dtype)+ bn_std = torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) + self.vae.config.batch_norm_eps).to(+ device=device, dtype=latents.dtype+ )

3. [P2] qk_layernorm=False Cannot Instantiate

qk_layernorm=False passes qk_norm=None, but ErnieImageAttention raises for None.

qk_norm="rms_norm"ifqk_layernormelseNone,

# QK Norm
ifqk_norm=="layer_norm":
self.norm_q=torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
self.norm_k=torch.nn.LayerNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
elifqk_norm=="rms_norm":
self.norm_q=torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
self.norm_k=torch.nn.RMSNorm(dim_head, eps=eps, elementwise_affine=elementwise_affine)
else:
raiseValueError(
f"unknown qk_norm: {qk_norm}. Should be one of None, 'layer_norm', 'fp32_layer_norm', 'layer_norm_across_heads', 'rms_norm', 'rms_norm_across_heads', 'l2'."
)

fromdiffusersimportErnieImageTransformer2DModeltry:
ErnieImageTransformer2DModel(
hidden_size=16, num_attention_heads=1, num_layers=1, ffn_hidden_size=16,
in_channels=16, out_channels=16, patch_size=1, text_in_dim=16,
rope_axes_dim=(8, 4, 4), qk_layernorm=False,
)
exceptExceptionaserror:
print(type(error).__name__+": "+str(error))

Small fix:

- if qk_norm == "layer_norm":+ if qk_norm is None:+ self.norm_q = None+ self.norm_k = None+ elif qk_norm == "layer_norm":

4. [P2] Direct prompt_embeds Are Not Expanded

Duplicate: already covered by PR #13532.

#13532

Direct prompt_embeds bypass encode_prompt(), so embeddings are not repeated for num_images_per_prompt, while latents are expanded.

ifprompt_embedsisnotNone:
text_hiddens=prompt_embeds
else:
text_hiddens=self.encode_prompt(prompt, device, num_images_per_prompt)

Reference:

prompt_embeds=prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds=prompt_embeds.view(batch_size*num_images_per_prompt, seq_len, -1)

importtorchfromdiffusersimportErnieImagePipeline, ErnieImageTransformer2DModel, FlowMatchEulerDiscreteSchedulerhidden_dim=16prompt_embeds= [torch.randn(3, hidden_dim), torch.randn(4, hidden_dim)]
pipe=ErnieImagePipeline(
transformer=ErnieImageTransformer2DModel(
hidden_size=hidden_dim, num_attention_heads=1, num_layers=1, ffn_hidden_size=hidden_dim,
in_channels=hidden_dim, out_channels=hidden_dim, patch_size=1, text_in_dim=hidden_dim,
rope_axes_dim=(8, 4, 4),
),
vae=None, text_encoder=None, tokenizer=None, scheduler=FlowMatchEulerDiscreteScheduler(),
)
try:
pipe(prompt=None, prompt_embeds=prompt_embeds, height=16, width=16, num_images_per_prompt=2,
num_inference_steps=1, guidance_scale=1.0, output_type="latent")
exceptExceptionaserror:
print(type(error).__name__+": "+str(error))

Small fix is the one in PR #13532: repeat provided positive and negative embed lists by num_images_per_prompt.

5. [P2] output_type="latent" Skips Cleanup And Ignores return_dict

The latent branch returns before maybe_free_model_hooks() and returns a raw tensor even when return_dict=True.

ifoutput_type=="latent":
returnlatents
# Decode latents to images
# Unnormalize latents using VAE's BN stats
bn_mean=self.vae.bn.running_mean.view(1, -1, 1, 1).to(device)
bn_std=torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) +1e-5).to(device)
latents=latents*bn_std+bn_mean
# Unpatchify
latents=self._unpatchify_latents(latents)
# Decode
images=self.vae.decode(latents, return_dict=False)[0]
# Post-process
images= (images.clamp(-1, 1) +1) /2
images=images.cpu().permute(0, 2, 3, 1).float().numpy()
ifoutput_type=="pil":
images= [Image.fromarray((img*255).astype("uint8")) forimginimages]
# Offload all models
self.maybe_free_model_hooks()
ifnotreturn_dict:
return (images,)
returnErnieImagePipelineOutput(images=images, revised_prompts=revised_prompts)

Reference:

ifoutput_type=="latent":
image=latents
else:
latents=self._unpack_latents(latents, height, width, self.vae_scale_factor)
latents=latents.to(self.vae.dtype)
latents_mean= (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std=1.0/torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
latents.device, latents.dtype
)
latents=latents/latents_std+latents_mean
image=self.vae.decode(latents, return_dict=False)[0][:, :, 0]
image=self.image_processor.postprocess(image, output_type=output_type)
# Offload all models
self.maybe_free_model_hooks()
ifnotreturn_dict:

ifoutput_type=="latent":
image=latents
else:
latents=self._unpack_latents_with_ids(latents, latent_ids)
latents_bn_mean=self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype)
latents_bn_std=torch.sqrt(self.vae.bn.running_var.view(1, -1, 1, 1) +self.vae.config.batch_norm_eps).to(
latents.device, latents.dtype
)
latents=latents*latents_bn_std+latents_bn_mean
latents=self._unpatchify_latents(latents)
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()
ifnotreturn_dict:

importtorchfromdiffusersimportErnieImagePipeline, ErnieImageTransformer2DModel, FlowMatchEulerDiscreteSchedulerpipe=ErnieImagePipeline(
transformer=ErnieImageTransformer2DModel(
hidden_size=16, num_attention_heads=1, num_layers=1, ffn_hidden_size=16,
in_channels=16, out_channels=16, patch_size=1, text_in_dim=16, rope_axes_dim=(8, 4, 4),
),
vae=None, text_encoder=None, tokenizer=None, scheduler=FlowMatchEulerDiscreteScheduler(),
)
pipe.maybe_free_model_hooks=lambda: print("cleanup called")
out=pipe(prompt=None, prompt_embeds=[torch.randn(3, 16)], height=16, width=16,
num_inference_steps=1, guidance_scale=1.0, output_type="latent", return_dict=True)
print("returned type:", type(out).__name__)

Small fix: replace the early return with images = latents, put the decode path under else, then call maybe_free_model_hooks() and honor return_dict.

6. [P3] Tiny Modular Test Model Is In A Personal Repo

The modular fast test uses akshan-main/tiny-ernie-image-modular-pipe. Test fixtures should live under hf-internal-testing/.

pretrained_model_name_or_path="akshan-main/tiny-ernie-image-modular-pipe"

Reference:

pretrained_model_name_or_path="hf-internal-testing/tiny-wan-modular-pipe"

pretrained_model_name_or_path="hf-internal-testing/tiny-zimage-modular-pipe"

Fix: move or copy the tiny fixture to hf-internal-testing/tiny-ernie-image-modular-pipe, then update the test constant.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions