#!/usr/bin/env python3 """Compute dLLM mask-ratio objective gradient alignment. This mirrors the Fig. 3 idea from arXiv:2409.15557v1: compute one gradient vector per denoising condition, then visualize pairwise cosine similarities. For dLLMs the denoising condition is represented by mask ratio. """ from __future__ import annotations import argparse import csv import html import json import math import random import sys from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path from typing import Iterable import torch import torch.nn.functional as F @dataclass class Sample: text: str input_ids: list[int] def parse_ratios(spec: str) -> list[float]: if ":" in spec: start, step, end = [float(x) for x in spec.split(":")] ratios = [] value = start while value <= end + 1e-9: ratios.append(round(value, 6)) value += step return ratios return [float(x) for x in spec.split(",") if x.strip()] def ensure_dir(path: Path) -> None: path.mkdir(parents=True, exist_ok=True) def flatten_grads(params: Iterable[torch.nn.Parameter], device: torch.device) -> torch.Tensor: pieces = [] for param in params: if param.grad is None: pieces.append(torch.zeros(param.numel(), device=device, dtype=torch.float32)) else: pieces.append(param.grad.detach().float().reshape(-1).to(device)) if not pieces: raise RuntimeError("No parameters selected for gradient vector") return torch.cat(pieces) def cosine_matrix(vectors: list[torch.Tensor]) -> list[list[float]]: normed = [] for vec in vectors: vec64 = vec.double() denom = vec64.norm().clamp_min(1e-12) normed.append(vec64 / denom) matrix = [] for a in normed: row = [] for b in normed: value = float(torch.dot(a, b).clamp(-1.0, 1.0).item()) row.append(value) matrix.append(row) return matrix def write_csv(path: Path, ratios: list[float], matrix: list[list[float]]) -> None: with path.open("w", newline="", encoding="utf-8") as f: writer = csv.writer(f) writer.writerow(["ratio"] + ratios) for ratio, row in zip(ratios, matrix): writer.writerow([ratio] + [f"{value:.8f}" for value in row]) def color_for(value: float) -> str: value = max(-1.0, min(1.0, value)) if value >= 0: t = value r = int(255 * (1 - 0.10 * t)) g = int(255 * (1 - 0.55 * t)) b = int(255 * (1 - 0.70 * t)) else: t = -value r = int(255 * (1 - 0.65 * t)) g = int(255 * (1 - 0.45 * t)) b = int(255 * (1 - 0.05 * t)) return f"#{r:02x}{g:02x}{b:02x}" def write_svg(path: Path, ratios: list[float], matrix: list[list[float]], title: str) -> None: cell = 34 label_w = 88 title_h = 42 legend_h = 34 n = len(ratios) width = label_w + n * cell + 24 height = title_h + n * cell + legend_h + 24 parts = [ f'', '', f'{html.escape(title)}', ] x0 = label_w y0 = title_h for i, ratio in enumerate(ratios): x = x0 + i * cell + cell / 2 parts.append(f'{ratio:g}') y = y0 + i * cell + cell / 2 + 4 parts.append(f'{ratio:g}') for row_idx, row in enumerate(matrix): for col_idx, value in enumerate(row): x = x0 + col_idx * cell y = y0 + row_idx * cell parts.append(f'') parts.append(f'{value:.2f}') ly = y0 + n * cell + 18 parts.append(f'blue/yellow: negative, white: 0, red: positive cosine similarity') parts.append("") path.write_text("\n".join(parts) + "\n", encoding="utf-8") def write_report(path: Path, payload: dict) -> None: ratios = payload["ratios"] matrix = payload["cosine_similarity"] lines = [ "# dLLM Mask-Ratio Gradient Alignment Run", "", f"- Updated: `{payload['updated_at']}`", f"- Mode: `{payload['mode']}`", f"- Dataset: `{payload.get('dataset', '-')}`", f"- Samples: `{payload['num_samples']}`", f"- Block size: `{payload['block_size']}`", f"- Gradient scope: `{payload['grad_scope']}`", f"- Output dir: `{payload['output_dir']}`", "", "## Cosine Similarity", "", "| ratio | " + " | ".join(f"{r:g}" for r in ratios) + " |", "| --- | " + " | ".join(["---:"] * len(ratios)) + " |", ] for ratio, row in zip(ratios, matrix): lines.append("| " + f"{ratio:g}" + " | " + " | ".join(f"{x:.3f}" for x in row) + " |") lines.extend([ "", "## Files", "", "- `alignment.json`", "- `alignment.csv`", "- `heatmap.svg`", "", "## Notes", "", "This run maps dLLM denoising timesteps to mask ratios. A visible block", "structure in the heatmap would support an interval/expert view similar", "to Fig. 3 in the image diffusion pruning paper.", ]) path.write_text("\n".join(lines) + "\n", encoding="utf-8") def deterministic_mask(length: int, ratio: float, seed: int, device: torch.device) -> torch.Tensor: count = max(1, min(length, int(round(length * ratio)))) generator = torch.Generator(device="cpu") generator.manual_seed(seed) perm = torch.randperm(length, generator=generator)[:count] mask = torch.zeros(length, dtype=torch.bool) mask[perm] = True return mask.to(device) def run_toy(args: argparse.Namespace, ratios: list[float], out_dir: Path) -> dict: torch.manual_seed(args.seed) device = torch.device(args.device) vocab_size = 257 mask_id = 0 model = torch.nn.Sequential( torch.nn.Embedding(vocab_size, args.toy_hidden), torch.nn.LayerNorm(args.toy_hidden), torch.nn.Linear(args.toy_hidden, vocab_size), ).to(device) params = [p for p in model.parameters() if p.requires_grad] samples = [] rng = random.Random(args.seed) for _ in range(args.num_samples): ids = [rng.randint(3, vocab_size - 1) for _ in range(args.block_size)] samples.append(torch.tensor(ids, device=device, dtype=torch.long)) vectors = [] losses_by_ratio = [] for ratio_idx, ratio in enumerate(ratios): model.zero_grad(set_to_none=True) total_loss = 0.0 for sample_idx, target in enumerate(samples): mask = deterministic_mask(args.block_size, ratio, args.seed + ratio_idx * 1009 + sample_idx, device) corrupted = target.clone() corrupted[mask] = mask_id logits = model(corrupted.unsqueeze(0))[0] loss = F.cross_entropy(logits[mask].float(), target[mask]) (loss / len(samples)).backward() total_loss += float(loss.detach().item()) vectors.append(flatten_grads(params, device=torch.device("cpu"))) losses_by_ratio.append(total_loss / len(samples)) matrix = cosine_matrix(vectors) return { "mode": "toy", "dataset": "synthetic", "num_samples": args.num_samples, "block_size": args.block_size, "grad_scope": "all_toy_params", "ratios": ratios, "loss_by_ratio": losses_by_ratio, "cosine_similarity": matrix, "output_dir": str(out_dir), } def load_text_samples(args: argparse.Namespace, tokenizer) -> list[Sample]: from datasets import load_dataset if args.dataset == "gsm8k": ds = load_dataset("openai/gsm8k", "main", split="test") texts = [row["question"] + "\n" + row["answer"] for row in ds] elif args.dataset == "math500": ds = load_dataset("HuggingFaceH4/MATH-500", split="test") texts = [row["problem"] + "\n" + row["answer"] for row in ds] elif args.dataset == "mbpp": ds = load_dataset("google-research-datasets/mbpp", "sanitized", split="test") texts = [row["prompt"] + "\n" + "\n".join(row.get("test_list", [])) for row in ds] else: raise ValueError(f"Unsupported dataset for dflash mode: {args.dataset}") samples = [] for text in texts: ids = tokenizer.encode(text, add_special_tokens=False) min_len = args.prefix_tokens + args.block_size if len(ids) >= min_len: samples.append(Sample(text=text, input_ids=ids[:min_len])) if len(samples) >= args.num_samples: break if len(samples) < args.num_samples: raise RuntimeError(f"Only found {len(samples)} usable samples; requested {args.num_samples}") return samples def select_dflash_params(draft, scope: str) -> list[torch.nn.Parameter]: for param in draft.parameters(): param.requires_grad_(False) selected = [] for name, param in draft.named_parameters(): if scope == "fc" and name.startswith("fc."): selected.append(param) elif scope == "last_layer" and name.startswith(f"layers.{len(draft.layers) - 1}."): selected.append(param) elif scope == "norms" and ("norm" in name): selected.append(param) elif scope == "all_draft": selected.append(param) if not selected: raise RuntimeError(f"No DFlash parameters selected for grad scope '{scope}'") for param in selected: param.requires_grad_(True) return selected def run_dflash(args: argparse.Namespace, ratios: list[float], out_dir: Path) -> dict: from transformers import AutoModelForCausalLM, AutoTokenizer from dflash.model import DFlashDraftModel, extract_context_feature torch.manual_seed(args.seed) device = torch.device(args.device) dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16 if args.dtype == "fp16" else torch.float32 tokenizer = AutoTokenizer.from_pretrained(args.model) target = AutoModelForCausalLM.from_pretrained( args.model, attn_implementation=args.attn_implementation, dtype=dtype, ).to(device).eval() draft = DFlashDraftModel.from_pretrained( args.draft_model, attn_implementation=args.attn_implementation, dtype=dtype, ).to(device).train() for param in target.parameters(): param.requires_grad_(False) selected_params = select_dflash_params(draft, args.grad_scope) mask_id = draft.mask_token_id if mask_id is None: raise RuntimeError("DFlash draft model did not expose mask_token_id") samples = load_text_samples(args, tokenizer) vectors = [] losses_by_ratio = [] selected_param_count = sum(p.numel() for p in selected_params) for ratio_idx, ratio in enumerate(ratios): draft.zero_grad(set_to_none=True) total_loss = 0.0 for sample_idx, sample in enumerate(samples): ids = torch.tensor(sample.input_ids, device=device, dtype=torch.long).unsqueeze(0) prefix = ids[:, : args.prefix_tokens] clean_block = ids[:, args.prefix_tokens : args.prefix_tokens + args.block_size] mask = deterministic_mask(args.block_size, ratio, args.seed + ratio_idx * 1009 + sample_idx, device) noisy_block = clean_block.clone() noisy_block[:, mask] = mask_id with torch.no_grad(): target_out = target( prefix, output_hidden_states=True, use_cache=False, ) target_hidden = extract_context_feature(target_out.hidden_states, draft.target_layer_ids) noise_embedding = target.model.embed_tokens(noisy_block) position_ids = torch.arange(args.prefix_tokens + args.block_size, device=device).unsqueeze(0) hidden = draft( target_hidden=target_hidden, noise_embedding=noise_embedding, position_ids=position_ids, use_cache=False, is_causal=False, ) logits = target.lm_head(hidden) loss = F.cross_entropy(logits[:, mask, :].float().reshape(-1, logits.shape[-1]), clean_block[:, mask].reshape(-1)) (loss / len(samples)).backward() total_loss += float(loss.detach().item()) del ids, prefix, clean_block, noisy_block, target_out, target_hidden, noise_embedding, hidden, logits, loss vectors.append(flatten_grads(selected_params, device=torch.device("cpu"))) losses_by_ratio.append(total_loss / len(samples)) if torch.cuda.is_available(): torch.cuda.empty_cache() matrix = cosine_matrix(vectors) return { "mode": "dflash", "model": args.model, "draft_model": args.draft_model, "dataset": args.dataset, "num_samples": args.num_samples, "block_size": args.block_size, "prefix_tokens": args.prefix_tokens, "grad_scope": args.grad_scope, "selected_param_count": selected_param_count, "ratios": ratios, "loss_by_ratio": losses_by_ratio, "cosine_similarity": matrix, "output_dir": str(out_dir), } def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--mode", choices=["toy", "dflash"], default="toy") parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--ratios", default="0.1:0.1:0.9") parser.add_argument("--num-samples", type=int, default=8) parser.add_argument("--block-size", type=int, default=16) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--device", default="cuda") parser.add_argument("--toy-hidden", type=int, default=64) parser.add_argument("--model", default="/home/l/liyj/shiying/hku_diffusion_dllm/models/qwen3-4b-target") parser.add_argument("--draft-model", default="/home/l/liyj/shiying/hku_diffusion_dllm/models/qwen3-4b-dflash-b16") parser.add_argument("--dataset", choices=["gsm8k", "math500", "mbpp"], default="gsm8k") parser.add_argument("--prefix-tokens", type=int, default=128) parser.add_argument("--grad-scope", choices=["fc", "last_layer", "norms", "all_draft"], default="fc") parser.add_argument("--dtype", choices=["bf16", "fp16", "fp32"], default="bf16") parser.add_argument("--attn-implementation", default="sdpa") args = parser.parse_args() ratios = parse_ratios(args.ratios) out_dir = args.output_dir ensure_dir(out_dir) if args.mode == "toy" and args.device == "cuda" and not torch.cuda.is_available(): args.device = "cpu" if args.mode == "toy": payload = run_toy(args, ratios, out_dir) else: payload = run_dflash(args, ratios, out_dir) payload["updated_at"] = datetime.now(timezone.utc).isoformat() payload["command"] = " ".join(sys.argv) json_path = out_dir / "alignment.json" csv_path = out_dir / "alignment.csv" svg_path = out_dir / "heatmap.svg" md_path = out_dir / "report.md" json_path.write_text(json.dumps(payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") write_csv(csv_path, payload["ratios"], payload["cosine_similarity"]) write_svg(svg_path, payload["ratios"], payload["cosine_similarity"], f"{payload['mode']} mask-ratio gradient alignment") write_report(md_path, payload) print(json.dumps({ "status": "ok", "json": str(json_path), "csv": str(csv_path), "svg": str(svg_path), "report": str(md_path), }, indent=2)) return 0 if __name__ == "__main__": raise SystemExit(main())