MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / ResConvBlock

Class ResConvBlock

diffusers/src/diffusers/models/unets/unet_1d_blocks.py:375–406  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

373
374
375class ResConvBlock(nn.Module):
376 def __init__(self, in_channels: int, mid_channels: int, out_channels: int, is_last: bool = False):
377 super().__init__()
378 self.is_last = is_last
379 self.has_conv_skip = in_channels != out_channels
380
381 if self.has_conv_skip:
382 self.conv_skip = nn.Conv1d(in_channels, out_channels, 1, bias=False)
383
384 self.conv_1 = nn.Conv1d(in_channels, mid_channels, 5, padding=2)
385 self.group_norm_1 = nn.GroupNorm(1, mid_channels)
386 self.gelu_1 = nn.GELU()
387 self.conv_2 = nn.Conv1d(mid_channels, out_channels, 5, padding=2)
388
389 if not self.is_last:
390 self.group_norm_2 = nn.GroupNorm(1, out_channels)
391 self.gelu_2 = nn.GELU()
392
393 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
394 residual = self.conv_skip(hidden_states) if self.has_conv_skip else hidden_states
395
396 hidden_states = self.conv_1(hidden_states)
397 hidden_states = self.group_norm_1(hidden_states)
398 hidden_states = self.gelu_1(hidden_states)
399 hidden_states = self.conv_2(hidden_states)
400
401 if not self.is_last:
402 hidden_states = self.group_norm_2(hidden_states)
403 hidden_states = self.gelu_2(hidden_states)
404
405 output = hidden_states + residual
406 return output
407
408
409class UNetMidBlock1D(nn.Module):

Callers 7

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected