File size: 3,862 Bytes
d8ed0eb
 
 
 
4e02ce9
d8ed0eb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4e02ce9
d8ed0eb
 
 
 
4e02ce9
d8ed0eb
 
 
 
 
 
 
 
 
 
4e02ce9
d8ed0eb
 
 
4e02ce9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8ed0eb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""
DINOv2 Encoder — frozen, extracts semantic features from LR images for DiT Cross-Attention.

Uses HF transformers AutoModel for DINOv2 (avoids torch.hub dependency on GitHub).
Uses output_hidden_states=True to get intermediate layer features.
"""

import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision.transforms import Normalize

IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD  = [0.229, 0.224, 0.225]

DINOV2_HF_NAMES = {
    "dinov2b": "facebook/dinov2-base",
    "dinov2l": "facebook/dinov2-large",
    "dinov2g": "facebook/dinov2-giant",
}


class Dinov2Encoder(nn.Module):
    """Frozen DINOv2 feature extractor, outputting features from specified intermediate layers."""

    def __init__(
        self,
        enc_type: str = "dinov2b",
        dinov2_size: int = 448,
        layer_indices: list[int] | None = None,
        device: str = "cuda",
    ):
        super().__init__()

        self.dinov2_size = dinov2_size
        self.layer_indices = layer_indices or [8]

        hf_name = DINOV2_HF_NAMES.get(enc_type)
        if hf_name is None:
            raise ValueError(
                f"Unknown DINOv2 type: {enc_type}, "
                f"expected one of {list(DINOV2_HF_NAMES)}"
            )

        print(f"Loading DINOv2 from HF transformers: {hf_name} ...")
        from transformers import AutoModel
        encoder = AutoModel.from_pretrained(hf_name)

        self.encoder = encoder.to(device).eval()
        for p in self.encoder.parameters():
            p.requires_grad_(False)

        print(f"DINOv2 encoder loaded, layers={self.layer_indices}")

    def preprocess(self, lr: torch.Tensor) -> torch.Tensor:
        """
        lr: [B, 3, H, W] float [0, 1]
        -> resize -> clamp -> ImageNet normalization
        """
        x = F.interpolate(lr, size=self.dinov2_size, mode="bicubic", align_corners=False)
        x = x.clamp(0, 1)
        x = Normalize(IMAGENET_MEAN, IMAGENET_STD)(x)
        return x

    @torch.no_grad()
    def forward(self, lr: torch.Tensor) -> list[torch.Tensor]:
        """
        lr: [B, 3, H, W] float [0, 1]
        -> list of [B, N_patches, enc_dim]
        """
        x = self.preprocess(lr)

        # Use output_hidden_states=True to get all layer outputs
        outputs = self.encoder(x, output_hidden_states=True)
        # hidden_states is a tuple of (num_layers + 1) tensors:
        #   hidden_states[0] = embeddings output
        #   hidden_states[1..N] = output of each encoder layer
        hidden_states = outputs.hidden_states  # tuple of (B, seq_len, hidden_size)

        # Build feature list: layer_i -> hidden_states[i+1] (skip embeddings output)
        # The original code extracts features from intermediate layers (without CLS token)
        z = []
        for idx in self.layer_indices:
            # hidden_states[idx+1] because hidden_states[0] is the embedding output
            # and hidden_states[1..N] are layer outputs
            feat = hidden_states[idx + 1][:, 1:]  # Remove CLS token
            z.append(feat)

        # Replace last with the final norm output (last_hidden_state)
        # The original code does: z[-1] = x_norm (the final normalized output without CLS)
        z[-1] = outputs.last_hidden_state[:, 1:]

        return z


def create_dinov2_encoder(config_path: str, device: str = "cuda") -> Dinov2Encoder | None:
    """Create DINOv2 encoder from YAML config, returns None if not configured."""
    import yaml
    with open(config_path, "r", encoding="utf-8") as f:
        cfg = yaml.safe_load(f)

    dv2 = cfg.get("dinov2", {}) or {}
    if not dv2:
        return None

    return Dinov2Encoder(
        enc_type=dv2.get("enc_type", "dinov2b"),
        dinov2_size=dv2.get("dinov2_size", 448),
        layer_indices=dv2.get("layer_dinov2b_list", [8]),
        device=device,
    )