MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / encode

Method encode

wan/models/wan_vae.py:520–550  ·  view source on GitHub ↗
(self, x, scale=None)

Source from the content-addressed store, hash-verified

518 return x_recon, mu, log_var
519
520 def encode(self, x, scale=None):
521 self.clear_cache()
522 ## cache
523 t = x.shape[2]
524 iter_ = 1 + (t - 1) // 4
525 if scale != None:
526 scale = [item.to(x.device, x.dtype) for item in scale]
527 ## 对encode输入的x,按时间拆分为1、4、4、4....
528 for i in range(iter_):
529 self._enc_conv_idx = [0]
530 if i == 0:
531 out = self.encoder(
532 x[:, :, :1, :, :],
533 feat_cache=self._enc_feat_map,
534 feat_idx=self._enc_conv_idx)
535 else:
536 out_ = self.encoder(
537 x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
538 feat_cache=self._enc_feat_map,
539 feat_idx=self._enc_conv_idx)
540 out = torch.cat([out, out_], 2)
541 mu, log_var = self.conv1(out).chunk(2, dim=1)
542 if scale != None:
543 if isinstance(scale[0], torch.Tensor):
544 mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
545 1, self.z_dim, 1, 1, 1)
546 else:
547 mu = (mu - scale[0]) * scale[1]
548 x = torch.cat([mu, log_var], dim = 1)
549 self.clear_cache()
550 return x
551
552 def decode(self, z, scale=None):
553 self.clear_cache()

Callers 2

forwardMethod · 0.95
sampleMethod · 0.95

Calls 1

clear_cacheMethod · 0.95

Tested by

no test coverage detected