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

Method decode

wan/models/wan_vae.py:552–578  ·  view source on GitHub ↗
(self, z, scale=None)

Source from the content-addressed store, hash-verified

550 return x
551
552 def decode(self, z, scale=None):
553 self.clear_cache()
554 # z: [b,c,t,h,w]
555 if scale != None:
556 scale = [item.to(z.device, z.dtype) for item in scale]
557 if isinstance(scale[0], torch.Tensor):
558 z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
559 1, self.z_dim, 1, 1, 1)
560 else:
561 z = z / scale[1] + scale[0]
562 iter_ = z.shape[2]
563 x = self.conv2(z)
564 for i in range(iter_):
565 self._conv_idx = [0]
566 if i == 0:
567 out = self.decoder(
568 x[:, :, i:i + 1, :, :],
569 feat_cache=self._feat_map,
570 feat_idx=self._conv_idx)
571 else:
572 out_ = self.decoder(
573 x[:, :, i:i + 1, :, :],
574 feat_cache=self._feat_map,
575 feat_idx=self._conv_idx)
576 out = torch.cat([out, out_], 2)
577 self.clear_cache()
578 return out
579
580 def reparameterize(self, mu, log_var):
581 std = torch.exp(0.5 * log_var)

Callers 1

forwardMethod · 0.95

Calls 1

clear_cacheMethod · 0.95

Tested by

no test coverage detected