Download dinac3/decoder.py from data-archetype/dinac3_96: direct link, hf CLI and curl.
- Browser
- Download file 1.59 kB
-
https://huggingface.co/data-archetype/dinac3_96/resolve/main/dinac3/decoder.py
- Command line
-
hf download hf://data-archetype/dinac3_96/dinac3/decoder.py
-
curl -L -o decoder.py https://huggingface.co/data-archetype/dinac3_96/resolve/main/dinac3/decoder.py
1.59 kB
| """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) | |