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)
|