File size: 11,873 Bytes
685e018
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""HuggingFace Qwen3 backbone served through the project's denoiser forward contract."""

from __future__ import annotations

import os

import torch
from torch import Tensor, nn

from diffusion_lm.config import ModelConfig

# The flex template's default 128x128 tiles need ~112 KiB of shared memory at head_dim 128,
# over Ada's 100 KiB per-block ceiling; halved tiles fit. The backward kernel budgets its
# tiles separately, hence the M1/N1/M2/N2 entries. Larger tiles are valid on Hopper.
_FLEX_KERNEL_OPTIONS = {
    'BLOCK_M': 64,
    'BLOCK_N': 64,
    'BLOCK_M1': 32,
    'BLOCK_N1': 64,
    'BLOCK_M2': 64,
    'BLOCK_N2': 32,
}


class Qwen3Denoiser(nn.Module):
    """Wrap ``Qwen3ForCausalLM`` behind the DiffusionTransformer forward contract.

    The backbone always receives a 4D attention mask so its stock causal masking never
    engages: block-diffusion objectives need bidirectional attention inside denoising
    windows, and an omitted mask must mean "attend everything", not "causal".
    ``use_flex_attention`` selects how the boolean blocking matrix reaches the backbone — a
    ``BlockMask`` for the flex kernel, which skips fully-masked blocks, or a materialized
    additive mask for SDPA. Hidden states are gathered at ``output_positions`` before the LM
    head so full-vocabulary logits are never materialized for visible tokens.
    """

    def __init__(
        self,
        config: ModelConfig,
        *,
        load_pretrained: bool = True,
        dtype: torch.dtype | None = None,
    ) -> None:
        super().__init__()
        try:
            from transformers import Qwen3Config, Qwen3ForCausalLM
        except ImportError as exc:
            raise RuntimeError(
                'the hf-qwen3 backbone requires transformers; install the [hf] extra'
            ) from exc
        if config.pretrained_path is None:
            raise ValueError('the hf-qwen3 backbone requires pretrained_path')
        self.config = config
        # Generation on Ada can hit the Triton shared-memory ceiling inside transformers' own
        # compiled flex kernel, where our kernel_options do not reach. MDLM_SDPA=1 sidesteps it
        # for serving and evaluation; flex earns its keep in training, not at q_len 1.
        use_flex = config.use_flex_attention and os.environ.get('MDLM_SDPA') != '1'
        attn_implementation = 'flex_attention' if use_flex else 'sdpa'
        if load_pretrained:
            kwargs = {} if dtype is None else {'dtype': dtype}
            self.backbone = Qwen3ForCausalLM.from_pretrained(
                config.pretrained_path, attn_implementation=attn_implementation, **kwargs
            )
        else:
            # Architecture-only construction: weights come from a later load_state_dict,
            # so checkpoint restore never re-reads the base model files.
            backbone_config = Qwen3Config.from_pretrained(
                config.pretrained_path, attn_implementation=attn_implementation
            )
            self.backbone = Qwen3ForCausalLM(backbone_config)
            if dtype is not None:
                self.backbone.to(dtype)
        self._use_flex = use_flex
        self._validate_backbone()
        self.backbone.config.use_cache = False
        if config.activation_checkpointing:
            self.backbone.gradient_checkpointing_enable(
                gradient_checkpointing_kwargs={'use_reentrant': False}
            )
        # A plain attribute keeps the compiled callable out of the module tree, so checkpoint
        # keys are identical whether or not compilation is on. Shapes reaching the backbone are
        # static; the varying count of output positions is gathered after it returns.
        # Serving generates at ever-changing lengths, which makes a compiled backbone
        # recompile per shape; MDLM_NO_COMPILE=1 turns it off without touching the checkpoint.
        compiling = config.compile_backbone and os.environ.get('MDLM_NO_COMPILE') != '1'
        self._compiled_backbone = (
            torch.compile(self.backbone.model.forward, dynamic=False) if compiling else None
        )
        self.register_buffer(
            '_forbidden_output_token_ids',
            torch.tensor(config.forbidden_output_token_ids, dtype=torch.long),
            persistent=False,
        )

    def _validate_backbone(self) -> None:
        backbone = self.backbone.config
        pairs = (
            ('vocab_size', self.config.vocab_size, backbone.vocab_size),
            ('d_model', self.config.d_model, backbone.hidden_size),
            ('n_layers', self.config.n_layers, backbone.num_hidden_layers),
            ('n_heads', self.config.n_heads, backbone.num_attention_heads),
            ('d_ff', self.config.d_ff, backbone.intermediate_size),
        )
        for name, expected, actual in pairs:
            if expected != actual:
                raise ValueError(f'config {name}={expected} but the backbone has {actual}')
        if self.config.max_seq_len > backbone.max_position_embeddings:
            raise ValueError(
                f'max_seq_len {self.config.max_seq_len} exceeds the backbone context '
                f'{backbone.max_position_embeddings}'
            )

    def _blocked_matrix(
        self,
        input_ids: Tensor,
        attention_mask: Tensor | None,
        attn_mask: Tensor | None,
    ) -> Tensor:
        """Boolean ``[batch, L, L]`` blocking matrix, ``True`` marking a key to suppress."""

        batch_size, seq_len = input_ids.shape
        device = input_ids.device
        if attn_mask is None:
            blocked = torch.zeros(
                batch_size, seq_len, seq_len, dtype=torch.bool, device=device
            )
        else:
            if attn_mask.dtype != torch.bool:
                raise ValueError('attn_mask must be boolean with True marking blocked positions')
            if attn_mask.shape == (seq_len, seq_len):
                blocked = attn_mask.unsqueeze(0).expand(batch_size, -1, -1)
            elif attn_mask.shape == (batch_size, seq_len, seq_len):
                blocked = attn_mask
            else:
                raise ValueError('attn_mask must have shape [L, L] or [batch, L, L]')
        if attention_mask is not None:
            if attention_mask.shape != input_ids.shape:
                raise ValueError('attention_mask must match input_ids')
            blocked = blocked | ~attention_mask.bool()[:, None, :]
        return blocked

    def _additive_mask(self, blocked: Tensor) -> Tensor:
        mask_dtype = self.backbone.model.embed_tokens.weight.dtype
        additive = torch.zeros(blocked.shape, dtype=mask_dtype, device=blocked.device)
        additive = additive.masked_fill(blocked, torch.finfo(mask_dtype).min)
        return additive.unsqueeze(1)

    def forward_cached(
        self,
        input_ids: Tensor,
        *,
        attn_mask: Tensor,
        past_key_values,
        output_positions: Tensor | None = None,
    ) -> tuple[Tensor, object]:
        """Run only the trailing ``input_ids`` against a populated key/value cache.

        ``attn_mask`` holds one row per NEW query and one column per position the query may
        see, cache included: ``[batch, query, cached + query]``. Block-causal masking is what
        makes caching sound here — prefix positions never attend forward into a block, so
        their keys and values stay valid while the block's own tokens keep changing.
        """

        cached = past_key_values.get_seq_length()
        query_len = input_ids.shape[1]
        blocked = attn_mask if attn_mask.dim() == 3 else attn_mask.unsqueeze(0)
        if blocked.shape[-2] != query_len or blocked.shape[-1] != cached + query_len:
            raise ValueError(
                f'attn_mask must be [batch, {query_len}, {cached + query_len}] for a cache of '
                f'{cached} positions, got {tuple(blocked.shape)}'
            )
        rows = blocked.expand(input_ids.shape[0], -1, -1)
        outputs = self.backbone.model(
            input_ids=input_ids,
            attention_mask=self._additive_mask(rows),
            past_key_values=past_key_values,
            use_cache=True,
            **self._backbone_kwargs(),
        )
        hidden = outputs.last_hidden_state
        if output_positions is not None:
            hidden = hidden[output_positions.bool()]
        logits = self.backbone.lm_head(hidden)
        self._forbid(logits)
        return logits, outputs.past_key_values

    def _backbone_kwargs(self) -> dict:
        """Extra backbone arguments the attention implementation needs.

        Every path into the backbone has to carry these, cached or not: without the reduced
        tiles the flex template asks Ada for more shared memory than it has and the kernel
        fails to compile at all.
        """

        if not self._use_flex:
            return {}
        return {'kernel_options': _FLEX_KERNEL_OPTIONS}

    def new_cache(self):
        """Empty key/value cache for :meth:`forward_cached`, kept here so callers stay
        independent of the transformers cache class."""

        from transformers import DynamicCache

        return DynamicCache()

    def _forbid(self, logits: Tensor) -> None:
        if self._forbidden_output_token_ids.numel():
            logits.index_fill_(
                -1, self._forbidden_output_token_ids, torch.finfo(logits.dtype).min
            )
        if self.config.forbidden_output_from is not None:
            logits[..., self.config.forbidden_output_from :] = torch.finfo(logits.dtype).min

    def _flex_mask(self, blocked: Tensor):
        """Compress the blocking matrix into a ``BlockMask`` the flex kernel can skip over.

        ``masking_utils`` forwards any mask reporting a 4D shape untouched, so a ``BlockMask``
        reaches ``flex_attention_forward`` as ``block_mask`` and the backbone never rebuilds a
        causal mask of its own.
        """

        from diffusion_lm.flexattn import build_block_mask

        batch_size, seq_len, _ = blocked.shape
        return build_block_mask(blocked, None, batch_size, seq_len, blocked.device)

    def forward(
        self,
        input_ids: Tensor,
        attention_mask: Tensor | None = None,
        output_positions: Tensor | None = None,
        attn_mask: Tensor | None = None,
    ) -> Tensor:
        if input_ids.ndim != 2:
            raise ValueError('input_ids must have shape [batch, sequence]')
        sequence_length = input_ids.shape[1]
        if sequence_length > self.config.max_seq_len:
            raise ValueError(
                f'sequence length {sequence_length} exceeds max_seq_len '
                f'{self.config.max_seq_len}'
            )
        blocked = self._blocked_matrix(input_ids, attention_mask, attn_mask)
        mask = (
            self._flex_mask(blocked) if self._use_flex else self._additive_mask(blocked)
        )
        call = self._compiled_backbone or self.backbone.model
        hidden = call(
            input_ids=input_ids, attention_mask=mask, use_cache=False,
            **self._backbone_kwargs(),
        ).last_hidden_state
        if output_positions is not None:
            if output_positions.shape != input_ids.shape:
                raise ValueError('output_positions must match input_ids')
            hidden = hidden[output_positions.bool()]
        logits = self.backbone.lm_head(hidden)
        self._forbid(logits)
        return logits

    @property
    def token_embedding(self) -> nn.Embedding:
        """Embedding module surfaced for the optimizer's 32-bit override (tied to the head)."""

        return self.backbone.model.embed_tokens

    @property
    def num_parameters(self) -> int:
        """Count unique trainable parameters (shared embeddings count once)."""

        return sum(parameter.numel() for parameter in self.parameters() if parameter.requires_grad)