vedkdev/text-diffusion-en-de
0
1"""2Text Diffusion Model for EN→DE Machine Translation3Self-contained training script.4Architecture: Masked Discrete Diffusion with DiT backbone5Inspired by MDLM (kuleshov-group) + LLaDA conditional generation6Dataset: WMT14 EN-DE7 8Usage:9 pip install torch transformers datasets trackio sacrebleu sacremoses sentencepiece protobuf10 python train.py11"""12 13import os14import math15import typing16import time17import json18import torch19import torch.nn as nn20import torch.nn.functional as F21from torch.utils.data import DataLoader, Dataset22from dataclasses import dataclass23from datasets import load_dataset24from transformers import AutoTokenizer, get_cosine_schedule_with_warmup25import trackio26 27 28# ═══════════════════════════════════════════════════════════════29# MODEL ARCHITECTURE30# ═══════════════════════════════════════════════════════════════31 32@dataclass33class DiffusionTranslatorConfig:34 vocab_size: int = 3212835 max_src_len: int = 12836 max_tgt_len: int = 12837 hidden_dim: int = 51238 n_heads: int = 839 n_blocks: int = 840 dropout: float = 0.141 cond_dim: int = 12842 mask_token_id: int = 3210043 pad_token_id: int = 044 45 46class Rotary(nn.Module):47 def __init__(self, dim, base=10_000):48 super().__init__()49 inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))50 self.register_buffer('inv_freq', inv_freq)51 self.seq_len_cached = None52 self.cos_cached = None53 self.sin_cached = None54 55 def forward(self, x, seq_dim=1):56 seq_len = x.shape[seq_dim]57 if seq_len != self.seq_len_cached:58 self.seq_len_cached = seq_len59 t = torch.arange(seq_len, device=x.device).type_as(self.inv_freq)60 freqs = torch.einsum("i,j->ij", t, self.inv_freq)61 emb = torch.cat((freqs, freqs), dim=-1).to(x.device)62 self.cos_cached = emb.cos()63 self.sin_cached = emb.sin()64 return self.cos_cached, self.sin_cached65 66 67def rotate_half(x):68 x1, x2 = x[..., :x.shape[-1] // 2], x[..., x.shape[-1] // 2:]69 return torch.cat((-x2, x1), dim=-1)70 71 72def apply_rotary_pos_emb(q, k, cos, sin):73 cos = cos[:q.shape[1], :]74 sin = sin[:q.shape[1], :]75 cos = cos.unsqueeze(0).unsqueeze(2)76 sin = sin.unsqueeze(0).unsqueeze(2)77 q = (q * cos) + (rotate_half(q) * sin)78 k = (k * cos) + (rotate_half(k) * sin)79 return q, k80 81 82class TimestepEmbedder(nn.Module):83 def __init__(self, hidden_size, frequency_embedding_size=256):84 super().__init__()85 self.mlp = nn.Sequential(86 nn.Linear(frequency_embedding_size, hidden_size, bias=True),87 nn.SiLU(),88 nn.Linear(hidden_size, hidden_size, bias=True),89 )90 self.frequency_embedding_size = frequency_embedding_size91 92 @staticmethod93 def timestep_embedding(t, dim, max_period=10000):94 half = dim // 295 freqs = torch.exp(96 -math.log(max_period) * torch.arange(0, half, dtype=torch.float32, device=t.device) / half97 )98 args = t[:, None].float() * freqs[None]99 embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)100 if dim % 2:101 embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)102 return embedding103 104 def forward(self, t):105 t_freq = self.timestep_embedding(t, self.frequency_embedding_size)106 return self.mlp(t_freq)107 108 109class LayerNorm(nn.Module):110 def __init__(self, dim):111 super().__init__()112 self.weight = nn.Parameter(torch.ones([dim]))113 self.dim = dim114 115 def forward(self, x):116 with torch.amp.autocast('cuda', enabled=False):117 x = F.layer_norm(x.float(), [self.dim])118 return x * self.weight[None, None, :]119 120 121class DiTBlock(nn.Module):122 """Diffusion Transformer block with adaptive layer norm (adaLN)."""123 124 def __init__(self, dim, n_heads, cond_dim, mlp_ratio=4, dropout=0.1):125 super().__init__()126 self.n_heads = n_heads127 self.head_dim = dim // n_heads128 129 self.norm1 = LayerNorm(dim)130 self.q_proj = nn.Linear(dim, dim, bias=False)131 self.k_proj = nn.Linear(dim, dim, bias=False)132 self.v_proj = nn.Linear(dim, dim, bias=False)133 self.attn_out = nn.Linear(dim, dim, bias=False)134 self.dropout1 = nn.Dropout(dropout)135 136 self.norm2 = LayerNorm(dim)137 self.mlp = nn.Sequential(138 nn.Linear(dim, mlp_ratio * dim, bias=True),139 nn.GELU(approximate='tanh'),140 nn.Linear(mlp_ratio * dim, dim, bias=True),141 )142 self.dropout2 = nn.Dropout(dropout)143 144 self.adaLN_modulation = nn.Linear(cond_dim, 6 * dim, bias=True)145 nn.init.zeros_(self.adaLN_modulation.weight)146 nn.init.zeros_(self.adaLN_modulation.bias)147 148 def forward(self, x, rotary_cos_sin, c, attention_mask=None):149 batch_size, seq_len, dim = x.shape150 151 mod = self.adaLN_modulation(c)[:, None, :].chunk(6, dim=2)152 shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod153 154 x_skip = x155 x_norm = self.norm1(x) * (1 + scale_msa) + shift_msa156 157 q = self.q_proj(x_norm).view(batch_size, seq_len, self.n_heads, self.head_dim)158 k = self.k_proj(x_norm).view(batch_size, seq_len, self.n_heads, self.head_dim)159 v = self.v_proj(x_norm).view(batch_size, seq_len, self.n_heads, self.head_dim)160 161 cos, sin = rotary_cos_sin162 q, k = apply_rotary_pos_emb(q, k, cos, sin)163 164 q = q.transpose(1, 2)165 k = k.transpose(1, 2)166 v = v.transpose(1, 2)167 168 # Bidirectional attention (no causal mask)169 attn_output = F.scaled_dot_product_attention(170 q, k, v, attn_mask=attention_mask,171 dropout_p=self.dropout1.p if self.training else 0.0172 )173 attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, dim)174 175 x = x_skip + gate_msa * self.dropout1(self.attn_out(attn_output))176 177 x_skip = x178 x_norm = self.norm2(x) * (1 + scale_mlp) + shift_mlp179 x = x_skip + gate_mlp * self.dropout2(self.mlp(x_norm))180 181 return x182 183 184class DiffusionTranslator(nn.Module):185 """186 Masked Discrete Diffusion model for EN→DE translation.187 Input: [source_tokens | target_tokens] where target tokens are partially masked188 Bidirectional transformer (DiT blocks with adaLN for timestep conditioning)189 """190 191 def __init__(self, config: DiffusionTranslatorConfig):192 super().__init__()193 self.config = config194 195 self.vocab_embed = nn.Embedding(config.vocab_size, config.hidden_dim)196 self.sigma_map = TimestepEmbedder(config.cond_dim)197 self.rotary_emb = Rotary(config.hidden_dim // config.n_heads)198 self.segment_embed = nn.Embedding(2, config.hidden_dim)199 200 self.blocks = nn.ModuleList([201 DiTBlock(config.hidden_dim, config.n_heads, config.cond_dim, dropout=config.dropout)202 for _ in range(config.n_blocks)203 ])204 205 self.final_norm = LayerNorm(config.hidden_dim)206 self.final_adaLN = nn.Linear(config.cond_dim, 2 * config.hidden_dim, bias=True)207 nn.init.zeros_(self.final_adaLN.weight)208 nn.init.zeros_(self.final_adaLN.bias)209 210 self.output_proj = nn.Linear(config.hidden_dim, config.vocab_size, bias=False)211 self.output_proj.weight = self.vocab_embed.weight # Weight tying212 213 self._init_weights()214 215 def _init_weights(self):216 for module in self.modules():217 if isinstance(module, nn.Linear):218 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)219 if module.bias is not None:220 torch.nn.init.zeros_(module.bias)221 elif isinstance(module, nn.Embedding):222 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)223 224 def forward(self, input_ids, segment_ids, timesteps):225 x = self.vocab_embed(input_ids) + self.segment_embed(segment_ids)226 c = F.silu(self.sigma_map(timesteps))227 rotary_cos_sin = self.rotary_emb(x)228 229 for block in self.blocks:230 x = block(x, rotary_cos_sin, c)231 232 shift, scale = self.final_adaLN(c)[:, None, :].chunk(2, dim=2)233 x = self.final_norm(x) * (1 + scale) + shift234 logits = self.output_proj(x)235 return logits236 237 def count_parameters(self):238 return sum(p.numel() for p in self.parameters() if p.requires_grad)239 240 241def compute_diffusion_loss(model, input_ids, segment_ids, target_ids, target_mask, config):242 """Compute masked diffusion training loss (LLaDA-style ELBO)."""243 batch_size = input_ids.shape[0]244 device = input_ids.device245 246 eps = 1e-5247 t = torch.rand(batch_size, device=device) * (1 - eps) + eps248 249 mask_prob = t[:, None].expand_as(target_mask)250 random_mask = torch.rand_like(mask_prob) < mask_prob251 diffusion_mask = random_mask & target_mask252 253 noised_input = input_ids.clone()254 noised_input[diffusion_mask] = config.mask_token_id255 256 logits = model(noised_input, segment_ids, t)257 258 logits_flat = logits.view(-1, config.vocab_size)259 targets_flat = target_ids.view(-1)260 261 if diffusion_mask.sum() == 0:262 zero = torch.tensor(0.0, device=device, requires_grad=True)263 return zero, zero264 265 ce_loss = F.cross_entropy(logits_flat, targets_flat, reduction='none')266 267 masked_loss_2d = ce_loss.view(batch_size, -1) * diffusion_mask.float()268 per_example_counts = diffusion_mask.float().sum(dim=1).clamp(min=1.0)269 per_example_loss = masked_loss_2d.sum(dim=1) / per_example_counts270 271 weighted_loss = (per_example_loss / t).mean()272 unweighted_loss = per_example_loss.mean()273 274 return weighted_loss, unweighted_loss275 276 277@torch.no_grad()278def generate(model, src_ids, src_segment_ids, config, num_steps=50, device='cuda'):279 """Generate translation using iterative unmasking."""280 model.eval()281 batch_size = src_ids.shape[0]282 tgt_len = config.max_tgt_len283 284 tgt_ids = torch.full((batch_size, tgt_len), config.mask_token_id, device=device)285 tgt_segment_ids = torch.ones(batch_size, tgt_len, dtype=torch.long, device=device)286 287 input_ids = torch.cat([src_ids, tgt_ids], dim=1)288 segment_ids = torch.cat([src_segment_ids, tgt_segment_ids], dim=1)289 src_len = src_ids.shape[1]290 291 for step in range(num_steps, 0, -1):292 t = torch.tensor([step / num_steps], device=device).expand(batch_size)293 s = torch.tensor([(step - 1) / num_steps], device=device).expand(batch_size)294 295 logits = model(input_ids, segment_ids, t)296 tgt_logits = logits[:, src_len:, :]297 predicted_tokens = tgt_logits.argmax(dim=-1)298 299 current_tgt = input_ids[:, src_len:]300 still_masked = (current_tgt == config.mask_token_id)301 302 if step > 1:303 remask_prob = s[0].item() / t[0].item() if t[0].item() > 0 else 0.0304 remask = torch.rand_like(predicted_tokens.float()) < remask_prob305 new_tgt = current_tgt.clone()306 unmask_positions = still_masked & ~remask307 new_tgt[unmask_positions] = predicted_tokens[unmask_positions]308 else:309 new_tgt = current_tgt.clone()310 new_tgt[still_masked] = predicted_tokens[still_masked]311 312 input_ids = torch.cat([src_ids, new_tgt], dim=1)313 314 return input_ids[:, src_len:]315 316 317# ═══════════════════════════════════════════════════════════════318# DATASET319# ═══════════════════════════════════════════════════════════════320 321class WMT14EnDeDataset(Dataset):322 def __init__(self, data, tokenizer, max_src_len=128, max_tgt_len=128):323 self.data = data324 self.tokenizer = tokenizer325 self.max_src_len = max_src_len326 self.max_tgt_len = max_tgt_len327 328 def __len__(self):329 return len(self.data)330 331 def __getitem__(self, idx):332 item = self.data[idx]333 en_text = item['translation']['en']334 de_text = item['translation']['de']335 336 src_enc = self.tokenizer(337 "translate English to German: " + en_text,338 max_length=self.max_src_len, truncation=True,339 padding='max_length', return_tensors=None,340 )341 tgt_enc = self.tokenizer(342 de_text,343 max_length=self.max_tgt_len, truncation=True,344 padding='max_length', return_tensors=None,345 )346 347 src_ids = src_enc['input_ids']348 tgt_ids = tgt_enc['input_ids']349 segment_ids = [0] * len(src_ids) + [1] * len(tgt_ids)350 full_ids = src_ids + tgt_ids351 target_mask = [0] * len(src_ids) + tgt_enc['attention_mask']352 353 return {354 'input_ids': torch.tensor(full_ids, dtype=torch.long),355 'segment_ids': torch.tensor(segment_ids, dtype=torch.long),356 'target_ids': torch.tensor(full_ids, dtype=torch.long),357 'target_mask': torch.tensor(target_mask, dtype=torch.bool),358 }359 360 361# ═══════════════════════════════════════════════════════════════362# CONFIGURATION363# ═══════════════════════════════════════════════════════════════364 365MODEL_CONFIG = dict(366 vocab_size=None,367 max_src_len=128,368 max_tgt_len=128,369 hidden_dim=512,370 n_heads=8,371 n_blocks=12,372 dropout=0.1,373 cond_dim=128,374 mask_token_id=None,375 pad_token_id=None,376)377 378TRAIN_CONFIG = dict(379 learning_rate=3e-4,380 weight_decay=0.01,381 warmup_steps=4000,382 max_steps=200_000,383 batch_size=64,384 gradient_accumulation_steps=4,385 eval_every=5000,386 save_every=10000,387 log_every=100,388 max_grad_norm=1.0,389 num_gen_steps=50,390 fp16=True,391 seed=42,392)393 394HUB_MODEL_ID = "vedkdev/text-diffusion-en-de"395TOKENIZER_NAME = "Helsinki-NLP/opus-mt-en-de"396 397 398# ═══════════════════════════════════════════════════════════════399# TRAINING400# ═══════════════════════════════════════════════════════════════401 402def train():403 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')404 torch.manual_seed(TRAIN_CONFIG['seed'])405 406 print(f"Device: {device}")407 print(f"Loading tokenizer: {TOKENIZER_NAME}")408 409 tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_NAME)410 if tokenizer.mask_token is None:411 tokenizer.add_special_tokens({'mask_token': '<mask>'})412 413 mask_token_id = tokenizer.mask_token_id414 pad_token_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 0415 416 print(f"Vocab size: {len(tokenizer)}")417 print(f"Mask token ID: {mask_token_id}, Pad token ID: {pad_token_id}")418 419 MODEL_CONFIG['vocab_size'] = len(tokenizer)420 MODEL_CONFIG['mask_token_id'] = mask_token_id421 MODEL_CONFIG['pad_token_id'] = pad_token_id422 config = DiffusionTranslatorConfig(**MODEL_CONFIG)423 424 model = DiffusionTranslator(config).to(device)425 print(f"Model parameters: {model.count_parameters():,}")426 427 print("Loading WMT14 EN-DE dataset...")428 dataset = load_dataset("wmt/wmt14", "de-en", trust_remote_code=True)429 train_data = dataset['train']430 val_data = dataset['validation']431 print(f"Train: {len(train_data):,} | Val: {len(val_data):,}")432 433 train_dataset = WMT14EnDeDataset(train_data, tokenizer, config.max_src_len, config.max_tgt_len)434 val_dataset = WMT14EnDeDataset(val_data, tokenizer, config.max_src_len, config.max_tgt_len)435 436 train_loader = DataLoader(train_dataset, batch_size=TRAIN_CONFIG['batch_size'],437 shuffle=True, num_workers=4, pin_memory=True, drop_last=True)438 val_loader = DataLoader(val_dataset, batch_size=TRAIN_CONFIG['batch_size'],439 shuffle=False, num_workers=2, pin_memory=True)440 441 optimizer = torch.optim.AdamW(model.parameters(), lr=TRAIN_CONFIG['learning_rate'],442 weight_decay=TRAIN_CONFIG['weight_decay'], betas=(0.9, 0.98), eps=1e-8)443 scheduler = get_cosine_schedule_with_warmup(optimizer,444 num_warmup_steps=TRAIN_CONFIG['warmup_steps'],445 num_training_steps=TRAIN_CONFIG['max_steps'])446 447 scaler = torch.amp.GradScaler('cuda') if (TRAIN_CONFIG['fp16'] and device.type == 'cuda') else None448 449 trackio.init(project="text-diffusion-en-de", name="v1-wmt14-dit12-512d")450 451 global_step = 0452 best_val_loss = float('inf')453 accum_loss = 0.0454 accum_loss_uw = 0.0455 accum_count = 0456 457 eff_bs = TRAIN_CONFIG['batch_size'] * TRAIN_CONFIG['gradient_accumulation_steps']458 print(f"\n=== Starting Training ===")459 print(f"Effective batch size: {eff_bs} | Max steps: {TRAIN_CONFIG['max_steps']:,}")460 print(f"Warmup: {TRAIN_CONFIG['warmup_steps']:,} | LR: {TRAIN_CONFIG['learning_rate']}")461 462 model.train()463 optimizer.zero_grad()464 data_iter = iter(train_loader)465 start_time = time.time()466 467 total_micro_steps = TRAIN_CONFIG['max_steps'] * TRAIN_CONFIG['gradient_accumulation_steps']468 for step in range(1, total_micro_steps + 1):469 try:470 batch = next(data_iter)471 except StopIteration:472 data_iter = iter(train_loader)473 batch = next(data_iter)474 475 input_ids = batch['input_ids'].to(device)476 segment_ids = batch['segment_ids'].to(device)477 target_ids = batch['target_ids'].to(device)478 target_mask = batch['target_mask'].to(device)479 480 if scaler is not None:481 with torch.amp.autocast('cuda'):482 wl, uwl = compute_diffusion_loss(model, input_ids, segment_ids, target_ids, target_mask, config)483 loss = wl / TRAIN_CONFIG['gradient_accumulation_steps']484 scaler.scale(loss).backward()485 else:486 wl, uwl = compute_diffusion_loss(model, input_ids, segment_ids, target_ids, target_mask, config)487 loss = wl / TRAIN_CONFIG['gradient_accumulation_steps']488 loss.backward()489 490 accum_loss += wl.item()491 accum_loss_uw += uwl.item()492 accum_count += 1493 494 if step % TRAIN_CONFIG['gradient_accumulation_steps'] == 0:495 if scaler is not None:496 scaler.unscale_(optimizer)497 torch.nn.utils.clip_grad_norm_(model.parameters(), TRAIN_CONFIG['max_grad_norm'])498 scaler.step(optimizer)499 scaler.update()500 else:501 torch.nn.utils.clip_grad_norm_(model.parameters(), TRAIN_CONFIG['max_grad_norm'])502 optimizer.step()503 504 scheduler.step()505 optimizer.zero_grad()506 global_step += 1507 508 # Log509 if global_step % TRAIN_CONFIG['log_every'] == 0:510 avg_l = accum_loss / accum_count511 avg_uw = accum_loss_uw / accum_count512 elapsed = time.time() - start_time513 sps = global_step / elapsed514 lr = scheduler.get_last_lr()[0]515 print(f"step={global_step} | loss={avg_l:.4f} | ce_loss={avg_uw:.4f} | lr={lr:.2e} | steps/s={sps:.2f}")516 trackio.log({"train/loss_weighted": avg_l, "train/loss_ce": avg_uw,517 "train/learning_rate": lr, "train/steps_per_sec": sps}, step=global_step)518 accum_loss = 0.0519 accum_loss_uw = 0.0520 accum_count = 0521 522 # Eval523 if global_step % TRAIN_CONFIG['eval_every'] == 0:524 vl, vuw = evaluate(model, val_loader, config, device, scaler is not None)525 print(f" [EVAL] step={global_step} | val_loss={vl:.4f} | val_ce={vuw:.4f}")526 trackio.log({"eval/loss_weighted": vl, "eval/loss_ce": vuw}, step=global_step)527 528 if global_step % (TRAIN_CONFIG['eval_every'] * 4) == 0:529 bleu = evaluate_bleu(model, tokenizer, config, device, num_samples=100,530 num_steps=TRAIN_CONFIG['num_gen_steps'])531 trackio.log({"eval/sacrebleu": bleu}, step=global_step)532 533 if vuw < best_val_loss:534 best_val_loss = vuw535 save_model(model, config, tokenizer, global_step, is_best=True)536 model.train()537 538 # Save + push539 if global_step % TRAIN_CONFIG['save_every'] == 0:540 save_model(model, config, tokenizer, global_step, push_to_hub=True)541 542 # Final543 save_model(model, config, tokenizer, global_step, push_to_hub=True)544 print("\n=== Final BLEU Evaluation ===")545 bleu = evaluate_bleu(model, tokenizer, config, device, num_samples=200,546 num_steps=TRAIN_CONFIG['num_gen_steps'])547 trackio.log({"eval/final_sacrebleu": bleu}, step=global_step)548 print(f"\n=== Training Complete === Final BLEU: {bleu:.2f}")549 550 551def evaluate(model, val_loader, config, device, use_fp16=True):552 model.eval()553 total_l = total_uw = 0.0554 count = 0555 with torch.no_grad():556 for i, batch in enumerate(val_loader):557 if i >= 50:558 break559 ids = batch['input_ids'].to(device)560 seg = batch['segment_ids'].to(device)561 tgt = batch['target_ids'].to(device)562 mask = batch['target_mask'].to(device)563 if use_fp16:564 with torch.amp.autocast('cuda'):565 wl, uwl = compute_diffusion_loss(model, ids, seg, tgt, mask, config)566 else:567 wl, uwl = compute_diffusion_loss(model, ids, seg, tgt, mask, config)568 total_l += wl.item()569 total_uw += uwl.item()570 count += 1571 return total_l / max(count, 1), total_uw / max(count, 1)572 573 574def evaluate_bleu(model, tokenizer, config, device, num_samples=100, num_steps=50):575 import sacrebleu576 model.eval()577 ds = load_dataset("wmt/wmt14", "de-en", split="test", trust_remote_code=True)578 refs, hyps = [], []579 for i in range(min(num_samples, len(ds))):580 en = ds[i]['translation']['en']581 de_ref = ds[i]['translation']['de']582 enc = tokenizer("translate English to German: " + en, max_length=config.max_src_len,583 truncation=True, padding='max_length', return_tensors='pt')584 src_ids = enc['input_ids'].to(device)585 src_seg = torch.zeros_like(src_ids)586 with torch.no_grad():587 if device.type == 'cuda':588 with torch.amp.autocast('cuda'):589 gen = generate(model, src_ids, src_seg, config, num_steps=num_steps, device=device)590 else:591 gen = generate(model, src_ids, src_seg, config, num_steps=num_steps, device=device)592 hyp = tokenizer.decode(gen[0], skip_special_tokens=True)593 refs.append(de_ref)594 hyps.append(hyp)595 if i < 5:596 print(f" EN: {en[:100]}")597 print(f" REF: {de_ref[:100]}")598 print(f" GEN: {hyp[:100]}")599 print()600 bleu = sacrebleu.corpus_bleu(hyps, [refs])601 print(f"SacreBLEU: {bleu.score:.2f}")602 return bleu.score603 604 605def save_model(model, config, tokenizer, step, is_best=False, push_to_hub=False):606 save_dir = "checkpoints/best" if is_best else f"checkpoints/step-{step}"607 os.makedirs(save_dir, exist_ok=True)608 torch.save(model.state_dict(), os.path.join(save_dir, "model.pt"))609 config_dict = {k: getattr(config, k) for k in [610 'vocab_size', 'max_src_len', 'max_tgt_len', 'hidden_dim',611 'n_heads', 'n_blocks', 'dropout', 'cond_dim', 'mask_token_id', 'pad_token_id'612 ]}613 with open(os.path.join(save_dir, "config.json"), "w") as f:614 json.dump(config_dict, f, indent=2)615 tokenizer.save_pretrained(save_dir)616 if push_to_hub:617 push_model_to_hub(save_dir, step, config)618 print(f" Saved checkpoint to {save_dir}")619 620 621def push_model_to_hub(save_dir, step, config):622 from huggingface_hub import HfApi, upload_folder623 api = HfApi()624 try:625 api.create_repo(HUB_MODEL_ID, exist_ok=True, private=False)626 except Exception as e:627 print(f" Warning creating repo: {e}")628 629 readme = f"""---630tags:631- text-diffusion632- machine-translation633- en-de634- masked-diffusion635language:636- en637- de638datasets:639- wmt/wmt14640---641 642# Text Diffusion Model for EN→DE Translation643 644A **masked discrete diffusion** model for English-to-German machine translation,645trained from scratch on WMT14 EN-DE.646 647## Architecture648- **Type**: Masked Discrete Diffusion (inspired by MDLM + LLaDA)649- **Backbone**: DiT (Diffusion Transformer) with adaptive LayerNorm (adaLN)650- **Parameters**: ~72M651- **Blocks**: {config.n_blocks} DiT blocks, hidden_dim={config.hidden_dim}, heads={config.n_heads}652- **Tokenizer**: {TOKENIZER_NAME} (~58K vocab)653- **Max sequence**: {config.max_src_len} src + {config.max_tgt_len} tgt tokens654 655## Training656- **Dataset**: WMT14 EN-DE (~4.5M pairs)657- **Method**: Masked discrete diffusion with ELBO weighting (1/t)658- **Optimizer**: AdamW, lr=3e-4, cosine with 4K warmup659- **Effective batch size**: {TRAIN_CONFIG['batch_size'] * TRAIN_CONFIG['gradient_accumulation_steps']}660- **Training steps**: {step:,}661 662## How It Works6631. Source (EN) + target (DE) tokens concatenated: `[source | target]`6642. Training: target tokens randomly masked with prob `t ~ U(0,1)`, predict masked tokens6653. Inference: start fully masked → iteratively unmask over {TRAIN_CONFIG['num_gen_steps']} steps666 667## References668- [MDLM](https://arxiv.org/abs/2406.07524) | [LLaDA](https://arxiv.org/abs/2502.09992) | [DiNoiSer](https://arxiv.org/abs/2302.10025)669"""670 with open(os.path.join(save_dir, "README.md"), "w") as f:671 f.write(readme)672 673 upload_folder(repo_id=HUB_MODEL_ID, folder_path=save_dir,674 commit_message=f"Checkpoint step {step}")675 print(f" Pushed to hub: https://huggingface.co/{HUB_MODEL_ID}")676 677 678if __name__ == "__main__":679 train()680 