MCPcopy Create free account
hub / github.com/TencentARC/BrushNet / forward

Method forward

src/diffusers/models/resnet.py:329–405  ·  view source on GitHub ↗
(
        self,
        input_tensor: torch.FloatTensor,
        temb: torch.FloatTensor,
        scale: float = 1.0,
    )

Source from the content-addressed store, hash-verified

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 )

Callers

nothing calls this directly

Calls 1

downsampleMethod · 0.80

Tested by

no test coverage detected