| name | memdlm-parametric-memory |
| title | MemDLM: Parametric Memory Enhancement for Diffusion Language Models |
| version | 0.0.3 |
| engine | skillxiv-v0.0.3-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2603.22241 |
| keywords | ["Memory Enhancement","Diffusion Language Models","Bi-level Optimization","Long Context","Needle-in-Haystack"] |
| description | Enhance diffusion language model performance on long-context tasks by embedding simulated denoising into training via bi-level optimization. Fast weights capture local trajectory experience; base model optimized with accumulated parametric memory. Achieves +17.0% on RULER Variable Tracking (8K) and +9.6% on BABILong with gains primarily from training-stage improvements. |
Component ID
Bi-level parametric memory optimization for diffusion language model training.
Motivation
Long-context language tasks require models to maintain accurate information over extended sequences. Diffusion language models struggle on needle-in-haystack tasks where critical information appears sparsely. Adding explicit parametric memory that captures local trajectory experience during training offloads memorization pressure from token representations, improving long-context capability.
What Was Modified
Bi-Level Optimization Framework
Insert a simulated denoising process into DLM training through two nested optimization loops:
class MemDLMTraining:
"""
Inner loop updates fast weights (parametric memory).
Outer loop updates base model conditioned on memory.
"""
def __init__(self, base_model, device="cuda"):
self.base_model = base_model
self.fast_weight_optimizer = None
def bi_level_step(self, batch_data, anchor_state, target_state):
"""
Two-stage trajectory: pre-anchor alignment → anchor-to-target prediction.
Fast weights act as Parametric Memory capturing local trajectory experience.
"""
current_state = batch_data
for denoising_step in range(self.inner_steps):
residual = current_state - anchor_state
loss_pre_anchor = self.fast_weights.predict_residual(
current_state, residual
)
loss_pre_anchor.backward()
.fast_weight_optimizer.step()
anchor_embedding = .base_model.encode(anchor_state)
loss_anchor_to_target = .fast_weights.predict_clean(
anchor_embedding, target_state
)
loss_anchor_to_target.backward()
.fast_weight_optimizer.step()
memory_repr = .fast_weights.get_memory()
base_loss = .base_model.compute_loss(
batch_data, target_state, memory_conditioning=memory_repr
)
base_loss.backward()
.base_model.optimizer.step()
loss_pre_anchor + loss_anchor_to_target + base_loss