(
self,
input_tensor: torch.FloatTensor,
temb: torch.FloatTensor,
scale: float = 1.0,
)
| 327 | ) |
| 328 | |
| 329 | def forward( |
| 330 | self, |
| 331 | input_tensor: torch.FloatTensor, |
| 332 | temb: torch.FloatTensor, |
| 333 | scale: float = 1.0, |
| 334 | ) -> torch.FloatTensor: |
| 335 | hidden_states = input_tensor |
| 336 | |
| 337 | hidden_states = self.norm1(hidden_states) |
| 338 | hidden_states = self.nonlinearity(hidden_states) |
| 339 | |
| 340 | if self.upsample is not None: |
| 341 | # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 |
| 342 | if hidden_states.shape[0] >= 64: |
| 343 | input_tensor = input_tensor.contiguous() |
| 344 | hidden_states = hidden_states.contiguous() |
| 345 | input_tensor = ( |
| 346 | self.upsample(input_tensor, scale=scale) |
| 347 | if isinstance(self.upsample, Upsample2D) |
| 348 | else self.upsample(input_tensor) |
| 349 | ) |
| 350 | hidden_states = ( |
| 351 | self.upsample(hidden_states, scale=scale) |
| 352 | if isinstance(self.upsample, Upsample2D) |
| 353 | else self.upsample(hidden_states) |
| 354 | ) |
| 355 | elif self.downsample is not None: |
| 356 | input_tensor = ( |
| 357 | self.downsample(input_tensor, scale=scale) |
| 358 | if isinstance(self.downsample, Downsample2D) |
| 359 | else self.downsample(input_tensor) |
| 360 | ) |
| 361 | hidden_states = ( |
| 362 | self.downsample(hidden_states, scale=scale) |
| 363 | if isinstance(self.downsample, Downsample2D) |
| 364 | else self.downsample(hidden_states) |
| 365 | ) |
| 366 | |
| 367 | hidden_states = self.conv1(hidden_states, scale) if not USE_PEFT_BACKEND else self.conv1(hidden_states) |
| 368 | |
| 369 | if self.time_emb_proj is not None: |
| 370 | if not self.skip_time_act: |
| 371 | temb = self.nonlinearity(temb) |
| 372 | temb = ( |
| 373 | self.time_emb_proj(temb, scale)[:, :, None, None] |
| 374 | if not USE_PEFT_BACKEND |
| 375 | else self.time_emb_proj(temb)[:, :, None, None] |
| 376 | ) |
| 377 | |
| 378 | if self.time_embedding_norm == "default": |
| 379 | if temb is not None: |
| 380 | hidden_states = hidden_states + temb |
| 381 | hidden_states = self.norm2(hidden_states) |
| 382 | elif self.time_embedding_norm == "scale_shift": |
| 383 | if temb is None: |
| 384 | raise ValueError( |
| 385 | f" `temb` should not be None when `time_embedding_norm` is {self.time_embedding_norm}" |
| 386 | ) |
nothing calls this directly
no test coverage detected