MCPcopy Create free account
hub / github.com/dek924/PerX2CT / forward

Method forward

taming/modules/diffusionmodules/model.py:652–698  ·  view source on GitHub ↗
(self, x, z)

Source from the content-addressed store, hash-verified

650
651
652 def forward(self, x, z):
653 #assert x.shape[2] == x.shape[3] == self.resolution
654
655 if self.use_timestep:
656 # timestep embedding
657 assert t is not None
658 temb = get_timestep_embedding(t, self.ch)
659 temb = self.temb.dense[0](temb)
660 temb = nonlinearity(temb)
661 temb = self.temb.dense[1](temb)
662 else:
663 temb = None
664
665 # downsampling
666 hs = [self.conv_in(x)]
667 for i_level in range(self.num_resolutions):
668 for i_block in range(self.num_res_blocks):
669 h = self.down[i_level].block[i_block](hs[-1], temb)
670 if len(self.down[i_level].attn) > 0:
671 h = self.down[i_level].attn[i_block](h)
672 hs.append(h)
673 if i_level != self.num_resolutions-1:
674 hs.append(self.down[i_level].downsample(hs[-1]))
675
676 # middle
677 h = hs[-1]
678 z = self.z_in(z)
679 h = torch.cat((h,z),dim=1)
680 h = self.mid.block_1(h, temb)
681 h = self.mid.attn_1(h)
682 h = self.mid.block_2(h, temb)
683
684 # upsampling
685 for i_level in reversed(range(self.num_resolutions)):
686 for i_block in range(self.num_res_blocks+1):
687 h = self.up[i_level].block[i_block](
688 torch.cat([h, hs.pop()], dim=1), temb)
689 if len(self.up[i_level].attn) > 0:
690 h = self.up[i_level].attn[i_block](h)
691 if i_level != 0:
692 h = self.up[i_level].upsample(h)
693
694 # end
695 h = self.norm_out(h)
696 h = nonlinearity(h)
697 h = self.conv_out(h)
698 return h
699
700
701class SimpleDecoder(nn.Module):

Callers

nothing calls this directly

Calls 2

get_timestep_embeddingFunction · 0.70
nonlinearityFunction · 0.70

Tested by

no test coverage detected