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