Instructions to use yresearch/MDLM-MMD-OWT with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use yresearch/MDLM-MMD-OWT with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="yresearch/MDLM-MMD-OWT", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("yresearch/MDLM-MMD-OWT", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
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
Model tree for yresearch/MDLM-MMD-OWT
Base model
kuleshov-group/mdlm-owt