Ouzhang's picture
Add files using upload-large-folder tool
d91766b verified
Raw
History Blame Contribute Delete
16.4 kB
#!/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'<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())