(self, x, latent)
| 159 | self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, kernel_size=3, padding=1) |
| 160 | |
| 161 | def forward(self, x, latent): |
| 162 | sample_latent = self.latent_conv_in(latent) |
| 163 | sample = self.conv_in(x) |
| 164 | emb = None |
| 165 | |
| 166 | down_block_res_samples = (sample,) |
| 167 | for i, downsample_block in enumerate(self.down_blocks): |
| 168 | if i == 3: |
| 169 | sample = sample + sample_latent |
| 170 | |
| 171 | sample, res_samples = downsample_block(hidden_states=sample, temb=emb) |
| 172 | down_block_res_samples += res_samples |
| 173 | |
| 174 | sample = self.mid_block(sample, emb) |
| 175 | |
| 176 | for upsample_block in self.up_blocks: |
| 177 | res_samples = down_block_res_samples[-len(upsample_block.resnets) :] |
| 178 | down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] |
| 179 | sample = upsample_block(sample, res_samples, emb) |
| 180 | |
| 181 | sample = self.conv_norm_out(sample) |
| 182 | sample = self.conv_act(sample) |
| 183 | sample = self.conv_out(sample) |
| 184 | return sample |
| 185 | |
| 186 | |
| 187 | def checkerboard(shape): |
nothing calls this directly
no outgoing calls
no test coverage detected