(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs)
| 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] |
no test coverage detected