"""Region-aware attention masks, objectives, and samplers for reasoning variants.
Three inference modes share one ``DiffusionTransformer``:
- ``ar``: prefix-LM. Bidirectional over the problem, causal generation after it.
- ``diffusion``: full-sequence denoising of the response region in parallel.
- ``hybrid``: thought slots denoised block-by-block, then the answer decoded
autoregressively under a bidirectional-prefix mask.
"""
from __future__ import annotations
import os
import time
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import Tensor
from diffusion_lm.diffusion import (
CorruptionBatch,
corrupt_tokens,
iterative_unmask,
_sample_categorical,
)
from diffusion_lm.model import DiffusionTransformer
def prefix_causal_blocked(
prefix_lens: Tensor, seq_len: int, *, causal_prefix: bool = False
) -> Tensor:
"""Blocked-attention mask: bidirectional before ``prefix_len``, causal after.
``allowed[b, i, j] = j < prefix_lens[b] or j <= i``; the returned tensor is the
inverse, matching the src_mask convention where ``True`` blocks attention.
``causal_prefix=True`` degenerates to the plain causal mask — the geometry
pretrained autoregressive backbones were trained under.
"""
positions = torch.arange(seq_len, device=prefix_lens.device)
causal = positions[None, :, None] >= positions[None, None, :]
if causal_prefix:
return (~causal).expand(prefix_lens.shape[0], seq_len, seq_len)
prefix = positions[None, None, :] < prefix_lens[:, None, None]
return ~(prefix | causal)
def window_blocked(window_ends: Tensor, seq_len: int) -> Tensor:
"""Blocked-attention mask hiding all keys at or beyond each sample's window end."""
positions = torch.arange(seq_len, device=window_ends.device)
allowed = positions[None, None, :] < window_ends[:, None, None]
return ~allowed.expand(window_ends.shape[0], seq_len, seq_len)
def slot_causal_blocked(
problem_len: Tensor, n_slots: Tensor, block: int, seq_len: int
) -> Tensor:
"""Block-causal mask over thought slots: each slot attends its prefix slots only.
Problem tokens (plus ````) attend the problem window; tokens of slot k
attend everything up to slot k's end; positions past the think region attend
the whole think window. Denoising slot k at inference with clean prefix slots
is the per-slot ``t -> 0`` limit of the training distribution.
"""
device = problem_len.device
positions = torch.arange(seq_len, device=device)[None, :]
prefix_end = (problem_len + 1)[:, None]
think_end = prefix_end + n_slots[:, None] * block
slot_index = torch.clamp((positions - prefix_end) // block, min=0)
slot_end = prefix_end + (slot_index + 1) * block
window = torch.where(positions < prefix_end, prefix_end, slot_end)
window = torch.where(positions >= think_end, think_end, window)
allowed = positions[:, None, :] < window[:, :, None]
return ~allowed
def adaptive_block_mask(
tokens: Tensor,
problem_len: Tensor,
answer_start: Tensor,
size_ids: Tensor,
end_think_id: int,
*,
causal_prefix: bool = False,
) -> Tensor:
"""Block-causal mask over variable-length thought blocks.
Block boundaries are read directly from the token stream: every ````
control token and the terminal ```` mark the start of the next region.
A thought-content token attends its clean prefix plus its own block
bidirectionally, never a following block; problem and ```` positions see
the problem window only — or, with ``causal_prefix=True``, only their causal
past, matching a pretrained autoregressive backbone. The returned tensor
follows the src_mask convention where ``True`` blocks a key.
"""
positions = torch.arange(tokens.shape[1], device=tokens.device)
prefix_end = (problem_len + 1)[:, None]
in_think = (positions[None, :] >= prefix_end) & (
positions[None, :] < answer_start[:, None]
)
is_size = (tokens.unsqueeze(-1) == size_ids).any(dim=-1)
boundary = (is_size | (tokens == end_think_id)) & in_think
return block_mask_from_boundaries(
boundary, problem_len + 1, answer_start, causal_prefix=causal_prefix
)
def block_mask_from_boundaries(
boundary: Tensor,
prefix_end: Tensor,
answer_start: Tensor,
*,
causal_prefix: bool = False,
) -> Tensor:
"""Blocked mask from per-position block-start marks (the adaptive-mask core).
``boundary[b, p]`` is True where a new block starts. A position attends every key
before the next boundary after it: its own block bidirectionally plus all preceding
context. Rows before ``prefix_end`` see the prefix window (or their causal past with
``causal_prefix=True``); rows at or past ``answer_start`` see up to ``answer_start``.
"""
device = boundary.device
batch_size, seq_len = boundary.shape
positions = torch.arange(seq_len, device=device)
prefix_end = prefix_end[:, None]
answer_start = answer_start[:, None]
boundary_index = torch.where(
boundary, positions[None, :].expand(batch_size, seq_len), seq_len
)
reverse_cummin = boundary_index.flip(1).cummin(dim=1).values.flip(1)
next_boundary = torch.full((batch_size, seq_len), seq_len, device=device)
next_boundary[:, :-1] = reverse_cummin[:, 1:]
if causal_prefix:
prefix_limit = (positions[None, :] + 1).expand(batch_size, seq_len)
else:
prefix_limit = prefix_end.expand(batch_size, seq_len)
attend_limit = next_boundary
attend_limit = torch.where(positions[None, :] < prefix_end, prefix_limit, attend_limit)
attend_limit = torch.where(
positions[None, :] >= answer_start,
answer_start.expand(batch_size, seq_len),
attend_limit,
)
allowed = positions[None, None, :] < attend_limit[:, :, None]
return ~allowed
@dataclass(frozen=True)
class HybridBatchLoss:
loss: Tensor
think_loss: float
answer_loss: float
think_accuracy: float
answer_accuracy: float
think_samples: int
answer_samples: int
# Accuracy of the block-boundary decision, restricted to the control menu, and of the
# stop-versus-continue half of it on its own. Zero outside the adaptive objective.
control_accuracy: float = 0.0
stop_accuracy: float = 0.0
def hybrid_objective(
model: DiffusionTransformer,
tokens: Tensor,
regions: Tensor,
*,
block: int,
think_probability: float,
mask_eps: float,
generator: torch.Generator | None = None,
) -> HybridBatchLoss:
"""Mixed objective: denoise all thought slots or predict the answer tokens.
Each sample is assigned one mode. Think samples corrupt every slot at an
independent noise level under a block-causal mask, so one forward trains all
slot conditionals; answer samples see the full clean reasoning prefix
bidirectionally and the answer causally.
"""
device = tokens.device
batch_size, seq_len = tokens.shape
problem_len, n_slots, answer_start, answer_end = regions.unbind(dim=1)
positions = torch.arange(seq_len, device=device)
think_sel = (
torch.rand(batch_size, device=device, generator=generator) < think_probability
)
if bool(think_sel.all()):
think_sel[-1] = False
if not bool(think_sel.any()):
think_sel[0] = True
prefix_end = problem_len + 1
think_end = prefix_end + n_slots * block
think_region = (
(positions[None, :] >= prefix_end[:, None])
& (positions[None, :] < think_end[:, None])
& think_sel[:, None]
)
max_slots = int(n_slots.max())
slot_noise = mask_eps + (1.0 - mask_eps) * torch.rand(
batch_size, max_slots, device=device, generator=generator
)
slot_index = torch.clamp(
(positions[None, :] - prefix_end[:, None]) // block, min=0, max=max_slots - 1
)
token_noise = slot_noise.gather(1, slot_index)
mask = (
torch.rand(tokens.shape, device=device, generator=generator) < token_noise
) & think_region
noisy_tokens = torch.where(mask, model.config.mask_token_id, tokens)
think_blocked = slot_causal_blocked(problem_len, n_slots, block, seq_len)
answer_blocked = prefix_causal_blocked(answer_start, seq_len)
blocked = torch.where(think_sel[:, None, None], think_blocked, answer_blocked)
predict_positions = (
(positions[None, :] >= answer_start[:, None] - 1)
& (positions[None, :] < answer_end[:, None] - 1)
& ~think_sel[:, None]
)
output_positions = mask | predict_positions
logits = model(noisy_tokens, output_positions=output_positions, attn_mask=blocked)
think_rows = mask[output_positions]
# Graph-connected zero: forbidden-output columns sit at finfo.min, so a raw sum
# overflows to -inf in low precision and would poison the scalar via -inf * 0.
zero = logits.sum().clamp(-1.0, 1.0) * 0.0
think_loss = zero
think_accuracy = 0.0
if bool(mask.any()):
think_logits = logits[think_rows]
think_targets = tokens[mask]
per_token = F.cross_entropy(think_logits.float(), think_targets, reduction='none')
weights = token_noise[mask]
normalizer = think_region.sum().clamp_min(1)
think_loss = (per_token / weights).sum() / normalizer
think_accuracy = float(
(think_logits.argmax(dim=-1) == think_targets).float().mean()
)
answer_loss = zero
answer_accuracy = 0.0
if bool(predict_positions.any()):
answer_logits = logits[~think_rows]
target_positions = torch.zeros_like(predict_positions)
target_positions[:, 1:] = predict_positions[:, :-1]
answer_targets = tokens[target_positions]
answer_loss = F.cross_entropy(answer_logits.float(), answer_targets)
answer_accuracy = float(
(answer_logits.argmax(dim=-1) == answer_targets).float().mean()
)
return HybridBatchLoss(
loss=think_loss + answer_loss,
think_loss=float(think_loss),
answer_loss=float(answer_loss),
think_accuracy=think_accuracy,
answer_accuracy=answer_accuracy,
think_samples=int(think_sel.sum()),
answer_samples=int((~think_sel).sum()),
)
def adaptive_hybrid_objective(
model: DiffusionTransformer,
tokens: Tensor,
regions: Tensor,
*,
size_ids: Tensor,
end_think_id: int,
think_probability: float,
mask_eps: float,
causal_prefix: bool = False,
control_context_noise: float = 0.0,
generator: torch.Generator | None = None,
) -> HybridBatchLoss:
"""Mixed objective for the adaptive layout.
Think samples denoise every thought block at an independent noise level under
the variable-boundary block-causal mask. The remaining samples run a causal
next-token objective over the whole reasoning-and-answer stream, which is where
the model learns the ```` block-size and ```` termination decisions.
Those causal samples read a pristine think region, while at inference the controller
reads thoughts the model just wrote, complete with sampling damage.
``control_context_noise`` closes that gap by swapping a fraction of their think tokens
for other tokens drawn from the batch: inputs degrade, targets stay clean.
"""
device = tokens.device
batch_size, seq_len = tokens.shape
problem_len, _, answer_start, answer_end = regions.unbind(dim=1)
positions = torch.arange(seq_len, device=device)
think_sel = (
torch.rand(batch_size, device=device, generator=generator) < think_probability
)
if bool(think_sel.all()):
think_sel[-1] = False
if not bool(think_sel.any()):
think_sel[0] = True
prefix_end = problem_len + 1
in_think = (positions[None, :] >= prefix_end[:, None]) & (
positions[None, :] < answer_start[:, None]
)
is_size = (tokens.unsqueeze(-1) == size_ids).any(dim=-1)
boundary = (is_size | (tokens == end_think_id)) & in_think
think_content = in_think & ~boundary & think_sel[:, None]
block_id = torch.cumsum((is_size & in_think).long(), dim=1)
max_blocks = int(block_id.max().clamp(min=1))
slot_noise = mask_eps + (1.0 - mask_eps) * torch.rand(
batch_size, max_blocks, device=device, generator=generator
)
gather_index = (block_id - 1).clamp(min=0, max=max_blocks - 1)
token_noise = slot_noise.gather(1, gather_index)
mask = (
torch.rand(tokens.shape, device=device, generator=generator) < token_noise
) & think_content
noisy_tokens = torch.where(mask, model.config.mask_token_id, tokens)
if control_context_noise > 0.0:
# Replacements come from the batch itself so the corrupted context keeps realistic
# token statistics; boundaries and targets are read from the clean tensor.
control_content = in_think & ~boundary & ~think_sel[:, None]
corrupt = (
torch.rand(tokens.shape, device=device, generator=generator) < control_context_noise
) & control_content
flat = tokens.reshape(-1)
picks = torch.randint(
0, flat.numel(), tokens.shape, device=device, generator=generator
)
noisy_tokens = torch.where(corrupt, flat[picks], noisy_tokens)
think_blocked = adaptive_block_mask(
tokens, problem_len, answer_start, size_ids, end_think_id,
causal_prefix=causal_prefix,
)
ar_blocked = prefix_causal_blocked(prefix_end, seq_len, causal_prefix=causal_prefix)
blocked = torch.where(think_sel[:, None, None], think_blocked, ar_blocked)
predict_positions = (
(positions[None, :] >= problem_len[:, None])
& (positions[None, :] < answer_end[:, None] - 1)
& ~think_sel[:, None]
)
output_positions = mask | predict_positions
logits = model(noisy_tokens, output_positions=output_positions, attn_mask=blocked)
think_rows = mask[output_positions]
# Graph-connected zero: forbidden-output columns sit at finfo.min, so a raw sum
# overflows to -inf in low precision and would poison the scalar via -inf * 0.
zero = logits.sum().clamp(-1.0, 1.0) * 0.0
think_loss = zero
think_accuracy = 0.0
if bool(mask.any()):
think_logits = logits[think_rows]
think_targets = tokens[mask]
per_token = F.cross_entropy(think_logits.float(), think_targets, reduction='none')
weights = token_noise[mask]
normalizer = think_content.sum().clamp_min(1)
think_loss = (per_token / weights).sum() / normalizer
think_accuracy = float(
(think_logits.argmax(dim=-1) == think_targets).float().mean()
)
answer_loss = zero
answer_accuracy = 0.0
control_accuracy = 0.0
stop_accuracy = 0.0
if bool(predict_positions.any()):
answer_logits = logits[~think_rows]
target_positions = torch.zeros_like(predict_positions)
target_positions[:, 1:] = predict_positions[:, :-1]
answer_targets = tokens[target_positions]
answer_loss = F.cross_entropy(answer_logits.float(), answer_targets)
control_ids = torch.cat(
[size_ids, torch.tensor([end_think_id], device=device, dtype=size_ids.dtype)]
)
is_control = (answer_targets.unsqueeze(-1) == control_ids).any(dim=-1)
answer_accuracy = float(
(answer_logits.argmax(dim=-1) == answer_targets).float().mean()
)
if bool(is_control.any()):
# Scored inside the control menu: an open-vocabulary argmax hides whether the
# boundary decision is right, since content tokens dominate those logits.
menu = answer_logits[is_control].index_select(-1, control_ids)
chosen = control_ids[menu.argmax(dim=-1)]
targets = answer_targets[is_control]
control_accuracy = float((chosen == targets).float().mean())
stop_accuracy = float(
((chosen == end_think_id) == (targets == end_think_id)).float().mean()
)
return HybridBatchLoss(
loss=think_loss + answer_loss,
think_loss=float(think_loss),
answer_loss=float(answer_loss),
think_accuracy=think_accuracy,
answer_accuracy=answer_accuracy,
think_samples=int(think_sel.sum()),
answer_samples=int((~think_sel).sum()),
control_accuracy=control_accuracy,
stop_accuracy=stop_accuracy,
)
def block_size_curriculum(
step: int | None, *, n_sizes: int, curriculum_steps: int
) -> Tensor:
"""Segment-size sampling weights: smallest-size-only ramping linearly to uniform.
``sizes`` are assumed ascending. ``step=None`` (evaluation) and a zero-length
curriculum both return the uniform end state so losses stay comparable across
checkpoints.
"""
uniform = torch.full((n_sizes,), 1.0 / n_sizes)
if step is None or curriculum_steps <= 0:
return uniform
progress = min(1.0, step / curriculum_steps)
smallest_only = torch.zeros(n_sizes)
smallest_only[0] = 1.0
return smallest_only * (1.0 - progress) + uniform * progress
def block_diffusion_objective(
model: DiffusionTransformer,
tokens: Tensor,
*,
sizes: tuple[int, ...],
size_weights: Tensor,
mask_eps: float,
ar_probability: float = 0.0,
generator: torch.Generator | None = None,
) -> HybridBatchLoss:
"""Variable-block denoising over plain packed text (the conversion objective).
Each sample is tiled with segments whose lengths are drawn from ``sizes`` under
``size_weights``; every segment is corrupted at an independent noise level and
denoised in one forward under the block-causal geometry (own segment bidirectional,
all preceding segments visible). No control tokens exist in the stream — pretraining
teaches variable-block denoising only; control decisions are learned in SFT. With
probability ``ar_probability`` a sample instead runs plain causal next-token loss,
retaining the autoregressive ability that control and answer decoding rely on.
Reported as ``think_*`` (denoising) and ``answer_*`` (causal retention) metrics.
"""
device = tokens.device
batch_size, seq_len = tokens.shape
positions = torch.arange(seq_len, device=device)
ar_sel = (
torch.rand(batch_size, device=device, generator=generator) < ar_probability
)
if bool(ar_sel.all()):
ar_sel[0] = False
diff_sel = ~ar_sel
size_tensor = torch.tensor(sizes, device=device, dtype=torch.long)
max_segments = -(-seq_len // int(min(sizes)))
drawn_index = torch.multinomial(
size_weights.to(device).expand(batch_size, -1),
max_segments,
replacement=True,
generator=generator,
)
drawn = size_tensor[drawn_index]
starts = torch.cumsum(drawn, dim=1) - drawn
valid = starts < seq_len
boundary_hits = torch.zeros(batch_size, seq_len, dtype=torch.long, device=device)
boundary_hits.scatter_add_(1, starts.clamp(max=seq_len - 1), valid.long())
boundary = boundary_hits > 0
segment_id = torch.cumsum(boundary.long(), dim=1) - 1
segment_noise = mask_eps + (1.0 - mask_eps) * torch.rand(
batch_size, max_segments, device=device, generator=generator
)
token_noise = segment_noise.gather(1, segment_id.clamp(max=max_segments - 1))
mask = (
torch.rand(tokens.shape, device=device, generator=generator) < token_noise
) & diff_sel[:, None]
noisy_tokens = torch.where(mask, model.config.mask_token_id, tokens)
zeros = torch.zeros(batch_size, dtype=torch.long, device=device)
full = torch.full((batch_size,), seq_len, dtype=torch.long, device=device)
block_blocked = block_mask_from_boundaries(boundary, zeros, full)
causal_blocked = prefix_causal_blocked(zeros, seq_len, causal_prefix=True)
blocked = torch.where(diff_sel[:, None, None], block_blocked, causal_blocked)
predict_positions = (positions[None, :] < seq_len - 1) & ar_sel[:, None]
output_positions = mask | predict_positions
logits = model(noisy_tokens, output_positions=output_positions, attn_mask=blocked)
think_rows = mask[output_positions]
# Graph-connected zero: forbidden-output columns sit at finfo.min, so a raw sum
# overflows to -inf in low precision and would poison the scalar via -inf * 0.
zero = logits.sum().clamp(-1.0, 1.0) * 0.0
think_loss = zero
think_accuracy = 0.0
if bool(mask.any()):
think_logits = logits[think_rows]
think_targets = tokens[mask]
per_token = F.cross_entropy(think_logits.float(), think_targets, reduction='none')
weights = token_noise[mask]
normalizer = (diff_sel.sum() * seq_len).clamp_min(1)
think_loss = (per_token / weights).sum() / normalizer
think_accuracy = float(
(think_logits.argmax(dim=-1) == think_targets).float().mean()
)
answer_loss = zero
answer_accuracy = 0.0
if bool(predict_positions.any()):
answer_logits = logits[~think_rows]
target_positions = torch.zeros_like(predict_positions)
target_positions[:, 1:] = predict_positions[:, :-1]
answer_targets = tokens[target_positions]
answer_loss = F.cross_entropy(answer_logits.float(), answer_targets)
answer_accuracy = float(
(answer_logits.argmax(dim=-1) == answer_targets).float().mean()
)
return HybridBatchLoss(
loss=think_loss + answer_loss,
think_loss=float(think_loss),
answer_loss=float(answer_loss),
think_accuracy=think_accuracy,
answer_accuracy=answer_accuracy,
think_samples=int(diff_sel.sum()),
answer_samples=int(ar_sel.sum()),
)
@dataclass(frozen=True)
class ARBatchLoss:
loss: Tensor
accuracy: float
token_count: int
def ar_objective(
model: DiffusionTransformer, tokens: Tensor, regions: Tensor
) -> ARBatchLoss:
"""Prefix-LM next-token objective over the think and answer regions."""
device = tokens.device
_, seq_len = tokens.shape
problem_len, _, _, answer_end = regions.unbind(dim=1)
positions = torch.arange(seq_len, device=device)
blocked = prefix_causal_blocked(problem_len + 1, seq_len)
predict_positions = (positions[None, :] >= problem_len[:, None]) & (
positions[None, :] < answer_end[:, None] - 1
)
logits = model(tokens, output_positions=predict_positions, attn_mask=blocked)
target_positions = torch.zeros_like(predict_positions)
target_positions[:, 1:] = predict_positions[:, :-1]
targets = tokens[target_positions]
loss = F.cross_entropy(logits.float(), targets)
accuracy = float((logits.argmax(dim=-1) == targets).float().mean())
return ARBatchLoss(loss=loss, accuracy=accuracy, token_count=int(targets.numel()))
def diffusion_objective(
model: DiffusionTransformer,
tokens: Tensor,
regions: Tensor,
*,
mask_eps: float,
mask_probability: Tensor | None = None,
generator: torch.Generator | None = None,
) -> tuple[Tensor, CorruptionBatch, Tensor]:
"""Whole-response denoising: corrupt everything after the problem, pads included."""
device = tokens.device
_, seq_len = tokens.shape
problem_len = regions[:, 0]
positions = torch.arange(seq_len, device=device)
response_mask = positions[None, :] >= (problem_len[:, None] + 1)
corruption = corrupt_tokens(
tokens,
model.config.mask_token_id,
valid_mask=response_mask,
mask_probability=mask_probability,
eps=mask_eps,
generator=generator,
)
logits = model(corruption.noisy_tokens, output_positions=corruption.mask)
return logits, corruption, response_mask
@dataclass
class GenerationResult:
tokens: list[int]
think_tokens: list[int]
answer_tokens: list[int]
think_seconds: float = 0.0
answer_seconds: float = 0.0
forward_passes: int = 0
slots_used: int = 0
block_sizes: tuple[int, ...] = ()
terminated: bool = False
@property
def total_seconds(self) -> float:
return self.think_seconds + self.answer_seconds
def _apply_repetition_penalty(
logits: Tensor, token_ids: list[int], penalty: float
) -> Tensor:
"""Divide logits of already-emitted tokens by ``penalty`` (CTRL convention)."""
if penalty == 1.0 or not token_ids:
return logits
index = torch.tensor(sorted(set(token_ids)), device=logits.device)
selected = logits.index_select(-1, index)
adjusted = torch.where(selected > 0, selected / penalty, selected * penalty)
return logits.index_copy(-1, index, adjusted)
def _apply_top_p(logits: Tensor, top_p: float) -> Tensor:
"""Restrict sampling to the smallest set of tokens whose mass reaches ``top_p``."""
if top_p >= 1.0:
return logits
ordered, indices = torch.sort(logits, descending=True, dim=-1)
cumulative = ordered.softmax(dim=-1).cumsum(dim=-1)
remove = cumulative - ordered.softmax(dim=-1) >= top_p
ordered = ordered.masked_fill(remove, torch.finfo(logits.dtype).min)
return ordered.gather(-1, indices.argsort(dim=-1))
def _kv_cache_enabled(part: str) -> bool:
"""Key/value caching per part, selected by ``MDLM_KV_CACHE``: ar, block, all or off.
Defaults to ``ar``, which measured 101s to 12.6s on the same prompt's answer: one query
token against an all-visible mask has no downside. The denoising prefix does — caching it
hands the backbone a dense mask, dropping flex attention onto its score_mod path and losing
the block skipping the uncached call gets, 10.6s per block against 5.3s on an L40S
(2026-07-27). It stays off until that path builds a rectangular BlockMask instead.
"""
setting = os.environ.get('MDLM_KV_CACHE', 'ar')
return setting in ('all', '1') or setting == part
def _cached_block_logits(model, blocked: Tensor, prefix_len: int):
"""Score a denoising block against a cached prefix, or ``None`` without cache support.
The prefix is encoded once per block; every denoising step then feeds only the block's
own positions. Its keys and values are dropped between steps because the block's tokens
keep changing as they are revealed, while the prefix behind them does not.
"""
if not hasattr(model, 'forward_cached') or not _kv_cache_enabled('block'):
return None
cache = model.new_cache()
def logits_fn(tokens: Tensor, masked: Tensor) -> Tensor:
if cache.get_seq_length() == 0:
with torch.inference_mode():
model.forward_cached(
tokens[:, :prefix_len],
attn_mask=blocked[:, :prefix_len, :prefix_len],
past_key_values=cache,
)
cache.crop(prefix_len)
with torch.inference_mode():
logits, _ = model.forward_cached(
tokens[:, prefix_len:],
attn_mask=blocked[:, prefix_len:, :],
past_key_values=cache,
output_positions=masked[:, prefix_len:],
)
return logits
return logits_fn
def _ar_decode_cached(
model: DiffusionTransformer,
sequence: list[int],
prefix_len: int,
*,
causal_prefix: bool,
eos_id: int,
max_new_tokens: int,
temperature: float,
repetition_penalty: float,
top_p: float,
device: torch.device,
generator: torch.Generator | None,
) -> tuple[list[int], int]:
"""Same decoding as :func:`_ar_decode` with the prefix kept in a key/value cache.
Masks come from :func:`prefix_causal_blocked` exactly as in the uncached path, sliced to
the rows of the queries being fed. Priming with an all-visible mask instead would look
right — the final row is identical, so a single generated token matches — while silently
computing every earlier position bidirectionally and poisoning the cached keys.
"""
generated: list[int] = []
cache = model.new_cache()
step_in = torch.tensor([sequence], dtype=torch.long, device=device)
forwards = 0
prefix = torch.tensor([prefix_len], device=device)
for _ in range(max_new_tokens):
cached = cache.get_seq_length()
length = cached + step_in.shape[1]
visible = prefix_causal_blocked(prefix, length, causal_prefix=causal_prefix)[:, cached:, :]
output_positions = torch.zeros_like(step_in, dtype=torch.bool)
output_positions[0, -1] = True
with torch.inference_mode():
logits, cache = model.forward_cached(
step_in, attn_mask=visible, past_key_values=cache,
output_positions=output_positions,
)
logits = _apply_repetition_penalty(logits, generated, repetition_penalty)
logits = _apply_top_p(logits, top_p)
token, _ = _sample_categorical(logits, temperature, generator)
forwards += 1
token_id = int(token.item())
generated.append(token_id)
if token_id == eos_id:
break
step_in = torch.tensor([[token_id]], dtype=torch.long, device=device)
return generated, forwards
def _ar_decode(
model: DiffusionTransformer,
sequence: list[int],
prefix_len: int,
*,
eos_id: int,
max_new_tokens: int,
temperature: float,
repetition_penalty: float,
top_p: float,
device: torch.device,
generator: torch.Generator | None,
causal_prefix: bool = False,
) -> tuple[list[int], int]:
"""Greedy/temperature decoding under the bidirectional-prefix causal mask.
``repetition_penalty`` (>1 discourages repeats) and ``top_p`` nucleus truncation
curb the degenerate loops small models fall into under plain temperature sampling.
"""
generated: list[int] = []
forwards = 0
max_new_tokens = min(max_new_tokens, model.config.max_seq_len - len(sequence))
if hasattr(model, 'forward_cached') and _kv_cache_enabled('ar'):
return _ar_decode_cached(
model, sequence, prefix_len, causal_prefix=causal_prefix, eos_id=eos_id,
max_new_tokens=max_new_tokens, temperature=temperature,
repetition_penalty=repetition_penalty, top_p=top_p, device=device,
generator=generator,
)
for _ in range(max_new_tokens):
current = torch.tensor([sequence + generated], dtype=torch.long, device=device)
seq_len = current.shape[1]
prefix = torch.tensor([prefix_len], device=device)
blocked = prefix_causal_blocked(prefix, seq_len, causal_prefix=causal_prefix)
output_positions = torch.zeros_like(current, dtype=torch.bool)
output_positions[0, -1] = True
with torch.inference_mode():
logits = model(current, output_positions=output_positions, attn_mask=blocked)
logits = _apply_repetition_penalty(logits, generated, repetition_penalty)
logits = _apply_top_p(logits, top_p)
token, _ = _sample_categorical(logits, temperature, generator)
forwards += 1
token_id = int(token.item())
generated.append(token_id)
if token_id == eos_id:
break
return generated, forwards
@torch.no_grad()
def generate_hybrid(
model: DiffusionTransformer,
prompt_ids: list[int],
*,
think_id: int,
end_think_id: int,
thought_pad_id: int,
eos_id: int,
block: int,
max_slots: int,
steps_per_block: int,
max_answer_tokens: int = 64,
temperature: float = 0.7,
repetition_penalty: float = 1.0,
top_p: float = 1.0,
strategy: str = 'confidence',
device: torch.device | str = 'cpu',
generator: torch.Generator | None = None,
) -> GenerationResult:
"""Denoise thought slots sequentially, then decode the answer autoregressively."""
device = torch.device(device)
mask_id = model.config.mask_token_id
sequence = [*prompt_ids, think_id]
forwards = 0
slots_used = 0
# Each new slot must leave room for itself plus at least a minimal answer.
slot_budget = model.config.max_seq_len - block - 8
think_started = time.perf_counter()
problem_tensor = torch.tensor([len(prompt_ids)], device=device)
for slot_index in range(max_slots):
if len(sequence) > slot_budget:
break
window = torch.tensor(
[sequence + [mask_id] * block], dtype=torch.long, device=device
)
blocked = slot_causal_blocked(
problem_tensor,
torch.tensor([slot_index + 1], device=device),
block,
window.shape[1],
)
filled = iterative_unmask(
model,
window,
mask_id,
steps=steps_per_block,
temperature=temperature,
strategy=strategy,
attn_mask=blocked,
generator=generator,
)
slot = [int(token) for token in filled[0, len(sequence):]]
forwards += steps_per_block
slots_used += 1
sequence.extend(slot)
if end_think_id in slot:
break
if end_think_id not in sequence[len(prompt_ids):]:
# Match the trained answer geometry when the model never closes its thinking.
terminal = [end_think_id] + [thought_pad_id] * (block - 1)
sequence.extend(terminal[: max(1, model.config.max_seq_len - 8 - len(sequence))])
think_seconds = time.perf_counter() - think_started
think_tokens = sequence[len(prompt_ids):]
answer_started = time.perf_counter()
answer, answer_forwards = _ar_decode(
model,
sequence,
prefix_len=len(sequence),
eos_id=eos_id,
max_new_tokens=max_answer_tokens,
temperature=temperature,
repetition_penalty=repetition_penalty,
top_p=top_p,
device=device,
generator=generator,
)
answer_seconds = time.perf_counter() - answer_started
return GenerationResult(
tokens=sequence + answer,
think_tokens=think_tokens,
answer_tokens=answer,
think_seconds=think_seconds,
answer_seconds=answer_seconds,
forward_passes=forwards + answer_forwards,
slots_used=slots_used,
)
def _ar_predict_control(
model: DiffusionTransformer,
sequence: list[int],
allowed_ids: list[int],
*,
prefix_len: int,
temperature: float,
device: torch.device,
generator: torch.Generator | None,
causal_prefix: bool = False,
) -> int:
"""Predict the next control token, restricted to the allowed size/stop ids."""
current = torch.tensor([sequence], dtype=torch.long, device=device)
seq_len = current.shape[1]
blocked = prefix_causal_blocked(
torch.tensor([prefix_len], device=device), seq_len, causal_prefix=causal_prefix
)
output_positions = torch.zeros_like(current, dtype=torch.bool)
output_positions[0, -1] = True
with torch.inference_mode():
logits = model(current, output_positions=output_positions, attn_mask=blocked)
restricted = torch.full_like(logits, torch.finfo(logits.dtype).min)
index = torch.tensor(allowed_ids, device=logits.device)
restricted.index_copy_(-1, index, logits.index_select(-1, index))
token, _ = _sample_categorical(restricted, temperature, generator)
return int(token.item())
@torch.no_grad()
def generate_hybrid_adaptive(
model: DiffusionTransformer,
prompt_ids: list[int],
*,
think_id: int,
end_think_id: int,
thought_pad_id: int,
eos_id: int,
size_ids: dict[int, int],
steps_per_block: int,
max_blocks: int,
max_answer_tokens: int = 96,
temperature: float = 0.7,
control_temperature: float = 0.0,
repetition_penalty: float = 1.0,
top_p: float = 1.0,
strategy: str = 'confidence',
causal_prefix: bool = False,
device: torch.device | str = 'cpu',
generator: torch.Generator | None = None,
) -> GenerationResult:
"""Interleave AR block-size decisions with in-block diffusion, then decode.
At each boundary the model predicts a ```` control token or ````.
A size token allocates that many masked positions denoised in parallel under the
variable-boundary block-causal mask; ```` ends thinking. The answer is
then decoded autoregressively with the same repetition and nucleus controls.
"""
device = torch.device(device)
mask_id = model.config.mask_token_id
size_by_id = {token_id: size for size, token_id in size_ids.items()}
size_ids_tensor = torch.tensor(sorted(size_ids.values()), device=device)
control_ids = [*size_by_id.keys(), end_think_id]
prefix_len = len(prompt_ids) + 1
problem_tensor = torch.tensor([len(prompt_ids)], device=device)
sequence = [*prompt_ids, think_id]
forwards = 0
blocks_used = 0
chosen: list[int] = []
terminated = False
think_started = time.perf_counter()
for _ in range(max_blocks):
control = _ar_predict_control(
model,
sequence,
control_ids,
prefix_len=prefix_len,
temperature=control_temperature,
device=device,
generator=generator,
causal_prefix=causal_prefix,
)
forwards += 1
if control == end_think_id:
terminated = True
break
size = size_by_id[control]
if len(sequence) + 1 + size > model.config.max_seq_len - 8:
break
sequence.append(control)
window_prefix = len(sequence)
window = torch.tensor(
[sequence + [mask_id] * size], dtype=torch.long, device=device
)
blocked = adaptive_block_mask(
window,
problem_tensor,
torch.tensor([window.shape[1]], device=device),
size_ids_tensor,
end_think_id,
causal_prefix=causal_prefix,
)
filled = iterative_unmask(
model,
window,
mask_id,
steps=steps_per_block,
temperature=temperature,
strategy=strategy,
attn_mask=blocked,
generator=generator,
logits_fn=_cached_block_logits(model, blocked, window_prefix),
)
sequence.extend(int(token) for token in filled[0, window_prefix:])
forwards += steps_per_block
blocks_used += 1
chosen.append(size)
sequence.append(end_think_id)
think_seconds = time.perf_counter() - think_started
think_tokens = sequence[len(prompt_ids):]
answer_started = time.perf_counter()
answer, answer_forwards = _ar_decode(
model,
sequence,
prefix_len=prefix_len,
eos_id=eos_id,
max_new_tokens=max_answer_tokens,
temperature=temperature,
repetition_penalty=repetition_penalty,
top_p=top_p,
device=device,
generator=generator,
causal_prefix=causal_prefix,
)
answer_seconds = time.perf_counter() - answer_started
return GenerationResult(
tokens=sequence + answer,
think_tokens=think_tokens,
answer_tokens=answer,
think_seconds=think_seconds,
answer_seconds=answer_seconds,
forward_passes=forwards + answer_forwards,
slots_used=blocks_used,
block_sizes=tuple(chosen),
terminated=terminated,
)
@torch.no_grad()
def generate_ar(
model: DiffusionTransformer,
prompt_ids: list[int],
*,
think_id: int,
end_think_id: int,
eos_id: int,
max_new_tokens: int = 384,
temperature: float = 0.7,
device: torch.device | str = 'cpu',
generator: torch.Generator | None = None,
) -> GenerationResult:
"""Classic sequential CoT baseline under the prefix-LM mask."""
device = torch.device(device)
sequence = [*prompt_ids, think_id]
started = time.perf_counter()
generated, forwards = _ar_decode(
model,
sequence,
prefix_len=len(sequence),
eos_id=eos_id,
max_new_tokens=max_new_tokens,
temperature=temperature,
device=device,
generator=generator,
)
elapsed = time.perf_counter() - started
if end_think_id in generated:
split = generated.index(end_think_id) + 1
else:
split = len(generated)
return GenerationResult(
tokens=sequence + generated,
think_tokens=generated[:split],
answer_tokens=generated[split:],
think_seconds=elapsed,
answer_seconds=0.0,
forward_passes=forwards,
)
@torch.no_grad()
def generate_diffusion(
model: DiffusionTransformer,
prompt_ids: list[int],
*,
think_id: int,
end_think_id: int,
eos_id: int,
response_budget: int,
steps: int,
temperature: float = 0.7,
blocked_token_ids: tuple[int, ...] = (),
device: torch.device | str = 'cpu',
generator: torch.Generator | None = None,
) -> GenerationResult:
"""Pure-diffusion baseline: denoise the entire response region at once.
Blocking the pad token here counters confidence-ordered pad collapse: pads are
the easiest predictions, so left unblocked they win every early reveal and
squeeze out the actual response text.
"""
device = torch.device(device)
mask_id = model.config.mask_token_id
budget = min(response_budget, model.config.max_seq_len - len(prompt_ids) - 1)
sequence = torch.tensor(
[[*prompt_ids, think_id] + [mask_id] * budget], dtype=torch.long, device=device
)
started = time.perf_counter()
filled = iterative_unmask(
model,
sequence,
mask_id,
steps=steps,
temperature=temperature,
strategy='confidence',
blocked_token_ids=blocked_token_ids,
generator=generator,
)
elapsed = time.perf_counter() - started
response = [int(token) for token in filled[0, len(prompt_ids) + 1:]]
if eos_id in response:
response = response[: response.index(eos_id) + 1]
if end_think_id in response:
split = response.index(end_think_id) + 1
else:
split = len(response)
return GenerationResult(
tokens=[*prompt_ids, think_id] + response,
think_tokens=response[:split],
answer_tokens=response[split:],
think_seconds=elapsed,
answer_seconds=0.0,
forward_passes=steps,
)