MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / __init__

Method __init__

architecture/autoencoder_kl_wan.py:520–584  ·  view source on GitHub ↗
(
        self,
        in_channels: int = 3,
        dim=128,
        z_dim=4,
        dim_mult=[1, 2, 4, 4],
        num_res_blocks=2,
        attn_scales=[],
        temperal_downsample=[True, True, False],
        dropout=0.0,
        non_linearity: str = "silu",
        is_residual: bool = False,  # wan 2.2 vae use a residual downblock
    )

Source from the content-addressed store, hash-verified

518 """
519
520 def __init__(
521 self,
522 in_channels: int = 3,
523 dim=128,
524 z_dim=4,
525 dim_mult=[1, 2, 4, 4],
526 num_res_blocks=2,
527 attn_scales=[],
528 temperal_downsample=[True, True, False],
529 dropout=0.0,
530 non_linearity: str = "silu",
531 is_residual: bool = False, # wan 2.2 vae use a residual downblock
532 ):
533 super().__init__()
534 self.dim = dim
535 self.z_dim = z_dim
536 self.dim_mult = dim_mult
537 self.num_res_blocks = num_res_blocks
538 self.attn_scales = attn_scales
539 self.temperal_downsample = temperal_downsample
540 self.nonlinearity = get_activation(non_linearity)
541
542 # dimensions
543 dims = [dim * u for u in [1] + dim_mult]
544 scale = 1.0
545
546 # init block
547 self.conv_in = WanCausalConv3d(in_channels, dims[0], 3, padding=1)
548
549 # downsample blocks
550 self.down_blocks = nn.ModuleList([])
551 for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
552 # residual (+attention) blocks
553 if is_residual:
554 self.down_blocks.append(
555 WanResidualDownBlock(
556 in_dim,
557 out_dim,
558 dropout,
559 num_res_blocks,
560 temperal_downsample=temperal_downsample[i] if i != len(dim_mult) - 1 else False,
561 down_flag=i != len(dim_mult) - 1,
562 )
563 )
564 else:
565 for _ in range(num_res_blocks):
566 self.down_blocks.append(WanResidualBlock(in_dim, out_dim, dropout))
567 if scale in attn_scales:
568 self.down_blocks.append(WanAttentionBlock(out_dim))
569 in_dim = out_dim
570
571 # downsample block
572 if i != len(dim_mult) - 1:
573 mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
574 self.down_blocks.append(WanResample(out_dim, mode=mode))
575 scale /= 2.0
576
577 # middle blocks

Callers

nothing calls this directly

Calls 8

WanCausalConv3dClass · 0.85
WanResidualBlockClass · 0.85
WanAttentionBlockClass · 0.85
WanResampleClass · 0.85
WanMidBlockClass · 0.85
WanRMS_normClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected