MCPcopy Create free account
hub / github.com/ByteDance-Seed/Bagel / AutoEncoder

Class AutoEncoder

modeling/autoencoder.py:290–325  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

288
289
290class AutoEncoder(nn.Module):
291 def __init__(self, params: AutoEncoderParams):
292 super().__init__()
293 self.encoder = Encoder(
294 resolution=params.resolution,
295 in_channels=params.in_channels,
296 ch=params.ch,
297 ch_mult=params.ch_mult,
298 num_res_blocks=params.num_res_blocks,
299 z_channels=params.z_channels,
300 )
301 self.decoder = Decoder(
302 resolution=params.resolution,
303 in_channels=params.in_channels,
304 ch=params.ch,
305 out_ch=params.out_ch,
306 ch_mult=params.ch_mult,
307 num_res_blocks=params.num_res_blocks,
308 z_channels=params.z_channels,
309 )
310 self.reg = DiagonalGaussian()
311
312 self.scale_factor = params.scale_factor
313 self.shift_factor = params.shift_factor
314
315 def encode(self, x: Tensor) -> Tensor:
316 z = self.reg(self.encoder(x))
317 z = self.scale_factor * (z - self.shift_factor)
318 return z
319
320 def decode(self, z: Tensor) -> Tensor:
321 z = z / self.scale_factor + self.shift_factor
322 return self.decoder(z)
323
324 def forward(self, x: Tensor) -> Tensor:
325 return self.decode(self.encode(x))
326
327
328def print_load_warning(missing: list[str], unexpected: list[str]) -> None:

Callers 1

load_aeFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected