[feature] support ada ln modulation, token concat and cfg (#14)

* [feature] support ada ln modulation

* [feature] token concat

* [feature] support cfg
This commit is contained in:
Hongxin Liu 2024-02-27 16:55:11 +08:00 committed by GitHub
parent 0c05cd2e9d
commit 14db4566e1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 213 additions and 52 deletions

View file

@ -46,7 +46,9 @@ class CrossAttention(nn.Module):
):
super().__init__()
self.hidden_size = head_dim * num_heads
cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
cross_attention_dim = (
cross_attention_dim if cross_attention_dim is not None else query_dim
)
self.scale = head_dim**-0.5
self.num_heads = num_heads
@ -57,7 +59,9 @@ class CrossAttention(nn.Module):
self.to_k = nn.Linear(cross_attention_dim, self.hidden_size, bias=bias)
self.to_v = nn.Linear(cross_attention_dim, self.hidden_size, bias=bias)
self.to_out = nn.Sequential(nn.Linear(self.hidden_size, query_dim), nn.Dropout(dropout))
self.to_out = nn.Sequential(
nn.Linear(self.hidden_size, query_dim), nn.Dropout(dropout)
)
def forward(self, hidden_states, context=None, mask=None):
bsz, q_len, _ = hidden_states.shape
@ -71,18 +75,24 @@ class CrossAttention(nn.Module):
# [B, S, H, D]
query = query.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
key = key.view(bsz, kv_seq_len, self.num_heads, self.head_dim).transpose(1, 2)
value = value.view(bsz, kv_seq_len, self.num_heads, self.head_dim).transpose(1, 2)
value = value.view(bsz, kv_seq_len, self.num_heads, self.head_dim).transpose(
1, 2
)
if mask is not None:
assert mask.shape == (bsz, 1, q_len, kv_seq_len)
if self.sdpa:
attn_output = F.scaled_dot_product_attention(query, key, value, attn_mask=mask, scale=self.scale)
attn_output = F.scaled_dot_product_attention(
query, key, value, attn_mask=mask, scale=self.scale
)
else:
attn_weights = torch.matmul(query, key.transpose(2, 3)) / self.scale
assert attn_weights.shape == (bsz, self.num_heads, q_len, kv_seq_len)
if mask is not None:
attn_weights = attn_weights + mask
attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(
query.dtype
)
attn_output = torch.matmul(attn_weights, value)
assert attn_output.shape == (bsz, self.num_heads, q_len, self.head_dim)
attn_output = attn_output.transpose(1, 2).contiguous()
@ -126,17 +136,23 @@ class TimestepEmbedder(nn.Module):
"""
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
device=t.device
)
freqs = torch.exp(
-math.log(max_period)
* torch.arange(start=0, end=half, dtype=torch.float32)
/ half
).to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
embedding = torch.cat(
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
)
return embedding
def forward(self, t):
t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(self.mlp[0].weight.dtype)
t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(
self.mlp[0].weight.dtype
)
t_emb = self.mlp(t_freq)
return t_emb
@ -149,7 +165,9 @@ class LabelEmbedder(nn.Module):
def __init__(self, num_classes, hidden_size, dropout_prob):
super().__init__()
use_cfg_embedding = dropout_prob > 0
self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size)
self.embedding_table = nn.Embedding(
num_classes + use_cfg_embedding, hidden_size
)
self.num_classes = num_classes
self.dropout_prob = dropout_prob
@ -158,7 +176,9 @@ class LabelEmbedder(nn.Module):
Drops labels to enable classifier-free guidance.
"""
if force_drop_ids is None:
drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
drop_ids = (
torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob
)
else:
drop_ids = force_drop_ids == 1
labels = torch.where(drop_ids, self.num_classes, labels)
@ -185,27 +205,50 @@ class PatchEmbedder(nn.Module):
) -> None:
super().__init__()
self.patch_size = patch_size
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
self.proj = nn.Conv2d(
in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias
)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
# [B, S, C, P, P] -> [B, S, C*P*P]
# FIXME: hack diffusion and use view
x = x.view(*x.shape[:2], -1)
out = F.linear(x, self.proj.weight.view(self.proj.weight.shape[0], -1), self.proj.bias)
out = F.linear(
x, self.proj.weight.view(self.proj.weight.shape[0], -1), self.proj.bias
)
out = self.norm(out)
# [B, S, H]
return out
class TextEmbedder(nn.Module):
def __init__(self, in_features: int, embed_dim: int = 768, bias: bool = True) -> None:
def __init__(
self,
in_features: int,
embed_dim: int = 768,
bias: bool = True,
dropout_prob: float = 0.0,
use_proj: bool = True,
) -> None:
super().__init__()
self.proj = nn.Linear(in_features, embed_dim, bias=bias)
self.dropout_prob = dropout_prob
self.use_proj = use_proj
if self.use_proj:
self.proj = nn.Linear(in_features, embed_dim, bias=bias)
def drop_sample(self, x: torch.Tensor) -> torch.Tensor:
drop_ids = torch.rand(x.shape[0], device=x.device) < self.dropout_prob
x = torch.where(drop_ids, torch.zeros_like(x), x)
return x
def forward(self, x: torch.Tensor) -> torch.Tensor:
# [B, S, C] -> [B, S, H]
return self.proj(x)
use_dropout = self.dropout_prob > 0
if self.training and use_dropout:
x = self.drop_sample(x)
if self.use_proj:
# [B, S, C] -> [B, S, H]
x = self.proj(x)
return x
class PositionEmbedding(nn.Module):
@ -241,7 +284,13 @@ class DiTBlock(nn.Module):
A DiT block with adaptive layer norm zero (adaLN-Zero) conditioning.
"""
def __init__(self, hidden_size, num_heads, cross_attention_dim, mlp_ratio=4.0, **block_kwargs):
def __init__(
self,
hidden_size,
num_heads,
cross_attention_dim=None,
mlp_ratio=4.0,
):
super().__init__()
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.attn = CrossAttention(
@ -262,11 +311,26 @@ class DiTBlock(nn.Module):
act_layer=approx_gelu,
drop=0,
)
self.mlp = Mlp(
in_features=hidden_size,
hidden_features=mlp_hidden_dim,
act_layer=approx_gelu,
drop=0,
)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)
)
def forward(self, x, context, attention_mask):
# TODO: use cross attn
x = x + self.attn(self.norm1(x), context, attention_mask)
x = x + self.mlp(self.norm2(x))
def forward(self, x, attention_mask, t, context=None):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.adaLN_modulation(t).chunk(6, dim=1)
)
x = x + gate_msa.unsqueeze(1) * self.attn(
modulate(self.norm1(x), shift_msa, scale_msa), context, attention_mask
)
x = x + gate_mlp.unsqueeze(1) * self.mlp(
modulate(self.norm2(x), shift_mlp, scale_mlp)
)
return x
@ -277,14 +341,22 @@ class FinalLayer(nn.Module):
def __init__(self, hidden_size, patch_size, out_channels):
super().__init__()
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.linear = nn.Linear(
hidden_size, patch_size * patch_size * out_channels, bias=True
)
self.patch_size = patch_size
self.adaLN_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)
)
def unpatchify(self, x):
b, s, h = x.shape
return x.view(b, s, -1, self.patch_size, self.patch_size)
def forward(self, x):
def forward(self, x, t):
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
x = self.unpatchify(x)
return x
@ -305,7 +377,9 @@ class DiT(nn.Module):
num_heads=16,
mlp_ratio=4.0,
max_num_embeddings=256 * 1024,
text_dropout_prob=0.1,
learn_sigma=True,
use_cross_attn=True,
):
super().__init__()
self.grad_checkpointing = False
@ -315,12 +389,29 @@ class DiT(nn.Module):
self.patch_size = patch_size
self.num_heads = num_heads
self.video_embedder = PatchEmbedder(patch_size, in_channels, hidden_size, bias=True)
self.t_embedder = TimestepEmbedder(text_embed_dim)
self.video_embedder = PatchEmbedder(
patch_size, in_channels, hidden_size, bias=True
)
self.t_embedder = TimestepEmbedder(hidden_size)
self.pos_embed = PositionEmbedding(hidden_size, max_num_embeddings)
self.text_embedder = TextEmbedder(
text_embed_dim,
hidden_size,
bias=True,
dropout_prob=text_dropout_prob,
use_proj=not use_cross_attn,
)
if not use_cross_attn:
cross_attn_dim = None
else:
cross_attn_dim = text_embed_dim
self.use_cross_attn = use_cross_attn
self.blocks = nn.ModuleList(
[DiTBlock(hidden_size, num_heads, text_embed_dim, mlp_ratio=mlp_ratio) for _ in range(depth)]
[
DiTBlock(hidden_size, num_heads, cross_attn_dim, mlp_ratio=mlp_ratio)
for _ in range(depth)
]
)
self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
self.initialize_weights()
@ -335,12 +426,22 @@ class DiT(nn.Module):
self.apply(_basic_init)
# TODO: update patch embed init
# Initialize text embedding layer
if self.text_embedder.use_proj:
nn.init.normal_(self.text_embedder.proj.weight, std=0.02)
# Initialize timestep embedding MLP:
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
# Zero-out adaLN modulation layers
for block in self.blocks:
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
# Zero-out output
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
nn.init.constant_(self.final_layer.linear.weight, 0)
nn.init.constant_(self.final_layer.linear.bias, 0)
@ -349,7 +450,9 @@ class DiT(nn.Module):
assert attention_mask.ndim == 4
attention_mask = attention_mask.to(dtype)
inverted_mask = 1.0 - attention_mask
return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)
return inverted_mask.masked_fill(
inverted_mask.to(torch.bool), torch.finfo(dtype).min
)
return attention_mask
def enable_gradient_checkpointing(self):
@ -369,29 +472,47 @@ class DiT(nn.Module):
video_latent_states: [B, S, C, P, P]
"""
video_latent_states = self.video_embedder(video_latent_states)
text_len = text_latent_states.shape[1]
text_latent_states = self.text_embedder(text_latent_states)
if not self.use_cross_attn:
video_latent_states = torch.cat(
[text_latent_states, video_latent_states], dim=1
)
text_latent_states = None
pos_embed = self.pos_embed(video_latent_states)
video_latent_states = video_latent_states + pos_embed
t = self.t_embedder(t) # (N, D)
text_latent_states = text_latent_states + t.unsqueeze(1)
attention_mask = self._prepare_mask(attention_mask, video_latent_states.dtype)
for block in self.blocks:
if self.grad_checkpointing and self.training:
video_latent_states = torch.utils.checkpoint.checkpoint(
block, video_latent_states, text_latent_states, attention_mask
block,
video_latent_states,
attention_mask,
t,
text_latent_states,
)
else:
video_latent_states = block(video_latent_states, text_latent_states, attention_mask)
video_latent_states = self.final_layer(video_latent_states)
video_latent_states = block(
video_latent_states, attention_mask, t, text_latent_states
)
if not self.use_cross_attn:
video_latent_states = video_latent_states[:, text_len:]
video_latent_states = self.final_layer(video_latent_states, t)
return video_latent_states
def forward_with_cfg(self, x, t, text_latent_states, cfg_scale, attention_mask=None):
def forward_with_cfg(
self, x, t, text_latent_states, cfg_scale, attention_mask=None
):
"""
Forward pass of DiT, but also batches the unconditional forward pass for classifier-free guidance.
"""
# https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb
half = x[: len(x) // 2]
combined = torch.cat([half, half], dim=0)
model_out = self.forward(combined, t, text_latent_states, attention_mask=attention_mask)
model_out = self.forward(
combined, t, text_latent_states, attention_mask=attention_mask
)
# For exact reproducibility reasons, we apply classifier-free guidance on only
# three channels by default. The standard approach to cfg applies it to all channels.
# This can be done by uncommenting the following line and commenting-out the line following that.
@ -425,7 +546,9 @@ def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=
grid = grid.reshape([2, 1, grid_size, grid_size])
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
if cls_token and extra_tokens > 0:
pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
pos_embed = np.concatenate(
[np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0
)
return pos_embed

View file

@ -37,7 +37,9 @@ def video2col(video_4d: torch.Tensor, patch_size: int) -> torch.Tensor:
return torch.stack(out, dim=1).view(-1, c, patch_size, patch_size)
def col2video(patches: torch.Tensor, video_shape: Tuple[int, int, int, int]) -> torch.Tensor:
def col2video(
patches: torch.Tensor, video_shape: Tuple[int, int, int, int]
) -> torch.Tensor:
"""
Convert a 2D tensor of patches to a 4D video tensor.
@ -74,7 +76,10 @@ def pad_sequences(sequences: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Te
"""
max_len = max([sequence.shape[0] for sequence in sequences])
padded_sequences = [
F.pad(sequence, [0] * (sequence.ndim - 1) * 2 + [0, max_len - sequence.shape[0]]) for sequence in sequences
F.pad(
sequence, [0] * (sequence.ndim - 1) * 2 + [0, max_len - sequence.shape[0]]
)
for sequence in sequences
]
padded_sequences = torch.stack(padded_sequences, dim=0)
padding_mask = torch.zeros(
@ -88,7 +93,9 @@ def pad_sequences(sequences: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Te
return padded_sequences, padding_mask
def patchify_batch(videos: List[torch.Tensor], patch_size: int) -> Tuple[torch.Tensor, torch.Tensor]:
def patchify_batch(
videos: List[torch.Tensor], patch_size: int
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Patchify a batch of videos.
Args:
@ -127,7 +134,10 @@ def make_batch(samples: List[dict], video_dir: str) -> dict:
Returns:
dict: A batch of samples.
"""
videos = [read_video(os.path.join(video_dir, sample["video_file"]), pts_unit="sec")[0] for sample in samples]
videos = [
read_video(os.path.join(video_dir, sample["video_file"]), pts_unit="sec")[0]
for sample in samples
]
texts = [sample["text_latent_states"] for sample in samples]
texts, text_padding_mask = pad_sequences(texts)
return {
@ -146,7 +156,13 @@ def unnormalize_video(video: torch.Tensor) -> torch.Tensor:
@torch.no_grad()
def preprocess_batch(batch: dict, patch_size: int, vqvae: Optional[nn.Module] = None, device=None) -> dict:
def preprocess_batch(
batch: dict,
patch_size: int,
vqvae: Optional[nn.Module] = None,
device=None,
use_cross_attn=True,
) -> dict:
if device is None:
device = get_current_device()
videos = []
@ -169,19 +185,27 @@ def preprocess_batch(batch: dict, patch_size: int, vqvae: Optional[nn.Module] =
batch["video_latent_states"] = video_latent_states
batch["video_padding_mask"] = video_padding_mask
text_padding_mask = batch.pop("text_padding_mask").to(device)
batch["attention_mask"] = expand_mask_4d(video_padding_mask, text_padding_mask)
if use_cross_attn:
batch["attention_mask"] = expand_mask_4d(video_padding_mask, text_padding_mask)
else:
attention_mask = torch.cat([text_padding_mask, video_padding_mask], dim=1)
batch["attention_mask"] = expand_mask_4d(attention_mask, attention_mask)
batch["text_latent_states"] = batch["text_latent_states"].to(device)
return batch
def load_datasets(dataset_paths: Union[PathType, List[PathType]], mode: str = "train") -> Optional[DatasetType]:
def load_datasets(
dataset_paths: Union[PathType, List[PathType]], mode: str = "train"
) -> Optional[DatasetType]:
"""
Load pre-tokenized dataset.
Each instance of dataset is a dictionary with
`{'input_ids': List[int], 'labels': List[int], sequence: str}` format.
"""
mode_map = {"train": "train", "dev": "validation", "test": "test"}
assert mode in tuple(mode_map), f"Unsupported mode {mode}, it must be in {tuple(mode_map)}"
assert mode in tuple(
mode_map
), f"Unsupported mode {mode}, it must be in {tuple(mode_map)}"
if isinstance(dataset_paths, (str, os.PathLike)):
dataset_paths = [dataset_paths]
@ -190,7 +214,9 @@ def load_datasets(dataset_paths: Union[PathType, List[PathType]], mode: str = "t
for ds_path in dataset_paths:
ds_path = os.path.abspath(ds_path)
assert os.path.exists(ds_path), f"Not existed file path {ds_path}"
ds_dict = load_from_disk(dataset_path=ds_path, keep_in_memory=False).with_format("torch")
ds_dict = load_from_disk(
dataset_path=ds_path, keep_in_memory=False
).with_format("torch")
if isinstance(ds_dict, HFDataset):
datasets.append(ds_dict)
else:

View file

@ -28,7 +28,11 @@ def main(args):
torch.set_grad_enabled(False)
device = get_current_device()
if len(args.vqvae) > 0:
vqvae = AutoModel.from_pretrained(args.vqvae, trust_remote_code=True).to(device).eval()
vqvae = (
AutoModel.from_pretrained(args.vqvae, trust_remote_code=True)
.to(device)
.eval()
)
in_channels = vqvae.embedding_dim
else:
# disable VQ-VAE if not provided, just use raw video frames
@ -50,7 +54,9 @@ def main(args):
num_frames = args.fps * args.sec
z = torch.randn(
1,
(args.height // patch_size // 4) * (args.width // patch_size // 4) * (num_frames // 2),
(args.height // patch_size // 4)
* (args.width // patch_size // 4)
* (num_frames // 2),
in_channels,
patch_size,
patch_size,
@ -60,7 +66,9 @@ def main(args):
# Setup classifier-free guidance:
model_kwargs = {}
z = torch.cat([z, z], 0)
model_kwargs["text_latent_states"] = torch.cat([text_latent_states, text_latent_states], 0)
model_kwargs["text_latent_states"] = torch.cat(
[text_latent_states, torch.zeros_like(text_latent_states)], 0
)
model_kwargs["cfg_scale"] = args.cfg_scale
model_kwargs["attention_mask"] = torch.ones(
2, 1, z.shape[1], text_latent_states.shape[1], device=device, dtype=torch.int
@ -96,7 +104,9 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, choices=list(DiT_models.keys()), default="DiT-S/8")
parser.add_argument(
"--model", type=str, choices=list(DiT_models.keys()), default="DiT-S/8"
)
parser.add_argument(
"--text",
type=str,
@ -112,7 +122,9 @@ if __name__ == "__main__":
help="Optional path to a DiT checkpoint (default: auto-download a pre-trained DiT-XL/2 model).",
)
parser.add_argument("--vqvae", default="hpcai-tech/vqvae")
parser.add_argument("--text_model", type=str, default="openai/clip-vit-base-patch32")
parser.add_argument(
"--text_model", type=str, default="openai/clip-vit-base-patch32"
)
parser.add_argument("--width", type=int, default=480)
parser.add_argument("--height", type=int, default=320)
parser.add_argument("--fps", type=int, default=15)