[diffusion] chore: improve z-image (#14104)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user