T

Diffusers
Safetensors
T
File size: 11,366 Bytes
964f845
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
"""
EditMapper: Maps LLaVA IMG token hidden states (4096) to SD text conditioning space (768).

Two implementations:
1. LightweightEditMapper: Simple MLP (recommended to start)
2. MGIEStyleEditMapper: Transformer-based with learnable queries (from MGIE paper)
"""

import torch
import torch.nn as nn
from typing import Optional


def match_dtype_to_model(mapper: nn.Module, model: nn.Module) -> nn.Module:
    """
    Helper function to match EditMapper dtype to LLaVA model dtype.
    
    Args:
        mapper: EditMapper instance
        model: LLaVA model
    
    Returns:
        mapper: EditMapper with matching dtype
    """
    model_dtype = next(model.parameters()).dtype
    mapper_dtype = next(mapper.parameters()).dtype
    
    if model_dtype != mapper_dtype:
        if model_dtype == torch.float16:
            mapper = mapper.half()
        elif model_dtype == torch.bfloat16:
            mapper = mapper.bfloat16()
        elif model_dtype == torch.float32:
            mapper = mapper.float()
        print(f"Converted EditMapper from {mapper_dtype} to {model_dtype}")
    
    return mapper


class LightweightEditMapper(nn.Module):
    """
    Lightweight EditMapper with simple MLP projection.
    
    Maps IMG token hidden states from LLaVA (4096) to SD text conditioning space (768).
    This is a stable, efficient baseline recommended for initial training.
    
    Architecture:
        LayerNorm -> Linear(4096->mid_dim) -> GELU -> Linear(mid_dim->768) -> LayerNorm
        Optional: L2 normalization + learnable scale (for CLIP stat matching)
    
    Args:
        in_dim: Input dimension (LLaVA hidden size, default: 4096)
        mid_dim: Hidden dimension (default: 1024)
        out_dim: Output dimension (SD CLIP text hidden, default: 768)
        k_tokens: Number of IMG tokens (default: 16)
        use_clip_norm: Apply L2 normalization + learnable scale to match CLIP (default: False)
    """
    
    def __init__(self, in_dim: int = 4096, mid_dim: int = 1024, 
                 out_dim: int = 768, k_tokens: int = 16,
                 use_clip_norm: bool = False):
        super().__init__()
        
        self.in_dim = in_dim
        self.mid_dim = mid_dim
        self.out_dim = out_dim
        self.k_tokens = k_tokens
        self.use_clip_norm = use_clip_norm
        
        # Simple MLP projection
        self.proj = nn.Sequential(
            nn.LayerNorm(in_dim),
            nn.Linear(in_dim, mid_dim),
            nn.GELU(),
            nn.Linear(mid_dim, out_dim),
        )
        
        # Output normalization
        self.out_norm = nn.LayerNorm(out_dim)
        
        # Optional: CLIP-style L2 normalization + learnable scale
        if use_clip_norm:
            self.scale = nn.Parameter(torch.ones(1) * 20.0)  # CLIP uses ~20-30
        
        # Initialize weights
        self._init_weights()
    
    def _init_weights(self):
        """
        Initialize weights with small values for stable training start.
        
        Using gain=0.02 intentionally to start with small outputs, preventing
        large initial perturbations to the frozen UNet cross-attention.
        Standard gain=1.0 can cause instability early in training.
        """
        for module in self.proj.modules():
            if isinstance(module, nn.Linear):
                # Small gain for gentle initialization
                nn.init.xavier_uniform_(module.weight, gain=0.02)
                if module.bias is not None:
                    nn.init.zeros_(module.bias)
    
    def forward(self, h_img: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
        """
        Args:
            h_img: (B, K, in_dim) - IMG token hidden states from LLaVA
            mask: (B, K) - Optional binary mask (1=valid, 0=missing/padding)
                  Used to zero out contributions from missing IMG tokens
        
        Returns:
            z_vlm: (B, K, out_dim) - Projected features for SD cross-attention
        
        Note:
            Input dtype should match mapper dtype. If using mixed precision:
            - Convert mapper to float16: mapper.half()
            - Or convert input to float32: h_img.float()
        """
        # Project: (B, K, 4096) -> (B, K, 768)
        z = self.proj(h_img)
        z = self.out_norm(z)
        
        # Optional: CLIP-style normalization
        if self.use_clip_norm:
            # L2 normalize then scale
            z = nn.functional.normalize(z, dim=-1)  # Unit vectors
            z = z * self.scale  # Learnable scaling (CLIP uses ~20-30)
        
        # Apply mask if provided (zero out missing tokens)
        if mask is not None:
            # Expand mask: (B, K) -> (B, K, 1) for broadcasting
            mask = mask.unsqueeze(-1).to(z.dtype)
            z = z * mask
        
        return z
    
    def extra_repr(self) -> str:
        return f'in_dim={self.in_dim}, mid_dim={self.mid_dim}, out_dim={self.out_dim}, k_tokens={self.k_tokens}'


class MGIEStyleEditMapper(nn.Module):
    """
    MGIE-style EditMapper with Transformer and learnable queries.
    
    This is more expressive but also more complex. Use if the lightweight mapper
    doesn't provide enough capacity.
    
    Architecture (from MGIE paper):
        1. Project LLaVA hidden: 4096 -> 512
        2. Transformer encoder-decoder with learnable queries
        3. Project to SD space: 512 -> 768
    
    Args:
        in_dim: Input dimension (LLaVA hidden size, default: 4096)
        hid_dim: Hidden dimension for transformer (default: 512)
        out_dim: Output dimension (SD CLIP text hidden, default: 768)
        num_queries: Number of output queries (default: 77, SD max sequence length)
        num_encoder_layers: Transformer encoder layers (default: 4)
        num_decoder_layers: Transformer decoder layers (default: 4)
        nhead: Number of attention heads (default: 4)
        dim_feedforward: FFN dimension (default: 2048)
        dropout: Dropout rate (default: 0.0)
        use_positional_encoding: Add learned positional encodings (default: True)
    """
    
    def __init__(self, 
                 in_dim: int = 4096,
                 hid_dim: int = 512,
                 out_dim: int = 768,
                 num_queries: int = 77,
                 num_encoder_layers: int = 4,
                 num_decoder_layers: int = 4,
                 nhead: int = 4,
                 dim_feedforward: int = 2048,
                 dropout: float = 0.0,
                 use_positional_encoding: bool = True):
        super().__init__()
        
        self.in_dim = in_dim
        self.hid_dim = hid_dim
        self.out_dim = out_dim
        self.num_queries = num_queries
        self.use_positional_encoding = use_positional_encoding
        
        # Project LLaVA hidden to transformer dimension
        self.llm2hid = nn.Linear(in_dim, hid_dim)
        
        # Learnable queries (output sequence)
        self.query = nn.Parameter(torch.randn(1, num_queries, hid_dim))
        
        # Optional: Positional encodings for better structure
        if use_positional_encoding:
            # Positional encoding for queries (decoder input)
            self.query_pos = nn.Parameter(torch.randn(1, num_queries, hid_dim))
            # Positional encoding for encoder input (source IMG tokens)
            self.src_pos = nn.Parameter(torch.randn(1, 16, hid_dim))  # Assuming max 16 IMG tokens
        
        # Transformer encoder-decoder
        # Note: src=encoder input (IMG hiddens), tgt=decoder input (queries)
        self.mapper = nn.Transformer(
            d_model=hid_dim,
            nhead=nhead,
            num_encoder_layers=num_encoder_layers,
            num_decoder_layers=num_decoder_layers,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            batch_first=True,
            norm_first=True
        )
        
        # Project to SD text conditioning space
        self.hid2feat = nn.Linear(hid_dim, out_dim)
        
        # Initialize
        self._init_weights()
    
    def _init_weights(self):
        """Initialize weights with small values for stable start"""
        # Initialize query and positional encodings with small values
        nn.init.normal_(self.query, mean=0.0, std=0.02)
        
        if self.use_positional_encoding:
            nn.init.normal_(self.query_pos, mean=0.0, std=0.02)
            nn.init.normal_(self.src_pos, mean=0.0, std=0.02)
        
        # Initialize linear layers (small gain for gentle start)
        nn.init.xavier_uniform_(self.llm2hid.weight, gain=0.02)
        nn.init.zeros_(self.llm2hid.bias)
        nn.init.xavier_uniform_(self.hid2feat.weight, gain=0.02)
        nn.init.zeros_(self.hid2feat.bias)
    
    def forward(self, llm: torch.Tensor, emb: Optional[torch.Tensor] = None) -> torch.Tensor:
        """
        Args:
            llm: (B, K, in_dim) - IMG token hidden states from LLaVA
            emb: (B, K, in_dim) - Optional, IMG token embeddings (MGIE adds these)
        
        Returns:
            feat: (B, num_queries, out_dim) - Features for SD cross-attention
        
        Note:
            Transformer forward signature: transformer(src, tgt)
            - src (encoder input): IMG hiddens + positional encoding
            - tgt (decoder input): learnable queries + positional encoding
        """
        batch_size = llm.shape[0]
        k_tokens = llm.shape[1]
        
        # Add embeddings if provided (MGIE does this)
        if emb is not None:
            hid = self.llm2hid(llm + emb)  # (B, K, hid_dim)
        else:
            hid = self.llm2hid(llm)  # (B, K, hid_dim)
        
        # Add positional encoding to source (encoder input)
        if self.use_positional_encoding:
            # Trim or pad src_pos to match actual k_tokens
            src_pos = self.src_pos[:, :k_tokens, :]  # (1, K, hid_dim)
            hid = hid + src_pos  # (B, K, hid_dim)
        
        # Prepare queries (decoder input)
        queries = self.query.repeat(batch_size, 1, 1)  # (B, num_queries, hid_dim)
        
        # Add positional encoding to queries
        if self.use_positional_encoding:
            queries = queries + self.query_pos  # (B, num_queries, hid_dim)
        
        # Transformer: src=hid (encoder), tgt=queries (decoder)
        # Returns (B, num_queries, hid_dim)
        hid = self.mapper(hid, queries)
        
        # Project to SD space
        feat = self.hid2feat(hid)  # (B, num_queries, out_dim)
        
        return feat
    
    def extra_repr(self) -> str:
        return (f'in_dim={self.in_dim}, hid_dim={self.hid_dim}, out_dim={self.out_dim}, '
                f'num_queries={self.num_queries}')


# Convenience factory
def create_edit_mapper(style: str = "lightweight", **kwargs) -> nn.Module:
    """
    Factory function to create EditMapper.
    
    Args:
        style: "lightweight" or "mgie"
        **kwargs: Arguments passed to the mapper constructor
    
    Returns:
        EditMapper instance
    """
    if style == "lightweight":
        return LightweightEditMapper(**kwargs)
    elif style == "mgie":
        return MGIEStyleEditMapper(**kwargs)
    else:
        raise ValueError(f"Unknown style: {style}. Choose 'lightweight' or 'mgie'")


# For backward compatibility and convenience
EditMapper = LightweightEditMapper