[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):
|
def post_denoising_loop(self, latents, batch):
|
||||||
bs, channels, num_frames, height, width = latents.shape
|
bs, channels, num_frames, height, width = latents.shape
|
||||||
return latents.view(bs, channels, height, width)
|
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
|
@dataclass
|
||||||
class ZImageSamplingParams(SamplingParams):
|
class ZImageSamplingParams(SamplingParams):
|
||||||
num_inference_steps: int = 9
|
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
|
guidance_scale: float = 0.0
|
||||||
|
|
||||||
teacache_params: TeaCacheParams = field(
|
teacache_params: TeaCacheParams = field(
|
||||||
|
|||||||
@@ -4,12 +4,12 @@ from typing import Any, List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch.nn.utils.rnn import pad_sequence
|
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
|
from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
@@ -150,13 +150,12 @@ class ZImageAttention(nn.Module):
|
|||||||
|
|
||||||
# Apply RoPE
|
# Apply RoPE
|
||||||
def apply_rotary_emb(
|
def apply_rotary_emb(
|
||||||
x_in: torch.Tensor, freqs_cis: torch.Tensor
|
x_in: torch.Tensor,
|
||||||
) -> torch.Tensor:
|
freqs_cis: torch.Tensor,
|
||||||
with torch.amp.autocast("cuda", enabled=False):
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2))
|
cos, sin = freqs_cis
|
||||||
freqs_cis = freqs_cis.unsqueeze(2)
|
x_out = _apply_rotary_emb(x_in, cos, sin, is_neox_style=False)
|
||||||
x_out = torch.view_as_real(x * freqs_cis).flatten(3)
|
return x_out
|
||||||
return x_out.type_as(x_in) # todo
|
|
||||||
|
|
||||||
if freqs_cis is not None:
|
if freqs_cis is not None:
|
||||||
q = apply_rotary_emb(q, freqs_cis)
|
q = apply_rotary_emb(q, freqs_cis)
|
||||||
@@ -211,7 +210,6 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
attn_mask: torch.Tensor,
|
|
||||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
|
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
|
||||||
adaln_input: Optional[torch.Tensor] = None,
|
adaln_input: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
@@ -288,12 +286,15 @@ class RopeEmbedder:
|
|||||||
assert len(axes_dims) == len(
|
assert len(axes_dims) == len(
|
||||||
axes_lens
|
axes_lens
|
||||||
), "axes_dims and axes_lens must have the same length"
|
), "axes_dims and axes_lens must have the same length"
|
||||||
self.freqs_cis = None
|
|
||||||
|
self.cos_cached = None
|
||||||
|
self.sin_cached = None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def precompute_freqs_cis(dim: List[int], end: List[int], theta: float = 256.0):
|
def precompute_freqs(dim: List[int], end: List[int], theta: float = 256.0):
|
||||||
with torch.device("cpu"):
|
with torch.device("cpu"):
|
||||||
freqs_cis = []
|
cos_list = []
|
||||||
|
sin_list = []
|
||||||
for i, (d, e) in enumerate(zip(dim, end)):
|
for i, (d, e) in enumerate(zip(dim, end)):
|
||||||
freqs = 1.0 / (
|
freqs = 1.0 / (
|
||||||
theta
|
theta
|
||||||
@@ -301,33 +302,43 @@ class RopeEmbedder:
|
|||||||
)
|
)
|
||||||
timestep = torch.arange(e, device=freqs.device, dtype=torch.float64)
|
timestep = torch.arange(e, device=freqs.device, dtype=torch.float64)
|
||||||
freqs = torch.outer(timestep, freqs).float()
|
freqs = torch.outer(timestep, freqs).float()
|
||||||
freqs_cis_i = torch.polar(torch.ones_like(freqs), freqs).to(
|
|
||||||
torch.complex64
|
|
||||||
) # complex64
|
|
||||||
freqs_cis.append(freqs_cis_i)
|
|
||||||
|
|
||||||
return freqs_cis
|
cos_list.append(torch.cos(freqs))
|
||||||
|
sin_list.append(torch.sin(freqs))
|
||||||
|
|
||||||
def __call__(self, ids: torch.Tensor):
|
return cos_list, sin_list
|
||||||
|
|
||||||
|
def __call__(self, ids: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
ids: [batch, len(axes_dims)] or [seq_len, len(axes_dims)]
|
||||||
|
Returns:
|
||||||
|
cos: [batch/seq, head_dim // 2]
|
||||||
|
sin: [batch/seq, head_dim // 2]
|
||||||
|
"""
|
||||||
assert ids.ndim == 2
|
assert ids.ndim == 2
|
||||||
assert ids.shape[-1] == len(self.axes_dims)
|
assert ids.shape[-1] == len(self.axes_dims)
|
||||||
device = ids.device
|
device = ids.device
|
||||||
|
|
||||||
if self.freqs_cis is None:
|
if self.cos_cached is None:
|
||||||
self.freqs_cis = self.precompute_freqs_cis(
|
self.cos_cached, self.sin_cached = self.precompute_freqs(
|
||||||
self.axes_dims, self.axes_lens, theta=self.theta
|
self.axes_dims, self.axes_lens, theta=self.theta
|
||||||
)
|
)
|
||||||
self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis]
|
self.cos_cached = [c.to(device) for c in self.cos_cached]
|
||||||
|
self.sin_cached = [s.to(device) for s in self.sin_cached]
|
||||||
else:
|
else:
|
||||||
# Ensure freqs_cis are on the same device as ids
|
if self.cos_cached[0].device != device:
|
||||||
if self.freqs_cis[0].device != device:
|
self.cos_cached = [c.to(device) for c in self.cos_cached]
|
||||||
self.freqs_cis = [freqs_cis.to(device) for freqs_cis in self.freqs_cis]
|
self.sin_cached = [s.to(device) for s in self.sin_cached]
|
||||||
|
|
||||||
result = []
|
cos_out = []
|
||||||
|
sin_out = []
|
||||||
for i in range(len(self.axes_dims)):
|
for i in range(len(self.axes_dims)):
|
||||||
index = ids[:, i]
|
index = ids[:, i]
|
||||||
result.append(self.freqs_cis[i][index])
|
cos_out.append(self.cos_cached[i][index])
|
||||||
return torch.cat(result, dim=-1)
|
sin_out.append(self.sin_cached[i][index])
|
||||||
|
|
||||||
|
return torch.cat(cos_out, dim=-1), torch.cat(sin_out, dim=-1)
|
||||||
|
|
||||||
|
|
||||||
class ZImageTransformer2DModel(CachableDiT):
|
class ZImageTransformer2DModel(CachableDiT):
|
||||||
@@ -435,7 +446,7 @@ class ZImageTransformer2DModel(CachableDiT):
|
|||||||
self.axes_dims = arch_config.axes_dims
|
self.axes_dims = arch_config.axes_dims
|
||||||
self.axes_lens = arch_config.axes_lens
|
self.axes_lens = arch_config.axes_lens
|
||||||
|
|
||||||
self.rope_embedder = RopeEmbedder(
|
self.rotary_emb = RopeEmbedder(
|
||||||
theta=self.rope_theta, axes_dims=self.axes_dims, axes_lens=self.axes_lens
|
theta=self.rope_theta, axes_dims=self.axes_dims, axes_lens=self.axes_lens
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -477,114 +488,64 @@ class ZImageTransformer2DModel(CachableDiT):
|
|||||||
patch_size: int,
|
patch_size: int,
|
||||||
f_patch_size: int,
|
f_patch_size: int,
|
||||||
):
|
):
|
||||||
|
assert len(all_image) == len(all_cap_feats) == 1
|
||||||
|
|
||||||
|
image = all_image[0] # C, F, H, W
|
||||||
|
cap_feat = all_cap_feats[0] # L, D
|
||||||
pH = pW = patch_size
|
pH = pW = patch_size
|
||||||
pF = f_patch_size
|
pF = f_patch_size
|
||||||
device = all_image[0].device
|
device = image.device
|
||||||
|
|
||||||
all_image_out = []
|
all_image_out = []
|
||||||
all_image_size = []
|
all_image_size = []
|
||||||
all_image_pos_ids = []
|
|
||||||
all_image_pad_mask = []
|
|
||||||
all_cap_pos_ids = []
|
|
||||||
all_cap_pad_mask = []
|
|
||||||
all_cap_feats_out = []
|
all_cap_feats_out = []
|
||||||
|
|
||||||
for i, (image, cap_feat) in enumerate(zip(all_image, all_cap_feats)):
|
# ------------ Process Caption ------------
|
||||||
### Process Caption
|
cap_ori_len = cap_feat.size(0)
|
||||||
cap_ori_len = len(cap_feat)
|
cap_padding_len = (-cap_ori_len) % SEQ_MULTI_OF
|
||||||
cap_padding_len = (-cap_ori_len) % SEQ_MULTI_OF
|
|
||||||
# padded position ids
|
|
||||||
cap_padded_pos_ids = self.create_coordinate_grid(
|
|
||||||
size=(cap_ori_len + cap_padding_len, 1, 1),
|
|
||||||
start=(1, 0, 0),
|
|
||||||
device=device,
|
|
||||||
).flatten(0, 2)
|
|
||||||
all_cap_pos_ids.append(cap_padded_pos_ids)
|
|
||||||
# pad mask
|
|
||||||
all_cap_pad_mask.append(
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.zeros((cap_ori_len,), dtype=torch.bool, device=device),
|
|
||||||
torch.ones((cap_padding_len,), dtype=torch.bool, device=device),
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
# padded feature
|
|
||||||
cap_padded_feat = torch.cat(
|
|
||||||
[cap_feat, cap_feat[-1:].repeat(cap_padding_len, 1)],
|
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
all_cap_feats_out.append(cap_padded_feat)
|
|
||||||
|
|
||||||
### Process Image
|
# padded feature
|
||||||
C, F, H, W = image.size()
|
cap_padded_feat = torch.cat(
|
||||||
all_image_size.append((F, H, W))
|
[cap_feat, cap_feat[-1:].repeat(cap_padding_len, 1)],
|
||||||
F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW
|
dim=0,
|
||||||
|
)
|
||||||
|
all_cap_feats_out.append(cap_padded_feat)
|
||||||
|
|
||||||
image = image.view(C, F_tokens, pF, H_tokens, pH, W_tokens, pW)
|
# ------------ Process Image ------------
|
||||||
# "c f pf h ph w pw -> (f h w) (pf ph pw c)"
|
C, F, H, W = image.size()
|
||||||
image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(
|
all_image_size.append((F, H, W))
|
||||||
F_tokens * H_tokens * W_tokens, pF * pH * pW * C
|
|
||||||
)
|
|
||||||
|
|
||||||
image_ori_len = len(image)
|
F_tokens, H_tokens, W_tokens = F // pF, H // pH, W // pW
|
||||||
image_padding_len = (-image_ori_len) % SEQ_MULTI_OF
|
image = image.view(C, F_tokens, pF, H_tokens, pH, W_tokens, pW)
|
||||||
|
# "c f pf h ph w pw -> (f h w) (pf ph pw c)"
|
||||||
|
image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(
|
||||||
|
F_tokens * H_tokens * W_tokens, pF * pH * pW * C
|
||||||
|
)
|
||||||
|
image_ori_len = image.size(0)
|
||||||
|
image_padding_len = (-image_ori_len) % SEQ_MULTI_OF
|
||||||
|
|
||||||
image_ori_pos_ids = self.create_coordinate_grid(
|
# padded feature
|
||||||
size=(F_tokens, H_tokens, W_tokens),
|
image_padded_feat = torch.cat(
|
||||||
start=(cap_ori_len + cap_padding_len + 1, 0, 0),
|
[image, image[-1:].repeat(image_padding_len, 1)],
|
||||||
device=device,
|
dim=0,
|
||||||
).flatten(0, 2)
|
)
|
||||||
image_padding_pos_ids = (
|
all_image_out.append(image_padded_feat)
|
||||||
self.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
|
|
||||||
)
|
|
||||||
all_image_pos_ids.append(image_padded_pos_ids)
|
|
||||||
# pad mask
|
|
||||||
all_image_pad_mask.append(
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.zeros((image_ori_len,), dtype=torch.bool, device=device),
|
|
||||||
torch.ones(
|
|
||||||
(image_padding_len,), dtype=torch.bool, device=device
|
|
||||||
),
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
# padded feature
|
|
||||||
image_padded_feat = torch.cat(
|
|
||||||
[image, image[-1:].repeat(image_padding_len, 1)], dim=0
|
|
||||||
)
|
|
||||||
all_image_out.append(image_padded_feat)
|
|
||||||
|
|
||||||
return (
|
return (
|
||||||
all_image_out,
|
all_image_out,
|
||||||
all_cap_feats_out,
|
all_cap_feats_out,
|
||||||
all_image_size,
|
all_image_size,
|
||||||
all_image_pos_ids,
|
|
||||||
all_cap_pos_ids,
|
|
||||||
all_image_pad_mask,
|
|
||||||
all_cap_pad_mask,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
hidden_states: List[torch.Tensor],
|
hidden_states: List[torch.Tensor],
|
||||||
timestep,
|
|
||||||
encoder_hidden_states: List[torch.Tensor],
|
encoder_hidden_states: List[torch.Tensor],
|
||||||
|
timestep,
|
||||||
guidance=0,
|
guidance=0,
|
||||||
patch_size=2,
|
patch_size=2,
|
||||||
f_patch_size=1,
|
f_patch_size=1,
|
||||||
|
freqs_cis=None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
assert patch_size in self.all_patch_size
|
assert patch_size in self.all_patch_size
|
||||||
@@ -594,129 +555,43 @@ class ZImageTransformer2DModel(CachableDiT):
|
|||||||
cap_feats = encoder_hidden_states
|
cap_feats = encoder_hidden_states
|
||||||
timestep = 1000.0 - timestep
|
timestep = 1000.0 - timestep
|
||||||
t = timestep
|
t = timestep
|
||||||
bsz = len(x)
|
bsz = 1
|
||||||
device = x[0].device
|
device = x[0].device
|
||||||
# t = t * self.t_scale
|
|
||||||
t = self.t_embedder(t)
|
t = self.t_embedder(t)
|
||||||
|
|
||||||
adaln_input = t.type_as(x)
|
adaln_input = t.type_as(x)
|
||||||
(
|
(
|
||||||
x,
|
x,
|
||||||
cap_feats,
|
cap_feats,
|
||||||
x_size,
|
x_size,
|
||||||
x_pos_ids,
|
|
||||||
cap_pos_ids,
|
|
||||||
x_inner_pad_mask,
|
|
||||||
cap_inner_pad_mask,
|
|
||||||
) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size)
|
) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size)
|
||||||
# x embed & refine
|
|
||||||
x_item_seqlens = [len(_) for _ in x]
|
|
||||||
assert all(_ % SEQ_MULTI_OF == 0 for _ in x_item_seqlens)
|
|
||||||
x_max_item_seqlen = max(x_item_seqlens)
|
|
||||||
|
|
||||||
x = torch.cat(x, dim=0)
|
x = torch.cat(x, dim=0)
|
||||||
x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x)
|
x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x)
|
||||||
x[torch.cat(x_inner_pad_mask)] = self.x_pad_token
|
x_freqs_cis = freqs_cis[1]
|
||||||
x = list(x.split(x_item_seqlens, dim=0))
|
|
||||||
x_freqs_cis = list(
|
|
||||||
self.rope_embedder(torch.cat(x_pos_ids, dim=0)).split(x_item_seqlens, dim=0)
|
|
||||||
)
|
|
||||||
|
|
||||||
# RoPE returns (cos, sin) now
|
x = x.unsqueeze(0)
|
||||||
# x_pos_ids_cat = torch.cat(x_pos_ids, dim=0)
|
x_freqs_cis = x_freqs_cis
|
||||||
# x_cos, x_sin = self.rope_embedder(x_pos_ids_cat)
|
|
||||||
|
|
||||||
# x_cos_list = list(x_cos.split(x_item_seqlens, dim=0))
|
|
||||||
# x_sin_list = list(x_sin.split(x_item_seqlens, dim=0))
|
|
||||||
|
|
||||||
x = pad_sequence(x, batch_first=True, padding_value=0.0)
|
|
||||||
x_freqs_cis = pad_sequence(x_freqs_cis, batch_first=True, padding_value=0.0)
|
|
||||||
# x_cos = pad_sequence(x_cos_list, batch_first=True, padding_value=0.0)
|
|
||||||
# x_sin = pad_sequence(x_sin_list, batch_first=True, padding_value=0.0)
|
|
||||||
# B, T, D_half = x_cos.shape # D_half = 64
|
|
||||||
|
|
||||||
# x_cos_triton = x_cos.reshape(B * T, D_half).contiguous() # [B*T, 64]
|
|
||||||
# x_sin_triton = x_sin.reshape(B * T, D_half).contiguous() # [B*T, 64]
|
|
||||||
|
|
||||||
x_attn_mask = torch.zeros(
|
|
||||||
(bsz, x_max_item_seqlen), dtype=torch.bool, device=device
|
|
||||||
)
|
|
||||||
for i, seq_len in enumerate(x_item_seqlens):
|
|
||||||
x_attn_mask[i, :seq_len] = 1
|
|
||||||
|
|
||||||
# Refiner logic
|
|
||||||
for layer in self.noise_refiner:
|
for layer in self.noise_refiner:
|
||||||
x = layer(x, x_attn_mask, x_freqs_cis, adaln_input)
|
x = layer(x, x_freqs_cis, adaln_input)
|
||||||
|
|
||||||
# cap embed & refine
|
|
||||||
cap_item_seqlens = [len(_) for _ in cap_feats]
|
|
||||||
assert all(_ % SEQ_MULTI_OF == 0 for _ in cap_item_seqlens)
|
|
||||||
cap_max_item_seqlen = max(cap_item_seqlens)
|
|
||||||
|
|
||||||
cap_feats = torch.cat(cap_feats, dim=0)
|
cap_feats = torch.cat(cap_feats, dim=0)
|
||||||
|
|
||||||
# cap_embedder is Sequential with ReplicatedLinear.
|
cap_feats, _ = self.cap_embedder(cap_feats)
|
||||||
# We need to handle this if ReplicatedLinear returns tuple.
|
|
||||||
# In __init__, cap_embedder = Sequential(RMSNorm, ReplicatedLinear).
|
|
||||||
# RMSNorm returns Tensor. ReplicatedLinear returns (Tensor, Gathered).
|
|
||||||
# Sequential returns (Tensor, Gathered).
|
|
||||||
# So we need to unpack.
|
|
||||||
cap_feats_out = self.cap_embedder(cap_feats)
|
|
||||||
if isinstance(cap_feats_out, tuple):
|
|
||||||
cap_feats = cap_feats_out[0]
|
|
||||||
else:
|
|
||||||
cap_feats = cap_feats_out
|
|
||||||
|
|
||||||
cap_feats[torch.cat(cap_inner_pad_mask)] = self.cap_pad_token
|
cap_freqs_cis = freqs_cis[0]
|
||||||
cap_feats = list(cap_feats.split(cap_item_seqlens, dim=0))
|
|
||||||
|
|
||||||
cap_freqs_cis = list(
|
|
||||||
self.rope_embedder(torch.cat(cap_pos_ids, dim=0)).split(
|
|
||||||
cap_item_seqlens, dim=0
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
cap_feats = pad_sequence(cap_feats, batch_first=True, padding_value=0.0)
|
|
||||||
cap_freqs_cis = pad_sequence(cap_freqs_cis, batch_first=True, padding_value=0.0)
|
|
||||||
|
|
||||||
cap_attn_mask = torch.zeros(
|
|
||||||
(bsz, cap_max_item_seqlen), dtype=torch.bool, device=device
|
|
||||||
)
|
|
||||||
for i, seq_len in enumerate(cap_item_seqlens):
|
|
||||||
cap_attn_mask[i, :seq_len] = 1
|
|
||||||
|
|
||||||
|
cap_feats = cap_feats.unsqueeze(0)
|
||||||
for layer in self.context_refiner:
|
for layer in self.context_refiner:
|
||||||
cap_feats = layer(cap_feats, cap_attn_mask, cap_freqs_cis)
|
cap_feats = layer(cap_feats, cap_freqs_cis)
|
||||||
|
|
||||||
# unified
|
unified = torch.cat([x, cap_feats], dim=1)
|
||||||
unified = []
|
unified_freqs_cis = (
|
||||||
unified_freqs_cis = []
|
torch.cat([x_freqs_cis[0], cap_freqs_cis[0]], dim=0),
|
||||||
|
torch.cat([x_freqs_cis[1], cap_freqs_cis[1]], dim=0),
|
||||||
for i in range(bsz):
|
|
||||||
x_len = x_item_seqlens[i]
|
|
||||||
cap_len = cap_item_seqlens[i]
|
|
||||||
unified.append(torch.cat([x[i][:x_len], cap_feats[i][:cap_len]]))
|
|
||||||
unified_freqs_cis.append(
|
|
||||||
torch.cat([x_freqs_cis[i][:x_len], cap_freqs_cis[i][:cap_len]])
|
|
||||||
)
|
|
||||||
|
|
||||||
unified_item_seqlens = [a + b for a, b in zip(cap_item_seqlens, x_item_seqlens)]
|
|
||||||
assert unified_item_seqlens == [len(_) for _ in unified]
|
|
||||||
unified_max_item_seqlen = max(unified_item_seqlens)
|
|
||||||
|
|
||||||
unified = pad_sequence(unified, batch_first=True, padding_value=0.0)
|
|
||||||
unified_freqs_cis = pad_sequence(
|
|
||||||
unified_freqs_cis, batch_first=True, padding_value=0.0
|
|
||||||
)
|
)
|
||||||
|
|
||||||
unified_attn_mask = torch.zeros(
|
|
||||||
(bsz, unified_max_item_seqlen), dtype=torch.bool, device=device
|
|
||||||
)
|
|
||||||
for i, seq_len in enumerate(unified_item_seqlens):
|
|
||||||
unified_attn_mask[i, :seq_len] = 1
|
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
unified = layer(unified, unified_attn_mask, unified_freqs_cis, adaln_input)
|
unified = layer(unified, unified_freqs_cis, adaln_input)
|
||||||
|
|
||||||
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](
|
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](
|
||||||
unified, adaln_input
|
unified, adaln_input
|
||||||
|
|||||||
@@ -850,18 +850,18 @@ class DenoisingStage(PipelineStage):
|
|||||||
# Predict noise residual
|
# Predict noise residual
|
||||||
attn_metadata = self._build_attn_metadata(i, batch, server_args)
|
attn_metadata = self._build_attn_metadata(i, batch, server_args)
|
||||||
noise_pred = self._predict_noise_with_cfg(
|
noise_pred = self._predict_noise_with_cfg(
|
||||||
current_model,
|
current_model=current_model,
|
||||||
latent_model_input,
|
latent_model_input=latent_model_input,
|
||||||
timestep,
|
timestep=timestep,
|
||||||
batch,
|
batch=batch,
|
||||||
i,
|
timestep_index=i,
|
||||||
attn_metadata,
|
attn_metadata=attn_metadata,
|
||||||
target_dtype,
|
target_dtype=target_dtype,
|
||||||
current_guidance_scale,
|
current_guidance_scale=current_guidance_scale,
|
||||||
image_kwargs,
|
image_kwargs=image_kwargs,
|
||||||
pos_cond_kwargs,
|
pos_cond_kwargs=pos_cond_kwargs,
|
||||||
neg_cond_kwargs,
|
neg_cond_kwargs=neg_cond_kwargs,
|
||||||
server_args,
|
server_args=server_args,
|
||||||
guidance=guidance,
|
guidance=guidance,
|
||||||
latents=latents,
|
latents=latents,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user