MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / forward

Method forward

diffsynth/models/sd_vae_decoder.py:92–117  ·  view source on GitHub ↗
(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs)

Source from the content-addressed store, hash-verified

90 return hidden_states
91
92 def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
93 original_dtype = sample.dtype
94 sample = sample.to(dtype=next(iter(self.parameters())).dtype)
95 # For VAE Decoder, we do not need to apply the tiler on each layer.
96 if tiled:
97 return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
98
99 # 1. pre-process
100 sample = sample / self.scaling_factor
101 hidden_states = self.post_quant_conv(sample)
102 hidden_states = self.conv_in(hidden_states)
103 time_emb = None
104 text_emb = None
105 res_stack = None
106
107 # 2. blocks
108 for i, block in enumerate(self.blocks):
109 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
110
111 # 3. output
112 hidden_states = self.conv_norm_out(hidden_states)
113 hidden_states = self.conv_act(hidden_states)
114 hidden_states = self.conv_out(hidden_states)
115 hidden_states = hidden_states.to(original_dtype)
116
117 return hidden_states
118
119 @staticmethod
120 def state_dict_converter():

Callers 1

tiled_forwardMethod · 0.95

Calls 2

tiled_forwardMethod · 0.95
toMethod · 0.45

Tested by

no test coverage detected