[diffusion] chore: improve z-image (#14104)

This commit is contained in:
Yuhao Yang
2025-12-01 12:26:17 +08:00
committed by GitHub
parent 982db4ebac
commit 0b9dbea593
4 changed files with 173 additions and 223 deletions
@@ -72,3 +72,71 @@ class ZImagePipelineConfig(ImagePipelineConfig):
def post_denoising_loop(self, latents, batch):
bs, channels, num_frames, height, width = latents.shape
return latents.view(bs, channels, height, width)
def get_freqs_cis(self, prompt_embeds, width, height, device, rotary_emb, batch):
def create_coordinate_grid(size, start=None, device=None):
if start is None:
start = (0 for _ in size)
axes = [
torch.arange(x0, x0 + span, dtype=torch.int32, device=device)
for x0, span in zip(start, size)
]
grids = torch.meshgrid(axes, indexing="ij")
return torch.stack(grids, dim=-1)
PATCH_SIZE = 2
F_PATCH_SIZE = 1
SEQ_MULTI_OF = 32
cap_ori_len = prompt_embeds.size(0)
cap_padding_len = (-cap_ori_len) % SEQ_MULTI_OF
cap_padded_pos_ids = create_coordinate_grid(
size=(cap_ori_len + cap_padding_len, 1, 1),
start=(1, 0, 0),
device=device,
).flatten(0, 2)
C = self.dit_config.num_channels_latents
F = 1
H = height // self.vae_config.arch_config.spatial_compression_ratio
W = width // self.vae_config.arch_config.spatial_compression_ratio
pH, pW = PATCH_SIZE, PATCH_SIZE
pF = F_PATCH_SIZE
F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW
image_ori_len = F_tokens * H_tokens * W_tokens
image_padding_len = (-image_ori_len) % SEQ_MULTI_OF
image_ori_pos_ids = create_coordinate_grid(
size=(F_tokens, H_tokens, W_tokens),
start=(cap_ori_len + cap_padding_len + 1, 0, 0),
device=device,
).flatten(0, 2)
image_padding_pos_ids = (
create_coordinate_grid(
size=(1, 1, 1),
start=(0, 0, 0),
device=device,
)
.flatten(0, 2)
.repeat(image_padding_len, 1)
)
image_padded_pos_ids = torch.cat(
[image_ori_pos_ids, image_padding_pos_ids], dim=0
)
cap_freqs_cis = rotary_emb(cap_padded_pos_ids)
x_freqs_cis = rotary_emb(image_padded_pos_ids)
return (cap_freqs_cis, x_freqs_cis)
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {
"freqs_cis": self.get_freqs_cis(
batch.prompt_embeds[0],
batch.width,
batch.height,
device,
rotary_emb,
batch,
),
}
@@ -10,6 +10,13 @@ from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
@dataclass
class ZImageSamplingParams(SamplingParams):
num_inference_steps: int = 9
num_frames: int = 1
negative_prompt: str = None
# height: int = 720
# width: int = 1280
# fps: int = 24
guidance_scale: float = 0.0
teacache_params: TeaCacheParams = field(