(self, x, xs, t_emb=None)
| 223 | self.upsample = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) |
| 224 | |
| 225 | def forward(self, x, xs, t_emb=None): |
| 226 | |
| 227 | for attn, net in zip(self.attns, self.nets): |
| 228 | res_x = xs[-1] |
| 229 | xs = xs[:-1] |
| 230 | x = torch.cat([x, res_x], dim=1) |
| 231 | x = net(x, t_emb) |
| 232 | if attn: |
| 233 | x = attn(x) |
| 234 | |
| 235 | if self.upsample: |
| 236 | x = F.interpolate(x, scale_factor=2.0, mode='nearest') |
| 237 | x = self.upsample(x) |
| 238 | |
| 239 | return x |
| 240 | |
| 241 | class UNet_timeimage_cond(nn.Module): |
| 242 | def __init__( |
nothing calls this directly
no outgoing calls
no test coverage detected