goldenfox's picture
Marimo Diffusion 0.6B: checkpoint, sampler, OpenAI server, ledger-needle bench
685e018 verified
Raw
History Blame Contribute Delete
9.84 kB
"""GPT-sized bidirectional Transformer used as a masked-token denoiser."""
from __future__ import annotations
import argparse
import math
from pathlib import Path
import torch
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint
from diffusion_lm.config import ModelConfig, load_config
class DiffusionTransformer(nn.Module):
"""A GPT-like Transformer with the causal mask deliberately removed.
The network predicts clean tokens from an input containing absorbing mask
tokens. Passing ``output_positions`` avoids materializing vocabulary logits
for already-visible tokens during training.
"""
def __init__(self, config: ModelConfig) -> None:
super().__init__()
self.config = config
self.tokenizer_sha256: str | None = None
self.token_embedding = nn.Embedding(config.vocab_size, config.d_model)
self.position_embedding = nn.Embedding(config.max_seq_len, config.d_model)
self.embedding_dropout = nn.Dropout(config.dropout)
if config.use_flex_attention:
from diffusion_lm.flexattn import FlexEncoder
self.transformer = FlexEncoder(
d_model=config.d_model,
n_heads=config.n_heads,
d_ff=config.d_ff,
dropout=config.dropout,
n_layers=config.n_layers,
activation_checkpointing=config.activation_checkpointing,
)
else:
layer = nn.TransformerEncoderLayer(
d_model=config.d_model,
nhead=config.n_heads,
dim_feedforward=config.d_ff,
dropout=config.dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.transformer = nn.TransformerEncoder(
layer,
num_layers=config.n_layers,
norm=nn.LayerNorm(config.d_model),
enable_nested_tensor=False,
)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
self.apply(self._init_weights)
self._init_residual_outputs()
if config.tie_embeddings:
self.lm_head.weight = self.token_embedding.weight
self.register_buffer(
"_forbidden_output_token_ids",
torch.tensor(config.forbidden_output_token_ids, dtype=torch.long),
persistent=False,
)
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def _init_residual_outputs(self) -> None:
"""Scale residual branch outputs as in GPT-2 for stable deep training."""
if self.config.use_flex_attention:
self.transformer.init_residual_outputs(self.config.n_layers)
return
residual_std = 0.02 / math.sqrt(2 * self.config.n_layers)
for layer in self.transformer.layers:
nn.init.normal_(layer.self_attn.out_proj.weight, mean=0.0, std=residual_std)
nn.init.normal_(layer.linear2.weight, mean=0.0, std=residual_std)
def _checkpointed_transformer(
self,
hidden: Tensor,
padding_mask: Tensor | None,
attn_mask: Tensor | None = None,
) -> Tensor:
for layer in self.transformer.layers:
def run_layer(layer_input: Tensor, *, current_layer: nn.Module = layer) -> Tensor:
return current_layer(
layer_input, src_mask=attn_mask, src_key_padding_mask=padding_mask
)
hidden = checkpoint(run_layer, hidden, use_reentrant=False)
if self.transformer.norm is not None:
hidden = self.transformer.norm(hidden)
return hidden
def _expand_attn_mask(self, attn_mask: Tensor | None, input_ids: Tensor) -> Tensor | None:
"""Broadcast a per-sample boolean blocking mask across attention heads.
Accepts ``[L, L]`` shared masks or ``[B, L, L]`` per-sample masks with
``True`` marking blocked key positions, matching the src_mask convention.
"""
if attn_mask is None:
return None
batch_size, sequence_length = input_ids.shape
if attn_mask.dtype != torch.bool:
raise ValueError("attn_mask must be boolean with True marking blocked positions")
if attn_mask.shape == (sequence_length, sequence_length):
return attn_mask
if attn_mask.shape != (batch_size, sequence_length, sequence_length):
raise ValueError("attn_mask must have shape [L, L] or [batch, L, L]")
return attn_mask.repeat_interleave(self.config.n_heads, dim=0)
def encode(
self,
input_ids: Tensor,
attention_mask: Tensor | None = None,
attn_mask: Tensor | None = None,
) -> Tensor:
"""Return contextual token states; ``attn_mask`` restricts attention topology."""
if input_ids.ndim != 2:
raise ValueError("input_ids must have shape [batch, sequence]")
batch_size, sequence_length = input_ids.shape
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}"
)
if attention_mask is not None and attention_mask.shape != input_ids.shape:
raise ValueError("attention_mask must match input_ids")
positions = torch.arange(sequence_length, device=input_ids.device)
hidden = self.token_embedding(input_ids) + self.position_embedding(positions)[None, :, :]
hidden = self.embedding_dropout(hidden)
# TransformerEncoder expects True for padding, the inverse of the common
# attention-mask convention. src_mask is only supplied by region-aware callers.
padding_mask = None if attention_mask is None else ~attention_mask.bool()
if self.config.use_flex_attention:
from diffusion_lm.flexattn import build_block_mask
if attn_mask is not None and attn_mask.dtype != torch.bool:
raise ValueError("attn_mask must be boolean with True marking blocked positions")
block_mask = build_block_mask(
attn_mask, padding_mask, batch_size, sequence_length, hidden.device
)
return self.transformer(hidden, block_mask)
expanded_attn_mask = self._expand_attn_mask(attn_mask, input_ids)
if (
self.config.activation_checkpointing
and self.training
and torch.is_grad_enabled()
):
return self._checkpointed_transformer(hidden, padding_mask, expanded_attn_mask)
return self.transformer(
hidden, mask=expanded_attn_mask, src_key_padding_mask=padding_mask
)
def forward(
self,
input_ids: Tensor,
attention_mask: Tensor | None = None,
output_positions: Tensor | None = None,
attn_mask: Tensor | None = None,
) -> Tensor:
"""Predict vocabulary logits for all tokens or selected positions only."""
hidden = self.encode(input_ids, attention_mask=attention_mask, attn_mask=attn_mask)
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.lm_head(hidden)
# Corruption/control tokens are never valid clean-token predictions. EOS
# deliberately remains available so generation can terminate naturally.
if self._forbidden_output_token_ids.numel():
logits.index_fill_(
-1,
self._forbidden_output_token_ids,
torch.finfo(logits.dtype).min,
)
return logits
@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)
def build_denoiser(
config: ModelConfig,
*,
load_pretrained: bool = True,
dtype: torch.dtype | None = None,
) -> nn.Module:
"""Construct the denoiser a config describes: project transformer or pretrained backbone.
``load_pretrained=False`` builds the architecture only, for callers that immediately
restore weights from a project checkpoint.
"""
if config.backbone == "hf-qwen3":
from diffusion_lm.hf_bridge import Qwen3Denoiser
return Qwen3Denoiser(config, load_pretrained=load_pretrained, dtype=dtype)
return DiffusionTransformer(config)
def format_parameter_count(count: int) -> str:
if count >= 1_000_000:
return f"{count / 1_000_000:.2f}M"
if count >= 1_000:
return f"{count / 1_000:.2f}K"
return str(count)
def main() -> None:
parser = argparse.ArgumentParser(description="Report the exact model parameter count")
parser.add_argument("--config", type=Path, required=True, help="experiment YAML")
args = parser.parse_args()
config = load_config(args.config)
# Parameter inspection should not allocate four gigabytes for the 1B preset.
with torch.device("meta"):
model = DiffusionTransformer(config.model)
print(f"parameters: {model.num_parameters:,} ({format_parameter_count(model.num_parameters)})")
if __name__ == "__main__":
main()