File size: 1,590 Bytes
75ff4df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
"""Deterministic one-pass decoder: latent projection, ViT trunk and conv up-path."""

from __future__ import annotations

from typing import TYPE_CHECKING

from torch import Tensor, nn

from .conv_up_head import ConvUpHead
from .trunk import DitBlock, rope_tables

if TYPE_CHECKING:
    from .config import Dinac3Config


class Decoder(nn.Module):
    """Raw latents ``[B, C, h, w]`` to RGB ``[B, 3, 16 h, 16 w]``."""

    def __init__(self, config: Dinac3Config) -> None:
        """Allocate the latent projection, trunk blocks and conv up-path head."""

        super().__init__()
        width = config.decoder_width
        self.head_dim = config.decoder_head_dim
        self.latent_up = nn.Conv2d(config.latent_channels, width, 1)
        self.trunk = nn.ModuleList(
            [
                DitBlock(width, config.decoder_head_dim, config.decoder_mlp_ratio)
                for _ in range(config.decoder_depth)
            ]
        )
        self.conv_up_head = ConvUpHead(config)

    def forward(self, latents: Tensor) -> Tensor:
        """Decode in one pass; latents are validated by the public API."""

        x = self.latent_up(latents)
        b, c, h, w = x.shape
        sin, cos = rope_tables(h, w, head_dim=self.head_dim, device=x.device)
        tokens = x.permute(0, 2, 3, 1).reshape(b, h * w, c)
        for block in self.trunk:
            tokens = block(tokens, sin, cos)
        # The head reads the trunk's tokens as a contiguous NCHW map.
        features = tokens.transpose(1, 2).reshape(b, c, h, w).contiguous()
        return self.conv_up_head(features)