MCPcopy Create free account
hub / github.com/MeiGen-AI/MultiTalk / decode

Method decode

wan/modules/vae.py:544–568  ·  view source on GitHub ↗
(self, z, scale)

Source from the content-addressed store, hash-verified

542 return mu
543
544 def decode(self, z, scale):
545 self.clear_cache()
546 # z: [b,c,t,h,w]
547 if isinstance(scale[0], torch.Tensor):
548 z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
549 1, self.z_dim, 1, 1, 1)
550 else:
551 z = z / scale[1] + scale[0]
552 iter_ = z.shape[2]
553 x = self.conv2(z)
554 for i in range(iter_):
555 self._conv_idx = [0]
556 if i == 0:
557 out = self.decoder(
558 x[:, :, i:i + 1, :, :],
559 feat_cache=self._feat_map,
560 feat_idx=self._conv_idx)
561 else:
562 out_ = self.decoder(
563 x[:, :, i:i + 1, :, :],
564 feat_cache=self._feat_map,
565 feat_idx=self._conv_idx)
566 out = torch.cat([out, out_], 2)
567 self.clear_cache()
568 return out
569
570 def reparameterize(self, mu, log_var):
571 std = torch.exp(0.5 * log_var)

Callers 1

forwardMethod · 0.95

Calls 1

clear_cacheMethod · 0.95

Tested by

no test coverage detected