Views
No views yet

pip install torch transformers accelerate[!NOTE] We recommend settingenable_thinking=Falsewhen using the model to ensure stable behavior and reproducible results.
1import torch
2import numpy as np
3import torch.nn.functional as F
4
5from transformers import AutoTokenizer, AutoModelForMaskedLM
6
7
8def add_gumbel_noise(logits, temperature):
9 if temperature == 0:
10 return logits
11 logits = logits.to(torch.float64)
12 noise = torch.rand_like(logits, dtype=torch.float64)
13 gumbel_noise = (- torch.log(noise)) ** temperature
14 return logits.exp() / gumbel_noise
15
16
17def get_num_transfer_tokens(mask_index, steps):
18 mask_num = mask_index.sum(dim=1, keepdim=True)
19 base = mask_num // steps
20 remainder = mask_num % steps
21 num_transfer_tokens = torch.zeros(mask_num.size(0), steps, device=mask_index.device, dtype=torch.int64) + base
22 for i in range(mask_num.size(0)):
23 num_transfer_tokens[i, :remainder[i]] += 1
24 return num_transfer_tokens
25
26
27@torch.no_grad()
28def generate(model, prompt, prompt_lens, pad_id, steps=128, max_new_tokens=128, block_size=64, temperature=0.0, cfg_scale=0.0, remasking="random"):
29 mask_id = tokenizer.mask_token_id
30 batch_size = prompt.size(0)
31 total_length = int(prompt_lens.max().item() + max_new_tokens)
32 x = torch.full((batch_size, total_length), pad_id, dtype=torch.long, device=model.device)
33 for i, length in enumerate(prompt_lens.tolist()):
34 x[i, :length] = prompt[i, :length]
35 x[i, length : length + max_new_tokens] = mask_id
36
37 prompt_index = torch.arange(total_length, device=x.device).unsqueeze(0) < prompt_lens.unsqueeze(1)
38 positions = torch.arange(total_length, device=x.device)
39
40 assert max_new_tokens % block_size == 0
41 num_blocks = max_new_tokens // block_size
42 assert steps % num_blocks == 0
43 steps_per_block = steps // num_blocks
44
45 for num_block in range(num_blocks):
46 block_start = prompt_lens + num_block * block_size
47 block_end = block_start + block_size
48 init_block_mask = (
49 (positions.unsqueeze(0) >= block_start.unsqueeze(1))
50 & (positions.unsqueeze(0) < block_end.unsqueeze(1))
51 & (x == mask_id)
52 )
53 num_transfer_tokens = get_num_transfer_tokens(init_block_mask, steps_per_block)
54
55 for i in range(steps_per_block):
56 block_mask = (
57 (positions.unsqueeze(0) >= block_start.unsqueeze(1))
58 & (positions.unsqueeze(0) < block_end.unsqueeze(1))
59 & (x == mask_id)
60 )
61
62 if cfg_scale > 0.0:
63 un_x = x.clone()
64 un_x[prompt_index] = mask_id
65 x_ = torch.cat([x, un_x], dim=0)
66 logits = model(x_).logits
67 logits, un_logits = torch.chunk(logits, 2, dim=0)
68 logits = un_logits + (cfg_scale + 1.0) * (logits - un_logits)
69 else:
70 logits = model(x).logits
71
72 logits_with_noise = add_gumbel_noise(logits, temperature=temperature)
73 x0 = torch.argmax(logits_with_noise, dim=-1)
74
75 if remasking == "low_confidence":
76 p = F.softmax(logits, dim=-1)
77 x0_p = torch.gather(p, dim=-1, index=x0.unsqueeze(-1)).squeeze(-1)
78 elif remasking == "random":
79 x0_p = torch.rand_like(x0, dtype=torch.float)
80 else:
81 raise NotImplementedError(remasking)
82
83 confidence = torch.full_like(x0_p, -np.inf)
84 confidence = torch.where(block_mask, x0_p, confidence)
85
86 x0 = torch.where(block_mask, x0, x)
87
88 transfer_index = torch.zeros_like(x0, dtype=torch.bool, device=x0.device)
89 for j in range(confidence.shape[0]):
90 k = int(num_transfer_tokens[j, i].item())
91 if k == 0:
92 continue
93 _, select_index = torch.topk(confidence[j], k=k)
94 transfer_index[j, select_index] = True
95 x[transfer_index] = x0[transfer_index]
96
97 return x
98
99device = "cuda" if torch.cuda.is_available() else "cpu"
100model = AutoModelForMaskedLM.from_pretrained("dllm-hub/Qwen3-0.6B-diffusion-mdlm-v0.1", dtype=torch.bfloat16, trust_remote_code=True).to(device).eval()
101tokenizer = AutoTokenizer.from_pretrained("dllm-hub/Qwen3-0.6B-diffusion-mdlm-v0.1")
102if tokenizer.pad_token_id is None and tokenizer.eos_token is not None:
103 tokenizer.pad_token = tokenizer.eos_token
104pad_id = tokenizer.pad_token_id or tokenizer.eos_token_id or tokenizer.mask_token_id
105
106messages = [
107 [
108 {"role": "system", "content": "You are a helpful AI assistant."},
109 {"role": "user", "content": "Implement a DFS traversal in Python with clear inline comments."},
110 ],
111 [
112 {"role": "system", "content": "You are a helpful AI assistant."},
113 {"role": "user", "content": "Lily can run 12 kilometers per hour for 4 hours. After that, she runs 10 kilometers per hour. How many kilometers can she run in 10 hours?"},
114 ],
115]
116
117encoded = [tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=True, enable_thinking=False) for m in messages]
118prompt_lens = torch.tensor([len(e) for e in encoded], dtype=torch.long)
119max_prompt_len = max(prompt_lens).item()
120prompt_tensor = torch.full((len(encoded), max_prompt_len), pad_id, dtype=torch.long)
121for i, ids in enumerate(encoded):
122 prompt_tensor[i, : len(ids)] = torch.tensor(ids, dtype=torch.long)
123
124prompt_tensor = prompt_tensor.to(device)
125prompt_lens = prompt_lens.to(device)
126max_new_tokens = 256
127
128text = generate(
129 model, prompt_tensor, prompt_lens, pad_id=pad_id, steps=256, max_new_tokens=max_new_tokens, block_size=64, temperature=0.0, cfg_scale=0.0, remasking="low_confidence"
130)
131
132new_tokens = [
133 text[i, prompt_lens[i] : prompt_lens[i] + max_new_tokens].tolist() for i in range(text.size(0))
134]
135for idx, decoded in enumerate(tokenizer.batch_decode(new_tokens, skip_special_tokens=False)):
136 print(f"
137[Sample {idx}]")
138 print(decoded)| Parameter | Description | Default |
|---|---|---|
max_new_tokens | Number of tokens to generate | 256 |
steps | Number of diffusion denoising iterations | 256 |
temperature | Sampling temperature; set to 0.0 for deterministic generation | 0.0 |
block_size | Token block size used during iterative denoising | 64 |
cfg_scale | Classifier-free guidance scale controlling instruction adherence (higher = more deterministic) | 0.0 |
remasking | Strategy for re-masking during each denoising step (random or low_confidence) | low_confidence |
1python -u examples/a2d/mdlm/chat.py \
2 --model_name_or_path dllm-hub/Qwen3-0.6B-diffusion-mdlm-v0.1 \
3 --chat_template True --block_size 64 --remasking low_confidence --steps 256 --max_new_tokens 256| Model | GSM8K | MATH | BBH | MMLU‑Pro | Hellaswag | MMLU | HumanEval | MBPP |
|---|---|---|---|---|---|---|---|---|
Qwen3-0.6B-diffusion-bd3lm-v0.1 (evaluated) | 46.6 | 13.9 | 27.0 | 14.1 | 40.0 | 38.8 | 47.6 | 32.0 |
Qwen3-0.6B-diffusion-mdlm-v0.1 (evaluated) | 29.8 | 8.8 | 27.0 | 17.6 | 42.1 | 40.0 | 30.5 | 29.2 |
Qwen3-0.6B-Base (reported) | 59.6 | 32.4 | 41.5 | 24.7 | 47.4 | 52.8 | 32.3 | 36.6 |
Qwen2.5-0.5B (reported) | 41.6 | 19.5 | 20.3 | 15.7 | 52.1 | 47.5 | 30.5 | 39.3 |
1bash examples/a2d/mdlm/eval.sh \
2 --model_name_or_path dllm-hub/Qwen3-0.6B-diffusion-mdlm-v0.11@misc{zhou2026dllm,
2 title={dLLM: Simple Diffusion Language Modeling},
3 author={Zhanhui Zhou and Lingjie Chen and Hanghang Tong and Dawn Song},
4 year={2026},
5 eprint={2602.22661},
6 archivePrefix={arXiv},
7 primaryClass={cs.CL},
8 url={https://arxiv.org/abs/2602.22661},
9}