(self, x, t=None, context=None)
| 379 | self.conv_out = torch.nn.Conv2d(block_in, out_ch, kernel_size=3, stride=1, padding=1) |
| 380 | |
| 381 | def forward(self, x, t=None, context=None): |
| 382 | # assert x.shape[2] == x.shape[3] == self.resolution |
| 383 | if context is not None: |
| 384 | # assume aligned context, cat along channel axis |
| 385 | x = torch.cat((x, context), dim=1) |
| 386 | if self.use_timestep: |
| 387 | # timestep embedding |
| 388 | assert t is not None |
| 389 | temb = get_timestep_embedding(t, self.ch) |
| 390 | temb = self.temb.dense[0](temb) |
| 391 | temb = nonlinearity(temb) |
| 392 | temb = self.temb.dense[1](temb) |
| 393 | else: |
| 394 | temb = None |
| 395 | |
| 396 | # downsampling |
| 397 | hs = [self.conv_in(x)] |
| 398 | for i_level in range(self.num_resolutions): |
| 399 | for i_block in range(self.num_res_blocks): |
| 400 | h = self.down[i_level].block[i_block](hs[-1], temb) |
| 401 | if len(self.down[i_level].attn) > 0: |
| 402 | h = self.down[i_level].attn[i_block](h) |
| 403 | hs.append(h) |
| 404 | if i_level != self.num_resolutions - 1: |
| 405 | hs.append(self.down[i_level].downsample(hs[-1])) |
| 406 | |
| 407 | # middle |
| 408 | h = hs[-1] |
| 409 | h = self.mid.block_1(h, temb) |
| 410 | h = self.mid.attn_1(h) |
| 411 | h = self.mid.block_2(h, temb) |
| 412 | |
| 413 | # upsampling |
| 414 | for i_level in reversed(range(self.num_resolutions)): |
| 415 | for i_block in range(self.num_res_blocks + 1): |
| 416 | h = self.up[i_level].block[i_block](torch.cat([h, hs.pop()], dim=1), temb) |
| 417 | if len(self.up[i_level].attn) > 0: |
| 418 | h = self.up[i_level].attn[i_block](h) |
| 419 | if i_level != 0: |
| 420 | h = self.up[i_level].upsample(h) |
| 421 | |
| 422 | # end |
| 423 | h = self.norm_out(h) |
| 424 | h = nonlinearity(h) |
| 425 | h = self.conv_out(h) |
| 426 | return h |
| 427 | |
| 428 | def get_last_layer(self): |
| 429 | return self.conv_out.weight |
no test coverage detected