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