Download src/diffusion_lm/hf_bridge.py from goldenfox/marimo-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/hf_bridge.py
- Command line
-
hf download hf://goldenfox/marimo-diffusion/src/diffusion_lm/hf_bridge.py
-
curl -L -o hf_bridge.py https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/hf_bridge.py
11.9 kB
| """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 | |
| 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 | |
| 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) | |