MCPcopy Create free account
hub / github.com/huggingface/diffusers / UpResnetBlock1D

Class UpResnetBlock1D

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

Source from the content-addressed store, hash-verified

84
85
86class UpResnetBlock1D(nn.Module):
87 def __init__(
88 self,
89 in_channels: int,
90 out_channels: int | None = None,
91 num_layers: int = 1,
92 temb_channels: int = 32,
93 groups: int = 32,
94 groups_out: int | None = None,
95 non_linearity: str | None = None,
96 time_embedding_norm: str = "default",
97 output_scale_factor: float = 1.0,
98 add_upsample: bool = True,
99 ):
100 super().__init__()
101 self.in_channels = in_channels
102 out_channels = in_channels if out_channels is None else out_channels
103 self.out_channels = out_channels
104 self.time_embedding_norm = time_embedding_norm
105 self.add_upsample = add_upsample
106 self.output_scale_factor = output_scale_factor
107
108 if groups_out is None:
109 groups_out = groups
110
111 # there will always be at least one resnet
112 resnets = [ResidualTemporalBlock1D(2 * in_channels, out_channels, embed_dim=temb_channels)]
113
114 for _ in range(num_layers):
115 resnets.append(ResidualTemporalBlock1D(out_channels, out_channels, embed_dim=temb_channels))
116
117 self.resnets = nn.ModuleList(resnets)
118
119 if non_linearity is None:
120 self.nonlinearity = None
121 else:
122 self.nonlinearity = get_activation(non_linearity)
123
124 self.upsample = None
125 if add_upsample:
126 self.upsample = Upsample1D(out_channels, use_conv_transpose=True)
127
128 def forward(
129 self,
130 hidden_states: torch.Tensor,
131 res_hidden_states_tuple: tuple[torch.Tensor, ...] | None = None,
132 temb: torch.Tensor | None = None,
133 ) -> torch.Tensor:
134 if res_hidden_states_tuple is not None:
135 res_hidden_states = res_hidden_states_tuple[-1]
136 hidden_states = torch.cat((hidden_states, res_hidden_states), dim=1)
137
138 hidden_states = self.resnets[0](hidden_states, temb)
139 for resnet in self.resnets[1:]:
140 hidden_states = resnet(hidden_states, temb)
141
142 if self.nonlinearity is not None:
143 hidden_states = self.nonlinearity(hidden_states)

Callers 1

get_up_blockFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…