| |
| """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'<svg xmlns="http://www.w3.org/2000/svg" width="{width}" height="{height}" viewBox="0 0 {width} {height}">', |
| '<rect width="100%" height="100%" fill="#ffffff"/>', |
| f'<text x="{label_w}" y="24" font-family="Arial, sans-serif" font-size="16" font-weight="700">{html.escape(title)}</text>', |
| ] |
| x0 = label_w |
| y0 = title_h |
| for i, ratio in enumerate(ratios): |
| x = x0 + i * cell + cell / 2 |
| parts.append(f'<text x="{x}" y="{y0 - 8}" text-anchor="middle" font-family="Arial, sans-serif" font-size="10">{ratio:g}</text>') |
| y = y0 + i * cell + cell / 2 + 4 |
| parts.append(f'<text x="{x0 - 10}" y="{y}" text-anchor="end" font-family="Arial, sans-serif" font-size="10">{ratio:g}</text>') |
| 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'<rect x="{x}" y="{y}" width="{cell}" height="{cell}" fill="{color_for(value)}" stroke="#ffffff" stroke-width="1"/>') |
| parts.append(f'<text x="{x + cell / 2}" y="{y + cell / 2 + 4}" text-anchor="middle" font-family="Arial, sans-serif" font-size="9" fill="#111111">{value:.2f}</text>') |
| ly = y0 + n * cell + 18 |
| parts.append(f'<text x="{x0}" y="{ly}" font-family="Arial, sans-serif" font-size="11">blue/yellow: negative, white: 0, red: positive cosine similarity</text>') |
| parts.append("</svg>") |
| 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()) |
|
|