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

Method forward

x2ct_nerf/modules/diffusionmodules/model.py:725–773  ·  view source on GitHub ↗
(self, z, temb=None)

Source from the content-addressed store, hash-verified

723 padding=1)
724
725 def forward(self, z, temb=None):
726 #assert z.shape[1:] == self.z_shape[1:]
727 outputs = {}
728 self.last_z_shape = z.shape
729
730 # timestep embedding
731 #temb = None
732
733 # z to block_in
734 h = self.conv_in(z)
735 if 'conv_in' in self.output_key:
736 outputs['conv_in'] = h.clone()
737
738 # middle
739 h = self.mid.block_1(h, temb)
740 # h = self.mid.attn_1(h)
741 h = self.mid.block_2(h, temb)
742
743 if 'mid' in self.output_key:
744 outputs['mid'] = h.clone()
745
746 # upsampling
747 for i_level in reversed(range(self.num_resolutions)):
748 for i_block in range(self.num_res_blocks+1):
749 h = self.up[i_level].block[i_block](h, temb)
750 if len(self.up[i_level].attn) > 0:
751 h = self.up[i_level].attn[i_block](h)
752 if i_level != 0:
753 h = self.up[i_level].upsample(h)
754
755 block_idx = self.num_resolutions - i_level - 1
756 if f'up_block{block_idx}' in self.output_key:
757 outputs[f'up_block{block_idx}'] = h.clone()
758
759 # end
760 if self.give_pre_end:
761 return h
762
763 h = self.norm_out(h)
764 if 'norm_out' in self.output_key:
765 outputs['norm_out'] = h.clone()
766
767 h = nonlinearity(h)
768 if 'nonlinear' in self.output_key:
769 outputs['nonlinear'] = h.clone()
770
771 h = self.conv_out(h)
772 # outputs['final'] = h
773 return h
774
775
776class UNetBlockDecoder(nn.Module):

Callers

nothing calls this directly

Calls 1

nonlinearityFunction · 0.70

Tested by

no test coverage detected