MCPcopy Create free account
hub / github.com/ToTheBeginning/PuLID / AutoEncoder

Class AutoEncoder

flux/modules/autoencoder.py:277–312  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

275
276
277class AutoEncoder(nn.Module):
278 def __init__(self, params: AutoEncoderParams):
279 super().__init__()
280 self.encoder = Encoder(
281 resolution=params.resolution,
282 in_channels=params.in_channels,
283 ch=params.ch,
284 ch_mult=params.ch_mult,
285 num_res_blocks=params.num_res_blocks,
286 z_channels=params.z_channels,
287 )
288 self.decoder = Decoder(
289 resolution=params.resolution,
290 in_channels=params.in_channels,
291 ch=params.ch,
292 out_ch=params.out_ch,
293 ch_mult=params.ch_mult,
294 num_res_blocks=params.num_res_blocks,
295 z_channels=params.z_channels,
296 )
297 self.reg = DiagonalGaussian()
298
299 self.scale_factor = params.scale_factor
300 self.shift_factor = params.shift_factor
301
302 def encode(self, x: Tensor) -> Tensor:
303 z = self.reg(self.encoder(x))
304 z = self.scale_factor * (z - self.shift_factor)
305 return z
306
307 def decode(self, z: Tensor) -> Tensor:
308 z = z / self.scale_factor + self.shift_factor
309 return self.decoder(z)
310
311 def forward(self, x: Tensor) -> Tensor:
312 return self.decode(self.encode(x))

Callers 1

load_aeFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected