(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_mask: torch.Tensor,
timestep: Optional[torch.Tensor],
ssl_hidden_states: Optional[List[torch.Tensor]] = None,
output_length: int = 0,
block_controlnet_hidden_states: Optional[
Union[List[torch.Tensor], torch.Tensor]
] = None,
controlnet_scale: Union[float, torch.Tensor] = 1.0,
return_dict: bool = True,
)
| 412 | return encoder_hidden_states, encoder_hidden_mask |
| 413 | |
| 414 | def decode( |
| 415 | self, |
| 416 | hidden_states: torch.Tensor, |
| 417 | attention_mask: torch.Tensor, |
| 418 | encoder_hidden_states: torch.Tensor, |
| 419 | encoder_hidden_mask: torch.Tensor, |
| 420 | timestep: Optional[torch.Tensor], |
| 421 | ssl_hidden_states: Optional[List[torch.Tensor]] = None, |
| 422 | output_length: int = 0, |
| 423 | block_controlnet_hidden_states: Optional[ |
| 424 | Union[List[torch.Tensor], torch.Tensor] |
| 425 | ] = None, |
| 426 | controlnet_scale: Union[float, torch.Tensor] = 1.0, |
| 427 | return_dict: bool = True, |
| 428 | ): |
| 429 | |
| 430 | embedded_timestep = self.timestep_embedder( |
| 431 | self.time_proj(timestep).to(dtype=hidden_states.dtype) |
| 432 | ) |
| 433 | temb = self.t_block(embedded_timestep) |
| 434 | |
| 435 | hidden_states = self.proj_in(hidden_states) |
| 436 | |
| 437 | # controlnet logic |
| 438 | if block_controlnet_hidden_states is not None: |
| 439 | control_condi = cross_norm(hidden_states, block_controlnet_hidden_states) |
| 440 | hidden_states = hidden_states + control_condi * controlnet_scale |
| 441 | |
| 442 | inner_hidden_states = [] |
| 443 | |
| 444 | rotary_freqs_cis = self.rotary_emb( |
| 445 | hidden_states, seq_len=hidden_states.shape[1] |
| 446 | ) |
| 447 | encoder_rotary_freqs_cis = self.rotary_emb( |
| 448 | encoder_hidden_states, seq_len=encoder_hidden_states.shape[1] |
| 449 | ) |
| 450 | |
| 451 | for index_block, block in enumerate(self.transformer_blocks): |
| 452 | |
| 453 | if self.training and self.gradient_checkpointing: |
| 454 | |
| 455 | hidden_states = torch.utils.checkpoint.checkpoint( |
| 456 | block, |
| 457 | hidden_states=hidden_states, |
| 458 | attention_mask=attention_mask, |
| 459 | encoder_hidden_states=encoder_hidden_states, |
| 460 | encoder_attention_mask=encoder_hidden_mask, |
| 461 | rotary_freqs_cis=rotary_freqs_cis, |
| 462 | rotary_freqs_cis_cross=encoder_rotary_freqs_cis, |
| 463 | temb=temb, |
| 464 | use_reentrant=False, |
| 465 | ) |
| 466 | |
| 467 | else: |
| 468 | hidden_states = block( |
| 469 | hidden_states=hidden_states, |
| 470 | attention_mask=attention_mask, |
| 471 | encoder_hidden_states=encoder_hidden_states, |
no test coverage detected