MDLM-MMD on OpenWebText

MDLM-MMD is an OpenWebText masked diffusion language model post-trained with representation-space MMD rewards. This repository provides a Transformers export of the EMA weights at checkpoint global step 3,250, together with the original model.ckpt for the training project's evaluation loader.

The base model is kuleshov-group/mdlm-owt. Training and diffusion sampling code is available in IDLM_orig.

Property Value
Parameters 169,627,218
Context length 1,024 tokens
Transformer blocks / attention heads 12 / 12
Hidden dimension 768
Weight precision FP32
Tokenizer GPT-2
Model vocabulary size 50,258
Absorbing mask token ID 50,257
BOS / EOS token ID 50,256
Time conditioning Disabled

Load with Transformers

The custom model implementation is the unchanged implementation from the base MDLM repository. It requires PyTorch, Transformers, Einops, and a compatible CUDA/FlashAttention installation for forward passes.

import torch
from transformers import AutoModelForMaskedLM, AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("gpt2")
model = AutoModelForMaskedLM.from_pretrained(
    "yresearch/MDLM_MMD_OWT",
    trust_remote_code=True,
    torch_dtype=torch.float32,
).cuda().eval()

# Example: one denoiser call on a fully masked sequence.
input_ids = torch.full(
    (1, 1024), model.config.mask_token_id,
    dtype=torch.long, device=model.device,
)
timesteps = torch.zeros(input_ids.shape[0], device=model.device)
with torch.inference_mode():
    logits = model(
        input_ids=input_ids,
        timesteps=timesteps,
        return_dict=True,
    ).logits

Use the GPT-2 tokenizer from gpt2; tokenizer files are not included here. The extra model vocabulary entry at ID 50,257 is the absorbing diffusion mask. The custom forward requires explicit timesteps even though time conditioning is disabled, and accepts token IDs rather than a tokenizer dictionary containing attention_mask.

The forward pass returns raw denoiser logits. Text generation requires the project's diffusion sampler, including suppression of the mask-token output and preservation of observed tokens. The standard autoregressive generate() API is not the sampling interface for this model.

Checkpoint and conversion

The source checkpoint is already an EMA-only evaluation checkpoint; EMA is not applied a second time during export. It contains no optimizer state for resuming training. Saved post-training settings include token-level RBF MMD, feature layer 3, RBF alpha 2e-5, group size 8, draw-kernel batch size 2, same-position exclusion, learning rate 1e-4, global batch size 512, and EMA decay 0.9999.

Export removes exactly one outer backbone. prefix from each checkpoint state key. All 131 state tensors, including the rotary-frequency buffer, are preserved exactly. No quantization or dtype conversion is performed.

Validation compared every tensor against the source checkpoint after Safetensors serialization and after a local AutoModelForMaskedLM.from_pretrained() reload. Both comparisons were exact, with no missing or unexpected keys. Loading was verified with PyTorch 2.2.1 and Transformers 4.38.2. A GPU forward pass and a new quality evaluation were not run as part of this export.

See export_metadata.json for source hashes, the pinned base revision, and conversion checks. The original checkpoint remains available as model.ckpt.

Attribution and license

The base weights and copied Hugging Face modeling/configuration code come from MDLM, by Subham Sekhar Sahoo and collaborators, under Apache 2.0. The Python model files are copied unchanged from base model revision d0958fa851335ece6c15260ce0025f030673c0fb; this repository contains post-trained weights and an updated configuration/model card. See LICENSE and NOTICE.

The training project builds on IDLM, whose training code is distributed under the MIT license. The underlying corpus is OpenWebText.

Downloads last month
6
Safetensors
Model size
0.2B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for yresearch/MDLM-MMD-OWT

Finetuned
(5)
this model

Dataset used to train yresearch/MDLM-MMD-OWT