#!/usr/bin/env python3 """Streaming adjacent-step dLLM mask-ratio gradient alignment. Use this for very fine ratio grids such as 0.001. It avoids storing hundreds or thousands of full gradient vectors by computing adjacent cosine online and only keeping a configurable downsampled subset for a small pairwise heatmap. """ from __future__ import annotations import argparse import json import sys from datetime import datetime, timezone from pathlib import Path import torch from compute_alignment import ( cosine_matrix, deterministic_mask, ensure_dir, flatten_grads, load_text_samples, parse_ratios, select_dflash_params, write_csv, write_report, write_svg, ) def grad_for_ratio( *, ratio: float, ratio_idx: int, args: argparse.Namespace, samples, tokenizer, target, draft, selected_params, mask_id: int, device: torch.device, ) -> tuple[torch.Tensor, float]: from dflash.model import extract_context_feature import torch.nn.functional as F 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 vec = flatten_grads(selected_params, device=torch.device("cpu")) if torch.cuda.is_available(): torch.cuda.empty_cache() return vec, total_loss / len(samples) def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--ratios", default="0.001:0.001:0.999") parser.add_argument("--num-samples", type=int, default=8) parser.add_argument("--block-size", type=int, default=64) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--device", default="cuda") 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") parser.add_argument("--downsample-every", type=int, default=10) args = parser.parse_args() from transformers import AutoModelForCausalLM, AutoTokenizer from dflash.model import DFlashDraftModel ratios = parse_ratios(args.ratios) out_dir = args.output_dir ensure_dir(out_dir) device = torch.device(args.device) dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16 if args.dtype == "fp16" else torch.float32 torch.manual_seed(args.seed) 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) adjacent = [] losses_by_ratio = [] downsample_ratios = [] downsample_vectors = [] prev_ratio = None prev_vec = None for idx, ratio in enumerate(ratios): vec, avg_loss = grad_for_ratio( ratio=ratio, ratio_idx=idx, args=args, samples=samples, tokenizer=tokenizer, target=target, draft=draft, selected_params=selected_params, mask_id=mask_id, device=device, ) losses_by_ratio.append(avg_loss) if prev_vec is not None: sim = cosine_matrix([prev_vec, vec])[0][1] adjacent.append({"left": prev_ratio, "right": ratio, "mid": (prev_ratio + ratio) / 2, "cosine": sim}) if idx % args.downsample_every == 0 or idx == len(ratios) - 1: downsample_ratios.append(ratio) downsample_vectors.append(vec) prev_ratio = ratio prev_vec = vec downsample_matrix = cosine_matrix(downsample_vectors) selected_param_count = sum(p.numel() for p in selected_params) candidate = min(adjacent, key=lambda x: x["cosine"]) if adjacent else None payload = { "updated_at": datetime.now(timezone.utc).isoformat(), "command": " ".join(sys.argv), "mode": "dflash_streaming_adjacent", "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, "adjacent_cosines": adjacent, "candidate_boundary": None if candidate is None else f"{candidate['left']:g} -> {candidate['right']:g}", "downsample_every": args.downsample_every, "downsampled_ratios": downsample_ratios, "downsampled_cosine_similarity": downsample_matrix, "output_dir": str(out_dir), } (out_dir / "adjacent_alignment.json").write_text(json.dumps(payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") with (out_dir / "adjacent_alignment.csv").open("w", encoding="utf-8") as f: f.write("left,right,mid,cosine\n") for row in adjacent: f.write(f"{row['left']},{row['right']},{row['mid']},{row['cosine']}\n") heat_payload = { "mode": payload["mode"], "dataset": payload["dataset"], "num_samples": payload["num_samples"], "block_size": payload["block_size"], "grad_scope": payload["grad_scope"], "ratios": downsample_ratios, "cosine_similarity": downsample_matrix, "output_dir": str(out_dir), "updated_at": payload["updated_at"], } (out_dir / "downsampled_alignment.json").write_text(json.dumps(heat_payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") write_csv(out_dir / "downsampled_alignment.csv", downsample_ratios, downsample_matrix) write_svg(out_dir / "downsampled_heatmap.svg", downsample_ratios, downsample_matrix, f"downsampled {args.block_size} adjacent run") write_report(out_dir / "downsampled_report.md", heat_payload) lines = [ "# Streaming Adjacent dLLM Alignment", "", f"- Updated: `{payload['updated_at']}`", f"- Dataset: `{args.dataset}`", f"- Samples: `{args.num_samples}`", f"- Block size: `{args.block_size}`", f"- Ratio grid: `{args.ratios}`", f"- Gradient scope: `{args.grad_scope}`", f"- Candidate boundary: `{payload['candidate_boundary']}`", "", "## Adjacent Cosines", "", "| Step pair | Cosine similarity |", "| --- | ---: |", ] for row in adjacent: lines.append(f"| {row['left']:g} -> {row['right']:g} | {row['cosine']:.3f} |") (out_dir / "adjacent_report.md").write_text("\n".join(lines) + "\n", encoding="utf-8") print(json.dumps({ "status": "ok", "adjacent_json": str(out_dir / "adjacent_alignment.json"), "adjacent_report": str(out_dir / "adjacent_report.md"), "downsampled_json": str(out_dir / "downsampled_alignment.json"), "downsampled_heatmap": str(out_dir / "downsampled_heatmap.svg"), "candidate_boundary": payload["candidate_boundary"], }, indent=2)) return 0 if __name__ == "__main__": raise SystemExit(main())