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

Method forward

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

Source from the content-addressed store, hash-verified

50 return hidden_states
51
52 def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
53 original_dtype = sample.dtype
54 sample = sample.to(dtype=next(iter(self.parameters())).dtype)
55 # For VAE Decoder, we do not need to apply the tiler on each layer.
56 if tiled:
57 return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
58
59 # 1. pre-process
60 hidden_states = self.conv_in(sample)
61 time_emb = None
62 text_emb = None
63 res_stack = None
64
65 # 2. blocks
66 for i, block in enumerate(self.blocks):
67 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
68
69 # 3. output
70 hidden_states = self.conv_norm_out(hidden_states)
71 hidden_states = self.conv_act(hidden_states)
72 hidden_states = self.conv_out(hidden_states)
73 hidden_states = self.quant_conv(hidden_states)
74 hidden_states = hidden_states[:, :4]
75 hidden_states *= self.scaling_factor
76 hidden_states = hidden_states.to(original_dtype)
77
78 return hidden_states
79
80 def encode_video(self, sample, batch_size=8):
81 B = sample.shape[0]

Callers 1

tiled_forwardMethod · 0.95

Calls 2

tiled_forwardMethod · 0.95
toMethod · 0.45

Tested by

no test coverage detected