| """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 |
|
|
| |
| |
| |
| _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 |
| |
| |
| |
| 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: |
| |
| |
| 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} |
| ) |
| |
| |
| |
| |
| |
| 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) |
|
|