| 530 | |
| 531 | |
| 532 | class InitialLayer(nn.Module): |
| 533 | def __init__(self, model, text_encoder, is_generic_llm): |
| 534 | super().__init__() |
| 535 | self.x_embedder = model.x_embedder |
| 536 | self.pos_embedder = model.pos_embedder |
| 537 | if model.extra_per_block_abs_pos_emb: |
| 538 | self.extra_pos_embedder = model.extra_pos_embedder |
| 539 | self.t_embedder = model.t_embedder |
| 540 | self.t_embedding_norm = model.t_embedding_norm |
| 541 | self.text_encoder = text_encoder |
| 542 | self.model = [model] |
| 543 | self.is_generic_llm = is_generic_llm |
| 544 | |
| 545 | @torch.autocast('cuda', dtype=AUTOCAST_DTYPE) |
| 546 | def forward(self, inputs): |
| 547 | x_B_C_T_H_W, timesteps_B_T, *prompt_embeds_or_batch_encoding = inputs |
| 548 | |
| 549 | if torch.is_floating_point(prompt_embeds_or_batch_encoding[0]): |
| 550 | crossattn_emb, attn_mask, t5_input_ids, t5_attn_mask = prompt_embeds_or_batch_encoding |
| 551 | else: |
| 552 | with torch.no_grad(): |
| 553 | input_ids, attn_mask, t5_input_ids, t5_attn_mask = prompt_embeds_or_batch_encoding |
| 554 | crossattn_emb = _compute_text_embeddings(self.text_encoder, input_ids, attn_mask, is_generic_llm=self.is_generic_llm) |
| 555 | |
| 556 | padding_mask = torch.zeros(x_B_C_T_H_W.shape[0], 1, x_B_C_T_H_W.shape[3], x_B_C_T_H_W.shape[4], dtype=x_B_C_T_H_W.dtype, device=x_B_C_T_H_W.device) |
| 557 | x_B_T_H_W_D, rope_emb_L_1_1_D, extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D = self.model[0].prepare_embedded_sequence( |
| 558 | x_B_C_T_H_W, |
| 559 | fps=None, |
| 560 | padding_mask=padding_mask, |
| 561 | ) |
| 562 | assert extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D is None |
| 563 | assert rope_emb_L_1_1_D is not None |
| 564 | |
| 565 | if timesteps_B_T.ndim == 1: |
| 566 | timesteps_B_T = timesteps_B_T.unsqueeze(1) |
| 567 | t_embedding_B_T_D, adaln_lora_B_T_3D = self.t_embedder(timesteps_B_T) |
| 568 | t_embedding_B_T_D = self.t_embedding_norm(t_embedding_B_T_D) |
| 569 | |
| 570 | outputs = make_contiguous(x_B_T_H_W_D, t_embedding_B_T_D, crossattn_emb, t5_input_ids, attn_mask, t5_attn_mask, rope_emb_L_1_1_D, adaln_lora_B_T_3D, timesteps_B_T) |
| 571 | for tensor in outputs: |
| 572 | if torch.is_floating_point(tensor): |
| 573 | tensor.requires_grad_(True) |
| 574 | return outputs |
| 575 | |
| 576 | |
| 577 | class LLMAdapterLayer(nn.Module): |