| 168 | |
| 169 | |
| 170 | class InitialLayer(nn.Module): |
| 171 | def __init__(self, model): |
| 172 | super().__init__() |
| 173 | self.pos_embed = model.pos_embed |
| 174 | self.time_text_embed = model.time_text_embed |
| 175 | self.context_embedder = model.context_embedder |
| 176 | self.model = [model] |
| 177 | |
| 178 | def __getattr__(self, name): |
| 179 | return getattr(self.model[0], name) |
| 180 | |
| 181 | @torch.autocast('cuda', dtype=AUTOCAST_DTYPE) |
| 182 | def forward(self, inputs): |
| 183 | for item in inputs: |
| 184 | if torch.is_floating_point(item): |
| 185 | item.requires_grad_(True) |
| 186 | hidden_states, timestep, encoder_hidden_states, pooled_projections = inputs |
| 187 | |
| 188 | height, width = hidden_states.shape[-2:] |
| 189 | latent_size = torch.tensor([height, width]).to(hidden_states.device) |
| 190 | |
| 191 | hidden_states = self.pos_embed(hidden_states) # takes care of adding positional embeddings too. |
| 192 | temb = self.time_text_embed(timestep, pooled_projections) |
| 193 | encoder_hidden_states = self.context_embedder(encoder_hidden_states) |
| 194 | |
| 195 | return make_contiguous(hidden_states, temb, latent_size, encoder_hidden_states) |
| 196 | |
| 197 | |
| 198 | class TransformerLayer(nn.Module): |