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